{% if useSubgroups is not defined %}{% set useSubgroups = true %}{% 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 %} {% set cacheSeqlensIndex = "b" %}{% set layer = 0 %} {% set cacheLen = 0 %} {% set scale = "0.0" %} {% if hasRotary is not defined %}{% set hasRotary = false %}{% endif %} {% set qHiddenV4 = qHiddenV4 | default(0) %} {% set kvHiddenV4 = kvHiddenV4 | default(0) %} {% set hasBias = hasBias is defined and hasBias %} {% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %} {% if useSubgroups %} enable subgroups; {% endif %} {{ env.wgsl.resourceDeclarations }} // Split-K flash attention. Single-partition direct output skips the merge pass. // Otherwise this is the first of two passes. This geometry // handles decode and short-query, long-context prefill inputs. // // 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 layout == "bhsd" %} // Layout: rank-4 [batch, heads, seq, headDim] for Q/K/V. {% elif layout == "layer_cache" %} // Layout: flat query [heads, headDim] plus a persistent KV cache laid out // [layer, cacheLen, kvHeads, headDim]. 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 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 layout == "bsh" %} const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u; {% elif 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; // 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) and (allowCooperativeQk if allowCooperativeQk is defined else true) %} {% 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 %} // Cross-subgroup merge that assumes nothing about which invocations share a // subgroup or how many subgroups there are: each subgroup's elected lane // publishes the subgroup pair in the slot at its OWN invocation index and sets // that index's bit in a workgroup bitmask; thread 0 then folds exactly the // published slots, in ascending index order (the online (m, d) merge is not // float-associative, so the order is fixed), and clears the mask for the next // call as it reads it. Workgroup memory starts zeroed, so the mask needs no // setup. Same three collectives as a single-subgroup reduce, two barriers. var partialM: array; var partialD: array; var leaderMask: array, (WG + 31u) / 32u>; var combinedMD: vec2; // When the whole workgroup is one subgroup the subgroup reduce already covers // it (no barriers, no shared state). `subgroup_size` is the size of the current // subgroup and uniform, so the test is exact and may guard the barriers below. 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); } if (subgroupElect()) { partialM[lidx] = sgM; partialD[lidx] = sgD; atomicOr(&leaderMask[lidx / 32u], 1u << (lidx % 32u)); } workgroupBarrier(); if (lidx == 0u) { var accM = -FLT_MAX; var accD = 0.0; for (var w = 0u; w < (WG + 31u) / 32u; w = w + 1u) { var bits = atomicExchange(&leaderMask[w], 0u); while (bits != 0u) { let slot = w * 32u + firstTrailingBit(bits); bits = bits & (bits - 1u); let mNew = max(accM, partialM[slot]); accD = accD * exp_shift(accM, mNew) + partialD[slot] * exp_shift(partialM[slot], mNew); accM = mNew; } } combinedMD = vec2(accM, accD); } workgroupBarrier(); return combinedMD; } {% else %} {% set mdExtent = "WG" %} var partialM: array; var partialD: array; 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 %} {% if layout == "layer_cache" %}{% set ATTN_SCALE_OVERRIDE = "ATTN_SCALE" %}{% endif %} {% set ATTN_SCALE_DIM = "HEAD_DIM" %}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 { return vec4({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk); }{% 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. {% set BW = "" %} {% set BC = "" %} fn load_bias4(base: u32, d4: u32) -> vec4 { let offset = base + d4 * 4u; return vec4({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }}); } {% 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 %} let split = wg.x; let h = wg.y; let b = wg.z; if (h >= Q_HEADS || split >= NUM_SPLITS{% if layout == "layer_cache" %} || params.past_len >= CACHE_LEN{% endif %}) { return; } let tid = lid.x; let hKv = h / (Q_HEADS / KV_HEADS); {% if 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[{{ cacheSeqlensIndex }}]) + 1u); {% else %} let kvSeq = cacheSeq; {% endif %} {% endif %} // Query row (decode uses token zero; short-query prefill folds the token into wg.x). {% if layout == "bsh" %} let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4; let kvBaseV4 = b * kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4; let kvTokenStrideV4 = KV_HIDDEN_V4; {% elif 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 %} let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4; 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) { let maskQuery = 0u; let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj; score = score + f32(attn_mask[maskIndex]); 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. let partialBase = ((b * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4; for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) { partial_out[partialBase + d4] = running_out[d4]; } if (tid == 0u) { let mdBase = (b * Q_HEADS + h) * NUM_SPLITS + split; // (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); } }