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
13 changes: 13 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,19 @@

## [Unreleased]

### Fixed

- **SDPA StableHLO export: explicit masks under grouped-query attention and onto a dynamic key
length** (#1302). A per-head mask `[b, H, Sq, Sk]` under GQA was broadcast onto the grouped
scores `[b, nKV, nRep, Sq, Sk]`, mapping `H` onto `nKV` (invalid IR even when static); it is now
reshaped into `[nKV, nRep]` in Q's own head order (`dynamic_reshape` when a dim is dynamic). A
mask that differs from a dynamic scores shape (the `?` key length of KV-cache chunk graphs) got a
static `broadcast_in_dim` to a dynamic type, which every backend rejects, with or without GQA; it
now uses `dynamic_broadcast_in_dim` with `known_expanding_dimensions` /
`known_nonexpanding_dimensions`, the form IREE 3.11 lowers. Static graphs are unchanged. Note:
IREE 3.11 does not lower `dynamic_reshape`, so for IREE prefer a head-shared `[b, 1, Sq, ?]` mask
when the key length is dynamic.

## [0.56.0] - 2026-09-20

Headline: **grouped-query attention is native to the engine, and the compiled leg of SKEEP-005 lands —
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -148,28 +148,113 @@ public class AttentionOperationsConverter : StableHloOperationConverter {
} else qScaled
ops += "$scores = stablehlo.dot_general $qForDot, ${operands[1]}, ${batchClause}contracting_dims = [$contractQ] x [$contractK] : ($qWorkType, $kType) -> $scoresType"

// When the scores shape is dynamic (the `?` key/cache dim of KV-cache decode), everything broadcast onto
// it — the explicit mask and the softmax's reduced max/sum — must use `stablehlo.dynamic_broadcast_in_dim`:
// a static `stablehlo.broadcast_in_dim` cannot target a dynamic shape. Its runtime `output_dimensions`
// operand is built once, here, from the scores tensor via `get_dimension_size`.
// Static graphs keep the original explicit `broadcast_in_dim` path (byte-for-byte unchanged).
val dyn = scoresShape.hasDynamic()
val shapeType = "tensor<${scoresShape.size}xi32>"
val scoresShapeOperand: String = if (!dyn) "" else run {
val parts = scoresShape.indices.map { d ->
if (Dim.isStatic(scoresShape[d])) {
val c = context.nextTempValue()
ops += "$c = stablehlo.constant dense<${scoresShape[d]}> : tensor<1xi32>"
c
} else {
val gd = context.nextTempValue(); val gr = context.nextTempValue()
ops += "$gd = stablehlo.get_dimension_size $scores, dim = $d : ($scoresType) -> tensor<i32>"
ops += "$gr = stablehlo.reshape $gd : (tensor<i32>) -> tensor<1xi32>"
gr
}
}
val sh = context.nextTempValue()
ops += "$sh = stablehlo.concatenate ${parts.joinToString(", ")}, dim = 0 : (${parts.joinToString(", ") { "tensor<1xi32>" }}) -> $shapeType"
sh
}

// Explicit additive mask (operands[3]) — e.g. a sliding-window+causal
// mask the caller built and passed with causal=false. It already
// encodes causality/window, so it takes priority over the built-in
// iota causal path. Broadcast (trailing-aligned) to the scores shape
// and add. Without this the masked layers run UNMASKED (attend to
// future tokens) — correct only at position 0.
// iota causal path. Brought to the scores shape and added. Without this
// the masked layers run UNMASKED (attend to future tokens) — correct
// only at position 0.
//
// Shape rules: trailing-aligned broadcast. Under GQA a rank-4 mask [b, M, Sq, Sk] keeps its batch on
// scores dim 0 and skips the nRep axis (dim 2); M may be 1 (head-shared) or nKV (one row per group).
// A per-head mask (M = H) is first VIEWED as [b, nKV, nRep, Sq, Sk] in Q's own h = kv * nRep + r order
// (reshape, or dynamic_reshape when a dim is dynamic) — broadcasting H onto nKV is invalid IR.
var softmaxIn = scores // scores are already scaled (scale folded into Q above)
val maskOperand = operands.getOrNull(3)
if (maskOperand != null) {
val maskShape = node.inputs.getOrNull(3)?.shape ?: scoresShape
val maskType = context.getValueType(maskOperand) ?: typeOf(maskShape)
var maskShape: List<Int> = node.inputs.getOrNull(3)?.shape ?: scoresShape
var maskVal = maskOperand
var maskType = context.getValueType(maskOperand) ?: typeOf(maskShape)
var maskDims: List<Int> = if (gqa && maskShape.size == 4) listOf(0, 1, 3, 4) else {
val offset = scoresShape.size - maskShape.size
maskShape.indices.map { it + offset }
}
if (gqa && maskShape.size == 4 && maskShape[1] != 1 && maskShape[1] != kShape[1]) {
val nKV = kShape[1]; val nH = qShape[1]
if (maskShape[1] != nH) {
return ConversionResult.Failure(
"SDPA grouped-query attention mask head dim must be 1, K/V heads ($nKV) or Q heads ($nH), got $maskShape",
"Unsupported GQA mask shape for ${node.id}",
)
}
val split = listOf(maskShape[0], nKV, nH / nKV, maskShape[2], maskShape[3])
val splitType = typeOf(split)
val r = context.nextTempValue()
if (!maskShape.hasDynamic()) {
ops += "$r = stablehlo.reshape $maskVal : ($maskType) -> $splitType"
} else {
// split dim -> source mask dim (the two group dims are static by construction)
val source = listOf(0, -1, -1, 2, 3)
val parts = split.indices.map { d ->
if (Dim.isStatic(split[d])) {
val c = context.nextTempValue()
ops += "$c = stablehlo.constant dense<${split[d]}> : tensor<1xi32>"
c
} else {
val gd = context.nextTempValue(); val gr = context.nextTempValue()
ops += "$gd = stablehlo.get_dimension_size $maskVal, dim = ${source[d]} : ($maskType) -> tensor<i32>"
ops += "$gr = stablehlo.reshape $gd : (tensor<i32>) -> tensor<1xi32>"
gr
}
}
val sh = context.nextTempValue()
ops += "$sh = stablehlo.concatenate ${parts.joinToString(", ")}, dim = 0 : (${parts.joinToString(", ") { "tensor<1xi32>" }}) -> tensor<5xi32>"
ops += "$r = stablehlo.dynamic_reshape $maskVal, $sh : ($maskType, tensor<5xi32>) -> $splitType"
}
maskVal = r; maskShape = split; maskType = splitType; maskDims = split.indices.toList()
}
val maskBc = if (maskShape == scoresShape) {
maskOperand
} else {
maskVal
} else if (!dyn) {
val mb = context.nextTempValue()
// Trailing-aligned. Under GQA a rank-4 mask [b, 1|H, Sq, Sk] keeps its batch on
// scores dim 0 and skips the nRep axis (dim 2): [0, 1, 3, 4].
val dims = if (gqa && maskShape.size == 4) "0, 1, 3, 4" else {
val offset = scoresShape.size - maskShape.size
maskShape.indices.joinToString(", ") { (it + offset).toString() }
ops += "$mb = stablehlo.broadcast_in_dim $maskVal, dims = [${maskDims.joinToString(", ")}] : ($maskType) -> $scoresType"
mb
} else {
// Dynamic target: dynamic_broadcast_in_dim, stating which operand dims expand. Without the
// hints a dynamic operand dim mapped onto a dynamic result dim is ambiguous (1 -> N or N -> N),
// and backends such as IREE refuse to lower it. A dynamic mask dim must match the scores dim
// it maps to (the key length), so it is non-expanding. Generic op syntax: the hint attributes
// are not part of every StableHLO version's pretty form.
val expanding = mutableListOf<Int>(); val nonExpanding = mutableListOf<Int>()
maskShape.indices.forEach { i ->
val o = maskShape[i]; val t = scoresShape[maskDims[i]]
when {
!Dim.isStatic(o) -> nonExpanding += i
Dim.isStatic(t) && o == t -> nonExpanding += i
Dim.isStatic(t) && o == 1 -> expanding += i
!Dim.isStatic(t) && o != 1 -> nonExpanding += i
}
}
ops += "$mb = stablehlo.broadcast_in_dim $maskOperand, dims = [$dims] : ($maskType) -> $scoresType"
fun arr(xs: List<Int>) = if (xs.isEmpty()) "array<i64>" else "array<i64: ${xs.joinToString(", ")}>"
val mb = context.nextTempValue()
ops += "$mb = \"stablehlo.dynamic_broadcast_in_dim\"($maskVal, $scoresShapeOperand) <{broadcast_dimensions = ${arr(maskDims)}, " +
"known_expanding_dimensions = ${arr(expanding)}, known_nonexpanding_dimensions = ${arr(nonExpanding)}}> : " +
"($maskType, $shapeType) -> $scoresType"
mb
}
val masked = context.nextTempValue()
Expand All @@ -193,30 +278,6 @@ public class AttentionOperationsConverter : StableHloOperationConverter {
softmaxIn = masked
}

// softmax(softmaxIn) over the key-length axis. When the scores shape is dynamic (the `?` key/cache dim
// of KV-cache decode), the reduced max/sum must broadcast back to the dynamic scores shape. A static
// `stablehlo.broadcast_in_dim` cannot target a dynamic shape, so we use `stablehlo.dynamic_broadcast_in_dim`
// with a runtime `output_dimensions` operand (built once from the scores tensor via `get_dimension_size`).
// Static graphs keep the original explicit `broadcast_in_dim` path (byte-for-byte unchanged).
val dyn = scoresShape.hasDynamic()
val shapeType = "tensor<${scoresShape.size}xi32>"
val scoresShapeOperand: String = if (!dyn) "" else run {
val parts = scoresShape.indices.map { d ->
if (Dim.isStatic(scoresShape[d])) {
val c = context.nextTempValue()
ops += "$c = stablehlo.constant dense<${scoresShape[d]}> : tensor<1xi32>"
c
} else {
val gd = context.nextTempValue(); val gr = context.nextTempValue()
ops += "$gd = stablehlo.get_dimension_size $scores, dim = $d : ($scoresType) -> tensor<i32>"
ops += "$gr = stablehlo.reshape $gd : (tensor<i32>) -> tensor<1xi32>"
gr
}
}
val sh = context.nextTempValue()
ops += "$sh = stablehlo.concatenate ${parts.joinToString(", ")}, dim = 0 : (${parts.joinToString(", ") { "tensor<1xi32>" }}) -> $shapeType"
sh
}
fun broadcastBack(src: String, dst: String) {
if (dyn) {
ops += "$dst = stablehlo.dynamic_broadcast_in_dim $src, $scoresShapeOperand, dims = [$bcastDims] : ($reducedType, $shapeType) -> $scoresType"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,57 @@ class SdpaGqaHloExportTest {
assertFalse(mlir.contains("stablehlo.concatenate %arg"), "K/V must not be materialised:\n$mlir")
}

@Test
fun perHeadMaskIsSplitIntoGroupsNotBroadcast() {
// [b, H, Sq, Sk] mask under GQA (H = 4, nKV = 2): the head dim must be viewed as [nKV, nRep]
// in the same h = kv * nRep + r order as Q. Broadcasting H onto nKV is invalid IR.
val mlir = StableHloConverterFactory.createBasic().convert(graph(listOf(1, 4, 8, 16), listOf(1, 2, 8, 16), causal = false, mask = listOf(1, 4, 8, 8)), "gqa_head_mask").content
assertTrue(mlir.contains("stablehlo.reshape %arg3 : (tensor<1x4x8x8xf32>) -> tensor<1x2x2x8x8xf32>"), "per-head mask is reshaped to [b, nKV, nRep, Sq, Sk]:\n$mlir")
assertFalse(mlir.contains("(tensor<1x4x8x8xf32>) -> tensor<1x2x2x8x8xf32>") && mlir.contains("broadcast_in_dim %arg3"), "per-head mask must not be broadcast onto nKV:\n$mlir")
}

@Test
fun headSharedMaskWithDynamicKeyLengthUsesAnAttributedDynamicBroadcast() {
// KV-cache chunk graph: mask [b, 1, Sq, past+Sq] with a dynamic key length. A static
// broadcast_in_dim cannot produce the dynamic scores type.
val d = TypeMapper.DYNAMIC_DIM
val mlir = StableHloConverterFactory.createBasic().convert(graph(listOf(1, 16, 32, 128), listOf(1, 8, d, 128), causal = false, mask = listOf(1, 1, 32, d)), "gqa_dyn_mask").content
assertFalse(mlir.contains(Regex("""stablehlo\.broadcast_in_dim %arg3[^\n]*\?""")), "no static broadcast to a dynamic type:\n$mlir")
assertTrue(
mlir.contains(
"\"stablehlo.dynamic_broadcast_in_dim\"(%arg3, ") &&
mlir.contains(
"<{broadcast_dimensions = array<i64: 0, 1, 3, 4>, known_expanding_dimensions = array<i64: 1>, " +
"known_nonexpanding_dimensions = array<i64: 0, 2, 3>}> : (tensor<1x1x32x?xf32>, tensor<5xi32>) -> tensor<1x8x2x32x?xf32>",
),
"head-shared dynamic mask uses dynamic_broadcast_in_dim with expansion hints:\n$mlir",
)
}

@Test
fun perHeadMaskWithDynamicKeyLengthUsesADynamicReshape() {
val d = TypeMapper.DYNAMIC_DIM
val mlir = StableHloConverterFactory.createBasic().convert(graph(listOf(1, 16, 32, 128), listOf(1, 8, d, 128), causal = false, mask = listOf(1, 16, 32, d)), "gqa_dyn_head_mask").content
assertTrue(mlir.contains("stablehlo.dynamic_reshape %arg3, "), "per-head dynamic mask is split with dynamic_reshape:\n$mlir")
assertTrue(mlir.contains(": (tensor<1x16x32x?xf32>, tensor<5xi32>) -> tensor<1x8x2x32x?xf32>"), "reshape target is [b, nKV, nRep, Sq, ?]:\n$mlir")
assertFalse(mlir.contains(Regex("""stablehlo\.broadcast_in_dim %arg3[^\n]*\?""")), "no static broadcast to a dynamic type:\n$mlir")
}

@Test
fun broadcastMaskWithDynamicKeyLengthIsDynamicSafeWithoutGqa() {
// Plain multi-head attention hits the same trap: [b, 1, Sq, ?] onto scores [b, H, Sq, ?].
val d = TypeMapper.DYNAMIC_DIM
val mlir = StableHloConverterFactory.createBasic().convert(graph(listOf(1, 4, 8, 16), listOf(1, 4, d, 16), causal = false, mask = listOf(1, 1, 8, d)), "mha_dyn_mask").content
assertFalse(mlir.contains(Regex("""stablehlo\.broadcast_in_dim %arg3[^\n]*\?""")), "no static broadcast to a dynamic type:\n$mlir")
assertTrue(
mlir.contains(
"<{broadcast_dimensions = array<i64: 0, 1, 2, 3>, known_expanding_dimensions = array<i64: 1>, " +
"known_nonexpanding_dimensions = array<i64: 0, 2, 3>}> : (tensor<1x1x8x?xf32>, tensor<4xi32>) -> tensor<1x4x8x?xf32>",
),
"mask uses dynamic_broadcast_in_dim with expansion hints:\n$mlir",
)
}

@Test
fun nonDividingHeadCountsAreRejected() {
val ex = kotlin.test.assertFailsWith<HloConversionException> {
Expand Down
Loading