From 4e586b24c4574fa74ea82fb04d9e164e3bfacea2 Mon Sep 17 00:00:00 2001 From: michalharakal Date: Thu, 24 Sep 2026 10:52:18 +0200 Subject: [PATCH] fix(hlo): SDPA explicit mask under GQA and onto a dynamic key length AttentionOperationsConverter's explicit-mask path emitted invalid StableHLO in two cases (#1302): - A per-head mask [b, H, Sq, Sk] under grouped-query attention was broadcast with dims = [0, 1, 3, 4] onto the grouped scores [b, nKV, nRep, Sq, Sk], mapping H onto nKV. It is now viewed as [b, nKV, nRep, Sq, Sk] by a reshape in the h = kv * nRep + r order Q already uses, or a dynamic_reshape when a dim is dynamic. A mask head dim other than 1, nKV or H is rejected. - A mask whose shape differs from a dynamic scores shape got a static broadcast_in_dim with a dynamic result type, with or without GQA. It now uses dynamic_broadcast_in_dim with the scores-shape operand the softmax builds (hoisted so both use it), stating known_expanding_dimensions / known_nonexpanding_dimensions. Without the hints a dynamic operand dim onto a dynamic result dim is ambiguous and IREE 3.11 fails to legalize it. The op is written in generic syntax because the hints are not in every StableHLO version's pretty form. Static graphs emit exactly what they did before. Tests: four new SdpaGqaHloExportTest cases (per-head static, head-shared dynamic, per-head dynamic, non-GQA dynamic) fail on develop and pass here; :skainet-compile:skainet-compile-hlo:jvmTest 240/240. The emitted modules were compiled with iree-compile 3.11.0: the static per-head, dynamic head-shared GQA and dynamic head-shared MHA cases compile for vulkan-spirv (valhall4) and llvm-cpu (arm32). The dynamic per-head case is valid StableHLO but IREE 3.11 does not lower dynamic_reshape; that limitation is noted in the changelog. Closes #1302 --- CHANGELOG.md | 13 ++ .../AttentionOperationsConverter.kt | 135 +++++++++++++----- .../ainet/compile/hlo/SdpaGqaHloExportTest.kt | 51 +++++++ 3 files changed, 162 insertions(+), 37 deletions(-) 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 {