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/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/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-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..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 @@ -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; @@ -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 { 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/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/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/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/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-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-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) + } +} 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