diff --git a/CHANGELOG.md b/CHANGELOG.md index a790c7f7..d5e044e3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 — 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 a3cb7385..98163a63 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 @@ -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" + ops += "$gr = stablehlo.reshape $gd : (tensor) -> 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 = node.inputs.getOrNull(3)?.shape ?: scoresShape + var maskVal = maskOperand + var maskType = context.getValueType(maskOperand) ?: typeOf(maskShape) + var maskDims: List = 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" + ops += "$gr = stablehlo.reshape $gd : (tensor) -> 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(); val nonExpanding = mutableListOf() + 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) = if (xs.isEmpty()) "array" else "array" + 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() @@ -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" - ops += "$gr = stablehlo.reshape $gd : (tensor) -> 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" diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/SdpaGqaHloExportTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/SdpaGqaHloExportTest.kt index d7375a01..1a01239e 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/SdpaGqaHloExportTest.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/SdpaGqaHloExportTest.kt @@ -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, known_expanding_dimensions = array, " + + "known_nonexpanding_dimensions = array}> : (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, known_expanding_dimensions = array, " + + "known_nonexpanding_dimensions = array}> : (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 {