Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,17 @@ public abstract interface class sk/ainet/compile/target/OpGranularityPolicy {
public abstract fun keepFused (Ljava/lang/String;)Z
}

public final class sk/ainet/lang/graph/BlockOrderLayout : sk/ainet/lang/graph/Layout {
public fun <init> (Ljava/lang/String;)V
public final fun component1 ()Ljava/lang/String;
public final fun copy (Ljava/lang/String;)Lsk/ainet/lang/graph/BlockOrderLayout;
public static synthetic fun copy$default (Lsk/ainet/lang/graph/BlockOrderLayout;Ljava/lang/String;ILjava/lang/Object;)Lsk/ainet/lang/graph/BlockOrderLayout;
public fun equals (Ljava/lang/Object;)Z
public final fun getBlockOrder ()Ljava/lang/String;
public fun hashCode ()I
public fun toString ()Ljava/lang/String;
}

public abstract interface class sk/ainet/lang/graph/ComputeGraph {
public abstract fun addEdge (Lsk/ainet/lang/graph/GraphEdge;)Lsk/ainet/lang/graph/GraphEdge;
public abstract fun addNode (Lsk/ainet/lang/graph/GraphNode;)Lsk/ainet/lang/graph/GraphNode;
Expand Down Expand Up @@ -319,6 +330,8 @@ public final class sk/ainet/lang/graph/OutputDesignatedGraph : sk/ainet/lang/gra
}

public final class sk/ainet/lang/graph/ResolvedComputeGraph {
public static final field BACKEND_ASSIGNMENT_METADATA_KEY Ljava/lang/String;
public static final field Companion Lsk/ainet/lang/graph/ResolvedComputeGraph$Companion;
public fun <init> (Lsk/ainet/lang/graph/ComputeGraph;)V
public final fun backendAssignment (Ljava/lang/String;)Ljava/lang/String;
public final fun getDelegate ()Lsk/ainet/lang/graph/ComputeGraph;
Expand All @@ -329,6 +342,9 @@ public final class sk/ainet/lang/graph/ResolvedComputeGraph {
public final fun validate ()Lsk/ainet/lang/graph/ResolvedGraphValidation;
}

public final class sk/ainet/lang/graph/ResolvedComputeGraph$Companion {
}

public final class sk/ainet/lang/graph/ResolvedGraphValidation {
public fun <init> (ZLjava/util/List;)V
public final fun component1 ()Z
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import sk.ainet.lang.tensor.withRequiresGrad
import sk.ainet.lang.tensor.ops.Operation
import sk.ainet.lang.tensor.ops.TensorSpec
import sk.ainet.lang.tensor.ops.withTensorEncoding
import sk.ainet.lang.tensor.ops.withBlockOrder
import sk.ainet.lang.tensor.ops.withTensorId
import sk.ainet.lang.types.DType
import sk.ainet.lang.trace.OpTrace
Expand Down Expand Up @@ -82,15 +83,15 @@ public open class DefaultExecutionTape(
name = ref.id,
shape = inputShapes?.getOrNull(i) ?: ref.shape.dimensions.toList(),
dtype = inputDTypes?.getOrNull(i) ?: ref.dtype.name,
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId)
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId).withBlockOrder(ref.blockOrder)
}
val outputs = List(trace.outputs.size) { i ->
val ref = trace.outputs[i]
TensorSpec(
name = ref.id,
shape = outputShapes?.getOrNull(i) ?: ref.shape.dimensions.toList(),
dtype = outputDTypes?.getOrNull(i) ?: ref.dtype.name,
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId)
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId).withBlockOrder(ref.blockOrder)
}

val op = object : sk.ainet.lang.tensor.ops.Operation {
Expand Down Expand Up @@ -128,7 +129,7 @@ public open class DefaultExecutionTape(
shape = tensor.shape.dimensions.toList(),
dtype = tensor.dtype.toString(),
requiresGrad = tensor.requiresGrad
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId)
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId).withBlockOrder(ref.blockOrder)
}

