{% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %} {% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %} {% if quantizedCache is not defined %}{% set quantizedCache = false %}{% endif %} {% if cacheSeqlens is not defined %}{% set cacheSeqlens = false %}{% endif %} {% if hasMask is not defined %}{% set hasMask = false %}{% endif %} {% if maskIsBool is not defined %}{% set maskIsBool = false %}{% endif %} {% set splitKWorkgroupSize = source.workgroupSize if source.workgroupSize is defined else tunables.WORKGROUP_SIZE %} {% if useSubgroups %} enable subgroups; {% endif %} {% if usesF16 %} enable f16; {% endif %} {{ env.wgsl.resourceDeclarations }} // Split-K flash attention, pass 1 of 2; the merge pass follows. Shared by // dense-attention decode and short-query/long-context prefill paths. // // The non-split flash decode launches only `batch * numHeads` workgroups, each // sweeping the whole KV sequence serially in WG-key tiles. This pass splits the // KV sequence into `NUM_SPLITS` contiguous ranges and gives each range its own // workgroup, so `batch * numHeads * NUM_SPLITS` workgroups run the tiled online // softmax in parallel. Each workgroup emits the *un-normalized* online state for // its range — the running (max, denom) and the softmax-weighted V sum before the // final divide — and the merge pass combines the per-split states with the online // rule. {% if source.layout == "bhsd" %} // Layout: rank-4 [batch, heads, seq, headDim] for Q/K/V. {% elif source.layout == "layer_cache" %} // Layout: flat query [heads, headDim] plus a persistent KV cache laid out // [layer, cacheLen, kvHeads, headDim]. This is the Qwen3.5 decode layout; the // dispatch has a single implicit batch. {% else %} // Layout: token-major [batch, seq, heads * headDim]; Q and KV hidden strides // are compiled constants. {% endif %} {% if (fusedQNormRope is defined and fusedQNormRope) or source.layout != "layer_cache" %}const HEAD_DIM: u32 = {{ headDim }}u; {% endif %} const HEAD_DIM_V4: u32 = {{ headDimV4 }}u; const Q_HEADS: u32 = {{ qNumHeads }}u; const KV_HEADS: u32 = {{ kvNumHeads }}u; {% if source.layout == "bsh" %} const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u; {% elif source.layout == "layer_cache" %} const LAYER: u32 = {{ layer }}u; const CACHE_LEN: u32 = {{ cacheLen }}u; const ATTN_SCALE: f32 = {{ scale }}; {% endif %} const WG: u32 = {{ splitKWorkgroupSize }}u; const NUM_SPLITS: u32 = {{ numSplits }}u; {% if splitQueries %} const Q_SEQ: u32 = {{ qSeq }}u; {% endif %} // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep // `m - m` finite so an empty lane / all--inf row contributes the exact // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret // a zero final denominator according to their public semantics. Using -inf // here changes +inf-row behavior. const FLT_MAX: f32 = 3.4028234663852886e38; fn is_finite_f32(value: f32) -> bool { return select(false, value <= FLT_MAX, value >= -FLT_MAX); } // x - m that is exactly 0 when x equals a finite m, so exp(shifted) == 1 // exactly at the row max. `x - x` on an infinite max is a legal fast-math // fold to 0, which would silently turn +inf rows finite — the explicit // equality test keeps the NaN propagation of the serial kernels. fn shifted_value(value: f32, maxValue: f32) -> f32 { let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue)); return select(value - maxValue, 0.0, equalFiniteMax); } fn exp_shift(value: f32, maxValue: f32) -> f32 { return exp(shifted_value(value, maxValue)); } var q_shared: array, HEAD_DIM_V4>; var running_out: array, HEAD_DIM_V4>; var probs: array; {% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) %} {% set jGroups = (splitKWorkgroupSize / headDimV4)|int %} {% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %} {% if coopQk %} var sval_sh: array; {% endif %} {% if jSplitV %} var vacc_sh: array, WG>; {% endif %} {% set combineSubgroups = useSubgroups %} // Workgroup-cooperative merge of per-thread online-softmax (m, d) partials: // mNew = max(m1, m2) // dNew = d1 * exp(m1 - mNew) + d2 * exp(m2 - mNew) // Both the subgroup and portable barrier-tree engines return the same merged // pair to every invocation. Repeated merges require a workgroup barrier between // calls before their shared partial storage is reused. {% set combineSubgroups = combineSubgroups is defined and combineSubgroups %} {% if combineSubgroups %} // Per-subgroup partials are published into a deterministic slot: the subgroup's // ordinal index within the workgroup (lidx / sgSize). The online (m, d) merge // is not float-associative, so thread 0 must fold partials in a fixed order. // Subgroups partition a workgroup into contiguous ordinal ranges on supported // backends, so the ordinal slot is unique per subgroup and every slot in // [0, subgroupCount) is written (each subgroup elects one leader). // Sized for the worst case of one partial per invocation. var partialM: array; var partialD: array; var combinedMD: vec2; // When the whole workgroup is one subgroup the subgroup reduce already covers // it (no barriers, no shared state); otherwise subgroup leaders publish // partials through shared memory and thread 0 folds them in ordinal order. fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2 { let sgM = subgroupMax(m); // A lane with no elements contributes d == 0 (exact identity). A +inf // element made exp(inf - inf) = NaN stick in that lane's d; a NaN element // landed in d via exp(NaN); both survive the merge and are detected by the // code after the reduction. let sgD = subgroupAdd(d * exp_shift(m, sgM)); if (sgSize == WG) { return vec2(sgM, sgD); } let subgroupCount = (WG + sgSize - 1u) / sgSize; // Pre-seed every fold slot with the (max, denom) identity. The fold below reads a // fixed subgroupCount slots in ordinal order (for determinism), but a slot whose // subgroup elects no leader this call — e.g. a fully out-of-window key tile in the // flash-attention loop that re-uses this shared memory each iteration — would // otherwise read stale shared memory. Identity makes such a slot a no-op. // (max identity = -FLT_MAX, denom identity = 0.) if (lidx < subgroupCount) { partialM[lidx] = -FLT_MAX; partialD[lidx] = 0.0; } workgroupBarrier(); if (subgroupElect()) { let slot = lidx / sgSize; partialM[slot] = sgM; partialD[slot] = sgD; } workgroupBarrier(); if (lidx == 0u) { var accM = -FLT_MAX; var accD = 0.0; for (var i = 0u; i < subgroupCount; i = i + 1u) { let mNew = max(accM, partialM[i]); accD = accD * exp_shift(accM, mNew) + partialD[i] * exp_shift(partialM[i], mNew); accM = mNew; } combinedMD = vec2(accM, accD); } workgroupBarrier(); return combinedMD; } {% else %} {% set mdStreamed = mdStreams is defined %} {% set mdStreams = mdStreams if mdStreams is defined else 1 %} {% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %} var partialM: array; var partialD: array; {% if mdStreamed %} // In-place fold of {{ mdStreams }} streams. The caller stores its per-thread // partials into partialM/partialD first and reads the merged pair of stream s // from slot s * WG afterwards. fn combine_partials_streams(lidx: u32) { workgroupBarrier(); var stride = WG / 2u; loop { if (stride == 0u) { break; } if (lidx < stride) { {% for s in range(mdStreams) %} { let slot = {{ s }}u * WG + lidx; let m1 = partialM[slot]; let d1 = partialD[slot]; let m2 = partialM[slot + stride]; let d2 = partialD[slot + stride]; let mNew = max(m1, m2); partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew); partialM[slot] = mNew; } {% endfor %} } workgroupBarrier(); stride = stride / 2u; } } {% else %} fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2 { partialM[lidx] = m; partialD[lidx] = d; workgroupBarrier(); var stride = WG / 2u; loop { if (stride == 0u) { break; } if (lidx < stride) { let m1 = partialM[lidx]; let d1 = partialD[lidx]; let m2 = partialM[lidx + stride]; let d2 = partialD[lidx + stride]; let mNew = max(m1, m2); partialD[lidx] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew); partialM[lidx] = mNew; } workgroupBarrier(); stride = stride / 2u; } let merged = vec2(partialM[0], partialD[0]); // Trailing barrier so back-to-back calls cannot race a next call's partial // stores against this call's reads of slot 0. workgroupBarrier(); return merged; } {% endif %} {% endif %} {% if source.layout == "layer_cache" %}{% set ATTN_SCALE_OVERRIDE = "ATTN_SCALE" %}{% endif %} {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %} fn scale_value() -> f32 { {% if ATTN_SCALE_OVERRIDE is defined %} return {{ ATTN_SCALE_OVERRIDE }}; {% else %} if (params.scale != 0.0) { return params.scale; } return inverseSqrt(f32({{ ATTN_SCALE_DIM }})); {% endif %} } {% if quantizedCache %} {% macro emit_quant_scale4(kind, scaleBuffer) %} fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4 { if (params.perChannel == 0u) { return vec4({{ scaleBuffer }}[0]); } let base = hk * HEAD_DIM + d4 * 4u; return vec4( {{ scaleBuffer }}[base], {{ scaleBuffer }}[base + 1u], {{ scaleBuffer }}[base + 2u], {{ scaleBuffer }}[base + 3u] ); } {%- endmacro %} {%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %} {{ emit_quant_scale4(kind, scaleBuffer) }} fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4 { {%- if format == "int8" %} return vec4({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk); {%- else %} // Two elements cover this vec4: each carries two +8-biased nibbles, low first. let rowBase = indexV4 - d4; let lo = {{ buffer }}[rowBase + d4 * 2u]; let hi = {{ buffer }}[rowBase + d4 * 2u + 1u]; let nibbles = vec4( i32(lo & 0xFu), i32((lo >> 4u) & 0xFu), i32(hi & 0xFu), i32((hi >> 4u) & 0xFu) ); let signed = nibbles - vec4(8); return vec4(signed) * {{ kind }}scale4(d4, hk); {%- endif %} } {%- endmacro %} {{ emit_quant_load4("int8", "key", "key", "k_scale") }} {{ emit_quant_load4("int8", "value", "value", "v_scale") }} {% else %} fn load_key4(indexV4: u32) -> vec4 { return vec4(key[indexV4]); } fn load_value4(indexV4: u32) -> vec4 { return vec4(value[indexV4]); } {% endif %} {% if hasBias %} // Packed [Q; K; V] bias rows (token-independent). The Q bias folds into the // query row before the Q.K dots; the K bias adds a constant to every key score // that softmax cancels, so it is skipped; the V bias is token-independent and // is applied once in the merge pass after the final normalize. fn load_bias4(base: u32, d4: u32) -> vec4 { let offset = base + d4 * 4u; return vec4(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]); } {% endif %} @compute @workgroup_size(WG, 1, 1) fn main( @builtin(workgroup_id) wg: vec3, @builtin(local_invocation_id) lid: vec3{% if useSubgroups %}, @builtin(subgroup_size) sgSize: u32{% endif %} ) { {% if useSubgroups %} // Subgroup tiles partition the fixed workgroup exactly. The advertised range // is validated before dispatch; retain this uniform guard for implementations that // choose an intermediate width at pipeline execution time. if (sgSize == 0u || sgSize > WG || WG % sgSize != 0u) { return; } {% endif %} {% if splitQueries %} let queryToken = wg.x / NUM_SPLITS; let split = wg.x % NUM_SPLITS; {% else %} let split = wg.x; {% endif %} let h = wg.y; let b = wg.z; if (h >= Q_HEADS || split >= NUM_SPLITS{% if splitQueries %} || queryToken >= Q_SEQ{% endif %}{% if source.layout == "layer_cache" %} || params.past_len >= CACHE_LEN{% endif %}) { return; } let tid = lid.x; let hKv = h / (Q_HEADS / KV_HEADS); {% if source.layout == "layer_cache" %} let kvSeq = params.past_len + 1u; {% else %} let cacheSeq = params.kvSeq; {% if cacheSeqlens %} // Buffer-sharing caches retain their capacity in the physical BNSH stride; // seqlens_k supplies the active end independently for each batch. let kvSeq = min(cacheSeq, u32(seqlens_k[b]) + 1u); {% else %} let kvSeq = cacheSeq; {% endif %} {% endif %} // Query row (decode uses token zero; short-query prefill folds the token into wg.x). {% if source.layout == "bsh" %} {% if splitQueries %} let qBaseV4 = (b * Q_SEQ + queryToken) * Q_HIDDEN_V4 + h * HEAD_DIM_V4; {% else %} let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4; {% endif %} let kvBaseV4 = b * kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4; let kvTokenStrideV4 = KV_HIDDEN_V4; {% elif source.layout == "layer_cache" %} let qBaseV4 = h * HEAD_DIM_V4; let kvBaseV4 = (LAYER * CACHE_LEN * KV_HEADS + hKv) * HEAD_DIM_V4; let kvTokenStrideV4 = KV_HEADS * HEAD_DIM_V4; {% else %} {% if splitQueries %} let qBaseV4 = ((b * Q_HEADS + h) * Q_SEQ + queryToken) * HEAD_DIM_V4; {% else %} let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4; {% endif %} let kvBaseV4 = (b * KV_HEADS + hKv) * cacheSeq * HEAD_DIM_V4; let kvTokenStrideV4 = HEAD_DIM_V4; {% endif %} // Contiguous KV range owned by this split. Ceil division lets the last split // absorb any remainder; empty ranges write identity partials and are ignored // by the merge pass. {% if hasWindow %} // Sliding window on the single decode query (absolute position // kvSeq-1): it attends only the last `windowSize` keys, so split the // contiguous [windowStart, kvSeq) range instead of the whole cache. var windowStart: u32 = 0u; if (kvSeq > params.windowSize) { windowStart = kvSeq - params.windowSize; } let activeKeys = kvSeq - windowStart; let keysPerSplit = (activeKeys + NUM_SPLITS - 1u) / NUM_SPLITS; let splitStart = windowStart + split * keysPerSplit; {% else %} let keysPerSplit = (kvSeq + NUM_SPLITS - 1u) / NUM_SPLITS; let splitStart = split * keysPerSplit; {% endif %} var splitEnd = splitStart + keysPerSplit; if (splitEnd > kvSeq) { splitEnd = kvSeq; } {% set hasBias = hasBias is defined and hasBias %} for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) { var qv = vec4(query[qBaseV4 + d4]); {% if hasBias %} qv = qv + load_bias4(h * HEAD_DIM, d4); {% endif %} q_shared[d4] = qv; running_out[d4] = vec4(0.0); } workgroupBarrier(); let scale = scale_value(); var runningMax = -FLT_MAX; var runningDenom = 0.0; var kjBase = splitStart; loop { if (kjBase >= splitEnd) { break; } let kj = kjBase + tid; var keyAllowed = kj < splitEnd; let tileCount = min(WG, splitEnd - kjBase); var score = -FLT_MAX; var m = -FLT_MAX; var dPart = 0.0; {% if coopQk %} // Cooperative Q.K: one subgroup per key, lanes splitting HEAD_DIM_V4, then a hardware // subgroupAdd — turns the per-thread HEAD_DIM_V4-long dependent dot chain into a few // strided vec4 dots + one reduce. Uniform trip count keeps subgroupAdd in uniform flow. let sgPerWg = WG / sgSize; let qkRounds = (tileCount + sgPerWg - 1u) / sgPerWg; let lane = tid % sgSize; let sgInWg = tid / sgSize; for (var rr: u32 = 0u; rr < qkRounds; rr = rr + 1u) { let j = rr * sgPerWg + sgInWg; var accS: f32 = 0.0; if (j < tileCount) { let kRowV4 = kvBaseV4 + (kjBase + j) * kvTokenStrideV4; for (var d4: u32 = lane; d4 < HEAD_DIM_V4; d4 = d4 + sgSize) { accS = accS + dot(q_shared[d4], load_key4(kRowV4 + d4{% if quantizedCache %}, d4, hKv{% endif %})); } } let sj = subgroupAdd(accS); if (lane == 0u && j < tileCount) { sval_sh[j] = sj; } } workgroupBarrier(); if (keyAllowed) { score = sval_sh[tid] * scale; m = score; dPart = 1.0; } {% else %} if (keyAllowed) { let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4; var acc: f32 = 0.0; for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) { acc = acc + dot(q_shared[d4], load_key4(kRowV4 + d4{% if quantizedCache %}, d4, hKv{% endif %})); } score = acc * scale; m = score; dPart = 1.0; } {% endif %} {% if hasMask %} if (keyAllowed) { {% if splitQueries %} let maskQuery = queryToken; {% else %} let maskQuery = 0u; {% endif %} let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj; {% if maskIsBool %} // A rejected bool-mask key contributes no probability mass. The merge // pass already maps a zero global denominator to an all-zero output row. if (attn_mask[maskIndex] == 0u) { keyAllowed = false; score = -FLT_MAX; dPart = 0.0; } {% else %} score = score + f32(attn_mask[maskIndex]); {% endif %} m = score; } {% endif %} let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %}); // Merge one key tile's online-softmax (maximum, denominator) partial into the // running state, then store the per-key probabilities consumed by V accumulation. let newMax = max(runningMax, tile.x); let correction = exp_shift(runningMax, newMax); runningDenom = runningDenom * correction + tile.y * exp_shift(tile.x, newMax); runningMax = newMax; var prob = 0.0; if (keyAllowed) { prob = exp_shift(score, newMax); } probs[tid] = prob; workgroupBarrier(); {% if jSplitV %} // j-split V accumulation: thread (jg, d4v) sums keys j == jg mod // J_GROUPS for dim block d4v into a register, then the groups combine // through shared memory so all lanes participate. const J_GROUPS: u32 = {{ jGroups }}u; let jg = tid / HEAD_DIM_V4; let d4v = tid % HEAD_DIM_V4; var vacc = vec4(0.0); var jj = jg; loop { if (jj >= tileCount) { break; } vacc = vacc + probs[jj] * load_value4( kvBaseV4 + (kjBase + jj) * kvTokenStrideV4 + d4v{% if quantizedCache %}, d4v, hKv{% endif %} ); jj = jj + J_GROUPS; } vacc_sh[tid] = vacc; workgroupBarrier(); for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) { var a4 = running_out[d4] * correction; for (var g: u32 = 0u; g < J_GROUPS; g = g + 1u) { a4 = a4 + vacc_sh[g * HEAD_DIM_V4 + d4]; } running_out[d4] = a4; } workgroupBarrier(); {% else %} for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) { var vSum = vec4(0.0); for (var i: u32 = 0u; i < tileCount; i = i + 1u) { vSum = vSum + probs[i] * load_value4( kvBaseV4 + (kjBase + i) * kvTokenStrideV4 + d4{% if quantizedCache %}, d4, hKv{% endif %} ); } running_out[d4] = running_out[d4] * correction + vSum; } workgroupBarrier(); {% endif %} kjBase = kjBase + WG; } // Emit un-normalized partials for (b, h, split): the merge pass divides. {% if splitQueries %} let partialBase = (((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4; {% else %} let partialBase = ((b * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4; {% endif %} for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) { partial_out[partialBase + d4] = running_out[d4]; } if (tid == 0u) { {% if splitQueries %} let mdBase = ((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + split; {% else %} let mdBase = (b * Q_HEADS + h) * NUM_SPLITS + split; {% endif %} // (max, denom) travel together to the merge, so they share one buffer as an // interleaved vec2 rather than costing two bindings. Interleaved, not two // halves, so the index needs no region size — and the merge reads both // fields of a split in a single load. partial_stats[mdBase] = vec2(runningMax, runningDenom); } }