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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -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<String>,
) : ComputeGraph by inner {
override fun getOutputNodes(): List<GraphNode> =
inner.nodes.filter { it.id in outputNodeIds }
}
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -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<String>): ComputeGraph {
require(outputNodeIds.isNotEmpty()) { "prunedToOutputs: outputNodeIds must not be empty" }
val designated = OutputDesignatedGraph(this, outputNodeIds)
return DeadCodeEliminationPass().apply(designated).graph
}
Original file line number Diff line number Diff line change
@@ -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<Int> = 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<IllegalArgumentException> { graph.prunedToOutputs(emptySet()) }
}
}
Loading