val outputSpecs = outputs.map { tensor ->
Expand All @@ -138,7 +139,7 @@ public open class DefaultExecutionTape(
shape = tensor.shape.dimensions.toList(),
dtype = tensor.dtype.toString(),
requiresGrad = tensor.requiresGrad
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId)
).withTensorEncoding(ref.encoding).withTensorId(ref.tensorId).withBlockOrder(ref.blockOrder)
}

val recordedOp = RecordedOperation(
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
package sk.ainet.lang.graph

import sk.ainet.lang.tensor.ops.blockOrder
import sk.ainet.lang.types.BF16
import sk.ainet.lang.types.DType
import sk.ainet.lang.types.FP16
Expand Down Expand Up @@ -48,16 +49,26 @@ public class ResolvedComputeGraph(public val delegate: ComputeGraph) {
}

/**
* Placeholder for the resolved memory layout. Returns `null`
* today — populated by future layout-planning passes.
* The resolved memory layout for [edgeId], derived from the block order a layout pass (or
* the tape) stamped on the edge's spec (#1180) — `null` where no decision was made.
*/
public fun resolvedLayout(edgeId: String): Layout? = null
public fun resolvedLayout(edgeId: String): Layout? {
val edge = edges.firstOrNull { it.id == edgeId } ?: return null
val order = edge.tensorSpec.blockOrder ?: return null
return BlockOrderLayout(order)
}

/**
* Placeholder for the backend assignment. Returns `null` today
* — populated by future multi-backend scheduling.
* The backend a layout/scheduling pass assigned [nodeId] to (#1180), or `null` where no
* pass made a decision. Read from [BACKEND_ASSIGNMENT_METADATA_KEY] node metadata.
*/
public fun backendAssignment(nodeId: String): String? = null
public fun backendAssignment(nodeId: String): String? =
nodes.firstOrNull { it.id == nodeId }?.metadata?.get(BACKEND_ASSIGNMENT_METADATA_KEY) as? String

public companion object {
/** Node metadata key a pass stamps with the target it decided for. */
public const val BACKEND_ASSIGNMENT_METADATA_KEY: String = "backendAssignment"
}

/**
* Precondition check for the resolved-DAG contract:
Expand Down Expand Up @@ -124,3 +135,9 @@ public data class ResolvedGraphValidation(
}
}
}

/**
* The one layout fact carried today (#1180): which way a packed weight's blocks run, as the
* stable string the whole carriage lane uses (`ROW_MAJOR` / `INPUT_BLOCK_MAJOR`).
*/
public data class BlockOrderLayout(public val blockOrder: String) : Layout
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,10 @@ import sk.ainet.lang.tensor.ops.Operation
import sk.ainet.lang.tensor.ops.TensorSpec
import sk.ainet.lang.tensor.ops.ValidationResult
import sk.ainet.lang.tensor.ops.inferTensorEncoding
import sk.ainet.lang.tensor.ops.blockOrder
import sk.ainet.lang.tensor.ops.tensorId
import sk.ainet.lang.tensor.ops.withTensorEncoding
import sk.ainet.lang.tensor.ops.withBlockOrder
import sk.ainet.lang.tensor.ops.withTensorId

/**
Expand Down Expand Up @@ -294,6 +296,7 @@ public class TraceToGraphBuilder(
dtype = weightDtype
).withTensorEncoding(encoding)
.withTensorId(refs.firstNotNullOfOrNull { it.spec.tensorId })
.withBlockOrder(refs.firstNotNullOfOrNull { it.spec.blockOrder })
syntheticNode = GraphNode(
id = nodeId,
operation = op,
Expand Down Expand Up @@ -351,6 +354,7 @@ public class TraceToGraphBuilder(
TensorSpec(name = name, shape = shape, dtype = dtype)
.withTensorEncoding(effectiveInputs[i].encoding)
.withTensorId(effectiveInputs[i].tensorId)
.withBlockOrder(effectiveInputs[i].blockOrder)
}
}

Expand All @@ -365,6 +369,7 @@ public class TraceToGraphBuilder(
TensorSpec(name = name, shape = shape, dtype = dtype)
.withTensorEncoding(trace.outputs[i].encoding)
.withTensorId(trace.outputs[i].tensorId)
.withBlockOrder(trace.outputs[i].blockOrder)
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -134,14 +134,17 @@ public final class sk/ainet/compile/hlo/DeadCodeEliminationPass : sk/ainet/compi
}

public final class sk/ainet/compile/hlo/ExternalParameterRef {
public fun <init> (Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;)V
public fun <init> (Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;Ljava/lang/String;)V
public synthetic fun <init> (Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
public final fun component1 ()Ljava/lang/String;
public final fun component2 ()Ljava/lang/String;
public final fun component3 ()Lsk/ainet/lang/tensor/storage/TensorEncoding;
public final fun component4 ()Lsk/ainet/lang/tensor/storage/BufferHandle;
public final fun copy (Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;)Lsk/ainet/compile/hlo/ExternalParameterRef;
public static synthetic fun copy$default (Lsk/ainet/compile/hlo/ExternalParameterRef;Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;ILjava/lang/Object;)Lsk/ainet/compile/hlo/ExternalParameterRef;
public final fun component5 ()Ljava/lang/String;
public final fun copy (Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;Ljava/lang/String;)Lsk/ainet/compile/hlo/ExternalParameterRef;
public static synthetic fun copy$default (Lsk/ainet/compile/hlo/ExternalParameterRef;Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;Ljava/lang/String;ILjava/lang/Object;)Lsk/ainet/compile/hlo/ExternalParameterRef;
public fun equals (Ljava/lang/Object;)Z
public final fun getBlockOrder ()Ljava/lang/String;
public final fun getEncoding ()Lsk/ainet/lang/tensor/storage/TensorEncoding;
public final fun getKey ()Ljava/lang/String;
public final fun getScope ()Ljava/lang/String;
Expand Down Expand Up @@ -490,8 +493,8 @@ public final class sk/ainet/compile/hlo/examples/ValidationDemonstrationResult {

public final class sk/ainet/compile/hlo/generate/HloGenerator {
public static final field INSTANCE Lsk/ainet/compile/hlo/generate/HloGenerator;
public final fun generate (Lsk/ainet/lang/model/Model;Lsk/ainet/lang/tensor/Tensor;Ljava/lang/String;Lkotlin/coroutines/Continuation;)Ljava/lang/Object;
public static synthetic fun generate$default (Lsk/ainet/compile/hlo/generate/HloGenerator;Lsk/ainet/lang/model/Model;Lsk/ainet/lang/tensor/Tensor;Ljava/lang/String;Lkotlin/coroutines/Continuation;ILjava/lang/Object;)Ljava/lang/Object;
public final fun generate (Lsk/ainet/lang/model/Model;Lsk/ainet/lang/tensor/Tensor;Ljava/lang/String;Ljava/lang/String;Lkotlin/coroutines/Continuation;)Ljava/lang/Object;
public static synthetic fun generate$default (Lsk/ainet/compile/hlo/generate/HloGenerator;Lsk/ainet/lang/model/Model;Lsk/ainet/lang/tensor/Tensor;Ljava/lang/String;Ljava/lang/String;Lkotlin/coroutines/Continuation;ILjava/lang/Object;)Ljava/lang/Object;
}

public final class sk/ainet/compile/hlo/generate/HloGeneratorJvmKt {
Expand Down
1 change: 1 addition & 0 deletions skainet-compile/skainet-compile-hlo/build.gradle.kts
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@ kotlin {
api(project(":skainet-lang:skainet-lang-models"))
api(project(":skainet-compile:skainet-compile-core"))
api(project(":skainet-compile:skainet-compile-dag"))
api(project(":skainet-compile:skainet-compile-opt"))
}

commonTest.dependencies {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -92,5 +92,13 @@ public data class ExternalParameterRef(
val scope: String,
val key: String,
val encoding: TensorEncoding,
val source: BufferHandle
val source: BufferHandle,
/**
* Physical block order of the packed bytes as a stable string (`ROW_MAJOR` /
* `INPUT_BLOCK_MAJOR`), or `null` when unknown / not packed (#1179). Carried so a
* kernel-feed-ordered weight can be blitted into the `.irpa` verbatim *and say so* —
* a string rather than the storage-model enum, by the same no-import rule the
* `TensorSpec` metadata follows.
*/
val blockOrder: String? = null,
)
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package sk.ainet.compile.hlo
import sk.ainet.lang.graph.ComputeGraph
import sk.ainet.lang.graph.GraphNode
import sk.ainet.lang.tensor.ops.TensorSpec
import sk.ainet.lang.tensor.ops.blockOrder
import sk.ainet.lang.tensor.ops.tensorEncoding
import sk.ainet.lang.tensor.storage.TensorEncoding

