From ed537158904a49627df4ae933d2dcaf2670b02c6 Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Wed, 26 Aug 2026 19:01:54 +0200 Subject: [PATCH 1/2] feat(hlo): structural layout attrs at the StableHLO and .irpa boundary MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Carriage slice 2 (#1147): the export boundary carries structure, not display names. The module attribute 'skainet.tensor_layouts' is machine-readable — block element counts, block bytes, bit widths, and block order — straight off the encoding objects' own constants, so a downstream consumer can size and address packed weights without a lookup table. TurboQuantPolar(4b, 128) used to degrade to a name whose block size was unrecoverable; now it is two integers in the header. The former 'skainet.tensor_encodings' names dictionary is REMOVED, not duplicated: an org-wide code search shows nothing outside this repo ever read it (its only consumers were the emitter and its own test), and 'kind' in the structural dictionary carries the name. The no-encodings case still emits a bare 'module {' header (test kept). Block order rides the whole way as a stable string (never the storage-model enum — the no-import rule): TensorRef.blockOrder is captured from PackedBlockStorage at the tape, flows through every spec site slice 1 opened, and lands both in the module attribute and on ExternalParameterRef, so a kernel-feed-ordered weight can be blitted into the .irpa verbatim and say so. StructuralLayoutsModuleAttributeTest pins the header facts for a feed-ordered Q4_K (256/144 + order) and a TurboQuant weight (bits + block), the absence of the removed dictionary, and the bare header. Closes #1179. Co-Authored-By: Claude Fable 5 --- .../ainet/lang/graph/DefaultExecutionTape.kt | 9 +- .../ainet/lang/trace/TraceToGraphBuilder.kt | 5 + .../api/jvm/skainet-compile-hlo.api | 9 +- .../compile/hlo/ConstantMaterialization.kt | 10 +- .../ainet/compile/hlo/StableHloConverter.kt | 59 +++++-- .../converters/ConstantOperationsConverter.kt | 7 +- .../hlo/ConstantMaterializationPolicyTest.kt | 2 +- .../StructuralLayoutsModuleAttributeTest.kt | 123 ++++++++++++++ .../hlo/TensorEncodingsModuleAttributeTest.kt | 158 ------------------ .../api/jvm/skainet-lang-core.api | 13 +- .../lang/tensor/ops/TensorSpecIdentity.kt | 21 +++ .../kotlin/sk/ainet/lang/trace/TensorRef.kt | 2 + .../sk/ainet/lang/trace/TraceSession.kt | 1 + 13 files changed, 230 insertions(+), 189 deletions(-) create mode 100644 skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/StructuralLayoutsModuleAttributeTest.kt delete mode 100644 skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/TensorEncodingsModuleAttributeTest.kt diff --git a/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/DefaultExecutionTape.kt b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/DefaultExecutionTape.kt index 11deb921a..ff69a5eea 100644 --- a/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/DefaultExecutionTape.kt +++ b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/DefaultExecutionTape.kt @@ -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 @@ -82,7 +83,7 @@ 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] @@ -90,7 +91,7 @@ public open class DefaultExecutionTape( 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 { @@ -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 -> @@ -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( diff --git a/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/trace/TraceToGraphBuilder.kt b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/trace/TraceToGraphBuilder.kt index 719ee1743..53b920daa 100644 --- a/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/trace/TraceToGraphBuilder.kt +++ b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/trace/TraceToGraphBuilder.kt @@ -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 /** @@ -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, @@ -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) } } @@ -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) } } diff --git a/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api b/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api index 921c5eebd..b86eff6a3 100644 --- a/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api +++ b/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api @@ -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 (Ljava/lang/String;Ljava/lang/String;Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/tensor/storage/BufferHandle;)V + public fun (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 (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; diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/ConstantMaterialization.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/ConstantMaterialization.kt index 7d86e8a9c..c12982215 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/ConstantMaterialization.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/ConstantMaterialization.kt @@ -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, ) diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/StableHloConverter.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/StableHloConverter.kt index 2415ad425..9681b4f08 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/StableHloConverter.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/StableHloConverter.kt @@ -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 @@ -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 @@ -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 {" } @@ -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): Map { - val result = linkedMapOf() + /** + * 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): Map { + val result = linkedMapOf() 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 */ diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ConstantOperationsConverter.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ConstantOperationsConverter.kt index 5c44df61f..3dac1ce8a 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ConstantOperationsConverter.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ConstantOperationsConverter.kt @@ -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 @@ -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 @@ -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( diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ConstantMaterializationPolicyTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ConstantMaterializationPolicyTest.kt index 8ccc9fb1c..d331563df 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ConstantMaterializationPolicyTest.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ConstantMaterializationPolicyTest.kt @@ -131,7 +131,7 @@ class ConstantMaterializationPolicyTest { @Test fun testModuleAttrsHeaderStillEmittedAboveUtilGlobal() { // When a tensor carries a tensorEncoding we already emit a - // `module attributes { skainet.tensor_encodings = {...} } {` + // `module attributes { skainet.tensor_layouts = {...} } {` // header. The new util.global decls must slot in AFTER that // header, before func.func — otherwise IREE's parser chokes. // This test is aspirational for now: we only assert both diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/StructuralLayoutsModuleAttributeTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/StructuralLayoutsModuleAttributeTest.kt new file mode 100644 index 000000000..5cea5558e --- /dev/null +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/StructuralLayoutsModuleAttributeTest.kt @@ -0,0 +1,123 @@ +package sk.ainet.compile.hlo + +import sk.ainet.lang.graph.DefaultComputeGraph +import sk.ainet.lang.graph.GraphEdge +import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.tensor.ops.AddOperation +import sk.ainet.lang.tensor.ops.InputOperation +import sk.ainet.lang.tensor.ops.TensorSpec +import sk.ainet.lang.tensor.ops.withBlockOrder +import sk.ainet.lang.tensor.ops.withTensorEncoding +import sk.ainet.lang.tensor.storage.TensorEncoding +import sk.ainet.lang.types.DType +import kotlin.test.Test +import kotlin.test.assertTrue + +/** + * #1179: the module header carries *structure*, not display names. `skainet.tensor_layouts` holds + * the machine-readable facts — block element counts, block bytes, bit widths, block order — that a + * downstream consumer needs to size and address packed weights; `kind` carries the name, so the + * former `skainet.tensor_encodings` names dictionary (which nothing outside this repo ever read) + * is gone rather than duplicated. `TurboQuantPolar(4b, 128)` used to degrade to a name whose block + * size was unrecoverable; now it is two integers in the header. + */ +class StructuralLayoutsModuleAttributeTest { + + @Test + fun structural_facts_reach_the_module_header() { + val graph = DefaultComputeGraph() + + val inputA = GraphNode( + id = "a", + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(TensorSpec("a", listOf(1, 4), "FP32")), + ) + // A feed-ordered Q4_K weight: the block order must ride along as a string. + val q4Spec = TensorSpec("w_q4", listOf(1, 4), "FP32") + .withTensorEncoding(TensorEncoding.Q4_K) + .withBlockOrder("INPUT_BLOCK_MAJOR") + val q4Node = GraphNode( + id = "w_q4", + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(q4Spec), + ) + // A TurboQuant weight: bits and block size must be recoverable from the header. + val tqSpec = TensorSpec("w_tq", listOf(1, 4), "FP32") + .withTensorEncoding(TensorEncoding.TurboQuantPolar(bitsPerElement = 4, blockSize = 128)) + val tqNode = GraphNode( + id = "w_tq", + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(tqSpec), + ) + val add1 = GraphNode( + id = "add1", + operation = AddOperation(), + inputs = listOf(inputA.outputs[0], q4Spec), + outputs = listOf(TensorSpec("sum1", listOf(1, 4), "FP32")), + ) + val add2 = GraphNode( + id = "add2", + operation = AddOperation(), + inputs = listOf(add1.outputs[0], tqSpec), + outputs = listOf(TensorSpec("sum2", listOf(1, 4), "FP32")), + ) + + graph.addNode(inputA); graph.addNode(q4Node); graph.addNode(tqNode) + graph.addNode(add1); graph.addNode(add2) + graph.addEdge(GraphEdge("e1", inputA, add1, 0, 0, inputA.outputs[0])) + graph.addEdge(GraphEdge("e2", q4Node, add1, 0, 1, q4Spec)) + graph.addEdge(GraphEdge("e3", add1, add2, 0, 0, add1.outputs[0])) + graph.addEdge(GraphEdge("e4", tqNode, add2, 0, 1, tqSpec)) + + val mlir = toStableHlo(graph, "structural_chain").content + + assertTrue(mlir.contains("skainet.tensor_layouts"), "structural dictionary must be emitted:\n$mlir") + assertTrue( + mlir.contains("w_q4 = {kind = \"Q4_K\", block_elems = 256, block_bytes = 144, block_order = \"INPUT_BLOCK_MAJOR\"}"), + "Q4_K facts + block order must be machine-readable:\n$mlir", + ) + assertTrue( + mlir.contains("w_tq = {kind = \"TurboQuant-Polar-4b\", bits = 4, block_elems = 128}") + || mlir.contains("bits = 4, block_elems = 128"), + "TurboQuant bits and block size must be recoverable:\n$mlir", + ) + // the redundant names dictionary is gone, not duplicated + assertTrue(!mlir.contains("skainet.tensor_encodings"), "legacy names dictionary must not be emitted:\n$mlir") + } + + @Test + fun dense_graph_keeps_bare_module_header() { + // A graph with no encoding metadata must emit the bare `module {` header with no + // `attributes` block. A null tensorEncoding is the unknown / not-carried state — not + // Dense — and the emitter must stay silent. + val graph = DefaultComputeGraph() + val a = GraphNode( + id = "a", + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(TensorSpec("a", listOf(1, 4), "FP32")), + ) + val b = GraphNode( + id = "b", + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(TensorSpec("b", listOf(1, 4), "FP32")), + ) + val add = GraphNode( + id = "add", + operation = AddOperation(), + inputs = listOf(a.outputs[0], b.outputs[0]), + outputs = listOf(TensorSpec("sum", listOf(1, 4), "FP32")), + ) + graph.addNode(a); graph.addNode(b); graph.addNode(add) + graph.addEdge(GraphEdge("e1", a, add, 0, 0, a.outputs[0])) + graph.addEdge(GraphEdge("e2", b, add, 0, 1, b.outputs[0])) + + val mlir = toStableHlo(graph, "dense_chain").content + assertTrue(mlir.contains("module {"), "bare module header expected:\n$mlir") + assertTrue(!mlir.contains("module attributes"), "no attributes block for a dense graph:\n$mlir") + } +} diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/TensorEncodingsModuleAttributeTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/TensorEncodingsModuleAttributeTest.kt deleted file mode 100644 index 634381ed6..000000000 --- a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/TensorEncodingsModuleAttributeTest.kt +++ /dev/null @@ -1,158 +0,0 @@ -package sk.ainet.compile.hlo - -import sk.ainet.lang.graph.DefaultComputeGraph -import sk.ainet.lang.graph.GraphEdge -import sk.ainet.lang.graph.GraphNode -import sk.ainet.lang.tensor.ops.AddOperation -import sk.ainet.lang.tensor.ops.InputOperation -import sk.ainet.lang.tensor.ops.TensorSpec -import sk.ainet.lang.tensor.ops.withTensorEncoding -import sk.ainet.lang.tensor.storage.TensorEncoding -import sk.ainet.lang.types.DType -import kotlin.test.Test -import kotlin.test.assertFalse -import kotlin.test.assertTrue - -/** - * Covers the structured module-level attribute emission for #477: - * every TensorSpec flowing through the graph with a non-null - * tensorEncoding must appear in a single `skainet.tensor_encodings` - * dictionary on the emitted `module attributes { ... }` header, so - * downstream tools can read it with one attribute lookup instead of - * string-matching against scattered comments. - */ -class TensorEncodingsModuleAttributeTest { - - @Test - fun encoded_weights_produce_module_attributes_block() { - val graph = DefaultComputeGraph() - - val inputA = GraphNode( - id = "a", - operation = InputOperation(), - inputs = emptyList(), - outputs = listOf(TensorSpec("a", listOf(1, 4), "FP32")) - ) - - // Two weight inputs with distinct encodings, exactly the shape - // TraceToGraphBuilder.finalize() produces post-#469 when a session - // resolves quantized weights. - val q8Spec = TensorSpec("w_q8", listOf(1, 4), "FP32") - .withTensorEncoding(TensorEncoding.Q8_0) - val q8Node = GraphNode( - id = "w_q8", - operation = InputOperation(), - inputs = emptyList(), - outputs = listOf(q8Spec) - ) - - val q4Spec = TensorSpec("w_q4", listOf(1, 4), "FP32") - .withTensorEncoding(TensorEncoding.Q4_K) - val q4Node = GraphNode( - id = "w_q4", - operation = InputOperation(), - inputs = emptyList(), - outputs = listOf(q4Spec) - ) - - val add1 = GraphNode( - id = "add1", - operation = AddOperation(), - inputs = listOf(TensorSpec("a", listOf(1, 4), "FP32"), q8Spec), - outputs = listOf(TensorSpec("sum1", listOf(1, 4), "FP32")) - ) - val add2 = GraphNode( - id = "add2", - operation = AddOperation(), - inputs = listOf(TensorSpec("sum1", listOf(1, 4), "FP32"), q4Spec), - outputs = listOf(TensorSpec("sum2", listOf(1, 4), "FP32")) - ) - - graph.addNode(inputA) - graph.addNode(q8Node) - graph.addNode(q4Node) - graph.addNode(add1) - graph.addNode(add2) - graph.addEdge(GraphEdge("e1", inputA, add1, 0, 0, inputA.outputs[0])) - graph.addEdge(GraphEdge("e2", q8Node, add1, 0, 1, q8Spec)) - graph.addEdge(GraphEdge("e3", add1, add2, 0, 0, add1.outputs[0])) - graph.addEdge(GraphEdge("e4", q4Node, add2, 0, 1, q4Spec)) - - val mlir = toStableHlo(graph, "quant_chain").content - println("[DEBUG_LOG] module-attribute export:\n$mlir") - - // The emitted module header must carry a structured attribute - // enumerating every encoded tensor in one place. - assertTrue( - mlir.contains("module attributes"), - "module header must be emitted with `module attributes { ... }` when encodings are present" - ) - assertTrue( - mlir.contains("skainet.tensor_encodings"), - "module attributes must include the `skainet.tensor_encodings` dictionary" - ) - - // Both encoded tensors must appear in the dictionary by name, - // each mapped to its TensorEncoding.name. - assertTrue( - mlir.contains("w_q8 = \"Q8_0\""), - "dictionary must map `w_q8` to `\"Q8_0\"`" - ) - assertTrue( - mlir.contains("w_q4 = \"Q4_K\""), - "dictionary must map `w_q4` to `\"Q4_K\"`" - ) - } - - @Test - fun dense_graph_keeps_bare_module_header() { - // A graph with no encoding metadata must emit the bare - // `module {` header with no `attributes` block. A `null` - // tensorEncoding is the unknown / not-carried state — not - // Dense — and the emitter must stay silent. - val graph = DefaultComputeGraph() - - val inputA = GraphNode( - id = "a", - operation = InputOperation(), - inputs = emptyList(), - outputs = listOf(TensorSpec("a", listOf(1, 4), "FP32")) - ) - val inputB = GraphNode( - id = "b", - operation = InputOperation(), - inputs = emptyList(), - outputs = listOf(TensorSpec("b", listOf(1, 4), "FP32")) - ) - val add = GraphNode( - id = "add1", - operation = AddOperation(), - inputs = listOf( - TensorSpec("a", listOf(1, 4), "FP32"), - TensorSpec("b", listOf(1, 4), "FP32") - ), - outputs = listOf(TensorSpec("c", listOf(1, 4), "FP32")) - ) - - graph.addNode(inputA) - graph.addNode(inputB) - graph.addNode(add) - graph.addEdge(GraphEdge("e1", inputA, add, 0, 0, inputA.outputs[0])) - graph.addEdge(GraphEdge("e2", inputB, add, 0, 1, inputB.outputs[0])) - - val mlir = toStableHlo(graph, "dense_add").content - - assertFalse( - mlir.contains("module attributes"), - "dense graph must not emit a `module attributes` block" - ) - assertFalse( - mlir.contains("skainet.tensor_encodings"), - "dense graph must not emit the `skainet.tensor_encodings` dictionary" - ) - assertTrue( - mlir.contains("module {"), - "dense graph must keep the bare `module {` header" - ) - } -} diff --git a/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api b/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api index 5ac226357..9ac495315 100644 --- a/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api +++ b/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api @@ -6564,8 +6564,11 @@ public final class sk/ainet/lang/tensor/ops/TensorSpec { } public final class sk/ainet/lang/tensor/ops/TensorSpecIdentities { + public static final field TENSOR_BLOCK_ORDER_METADATA_KEY Ljava/lang/String; public static final field TENSOR_ID_METADATA_KEY Ljava/lang/String; + public static final fun getBlockOrder (Lsk/ainet/lang/tensor/ops/TensorSpec;)Ljava/lang/String; public static final fun getTensorId (Lsk/ainet/lang/tensor/ops/TensorSpec;)Lsk/ainet/lang/tensor/TensorId; + public static final fun withBlockOrder (Lsk/ainet/lang/tensor/ops/TensorSpec;Ljava/lang/String;)Lsk/ainet/lang/tensor/ops/TensorSpec; public static final fun withTensorId (Lsk/ainet/lang/tensor/ops/TensorSpec;Lsk/ainet/lang/tensor/TensorId;)Lsk/ainet/lang/tensor/ops/TensorSpec; } @@ -7802,16 +7805,18 @@ public final class sk/ainet/lang/trace/OpTrace { } public final class sk/ainet/lang/trace/TensorRef { - public fun (Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;)V - public synthetic fun (Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public fun (Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;Ljava/lang/String;)V + public synthetic fun (Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V public final fun component1 ()Ljava/lang/String; public final fun component2 ()Lsk/ainet/lang/tensor/Shape; public final fun component3 ()Lsk/ainet/lang/types/DType; public final fun component4 ()Lsk/ainet/lang/tensor/TensorId; public final fun component5 ()Lsk/ainet/lang/tensor/storage/TensorEncoding; - public final fun copy (Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;)Lsk/ainet/lang/trace/TensorRef; - public static synthetic fun copy$default (Lsk/ainet/lang/trace/TensorRef;Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;ILjava/lang/Object;)Lsk/ainet/lang/trace/TensorRef; + public final fun component6 ()Ljava/lang/String; + public final fun copy (Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;Ljava/lang/String;)Lsk/ainet/lang/trace/TensorRef; + public static synthetic fun copy$default (Lsk/ainet/lang/trace/TensorRef;Ljava/lang/String;Lsk/ainet/lang/tensor/Shape;Lsk/ainet/lang/types/DType;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/tensor/storage/TensorEncoding;Ljava/lang/String;ILjava/lang/Object;)Lsk/ainet/lang/trace/TensorRef; public fun equals (Ljava/lang/Object;)Z + public final fun getBlockOrder ()Ljava/lang/String; public final fun getDtype ()Lsk/ainet/lang/types/DType; public final fun getEncoding ()Lsk/ainet/lang/tensor/storage/TensorEncoding; public final fun getId ()Ljava/lang/String; diff --git a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/tensor/ops/TensorSpecIdentity.kt b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/tensor/ops/TensorSpecIdentity.kt index 5252b682a..20fc87dc6 100644 --- a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/tensor/ops/TensorSpecIdentity.kt +++ b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/tensor/ops/TensorSpecIdentity.kt @@ -34,3 +34,24 @@ public fun TensorSpec.withTensorId(id: TensorId?): TensorSpec { } return copy(metadata = newMetadata) } + +/** + * Metadata key for the physical block order of packed storage, carried as a **stable string** + * (`ROW_MAJOR` / `INPUT_BLOCK_MAJOR`) rather than the `sk.ainet.lang.memory.BlockOrder` enum — + * deliberately, so the compile pipeline reads it without importing storage-model types (#1179). + */ +public const val TENSOR_BLOCK_ORDER_METADATA_KEY: String = "tensorBlockOrder" + +/** The packed block order carried on this spec, or `null` when unknown / not packed. */ +public val TensorSpec.blockOrder: String? + get() = metadata[TENSOR_BLOCK_ORDER_METADATA_KEY] as? String + +/** Return a copy of this spec with [order] stored in its metadata map (`null` removes the entry). */ +public fun TensorSpec.withBlockOrder(order: String?): TensorSpec { + val newMetadata: Map = if (order == null) { + metadata - TENSOR_BLOCK_ORDER_METADATA_KEY + } else { + metadata + (TENSOR_BLOCK_ORDER_METADATA_KEY to order) + } + return copy(metadata = newMetadata) +} diff --git a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TensorRef.kt b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TensorRef.kt index 6b1387a54..ed8d8adde 100644 --- a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TensorRef.kt +++ b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TensorRef.kt @@ -21,4 +21,6 @@ public data class TensorRef( val dtype: DType, val tensorId: TensorId? = null, val encoding: TensorEncoding? = null, + /** Physical block order of packed storage as a stable string (`ROW_MAJOR`/`INPUT_BLOCK_MAJOR`), or null. */ + val blockOrder: String? = null, ) diff --git a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TraceSession.kt b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TraceSession.kt index 3a64150b4..43fb464a7 100644 --- a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TraceSession.kt +++ b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/trace/TraceSession.kt @@ -53,6 +53,7 @@ public open class TraceSession { dtype = dtypeInstance, tensorId = identities[key], encoding = tensor.data.inferTensorEncoding(), + blockOrder = (tensor.data as? sk.ainet.lang.tensor.storage.PackedBlockStorage)?.blockOrder?.name, ) refToId[ref.id] = tensor ref From 6f0cb61d47e773d194f8f73253cb6752ba8225ab Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Wed, 26 Aug 2026 20:45:54 +0200 Subject: [PATCH 2/2] =?UTF-8?q?feat(opt):=20LayoutAssignmentPass=20on=20th?= =?UTF-8?q?e=20production=20path=20=E2=80=94=20one=20target,=20one=20decis?= =?UTF-8?q?ion?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Carriage slice 3 (#1147): the optimizer pipeline joins the production tape→HLO path, and the first real layout decision runs on it. - HloGenerator.generate gains target: String? = null. Null — the default — runs no pipeline and emits byte-identically to before; a named target runs dagPipelineFor(target) with LayoutAssignmentPass as a core pass, plus whatever the target registered through TargetOptimizers. compile-hlo gains the compile-opt dependency this wiring always implied (opt depends on dag; no cycle). - LayoutAssignmentPass, deliberately narrow: rank-2 block-quantized weights (the GGML block formats) get kernel-feed order (INPUT_BLOCK_MAJOR) for the target — mirroring the decision the eager form resolver makes at load (#1120) — and never overrides an order the tape already carried: the loader's fact outranks the pass's preference. Touched nodes are stamped with the backend assignment. - The pre-built ResolvedComputeGraph seams stop returning hardcoded null: resolvedLayout derives BlockOrderLayout from the spec the passes (or the tape) stamped; backendAssignment reads the node metadata key, which lives on ResolvedComputeGraph so the pass imports it along the legal opt→dag direction. - Tests: pass semantics (decides / never-overrides / leaves dense and rank-1 alone; seams surface exactly the decisions made) and the pipeline→emission composition (a decided order reaches the module header as block_order; with no pipeline, no order is invented). Closes #1180. Co-Authored-By: Claude Fable 5 --- .../api/jvm/skainet-compile-dag.api | 16 ++++ .../ainet/lang/graph/ResolvedComputeGraph.kt | 29 +++++-- .../api/jvm/skainet-compile-hlo.api | 4 +- .../skainet-compile-hlo/build.gradle.kts | 1 + .../compile/hlo/generate/HloGenerator.kt | 24 +++++- .../compile/hlo/LayoutPassToHeaderTest.kt | 58 +++++++++++++ .../api/jvm/skainet-compile-opt.api | 10 +++ .../opt/passes/LayoutAssignmentPass.kt | 81 ++++++++++++++++++ .../opt/passes/LayoutAssignmentPassTest.kt | 82 +++++++++++++++++++ 9 files changed, 295 insertions(+), 10 deletions(-) create mode 100644 skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/LayoutPassToHeaderTest.kt create mode 100644 skainet-compile/skainet-compile-opt/src/commonMain/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPass.kt create mode 100644 skainet-compile/skainet-compile-opt/src/commonTest/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPassTest.kt diff --git a/skainet-compile/skainet-compile-dag/api/jvm/skainet-compile-dag.api b/skainet-compile/skainet-compile-dag/api/jvm/skainet-compile-dag.api index bf3c5b8c5..08c8017af 100644 --- a/skainet-compile/skainet-compile-dag/api/jvm/skainet-compile-dag.api +++ b/skainet-compile/skainet-compile-dag/api/jvm/skainet-compile-dag.api @@ -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 (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; @@ -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 (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; @@ -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 (ZLjava/util/List;)V public final fun component1 ()Z diff --git a/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/ResolvedComputeGraph.kt b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/ResolvedComputeGraph.kt index 826c5b751..bc1e7e4e5 100644 --- a/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/ResolvedComputeGraph.kt +++ b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/ResolvedComputeGraph.kt @@ -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 @@ -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: @@ -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 diff --git a/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api b/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api index b86eff6a3..424ef2c94 100644 --- a/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api +++ b/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api @@ -493,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 { diff --git a/skainet-compile/skainet-compile-hlo/build.gradle.kts b/skainet-compile/skainet-compile-hlo/build.gradle.kts index 01fd77c50..d2e08e072 100644 --- a/skainet-compile/skainet-compile-hlo/build.gradle.kts +++ b/skainet-compile/skainet-compile-hlo/build.gradle.kts @@ -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 { diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/generate/HloGenerator.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/generate/HloGenerator.kt index 47c06f4a9..ae90937bc 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/generate/HloGenerator.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/generate/HloGenerator.kt @@ -29,7 +29,15 @@ public object HloGenerator { public suspend fun generate( model: Model, Tensor>, sampleInput: Tensor, - functionName: String = "main" + functionName: String = "main", + /** + * The compile target (an IREE device name). `null` — the default, byte-identical to the + * pre-#1180 behaviour — runs no optimization pipeline. A named target runs + * `dagPipelineFor(target)` with [sk.ainet.compile.opt.passes.LayoutAssignmentPass] as a + * core pass, plus whatever the target registered via + * [sk.ainet.compile.opt.TargetOptimizers]. + */ + target: String? = null, ): StableHloModule { val ctx = DefaultGraphExecutionContext.tape(baseOps = VoidTensorOps()) val traceInput = sampleInput.bind(ctx) @@ -52,8 +60,20 @@ public object HloGenerator { inputTensorIds = setOf(inputRefId) ) ?: error("Failed to create compute graph: no execution tape was recorded") + // The optimizer pipeline joins the production path here (#1180): with no target named, + // nothing runs and the emitted module is what it always was; with one, the layout pass + // decides and every registered target optimizer gets its say. + val optimizedGraph = if (target == null) { + computeGraph + } else { + sk.ainet.compile.opt.dagPipelineFor( + target, + corePasses = listOf(sk.ainet.compile.opt.passes.LayoutAssignmentPass(target)), + ).optimize(computeGraph).graph + } + val converter = StableHloConverterFactory.createExtended() - return converter.convert(computeGraph, functionName) + return converter.convert(optimizedGraph, functionName) } private suspend fun traceForwardPass( diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/LayoutPassToHeaderTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/LayoutPassToHeaderTest.kt new file mode 100644 index 000000000..6c0cc5269 --- /dev/null +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/LayoutPassToHeaderTest.kt @@ -0,0 +1,58 @@ +package sk.ainet.compile.hlo + +import sk.ainet.compile.opt.dagPipelineFor +import sk.ainet.compile.opt.passes.LayoutAssignmentPass +import sk.ainet.lang.graph.DefaultComputeGraph +import sk.ainet.lang.graph.GraphEdge +import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.tensor.ops.AddOperation +import sk.ainet.lang.tensor.ops.InputOperation +import sk.ainet.lang.tensor.ops.TensorSpec +import sk.ainet.lang.tensor.ops.withTensorEncoding +import sk.ainet.lang.tensor.storage.TensorEncoding +import sk.ainet.lang.types.DType +import kotlin.test.Test +import kotlin.test.assertFalse +import kotlin.test.assertTrue + +/** + * #1180 end to end at the graph level: the layout pass decides, and slice 2's structural module + * attribute carries the decision — a rank-2 Q4_K weight with no carried order enters the pipeline + * and leaves the emitter declaring `block_order = "INPUT_BLOCK_MAJOR"` in the header. With no + * pipeline run, the header carries no order — proving the default path is untouched. + */ +class LayoutPassToHeaderTest { + + private fun graphWithPackedWeight(): DefaultComputeGraph { + val graph = DefaultComputeGraph() + val q4 = TensorSpec("w_q4", listOf(64, 256), "FP32").withTensorEncoding(TensorEncoding.Q4_K) + val w = GraphNode("w_q4", InputOperation(), emptyList(), listOf(q4)) + val x = GraphNode("x", InputOperation(), emptyList(), listOf(TensorSpec("x", listOf(1, 64), "FP32"))) + val add = GraphNode("add", AddOperation(), listOf(x.outputs[0], q4), listOf(TensorSpec("y", listOf(1, 256), "FP32"))) + graph.addNode(w); graph.addNode(x); graph.addNode(add) + graph.addEdge(GraphEdge("e_w", w, add, 0, 1, q4)) + graph.addEdge(GraphEdge("e_x", x, add, 0, 0, x.outputs[0])) + return graph + } + + @Test + fun decidedOrderReachesTheModuleHeader() { + val optimized = dagPipelineFor( + "test-target", + corePasses = listOf(LayoutAssignmentPass("test-target")), + ).optimize(graphWithPackedWeight()).graph + + val mlir = toStableHlo(optimized, "layout_chain").content + assertTrue( + mlir.contains("block_order = \"INPUT_BLOCK_MAJOR\""), + "the pass's decision must be declared in the header:\n$mlir", + ) + } + + @Test + fun noPipelineMeansNoOrderInTheHeader() { + val mlir = toStableHlo(graphWithPackedWeight(), "layout_chain").content + assertTrue(mlir.contains("skainet.tensor_layouts"), "structural facts still emitted:\n$mlir") + assertFalse(mlir.contains("block_order"), "no pass ran, so no order may be invented:\n$mlir") + } +} diff --git a/skainet-compile/skainet-compile-opt/api/jvm/skainet-compile-opt.api b/skainet-compile/skainet-compile-opt/api/jvm/skainet-compile-opt.api index 82e878350..5e3a3b26b 100644 --- a/skainet-compile/skainet-compile-opt/api/jvm/skainet-compile-opt.api +++ b/skainet-compile/skainet-compile-opt/api/jvm/skainet-compile-opt.api @@ -127,6 +127,16 @@ public final class sk/ainet/compile/opt/passes/LLMFusionPass : sk/ainet/compile/ public fun getName ()Ljava/lang/String; } +public final class sk/ainet/compile/opt/passes/LayoutAssignmentPass : sk/ainet/compile/opt/GraphOptimizationPass { + public static final field Companion Lsk/ainet/compile/opt/passes/LayoutAssignmentPass$Companion; + public fun (Ljava/lang/String;)V + public fun apply (Lsk/ainet/lang/graph/ComputeGraph;)Lsk/ainet/compile/opt/GraphOptimizationResult; + public fun getName ()Ljava/lang/String; +} + +public final class sk/ainet/compile/opt/passes/LayoutAssignmentPass$Companion { +} + public final class sk/ainet/compile/opt/passes/OperationFusionPass : sk/ainet/compile/opt/GraphOptimizationPass { public fun ()V public fun apply (Lsk/ainet/lang/graph/ComputeGraph;)Lsk/ainet/compile/opt/GraphOptimizationResult; diff --git a/skainet-compile/skainet-compile-opt/src/commonMain/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPass.kt b/skainet-compile/skainet-compile-opt/src/commonMain/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPass.kt new file mode 100644 index 000000000..407b8f8ca --- /dev/null +++ b/skainet-compile/skainet-compile-opt/src/commonMain/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPass.kt @@ -0,0 +1,81 @@ +package sk.ainet.compile.opt.passes + +import sk.ainet.compile.opt.GraphOptimizationPass +import sk.ainet.compile.opt.GraphOptimizationResult +import sk.ainet.lang.graph.ComputeGraph +import sk.ainet.lang.graph.DefaultComputeGraph +import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.graph.ResolvedComputeGraph +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.ops.withBlockOrder +import sk.ainet.lang.tensor.storage.TensorEncoding + +/** + * The first real layout decision on the compile path (#1180): rank-2 block-quantized weights are + * assigned kernel-feed block order (`INPUT_BLOCK_MAJOR`) for [target], mirroring the decision the + * eager side's form resolver already makes at load — packed matmul kernels read input-block-major, + * and a weight delivered in that order has nothing to convert on first use (#1120). + * + * Decision rules, deliberately narrow: + * - only block-quantized encodings whose kernels feed them (the GGML block formats); + * - only rank-2 specs (feed order is defined relative to an `[out, in]` weight); + * - never overrides an order the tape already carried — the loader's fact outranks a preference. + * + * The pass stamps `backendAssignment` metadata on the nodes it touched, which + * `ResolvedComputeGraph.backendAssignment` surfaces — so the pre-built seams stop returning null + * exactly where a decision was actually made, and nowhere else. + * + * Mechanically target-parameterized rather than registry-registered: `HloGenerator` runs it as a + * core pass of `dagPipelineFor(target, …)` whenever a target is named, and target-specific + * optimizers can still contribute their own passes through [sk.ainet.compile.opt.TargetOptimizers]. + */ +public class LayoutAssignmentPass(private val target: String) : GraphOptimizationPass { + + override val name: String = "layout-assignment($target)" + + public companion object { + /** Block-quantized encodings whose matmul kernels read input-block-major (#973/#1120). */ + private val FEED_ORDERED = setOf( + TensorEncoding.Q4_K, TensorEncoding.Q5_K, TensorEncoding.Q6_K, + TensorEncoding.Q4_0, TensorEncoding.Q5_0, TensorEncoding.Q5_1, TensorEncoding.Q8_0, + ) + } + + private fun decide(spec: TensorSpec): TensorSpec? { + val encoding = spec.tensorEncoding ?: return null + if (encoding !in FEED_ORDERED) return null + if (spec.shape?.size != 2) return null + if (spec.blockOrder != null) return null // the tape's fact outranks this preference + return spec.withBlockOrder("INPUT_BLOCK_MAJOR") + } + + override fun apply(graph: ComputeGraph): GraphOptimizationResult { + var changed = false + val newNodes = graph.nodes.map { node -> + var touched = false + val outputs = node.outputs.map { spec -> decide(spec)?.also { touched = true } ?: spec } + val inputs = node.inputs.map { spec -> decide(spec)?.also { touched = true } ?: spec } + if (!touched) return@map node + changed = true + node.copy( + inputs = inputs, + outputs = outputs, + metadata = node.metadata + (ResolvedComputeGraph.BACKEND_ASSIGNMENT_METADATA_KEY to target), + ) + } + if (!changed) return GraphOptimizationResult(graph, changed = false) + + val byId = newNodes.associateBy { it.id } + val newGraph = DefaultComputeGraph() + for (node in newNodes) newGraph.addNode(node) + for (edge in graph.edges) { + val src = byId.getValue(edge.source.id) + val dst = byId.getValue(edge.destination.id) + val spec = decide(edge.tensorSpec) ?: edge.tensorSpec + newGraph.addEdge(edge.copy(source = src, destination = dst, tensorSpec = spec)) + } + return GraphOptimizationResult(newGraph, changed = true) + } +} diff --git a/skainet-compile/skainet-compile-opt/src/commonTest/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPassTest.kt b/skainet-compile/skainet-compile-opt/src/commonTest/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPassTest.kt new file mode 100644 index 000000000..80c35d7d0 --- /dev/null +++ b/skainet-compile/skainet-compile-opt/src/commonTest/kotlin/sk/ainet/compile/opt/passes/LayoutAssignmentPassTest.kt @@ -0,0 +1,82 @@ +package sk.ainet.compile.opt.passes + +import sk.ainet.lang.graph.BlockOrderLayout +import sk.ainet.lang.graph.DefaultComputeGraph +import sk.ainet.lang.graph.GraphEdge +import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.graph.ResolvedComputeGraph +import sk.ainet.lang.tensor.ops.AddOperation +import sk.ainet.lang.tensor.ops.InputOperation +import sk.ainet.lang.tensor.ops.TensorSpec +import sk.ainet.lang.tensor.ops.blockOrder +import sk.ainet.lang.tensor.ops.withBlockOrder +import sk.ainet.lang.tensor.ops.withTensorEncoding +import sk.ainet.lang.tensor.storage.TensorEncoding +import sk.ainet.lang.types.DType +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFalse +import kotlin.test.assertNull +import kotlin.test.assertTrue + +/** + * #1180: the first layout decision, narrowly scoped — rank-2 block-quantized weights get + * kernel-feed order for the target; the tape's own facts are never overridden; nothing else is + * touched. The pre-built `ResolvedComputeGraph` seams surface exactly the decisions made. + */ +class LayoutAssignmentPassTest { + + private fun weightNode(id: String, spec: TensorSpec) = GraphNode( + id = id, + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(spec), + ) + + @Test + fun rank2PackedWeightGetsFeedOrderAndBackendAssignment() { + val graph = DefaultComputeGraph() + val q4 = TensorSpec("w", listOf(64, 256), "FP32").withTensorEncoding(TensorEncoding.Q4_K) + val w = weightNode("w", q4) + val x = weightNode("x", TensorSpec("x", listOf(1, 64), "FP32")) + val add = GraphNode("add", AddOperation(), listOf(x.outputs[0], q4), listOf(TensorSpec("y", listOf(1, 256), "FP32"))) + graph.addNode(w); graph.addNode(x); graph.addNode(add) + graph.addEdge(GraphEdge("e_w", w, add, 0, 1, q4)) + graph.addEdge(GraphEdge("e_x", x, add, 0, 0, x.outputs[0])) + + val result = LayoutAssignmentPass("test-target").apply(graph) + assertTrue(result.changed) + + val resolved = ResolvedComputeGraph(result.graph) + val layout = resolved.resolvedLayout("e_w") + assertEquals(BlockOrderLayout("INPUT_BLOCK_MAJOR"), layout, "the seam surfaces the decision") + assertEquals("test-target", resolved.backendAssignment("w"), "touched node carries the assignment") + assertNull(resolved.resolvedLayout("e_x"), "undecided edges stay null") + assertNull(resolved.backendAssignment("x"), "untouched nodes stay null") + } + + @Test + fun theTapesOwnFactIsNeverOverridden() { + val graph = DefaultComputeGraph() + val loaded = TensorSpec("w", listOf(64, 256), "FP32") + .withTensorEncoding(TensorEncoding.Q8_0) + .withBlockOrder("ROW_MAJOR") // the loader delivered row-major and said so + graph.addNode(weightNode("w", loaded)) + + val result = LayoutAssignmentPass("test-target").apply(graph) + assertFalse(result.changed, "a carried fact outranks the pass's preference") + assertEquals("ROW_MAJOR", result.graph.nodes.single().outputs.single().blockOrder) + } + + @Test + fun denseAndRank1SpecsAreUntouched() { + val graph = DefaultComputeGraph() + graph.addNode(weightNode("dense", TensorSpec("dense", listOf(4, 4), "FP32"))) + graph.addNode( + weightNode("bias", TensorSpec("bias", listOf(256), "FP32").withTensorEncoding(TensorEncoding.Q8_0)), + ) + val result = LayoutAssignmentPass("test-target").apply(graph) + assertFalse(result.changed) + assertNull(result.graph.nodes.first { it.id == "bias" }.outputs.single().blockOrder) + } +}