From 6f0cb61d47e773d194f8f73253cb6752ba8225ab Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Wed, 26 Aug 2026 20:45:54 +0200 Subject: [PATCH] =?UTF-8?q?feat(opt):=20LayoutAssignmentPass=20on=20the=20?= =?UTF-8?q?production=20path=20=E2=80=94=20one=20target,=20one=20decision?= 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) + } +}