Expand Down Expand Up @@ -61,11 +62,11 @@ public class StableHloConverter @kotlin.jvm.JvmOverloads constructor(
val functionSignature = buildFunctionSignature(inputNodes, outputSpecs, functionName)

// Collect every TensorSpec with a non-null tensorEncoding into a
// single name -> encoding map. Emitting this as a structured
// single name -> structural-facts map. Emitting this as a structured
// MLIR attribute on the module header lets downstream tools
// enumerate every encoded tensor via one attribute lookup
// instead of string-matching against scattered comments.
val tensorEncodings = collectTensorEncodings(topo)
// enumerate every encoded tensor via one attribute lookup —
// block sizes and bit widths as integers, not display names (#1179).
val structuralLayouts = collectStructuralLayouts(topo)

// Process nodes first, then assemble the final content.
// Converters populate two buffers on the context — op emissions
Expand All @@ -78,11 +79,11 @@ public class StableHloConverter @kotlin.jvm.JvmOverloads constructor(
processNodes(topo, context)
generateReturnStatement(outputNodes, context)

val moduleHeader = if (tensorEncodings.isNotEmpty()) {
val dictEntries = tensorEncodings.entries
val moduleHeader = if (structuralLayouts.isNotEmpty()) {
val layoutEntries = structuralLayouts.entries
.sortedBy { it.key }
.joinToString(", ") { (name, encoding) -> "$name = \"${encoding.name}\"" }
"module attributes {skainet.tensor_encodings = {$dictEntries}} {"
.joinToString(", ") { (name, attr) -> "$name = $attr" }
"module attributes {skainet.tensor_layouts = {$layoutEntries}} {"
} else {
"module {"
}
Expand Down Expand Up @@ -248,21 +249,47 @@ public class StableHloConverter @kotlin.jvm.JvmOverloads constructor(
* extension and does not exist in Kotlin common-stdlib for WasmJS
* / JS / Native targets.
*/
private fun collectTensorEncodings(nodes: List<GraphNode>): Map<String, TensorEncoding> {
val result = linkedMapOf<String, TensorEncoding>()
/**
* Structural per-tensor storage facts for the module header (#1179): unlike
* `skainet.tensor_encodings` (display names, kept for compatibility), these entries are
* machine-readable — block element counts, block bytes, bit widths, block order — so a
* downstream consumer can size and address packed weights without a lookup table of names.
*/
private fun collectStructuralLayouts(nodes: List<GraphNode>): Map<String, String> {
val result = linkedMapOf<String, String>()
for (node in nodes) {
for (spec in node.outputs) {
val encoding = spec.tensorEncoding ?: continue
if (spec.name !in result) result[spec.name] = encoding
}
for (spec in node.inputs) {
for (spec in node.outputs + node.inputs) {
val encoding = spec.tensorEncoding ?: continue
if (spec.name !in result) result[spec.name] = encoding
if (spec.name in result) continue
result[spec.name] = structuralAttr(encoding, spec.blockOrder)
}
}
return result
}

private fun structuralAttr(encoding: TensorEncoding, blockOrder: String?): String {
val facts = buildList {
add("kind = \"" + encoding.name + "\"")
when (encoding) {
is TensorEncoding.Dense -> add("bytes_per_element = " + encoding.bytesPerElement)
TensorEncoding.Q4_K -> { add("block_elems = " + TensorEncoding.Q4_K.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.Q4_K.BYTES_PER_BLOCK) }
TensorEncoding.Q5_K -> { add("block_elems = " + TensorEncoding.Q5_K.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.Q5_K.BYTES_PER_BLOCK) }
TensorEncoding.Q6_K -> { add("block_elems = " + TensorEncoding.Q6_K.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.Q6_K.BYTES_PER_BLOCK) }
TensorEncoding.Q4_0 -> { add("block_elems = " + TensorEncoding.Q4_0.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.Q4_0.BYTES_PER_BLOCK) }
TensorEncoding.Q8_0 -> { add("block_elems = " + TensorEncoding.Q8_0.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.Q8_0.BYTES_PER_BLOCK) }
TensorEncoding.Q5_0 -> { add("block_elems = " + TensorEncoding.Q5_0.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.Q5_0.BYTES_PER_BLOCK) }
TensorEncoding.Q5_1 -> { add("block_elems = " + TensorEncoding.Q5_1.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.Q5_1.BYTES_PER_BLOCK) }
TensorEncoding.TQ1_0 -> { add("block_elems = " + TensorEncoding.TQ1_0.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.TQ1_0.BYTES_PER_BLOCK) }
TensorEncoding.TQ2_0 -> { add("block_elems = " + TensorEncoding.TQ2_0.BLOCK_SIZE); add("block_bytes = " + TensorEncoding.TQ2_0.BYTES_PER_BLOCK) }
is TensorEncoding.TurboQuantPolar -> { add("bits = " + encoding.bitsPerElement); add("block_elems = " + encoding.blockSize) }
TensorEncoding.TernaryPacked -> add("bits = 2")
else -> { /* kind alone: unknown structure is stated, not invented */ }
}
if (blockOrder != null) add("block_order = \"" + blockOrder + "\"")
}
return "{" + facts.joinToString(", ") + "}"
}

/**
* Determine output specifications from output nodes
*/
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import sk.ainet.compile.hlo.elementCountFromShape
import sk.ainet.compile.hlo.floatArrayToLittleEndianBytes
import sk.ainet.compile.hlo.numberListToLittleEndianBytes
import sk.ainet.lang.graph.GraphNode
import sk.ainet.lang.tensor.ops.blockOrder
import sk.ainet.lang.tensor.ops.TensorSpec
import sk.ainet.lang.tensor.ops.tensorEncoding
import sk.ainet.lang.tensor.storage.BufferHandle
Expand Down Expand Up @@ -439,7 +440,8 @@ public class ConstantOperationsConverter : StableHloOperationConverter {
scope = scope,
key = key,
encoding = encoding,
source = BufferHandle.Owned(bytes)
source = BufferHandle.Owned(bytes),
blockOrder = outputSpec.blockOrder,
)
)
// Bind the global to an archive entry via the IREE flow-dialect
Expand Down Expand Up @@ -514,7 +516,8 @@ public class ConstantOperationsConverter : StableHloOperationConverter {
scope = scope,
key = key,
encoding = encoding,
source = BufferHandle.Owned(bytes)
source = BufferHandle.Owned(bytes),
blockOrder = outputSpec.blockOrder,
)
)
context.emitModuleDeclaration(
Expand Down
Loading
Loading