diff --git a/CHANGELOG.md b/CHANGELOG.md index c83e2e07f..e237bfd97 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,23 @@ ## [Unreleased] +### Added + +- **Graph-output pruning for export (`ComputeGraph.prunedToOutputs`).** A traced decoder surfaces + every leaf tensor (e.g. per-layer intermediates) as a graph output, so a StableHLO/IREE export + returns the logits plus dozens of dangling tensors — extra `func` returns and dead op subgraphs. + Adds `OutputDesignatedGraph` (skainet-compile-dag) to override the output set by node id, and + `ComputeGraph.prunedToOutputs(outputNodeIds)` (skainet-compile-opt) which designates those outputs + and runs `DeadCodeEliminationPass` so only the nodes feeding them survive. Exporters can now keep + just the logits. Adds `GraphPruningTest` (commonTest). + +### Changed + +- **SDPA causal mask uses a large finite fill (`-1e30`) instead of `-inf`.** The attention HLO + converter's causal path emitted `dense<0xFF800000>` (`-inf`) for the masked-fill select; it now + emits `-1.000000e+30`, matching `MultiHeadAttention.buildSlidingCausalMask` and avoiding a `-inf` + splat in the IR (numerically equivalent after softmax). (`AttentionOperationsConverter`) + ## [0.32.2] - 2026-06-24 ### Added diff --git a/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/OutputDesignatedGraph.kt b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/OutputDesignatedGraph.kt new file mode 100644 index 000000000..ac51c3e8a --- /dev/null +++ b/skainet-compile/skainet-compile-dag/src/commonMain/kotlin/sk/ainet/lang/graph/OutputDesignatedGraph.kt @@ -0,0 +1,22 @@ +package sk.ainet.lang.graph + +/** + * A [ComputeGraph] view that overrides the set of *output* nodes to a caller-designated subset + * (by node id), delegating every other operation to [inner]. + * + * By default a graph's outputs are inferred as the nodes with no outgoing edges + * ([ComputeGraph.getOutputNodes]). For a traced decoder that surfaces dangling intermediates + * (e.g. per-layer post-RoPE q/k tensors) as extra outputs, which then get emitted as additional + * `func` returns and dead subgraphs. Wrapping the graph here and feeding it to a dead-code pass + * lets an exporter keep only what's reachable from a chosen output (e.g. the decoder logits). + * + * See `ComputeGraph.prunedToOutputs` in skainet-compile-opt, which combines this with + * `DeadCodeEliminationPass` to physically remove the now-unreachable nodes before conversion. + */ +public class OutputDesignatedGraph( + private val inner: ComputeGraph, + private val outputNodeIds: Set, +) : ComputeGraph by inner { + override fun getOutputNodes(): List = + inner.nodes.filter { it.id in outputNodeIds } +} diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/AttentionOperationsConverter.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/AttentionOperationsConverter.kt index 9f35ad021..b221e6764 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/AttentionOperationsConverter.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/AttentionOperationsConverter.kt @@ -134,7 +134,10 @@ public class AttentionOperationsConverter : StableHloOperationConverter { ops += "$iotaK = stablehlo.iota dim = $sdAxis : $scoresI32Type" ops += "$keep = stablehlo.compare GE, $iotaQ, $iotaK : ($scoresI32Type, $scoresI32Type) -> $scoresI1Type" ops += "$zeros = stablehlo.constant dense<0.0> : $scoresType" - ops += "$ninf = stablehlo.constant dense<0xFF800000> : $scoresType" + // Masked-fill with a large finite negative (not -inf): matches + // MultiHeadAttention.buildSlidingCausalMask (-1e30) and avoids a -inf splat in the + // masked-fill select, which can trip downstream greedy constant-folding. + ops += "$ninf = stablehlo.constant dense<-1.000000e+30> : $scoresType" ops += "$maskAdd = stablehlo.select $keep, $zeros, $ninf : $scoresI1Type, $scoresType" ops += "$masked = stablehlo.add $scaled, $maskAdd : $scoresType" softmaxIn = masked diff --git a/skainet-compile/skainet-compile-opt/src/commonMain/kotlin/sk/ainet/compile/opt/GraphPruning.kt b/skainet-compile/skainet-compile-opt/src/commonMain/kotlin/sk/ainet/compile/opt/GraphPruning.kt new file mode 100644 index 000000000..a4b3e7cb7 --- /dev/null +++ b/skainet-compile/skainet-compile-opt/src/commonMain/kotlin/sk/ainet/compile/opt/GraphPruning.kt @@ -0,0 +1,24 @@ +package sk.ainet.compile.opt + +import sk.ainet.compile.opt.passes.DeadCodeEliminationPass +import sk.ainet.lang.graph.ComputeGraph +import sk.ainet.lang.graph.OutputDesignatedGraph + +/** + * Return a graph containing only the nodes that contribute to [outputNodeIds], with those nodes as + * the sole outputs. + * + * Designates [outputNodeIds] as the graph outputs (via [OutputDesignatedGraph]) and runs + * [DeadCodeEliminationPass] to drop every node not reachable backward from them. Use before a + * StableHLO / IREE export to keep only the desired result (e.g. a decoder's logits) and discard + * dangling intermediates the trace leaves as extra outputs — these would otherwise be emitted as + * additional `func` returns and dead op subgraphs that can crash downstream compilers + * (observed: an `iree-compile` constant-folding null-deref on a multi-position decoder graph). + * + * @throws IllegalArgumentException if [outputNodeIds] is empty. + */ +public fun ComputeGraph.prunedToOutputs(outputNodeIds: Set): ComputeGraph { + require(outputNodeIds.isNotEmpty()) { "prunedToOutputs: outputNodeIds must not be empty" } + val designated = OutputDesignatedGraph(this, outputNodeIds) + return DeadCodeEliminationPass().apply(designated).graph +} diff --git a/skainet-compile/skainet-compile-opt/src/commonTest/kotlin/sk/ainet/compile/opt/GraphPruningTest.kt b/skainet-compile/skainet-compile-opt/src/commonTest/kotlin/sk/ainet/compile/opt/GraphPruningTest.kt new file mode 100644 index 000000000..5c0a32027 --- /dev/null +++ b/skainet-compile/skainet-compile-opt/src/commonTest/kotlin/sk/ainet/compile/opt/GraphPruningTest.kt @@ -0,0 +1,83 @@ +package sk.ainet.compile.opt + +import sk.ainet.lang.graph.DefaultComputeGraph +import sk.ainet.lang.graph.GraphEdge +import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.tensor.ops.GenericOperation +import sk.ainet.lang.tensor.ops.TensorSpec +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFailsWith +import kotlin.test.assertTrue + +/** + * Tests for [prunedToOutputs] — keeping only the nodes that feed a designated output and dropping + * every other leaf. This is the capability a decoder export needs: a traced graph surfaces + * dangling intermediates as extra outputs (every leaf with no outgoing edge), and only the logits + * should survive. (See DeadCodeEliminationPassTest's own comments: in a DAG, dead code only + * manifests once you can mark explicit outputs — which is exactly what this provides.) + */ +class GraphPruningTest { + + private fun spec(name: String = "t", shape: List = listOf(1)) = + TensorSpec(name = name, shape = shape, dtype = "float32") + + private fun constNode(id: String) = GraphNode( + id = id, + operation = GenericOperation("constant", mapOf("values" to listOf(1.0f)), "constant"), + inputs = emptyList(), + outputs = listOf(spec()), + ) + + private fun opNode(id: String, opName: String = "add") = GraphNode( + id = id, + operation = GenericOperation(opName), + inputs = listOf(spec()), + outputs = listOf(spec()), + ) + + @Test + fun keepsOnlyDesignatedOutputAndItsAncestors() { + // a → b (the logits we keep); a → c (a dangling sibling leaf to drop). + val graph = DefaultComputeGraph() + val a = graph.addNode(constNode("a")) + val b = graph.addNode(opNode("b")) + val c = graph.addNode(opNode("c")) + graph.addEdge(GraphEdge("e1", a, b, tensorSpec = spec())) + graph.addEdge(GraphEdge("e2", a, c, tensorSpec = spec())) + // Default outputs are both leaves b and c. + assertEquals(setOf("b", "c"), graph.getOutputNodes().map { it.id }.toSet()) + + val pruned = graph.prunedToOutputs(setOf("b")) + + // c (and only c) is gone; a and b remain; b is now the sole output. + assertEquals(setOf("a", "b"), pruned.nodes.map { it.id }.toSet()) + assertEquals(listOf("b"), pruned.getOutputNodes().map { it.id }) + assertEquals(1, pruned.edges.size) + } + + @Test + fun keepsSharedAncestorsAcrossLiveAndDeadBranches() { + // a → b → out (keep); a → c (drop). 'a' is a shared ancestor → must survive. + val graph = DefaultComputeGraph() + graph.addNode(constNode("a")) + graph.addNode(opNode("b")) + graph.addNode(opNode("out")) + graph.addNode(opNode("c")) + graph.addEdge(GraphEdge("e1", graph.nodes[0], graph.nodes[1], tensorSpec = spec())) + graph.addEdge(GraphEdge("e2", graph.nodes[1], graph.nodes[2], tensorSpec = spec())) + graph.addEdge(GraphEdge("e3", graph.nodes[0], graph.nodes[3], tensorSpec = spec())) + + val pruned = graph.prunedToOutputs(setOf("out")) + + assertEquals(setOf("a", "b", "out"), pruned.nodes.map { it.id }.toSet()) + assertTrue(pruned.nodes.none { it.id == "c" }) + } + + @Test + fun emptyOutputSetIsRejected() { + val graph = DefaultComputeGraph() + graph.addNode(constNode("a")) + assertFailsWith { graph.prunedToOutputs(emptySet()) } + } +}