Xenova HF Staff commited on
Commit
ba9df6e
·
verified ·
1 Parent(s): 78d4fcb

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -87,19 +87,16 @@ One implementation is selected per call from the device capabilities, the reques
87
  - `window_shift_materialized_sgmat_f32` — Materializes causal chunked-prefill attention with a windowed cache. The shift pass compacts surviving rows and appends the chunk; score and apply passes enforce both the sliding-window floor and causal bound.
88
  - `share_append_materialized_sgmat_f32` — Materializes causal chunked-prefill attention while updating a shared-capacity cache in place. Score and apply passes bound every causal tile by the live length from `seqlens_k`, not the buffer capacity, so right-padded batches remain left-aligned.
89
  - `new_kv_share_append_split` — Retains the existing cache and appends new key/value rows in separate passes before portable attention. It avoids rebuilding unchanged cache positions when the past and present allocations share capacity.
 
90
  - `qkv_present_tiled_nosg` — Portable tiled prefill route that computes attention online and writes the present cache separately. It is used when the flash shape is valid but no suitable subgroup route is admissible.
 
 
91
  - `quant_int8_decode_splitk` — Appends int8-quantized key/value rows, partitions cached decode across the key axis, and merges partial online-softmax results. Used when one workgroup per head exposes too little independent work.
92
  - `qkv_present_flash_splitk` — Partitions direct Q/K/V attention across the key axis, merges partial online-softmax results, and writes the present cache separately. It serves short-query shapes needing more key-axis parallelism.
93
  - `qkv_present_flash_cluster` — Computes clustered online-softmax prefill directly from Q/K/V and writes the present cache separately. The family provides subgroup and portable reductions for the same tiled algorithm.
94
- - `quant_int8_decode_splitk_nosg` — Appends int8-quantized key/value rows, partitions cached decode across the key axis, and merges partial online-softmax results. Used when one workgroup per head exposes too little independent work.
95
- - `qkv_present_flash_splitk_nosg` — Partitions direct Q/K/V attention across the key axis, merges partial online-softmax results, and writes the present cache separately. It serves short-query shapes needing more key-axis parallelism.
96
- - `qkv_present_flash_cluster_nosg` — Computes clustered online-softmax prefill directly from Q/K/V and writes the present cache separately. The family provides subgroup and portable reductions for the same tiled algorithm.
97
  - `past_kv_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
98
  - `new_kv_past_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
99
  - `window_shift_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
100
- - `past_kv_decode_splitk_nosg` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
101
- - `new_kv_past_decode_splitk_nosg` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
102
- - `window_shift_decode_splitk_nosg` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
103
  - `past_kv_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
104
  - `new_kv_past_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
105
  - `past_kv_rotary_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
@@ -107,21 +104,14 @@ One implementation is selected per call from the device capabilities, the reques
107
  - `past_kv_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
108
  - `past_kv_bias_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
109
  - `window_shift_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
110
- - `past_kv_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
111
- - `new_kv_past_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
112
- - `past_kv_rotary_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
113
- - `past_kv_softcap_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
114
- - `past_kv_headsink_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
115
- - `past_kv_bias_headsink_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
116
- - `window_shift_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
117
  - `quant_int8_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
118
  - `quant_int4_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
119
- - `quant_int8_flash_prefill_nosg` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
120
- - `quant_int4_flash_prefill_nosg` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
121
  - `share_append_split_decode_splitk` — Retains shared-capacity cache rows, appends new rows separately, then partitions decode across the key axis. It avoids rebuilding unchanged cache positions while exposing split-key parallelism.
122
- - `share_append_split_decode_splitk_nosg` — Retains shared-capacity cache rows, appends new rows separately, then partitions decode across the key axis. It avoids rebuilding unchanged cache positions while exposing split-key parallelism.
123
  - `share_append_split_flash_prefill` — Retains shared-capacity cache rows, appends new rows separately, then applies clustered causal prefill. It avoids rebuilding unchanged cache positions and includes subgroup and portable forms.
124
- - `share_append_split_flash_prefill_nosg` — Retains shared-capacity cache rows, appends new rows separately, then applies clustered causal prefill. It avoids rebuilding unchanged cache positions and includes subgroup and portable forms.
125
 
126
  ## Device requirements
127
 
@@ -132,7 +122,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
132
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
133
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
134
  - [`test.json`](build/webgpu/test.json) — correctness cases
135
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
136
  - [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
137
  - [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
138
  - [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
@@ -149,7 +139,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
149
  ## Use with `@huggingface/kernels`
150
 
151
  ```sh
152
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
153
  ```
154
 
155
  Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
@@ -170,18 +160,16 @@ import { getKernel } from "@huggingface/kernels";
170
  const kernel = await getKernel("webgpu-kernels/com.microsoft.GroupQueryAttention", { version: 1 });
171
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
172
  const { outputT, presentKeyT, presentValueT } = await kernel({
173
- queryT: { data: queryTData, shape: [2, 1, 16] },
174
- keyT: { data: keyTData, shape: [2, 1, 8] },
175
- valueT: { data: valueTData, shape: [2, 1, 8] },
176
- pastKeyT: { data: pastKeyTData, shape: [2, 1, 8, 8] },
177
- pastValueT: { data: pastValueTData, shape: [2, 1, 8, 8] },
178
- seqlensKT: { data: seqlensKTData, shape: [2] },
179
  totalSequenceLengthT: { data: totalSequenceLengthTData, shape: [1] },
180
  }, {
181
- attrs: { num_heads: 2, kv_num_heads: 1 },
182
  outputs: {
183
- presentKeyT: { shape: [2, 1, 8, 8], dtype: "float32" },
184
- presentValueT: { shape: [2, 1, 8, 8], dtype: "float32" },
185
  },
186
  });
187
  ```
 
87
  - `window_shift_materialized_sgmat_f32` — Materializes causal chunked-prefill attention with a windowed cache. The shift pass compacts surviving rows and appends the chunk; score and apply passes enforce both the sliding-window floor and causal bound.
88
  - `share_append_materialized_sgmat_f32` — Materializes causal chunked-prefill attention while updating a shared-capacity cache in place. Score and apply passes bound every causal tile by the live length from `seqlens_k`, not the buffer capacity, so right-padded batches remain left-aligned.
89
  - `new_kv_share_append_split` — Retains the existing cache and appends new key/value rows in separate passes before portable attention. It avoids rebuilding unchanged cache positions when the past and present allocations share capacity.
90
+ - `new_kv_share_append_bidirectional_split` — Bidirectional (`causal = 0`) attention over a shared-capacity cache: the new key/value rows are appended exactly as in the causal sibling, then every query attends to the whole live range `[0, seqlens_k[b] + 1)` instead of stopping at its own position. ONNX Runtime forbids `causal = 0` with a local window, so the window floor never applies here.
91
  - `qkv_present_tiled_nosg` — Portable tiled prefill route that computes attention online and writes the present cache separately. It is used when the flash shape is valid but no suitable subgroup route is admissible.
92
+ - `qkv_present_flash_causal` — Causal prefill without a past cache and without a sliding window (`causal=1`, `local_window_size=-1`): the spelling onnxruntime's GroupQueryAttention tests use for a first prompt. The twin of `qkv_present_flash` with the causal ceiling enabled; the key sequence equals the query sequence, so the upper-left ceiling is the right one.
93
+ - `qkv_present_causal` — Causal prefill without a past cache and without a sliding window (`causal=1`, `local_window_size=-1`): the spelling onnxruntime's GroupQueryAttention tests use for a first prompt. The twin of `qkv_present` with the causal ceiling enabled; the key sequence equals the query sequence, so the upper-left ceiling is the right one.
94
  - `quant_int8_decode_splitk` — Appends int8-quantized key/value rows, partitions cached decode across the key axis, and merges partial online-softmax results. Used when one workgroup per head exposes too little independent work.
95
  - `qkv_present_flash_splitk` — Partitions direct Q/K/V attention across the key axis, merges partial online-softmax results, and writes the present cache separately. It serves short-query shapes needing more key-axis parallelism.
96
  - `qkv_present_flash_cluster` — Computes clustered online-softmax prefill directly from Q/K/V and writes the present cache separately. The family provides subgroup and portable reductions for the same tiled algorithm.
 
 
 
97
  - `past_kv_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
98
  - `new_kv_past_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
99
  - `window_shift_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
 
 
 
100
  - `past_kv_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
101
  - `new_kv_past_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
102
  - `past_kv_rotary_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
 
104
  - `past_kv_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
105
  - `past_kv_bias_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
106
  - `window_shift_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
107
+ - `window_shift_rotary_append` — Windowed cache with rotary embeddings. The shift pass slides surviving rows down by the eviction count and clears what left the window; the append pass writes the new key/value rows and rotates K at its ABSOLUTE position (window origin plus row, not the cache row); attention rotates Q at the same absolute position and masks cache-relative.
108
+ - `window_shift_rotary_headsink_append` — Windowed rotary cache with a per-head softmax sink (the gpt-oss attention layer). The sink contributes `exp(sink - m)` to the denominator once per query row and nothing to the numerator.
109
+ - `window_shift_rotary_bias_append` — Windowed rotary cache with an additive attention bias. The bias row is indexed by ABSOLUTE key position with the tensor's own last dimension as the stride (ONNX spells it `total_sequence_length`, which keeps counting past the cache capacity), so resident cache row j reads bias column `origin + j`.
 
 
 
 
110
  - `quant_int8_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
111
  - `quant_int4_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
 
 
112
  - `share_append_split_decode_splitk` — Retains shared-capacity cache rows, appends new rows separately, then partitions decode across the key axis. It avoids rebuilding unchanged cache positions while exposing split-key parallelism.
 
113
  - `share_append_split_flash_prefill` — Retains shared-capacity cache rows, appends new rows separately, then applies clustered causal prefill. It avoids rebuilding unchanged cache positions and includes subgroup and portable forms.
114
+ - `share_append_bidirectional_flash_prefill` — Retains and appends a shared-capacity cache, then applies clustered bidirectional prefill to the full live key range.
115
 
116
  ## Device requirements
117
 
 
122
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
123
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
124
  - [`test.json`](build/webgpu/test.json) — correctness cases
125
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
126
  - [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
127
  - [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
128
  - [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
 
139
  ## Use with `@huggingface/kernels`
140
 
141
  ```sh
142
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
143
  ```
144
 
145
  Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
 
160
  const kernel = await getKernel("webgpu-kernels/com.microsoft.GroupQueryAttention", { version: 1 });
161
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
162
  const { outputT, presentKeyT, presentValueT } = await kernel({
163
+ queryT: { data: queryTData, shape: [1, 2, 8] },
164
+ keyT: { data: keyTData, shape: [1, 2, 8] },
165
+ valueT: { data: valueTData, shape: [1, 2, 8] },
166
+ seqlensKT: { data: seqlensKTData, shape: [1] },
 
 
167
  totalSequenceLengthT: { data: totalSequenceLengthTData, shape: [1] },
168
  }, {
169
+ attrs: { num_heads: 1, kv_num_heads: 1 },
170
  outputs: {
171
+ presentKeyT: { shape: [1, 1, 2, 8], dtype: "float32" },
172
+ presentValueT: { shape: [1, 1, 2, 8], dtype: "float32" },
173
  },
174
  });
175
  ```
build/webgpu/attention-rank4-tiled.wgsl.jinja CHANGED
@@ -1,4 +1,4 @@
1
- {% set layout = layout if layout is defined else "bnsh" %}
2
  {{ env.wgsl.resourceDeclarations }}
3
 
4
  // Thread-per-query online-softmax flash fallback for rank-4 (BNSH) attention with
@@ -16,13 +16,11 @@ const BLOCK_M: u32 = {{ blockM }}u;
16
  var<workgroup> acc: array<f32, {{ vHeadCap * blockM }}>;
17
 
18
  {% set ATTN_SCALE_DIM = "params.headSize" %}
19
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
20
  fn scale_value() -> f32 {
21
  if (params.scale != 0.0) { return params.scale; }
22
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
23
  }
24
 
25
-
26
  fn kv_head(q_head: u32) -> u32 {
27
  return q_head / (params.qHeads / params.kvHeads);
28
  }
@@ -51,14 +49,8 @@ fn main(
51
  if (qs >= params.qSeq) { return; }
52
  let kh = kv_head(qh);
53
 
54
- {% if layout == "bsh" %}
55
  // Token-major packed QKV: row = (batch*seq + token)*HIDDEN + head*head_dim.
56
  let qBase = (batch * params.qSeq + qs) * params.qHidden + qh * params.headSize;
57
- {% else %}
58
- let qBase = ((batch * params.qHeads + qh) * params.qSeq + qs) * params.headSize;
59
- let kBase = (batch * params.kvHeads + kh) * params.kvSeq;
60
- let vBase = (batch * params.kvHeads + kh) * params.kvSeq;
61
- {% endif %}
62
  let scale = scale_value();
63
 
64
  for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
@@ -76,21 +68,13 @@ fn main(
76
 
77
  for (var ks: u32 = 0u; ks < maxKj; ks = ks + 1u) {
78
  var masked = false;
79
- {% if hasMask and maskIsBool %}
80
- let mIdxB = batch * params.maskBatchStride + qh * params.maskHeadStride + qs * params.maskSeqStride + ks;
81
- masked = attn_mask[mIdxB] == 0u;
82
- {% endif %}
83
  // A rejected bool-mask key has zero softmax mass. Skipping it is safe here:
84
  // this thread-per-query kernel has no barriers inside the key loop.
85
  if (masked) { continue; }
86
  // Adjacent query threads load the same K/V row for this key.
87
  var score: f32 = -3.4028234663852886e38;
88
  if (!masked) {
89
- {% if layout == "bsh" %}
90
  let kRow = (batch * params.kvSeq + ks) * params.kvHidden + kh * params.headSize;
91
- {% else %}
92
- let kRow = (kBase + ks) * params.headSize;
93
- {% endif %}
94
  var dot: f32 = 0.0;
95
  for (var d: u32 = 0u; d < params.headSize; d = d + 1u) {
96
  dot = dot + f32(q[qBase + d]) * f32(k[kRow + d]);
@@ -110,22 +94,14 @@ fn main(
110
  let weight = exp(score - next_max);
111
  running_max = next_max;
112
  running_denom = running_denom * prev_scale + weight;
113
- {% if layout == "bsh" %}
114
  let vRow = (batch * params.kvSeq + ks) * params.vHidden + kh * params.vHeadSize;
115
- {% else %}
116
- let vRow = (vBase + ks) * params.vHeadSize;
117
- {% endif %}
118
  for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
119
  acc[d * BLOCK_M + tid] = acc[d * BLOCK_M + tid] * prev_scale + weight * f32(v[vRow + d]);
120
  }
121
  }
122
 
123
  let inv_denom = select(0.0, 1.0 / running_denom, running_denom > 0.0);
124
- {% if layout == "bsh" %}
125
  let yBase = (batch * params.qSeq + qs) * (params.qHeads * params.vHeadSize) + qh * params.vHeadSize;
126
- {% else %}
127
- let yBase = ((batch * params.qHeads + qh) * params.qSeq + qs) * params.vHeadSize;
128
- {% endif %}
129
  for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
130
  y[yBase + d] = {{ scalar }}(acc[d * BLOCK_M + tid] * inv_denom);
131
  }
 
1
+ {% set maskIsBool = false %}
2
  {{ env.wgsl.resourceDeclarations }}
3
 
4
  // Thread-per-query online-softmax flash fallback for rank-4 (BNSH) attention with
 
16
  var<workgroup> acc: array<f32, {{ vHeadCap * blockM }}>;
17
 
18
  {% set ATTN_SCALE_DIM = "params.headSize" %}
 
19
  fn scale_value() -> f32 {
20
  if (params.scale != 0.0) { return params.scale; }
21
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
22
  }
23
 
 
24
  fn kv_head(q_head: u32) -> u32 {
25
  return q_head / (params.qHeads / params.kvHeads);
26
  }
 
49
  if (qs >= params.qSeq) { return; }
50
  let kh = kv_head(qh);
51
 
 
52
  // Token-major packed QKV: row = (batch*seq + token)*HIDDEN + head*head_dim.
53
  let qBase = (batch * params.qSeq + qs) * params.qHidden + qh * params.headSize;
 
 
 
 
 
54
  let scale = scale_value();
55
 
56
  for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
 
68
 
69
  for (var ks: u32 = 0u; ks < maxKj; ks = ks + 1u) {
70
  var masked = false;
 
 
 
 
71
  // A rejected bool-mask key has zero softmax mass. Skipping it is safe here:
72
  // this thread-per-query kernel has no barriers inside the key loop.
73
  if (masked) { continue; }
74
  // Adjacent query threads load the same K/V row for this key.
75
  var score: f32 = -3.4028234663852886e38;
76
  if (!masked) {
 
77
  let kRow = (batch * params.kvSeq + ks) * params.kvHidden + kh * params.headSize;
 
 
 
78
  var dot: f32 = 0.0;
79
  for (var d: u32 = 0u; d < params.headSize; d = d + 1u) {
80
  dot = dot + f32(q[qBase + d]) * f32(k[kRow + d]);
 
94
  let weight = exp(score - next_max);
95
  running_max = next_max;
96
  running_denom = running_denom * prev_scale + weight;
 
97
  let vRow = (batch * params.kvSeq + ks) * params.vHidden + kh * params.vHeadSize;
 
 
 
98
  for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
99
  acc[d * BLOCK_M + tid] = acc[d * BLOCK_M + tid] * prev_scale + weight * f32(v[vRow + d]);
100
  }
101
  }
102
 
103
  let inv_denom = select(0.0, 1.0 / running_denom, running_denom > 0.0);
 
104
  let yBase = (batch * params.qSeq + qs) * (params.qHeads * params.vHeadSize) + qh * params.vHeadSize;
 
 
 
105
  for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
106
  y[yBase + d] = {{ scalar }}(acc[d * BLOCK_M + tid] * inv_denom);
107
  }
build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja CHANGED
@@ -1,9 +1,6 @@
1
- {% if scalar == "f16" %}
2
- enable f16;
3
- {% endif %}
4
- {% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
5
- {% if hasGate is not defined %}{% set hasGate = false %}{% endif %}
6
- {% set qSeq = qSeq | default(0) %}
7
  {{ env.wgsl.resourceDeclarations }}
8
 
9
  // Split-K flash decode, pass 2 of 2. Combines the per-split un-normalized online
@@ -25,9 +22,6 @@ enable f16;
25
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
26
  const Q_HEADS: u32 = {{ qNumHeads }}u;
27
  const NUM_SPLITS: u32 = {{ numSplits }}u;
28
- {% if splitQueries %}
29
- const Q_SEQ: u32 = {{ qSeq }}u;
30
- {% endif %}
31
  {% if layout == "bsh" %}
32
  const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u;
33
  {% endif %}
@@ -57,6 +51,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
57
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
58
  return select(value - maxValue, 0.0, equalFiniteMax);
59
  }
 
60
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
61
  return exp(shifted_value(value, maxValue));
62
  }
@@ -65,21 +60,14 @@ fn main(
65
  @builtin(workgroup_id) wg: vec3<u32>,
66
  @builtin(local_invocation_id) lid: vec3<u32>
67
  ) {
68
- {% if splitQueries %}
69
- let queryToken = wg.x;
70
- {% endif %}
71
  let h = wg.y;
72
  let b = wg.z;
73
  let d4 = lid.x;
74
- if (h >= Q_HEADS || d4 >= HEAD_DIM_V4{% if splitQueries %} || queryToken >= Q_SEQ{% endif %}) {
75
  return;
76
  }
77
 
78
- {% if splitQueries %}
79
- let mdBase = ((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS;
80
- {% else %}
81
  let mdBase = (b * Q_HEADS + h) * NUM_SPLITS;
82
- {% endif %}
83
  var globalMax = -FLT_MAX;
84
  for (var s: u32 = 0u; s < NUM_SPLITS; s = s + 1u) {
85
  globalMax = max(globalMax, partial_stats[mdBase + s].x);
@@ -90,28 +78,16 @@ fn main(
90
  let stats = partial_stats[mdBase + s];
91
  let w = exp_shift(stats.x, globalMax);
92
  globalDenom = globalDenom + stats.y * w;
93
- {% if splitQueries %}
94
- let pBase = (((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + s) * HEAD_DIM_V4;
95
- {% else %}
96
  let pBase = ((b * Q_HEADS + h) * NUM_SPLITS + s) * HEAD_DIM_V4;
97
- {% endif %}
98
  outV = outV + partial_out[pBase + d4] * w;
99
  }
100
 
101
  {% if layout == "bsh" %}
102
- {% if splitQueries %}
103
- let qBaseV4 = (b * Q_SEQ + queryToken) * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
104
- {% else %}
105
  let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
106
- {% endif %}
107
  {% elif layout == "layer_cache" %}
108
  let qBaseV4 = h * HEAD_DIM_V4;
109
- {% else %}
110
- {% if splitQueries %}
111
- let qBaseV4 = ((b * Q_HEADS + h) * Q_SEQ + queryToken) * HEAD_DIM_V4;
112
  {% else %}
113
  let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
114
- {% endif %}
115
  {% endif %}
116
  // A split that covers no keys writes vec2(-FLT_MAX, 0), so when every split is empty —
117
  // a query whose window admits nothing — globalMax stays -FLT_MAX, exp_shift(x, x) is 1,
@@ -122,11 +98,6 @@ fn main(
122
  // V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
123
  let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
124
  outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
125
- {% endif %}
126
- {% if hasGate %}
127
- // The gated route multiplies the normalized attention output elementwise by its gate.
128
- let gateV = vec4<f32>(gate[qBaseV4 + d4]);
129
- outValue = outValue * (vec4<f32>(1.0) / (vec4<f32>(1.0) + exp(-gateV)));
130
  {% endif %}
131
  output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
132
  }
 
1
+ {% set qHiddenV4 = qHiddenV4 | default(0) %}
2
+ {% set qHidden = qHidden | default(0) %}
3
+ {% set hasBias = hasBias is defined and hasBias %}
 
 
 
4
  {{ env.wgsl.resourceDeclarations }}
5
 
6
  // Split-K flash decode, pass 2 of 2. Combines the per-split un-normalized online
 
22
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
23
  const Q_HEADS: u32 = {{ qNumHeads }}u;
24
  const NUM_SPLITS: u32 = {{ numSplits }}u;
 
 
 
25
  {% if layout == "bsh" %}
26
  const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u;
27
  {% endif %}
 
51
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
52
  return select(value - maxValue, 0.0, equalFiniteMax);
53
  }
54
+
55
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
56
  return exp(shifted_value(value, maxValue));
57
  }
 
60
  @builtin(workgroup_id) wg: vec3<u32>,
61
  @builtin(local_invocation_id) lid: vec3<u32>
62
  ) {
 
 
 
63
  let h = wg.y;
64
  let b = wg.z;
65
  let d4 = lid.x;
66
+ if (h >= Q_HEADS || d4 >= HEAD_DIM_V4) {
67
  return;
68
  }
69
 
 
 
 
70
  let mdBase = (b * Q_HEADS + h) * NUM_SPLITS;
 
71
  var globalMax = -FLT_MAX;
72
  for (var s: u32 = 0u; s < NUM_SPLITS; s = s + 1u) {
73
  globalMax = max(globalMax, partial_stats[mdBase + s].x);
 
78
  let stats = partial_stats[mdBase + s];
79
  let w = exp_shift(stats.x, globalMax);
80
  globalDenom = globalDenom + stats.y * w;
 
 
 
81
  let pBase = ((b * Q_HEADS + h) * NUM_SPLITS + s) * HEAD_DIM_V4;
 
82
  outV = outV + partial_out[pBase + d4] * w;
83
  }
84
 
85
  {% if layout == "bsh" %}
 
 
 
86
  let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
 
87
  {% elif layout == "layer_cache" %}
88
  let qBaseV4 = h * HEAD_DIM_V4;
 
 
 
89
  {% else %}
90
  let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
 
91
  {% endif %}
92
  // A split that covers no keys writes vec2(-FLT_MAX, 0), so when every split is empty —
93
  // a query whose window admits nothing — globalMax stays -FLT_MAX, exp_shift(x, x) is 1,
 
98
  // V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
99
  let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
100
  outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
 
 
 
 
 
101
  {% endif %}
102
  output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
103
  }
build/webgpu/attn-flash-decode-splitk.wgsl.jinja CHANGED
@@ -1,21 +1,22 @@
1
  {% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
2
- {% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
3
  {% if quantizedCache is not defined %}{% set quantizedCache = false %}{% endif %}
4
  {% if cacheSeqlens is not defined %}{% set cacheSeqlens = false %}{% endif %}
5
  {% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
6
- {% if maskIsBool is not defined %}{% set maskIsBool = false %}{% endif %}
7
- {% set layer = layer | default(0) %}
8
- {% set cacheLen = cacheLen | default(0) %}
9
- {% set scale = scale | default("0.0") %}
10
  {% if hasRotary is not defined %}{% set hasRotary = false %}{% endif %}
11
- {% set qSeq = qSeq | default(0) %}
 
 
12
  {% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
13
  {% if useSubgroups %}
14
  enable subgroups;
15
  {% endif %}
16
  {{ env.wgsl.resourceDeclarations }}
17
 
18
- // Split-K flash attention, pass 1 of 2; the merge pass follows. This geometry
 
19
  // handles decode and short-query, long-context prefill inputs.
20
  //
21
  // The non-split flash decode launches only `batch * numHeads` workgroups, each
@@ -50,9 +51,6 @@ const ATTN_SCALE: f32 = {{ scale }};
50
  {% endif %}
51
  const WG: u32 = {{ splitKWorkgroupSize }}u;
52
  const NUM_SPLITS: u32 = {{ numSplits }}u;
53
- {% if splitQueries %}
54
- const Q_SEQ: u32 = {{ qSeq }}u;
55
- {% endif %}
56
  // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
57
  // `m - m` finite so an empty lane / all--inf row contributes the exact
58
  // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
@@ -72,6 +70,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
72
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
73
  return select(value - maxValue, 0.0, equalFiniteMax);
74
  }
 
75
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
76
  return exp(shifted_value(value, maxValue));
77
  }
@@ -79,7 +78,7 @@ fn exp_shift(value: f32, maxValue: f32) -> f32 {
79
  var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
80
  var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
81
  var<workgroup> probs: array<f32, WG>;
82
- {% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) %}
83
  {% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
84
  {% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
85
  {% if coopQk %}
@@ -148,41 +147,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
148
  return combinedMD;
149
  }
150
  {% else %}
151
- {% set mdStreamed = mdStreams is defined %}
152
- {% set mdStreams = mdStreams if mdStreams is defined else 1 %}
153
- {% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
154
  var<workgroup> partialM: array<f32, {{ mdExtent }}>;
155
  var<workgroup> partialD: array<f32, {{ mdExtent }}>;
156
- {% if mdStreamed %}
157
-
158
- // In-place fold of {{ mdStreams }} streams. Input partials occupy
159
- // partialM/partialD; stream s returns its merged pair in slot s * WG.
160
- fn combine_partials_streams(lidx: u32) {
161
- workgroupBarrier();
162
- var stride = WG / 2u;
163
- loop {
164
- if (stride == 0u) {
165
- break;
166
- }
167
- if (lidx < stride) {
168
- {% for s in range(mdStreams) %}
169
- {
170
- let slot = {{ s }}u * WG + lidx;
171
- let m1 = partialM[slot];
172
- let d1 = partialD[slot];
173
- let m2 = partialM[slot + stride];
174
- let d2 = partialD[slot + stride];
175
- let mNew = max(m1, m2);
176
- partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
177
- partialM[slot] = mNew;
178
- }
179
- {% endfor %}
180
- }
181
- workgroupBarrier();
182
- stride = stride / 2u;
183
- }
184
- }
185
- {% else %}
186
 
187
  fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
188
  partialM[lidx] = m;
@@ -212,12 +179,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
212
  return merged;
213
  }
214
  {% endif %}
215
- {% endif %}
216
-
217
 
218
  {% if layout == "layer_cache" %}{% set ATTN_SCALE_OVERRIDE = "ATTN_SCALE" %}{% endif %}
219
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
220
- fn scale_value() -> f32 {
221
  {% if ATTN_SCALE_OVERRIDE is defined %}
222
  return {{ ATTN_SCALE_OVERRIDE }};
223
  {% else %}
@@ -226,7 +190,6 @@ fn scale_value() -> f32 {
226
  {% endif %}
227
  }
228
 
229
-
230
  {% if quantizedCache %}
231
  {% macro emit_quant_scale4(kind, scaleBuffer) %}
232
  fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
@@ -240,28 +203,12 @@ fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
240
  {{ scaleBuffer }}[base + 2u],
241
  {{ scaleBuffer }}[base + 3u]
242
  );
243
- }
244
- {%- endmacro %}
245
- {%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
246
  {{ emit_quant_scale4(kind, scaleBuffer) }}
247
  fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
248
- {%- if format == "int8" %}
249
  return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
250
- {%- else %}
251
- // Two elements cover this vec4: each carries two +8-biased nibbles, low first.
252
- let rowBase = indexV4 - d4;
253
- let lo = {{ buffer }}[rowBase + d4 * 2u];
254
- let hi = {{ buffer }}[rowBase + d4 * 2u + 1u];
255
- let nibbles = vec4<i32>(
256
- i32(lo & 0xFu), i32((lo >> 4u) & 0xFu),
257
- i32(hi & 0xFu), i32((hi >> 4u) & 0xFu)
258
- );
259
- let signed = nibbles - vec4<i32>(8);
260
- return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
261
- {%- endif %}
262
- }
263
- {%- endmacro %}
264
-
265
  {{ emit_quant_load4("int8", "key", "key", "k_scale") }}
266
  {{ emit_quant_load4("int8", "value", "value", "v_scale") }}
267
  {% else %}
@@ -279,9 +226,11 @@ fn load_value4(indexV4: u32) -> vec4<f32> {
279
  // query row before the Q.K dots; the K bias adds a constant to every key score
280
  // that softmax cancels, so it is skipped; the V bias is token-independent and
281
  // is applied once in the merge pass after the final normalize.
 
 
282
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
283
  let offset = base + d4 * 4u;
284
- return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
285
  }
286
 
287
  {% endif %}
@@ -297,15 +246,10 @@ fn main(
297
  // choose an intermediate width at pipeline execution time.
298
  if (sgSize == 0u || sgSize > WG || WG % sgSize != 0u) { return; }
299
  {% endif %}
300
- {% if splitQueries %}
301
- let queryToken = wg.x / NUM_SPLITS;
302
- let split = wg.x % NUM_SPLITS;
303
- {% else %}
304
  let split = wg.x;
305
- {% endif %}
306
  let h = wg.y;
307
  let b = wg.z;
308
- if (h >= Q_HEADS || split >= NUM_SPLITS{% if splitQueries %} || queryToken >= Q_SEQ{% endif %}{% if layout == "layer_cache" %} || params.past_len >= CACHE_LEN{% endif %}) {
309
  return;
310
  }
311
  let tid = lid.x;
@@ -317,7 +261,7 @@ fn main(
317
  {% if cacheSeqlens %}
318
  // Buffer-sharing caches retain their capacity in the physical BNSH stride;
319
  // seqlens_k supplies the active end independently for each batch.
320
- let kvSeq = min(cacheSeq, u32(seqlens_k[b]) + 1u);
321
  {% else %}
322
  let kvSeq = cacheSeq;
323
  {% endif %}
@@ -325,23 +269,15 @@ fn main(
325
 
326
  // Query row (decode uses token zero; short-query prefill folds the token into wg.x).
327
  {% if layout == "bsh" %}
328
- {% if splitQueries %}
329
- let qBaseV4 = (b * Q_SEQ + queryToken) * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
330
- {% else %}
331
  let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
332
- {% endif %}
333
  let kvBaseV4 = b * kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
334
  let kvTokenStrideV4 = KV_HIDDEN_V4;
335
  {% elif layout == "layer_cache" %}
336
  let qBaseV4 = h * HEAD_DIM_V4;
337
  let kvBaseV4 = (LAYER * CACHE_LEN * KV_HEADS + hKv) * HEAD_DIM_V4;
338
  let kvTokenStrideV4 = KV_HEADS * HEAD_DIM_V4;
339
- {% else %}
340
- {% if splitQueries %}
341
- let qBaseV4 = ((b * Q_HEADS + h) * Q_SEQ + queryToken) * HEAD_DIM_V4;
342
  {% else %}
343
  let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
344
- {% endif %}
345
  let kvBaseV4 = (b * KV_HEADS + hKv) * cacheSeq * HEAD_DIM_V4;
346
  let kvTokenStrideV4 = HEAD_DIM_V4;
347
  {% endif %}
@@ -437,27 +373,12 @@ fn main(
437
  {% endif %}
438
  {% if hasMask %}
439
  if (keyAllowed) {
440
- {% if splitQueries %}
441
- let maskQuery = queryToken;
442
- {% else %}
443
  let maskQuery = 0u;
444
- {% endif %}
445
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
446
- {% if maskIsBool %}
447
- // A rejected bool-mask key contributes no probability mass. The merge
448
- // pass already maps a zero global denominator to an all-zero output row.
449
- if (attn_mask[maskIndex] == 0u) {
450
- keyAllowed = false;
451
- score = -FLT_MAX;
452
- dPart = 0.0;
453
- }
454
- {% else %}
455
  score = score + f32(attn_mask[maskIndex]);
456
- {% endif %}
457
  m = score;
458
  }
459
- {% endif %}
460
- let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
461
 
462
  // Merge one key tile's online-softmax (maximum, denominator) partial into the
463
  // running state, then store the per-key probabilities consumed by V accumulation.
@@ -473,7 +394,6 @@ fn main(
473
  probs[tid] = prob;
474
  workgroupBarrier();
475
 
476
-
477
  {% if jSplitV %}
478
  // j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
479
  // J_GROUPS for dim block d4v into a register, then the groups combine
@@ -521,20 +441,12 @@ fn main(
521
  }
522
 
523
  // Emit un-normalized partials for (b, h, split): the merge pass divides.
524
- {% if splitQueries %}
525
- let partialBase = (((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4;
526
- {% else %}
527
  let partialBase = ((b * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4;
528
- {% endif %}
529
  for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
530
  partial_out[partialBase + d4] = running_out[d4];
531
  }
532
  if (tid == 0u) {
533
- {% if splitQueries %}
534
- let mdBase = ((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + split;
535
- {% else %}
536
  let mdBase = (b * Q_HEADS + h) * NUM_SPLITS + split;
537
- {% endif %}
538
  // (max, denom) travel together to the merge, so they share one buffer as an
539
  // interleaved vec2 rather than costing two bindings. Interleaved, not two
540
  // halves, so the index needs no region size — and the merge reads both
 
1
  {% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
 
2
  {% if quantizedCache is not defined %}{% set quantizedCache = false %}{% endif %}
3
  {% if cacheSeqlens is not defined %}{% set cacheSeqlens = false %}{% endif %}
4
  {% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
5
+ {% set cacheSeqlensIndex = "b" %}{% set layer = 0 %}
6
+ {% set cacheLen = 0 %}
7
+ {% set scale = "0.0" %}
 
8
  {% if hasRotary is not defined %}{% set hasRotary = false %}{% endif %}
9
+ {% set qHiddenV4 = qHiddenV4 | default(0) %}
10
+ {% set kvHiddenV4 = kvHiddenV4 | default(0) %}
11
+ {% set hasBias = hasBias is defined and hasBias %}
12
  {% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
13
  {% if useSubgroups %}
14
  enable subgroups;
15
  {% endif %}
16
  {{ env.wgsl.resourceDeclarations }}
17
 
18
+ // Split-K flash attention. Single-partition direct output skips the merge pass.
19
+ // Otherwise this is the first of two passes. This geometry
20
  // handles decode and short-query, long-context prefill inputs.
21
  //
22
  // The non-split flash decode launches only `batch * numHeads` workgroups, each
 
51
  {% endif %}
52
  const WG: u32 = {{ splitKWorkgroupSize }}u;
53
  const NUM_SPLITS: u32 = {{ numSplits }}u;
 
 
 
54
  // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
55
  // `m - m` finite so an empty lane / all--inf row contributes the exact
56
  // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
 
70
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
71
  return select(value - maxValue, 0.0, equalFiniteMax);
72
  }
73
+
74
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
75
  return exp(shifted_value(value, maxValue));
76
  }
 
78
  var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
79
  var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
80
  var<workgroup> probs: array<f32, WG>;
81
+ {% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) and (allowCooperativeQk if allowCooperativeQk is defined else true) %}
82
  {% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
83
  {% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
84
  {% if coopQk %}
 
147
  return combinedMD;
148
  }
149
  {% else %}
150
+ {% set mdExtent = "WG" %}
 
 
151
  var<workgroup> partialM: array<f32, {{ mdExtent }}>;
152
  var<workgroup> partialD: array<f32, {{ mdExtent }}>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
153
 
154
  fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
155
  partialM[lidx] = m;
 
179
  return merged;
180
  }
181
  {% endif %}
 
 
182
 
183
  {% if layout == "layer_cache" %}{% set ATTN_SCALE_OVERRIDE = "ATTN_SCALE" %}{% endif %}
184
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
185
  {% if ATTN_SCALE_OVERRIDE is defined %}
186
  return {{ ATTN_SCALE_OVERRIDE }};
187
  {% else %}
 
190
  {% endif %}
191
  }
192
 
 
193
  {% if quantizedCache %}
194
  {% macro emit_quant_scale4(kind, scaleBuffer) %}
195
  fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
 
203
  {{ scaleBuffer }}[base + 2u],
204
  {{ scaleBuffer }}[base + 3u]
205
  );
206
+ }{% endmacro %}
207
+ {% macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
 
208
  {{ emit_quant_scale4(kind, scaleBuffer) }}
209
  fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
 
210
  return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
211
+ }{% endmacro %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
212
  {{ emit_quant_load4("int8", "key", "key", "k_scale") }}
213
  {{ emit_quant_load4("int8", "value", "value", "v_scale") }}
214
  {% else %}
 
226
  // query row before the Q.K dots; the K bias adds a constant to every key score
227
  // that softmax cancels, so it is skipped; the V bias is token-independent and
228
  // is applied once in the merge pass after the final normalize.
229
+ {% set BW = "" %}
230
+ {% set BC = "" %}
231
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
232
  let offset = base + d4 * 4u;
233
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
234
  }
235
 
236
  {% endif %}
 
246
  // choose an intermediate width at pipeline execution time.
247
  if (sgSize == 0u || sgSize > WG || WG % sgSize != 0u) { return; }
248
  {% endif %}
 
 
 
 
249
  let split = wg.x;
 
250
  let h = wg.y;
251
  let b = wg.z;
252
+ if (h >= Q_HEADS || split >= NUM_SPLITS{% if layout == "layer_cache" %} || params.past_len >= CACHE_LEN{% endif %}) {
253
  return;
254
  }
255
  let tid = lid.x;
 
261
  {% if cacheSeqlens %}
262
  // Buffer-sharing caches retain their capacity in the physical BNSH stride;
263
  // seqlens_k supplies the active end independently for each batch.
264
+ let kvSeq = min(cacheSeq, u32(seqlens_k[{{ cacheSeqlensIndex }}]) + 1u);
265
  {% else %}
266
  let kvSeq = cacheSeq;
267
  {% endif %}
 
269
 
270
  // Query row (decode uses token zero; short-query prefill folds the token into wg.x).
271
  {% if layout == "bsh" %}
 
 
 
272
  let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
 
273
  let kvBaseV4 = b * kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
274
  let kvTokenStrideV4 = KV_HIDDEN_V4;
275
  {% elif layout == "layer_cache" %}
276
  let qBaseV4 = h * HEAD_DIM_V4;
277
  let kvBaseV4 = (LAYER * CACHE_LEN * KV_HEADS + hKv) * HEAD_DIM_V4;
278
  let kvTokenStrideV4 = KV_HEADS * HEAD_DIM_V4;
 
 
 
279
  {% else %}
280
  let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
 
281
  let kvBaseV4 = (b * KV_HEADS + hKv) * cacheSeq * HEAD_DIM_V4;
282
  let kvTokenStrideV4 = HEAD_DIM_V4;
283
  {% endif %}
 
373
  {% endif %}
374
  {% if hasMask %}
375
  if (keyAllowed) {
 
 
 
376
  let maskQuery = 0u;
 
377
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
 
 
 
 
 
 
 
 
 
378
  score = score + f32(attn_mask[maskIndex]);
 
379
  m = score;
380
  }
381
+ {% endif %} let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
 
382
 
383
  // Merge one key tile's online-softmax (maximum, denominator) partial into the
384
  // running state, then store the per-key probabilities consumed by V accumulation.
 
394
  probs[tid] = prob;
395
  workgroupBarrier();
396
 
 
397
  {% if jSplitV %}
398
  // j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
399
  // J_GROUPS for dim block d4v into a register, then the groups combine
 
441
  }
442
 
443
  // Emit un-normalized partials for (b, h, split): the merge pass divides.
 
 
 
444
  let partialBase = ((b * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4;
 
445
  for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
446
  partial_out[partialBase + d4] = running_out[d4];
447
  }
448
  if (tid == 0u) {
 
 
 
449
  let mdBase = (b * Q_HEADS + h) * NUM_SPLITS + split;
 
450
  // (max, denom) travel together to the merge, so they share one buffer as an
451
  // interleaved vec2 rather than costing two bindings. Interleaved, not two
452
  // halves, so the index needs no region size — and the merge reads both
build/webgpu/attn-flash-online.wgsl.jinja CHANGED
@@ -1,6 +1,4 @@
1
- {% if combineSubgroups %}
2
  enable subgroups;
3
- {% endif %}
4
  {{ env.wgsl.resourceDeclarations }}
5
 
6
  // Flash-style tiled online-softmax attention for vec4-aligned head dimensions.
@@ -28,8 +26,8 @@ const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u;
28
  const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
29
  const Q_HEADS: u32 = {{ qNumHeads }}u;
30
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
31
- {% set qHeads = "params.qHeads" if headsFromParams else "Q_HEADS" %}
32
- {% set kvHeads = "params.kvHeads" if headsFromParams else "KV_HEADS" %}
33
  const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
34
  // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
35
  // `m - m` finite so an empty lane / all--inf row contributes the exact
@@ -50,6 +48,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
50
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
51
  return select(value - maxValue, 0.0, equalFiniteMax);
52
  }
 
53
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
54
  return exp(shifted_value(value, maxValue));
55
  }
@@ -63,8 +62,6 @@ var<workgroup> probs: array<f32, WG>;
63
  // Both the subgroup and portable barrier-tree engines return the same merged
64
  // pair to every invocation. Repeated merges require a workgroup barrier between
65
  // calls before their shared partial storage is reused.
66
- {% set combineSubgroups = combineSubgroups is defined and combineSubgroups %}
67
- {% if combineSubgroups %}
68
  // Cross-subgroup merge that assumes nothing about which invocations share a
69
  // subgroup or how many subgroups there are: each subgroup's elected lane
70
  // publishes the subgroup pair in the slot at its OWN invocation index and sets
@@ -115,96 +112,28 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
115
  workgroupBarrier();
116
  return combinedMD;
117
  }
118
- {% else %}
119
- {% set mdStreamed = mdStreams is defined %}
120
- {% set mdStreams = mdStreams if mdStreams is defined else 1 %}
121
- {% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
122
- var<workgroup> partialM: array<f32, {{ mdExtent }}>;
123
- var<workgroup> partialD: array<f32, {{ mdExtent }}>;
124
- {% if mdStreamed %}
125
-
126
- // In-place fold of {{ mdStreams }} streams. Input partials occupy
127
- // partialM/partialD; stream s returns its merged pair in slot s * WG.
128
- fn combine_partials_streams(lidx: u32) {
129
- workgroupBarrier();
130
- var stride = WG / 2u;
131
- loop {
132
- if (stride == 0u) {
133
- break;
134
- }
135
- if (lidx < stride) {
136
- {% for s in range(mdStreams) %}
137
- {
138
- let slot = {{ s }}u * WG + lidx;
139
- let m1 = partialM[slot];
140
- let d1 = partialD[slot];
141
- let m2 = partialM[slot + stride];
142
- let d2 = partialD[slot + stride];
143
- let mNew = max(m1, m2);
144
- partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
145
- partialM[slot] = mNew;
146
- }
147
- {% endfor %}
148
- }
149
- workgroupBarrier();
150
- stride = stride / 2u;
151
- }
152
- }
153
- {% else %}
154
-
155
- fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
156
- partialM[lidx] = m;
157
- partialD[lidx] = d;
158
- workgroupBarrier();
159
- var stride = WG / 2u;
160
- loop {
161
- if (stride == 0u) {
162
- break;
163
- }
164
- if (lidx < stride) {
165
- let m1 = partialM[lidx];
166
- let d1 = partialD[lidx];
167
- let m2 = partialM[lidx + stride];
168
- let d2 = partialD[lidx + stride];
169
- let mNew = max(m1, m2);
170
- partialD[lidx] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
171
- partialM[lidx] = mNew;
172
- }
173
- workgroupBarrier();
174
- stride = stride / 2u;
175
- }
176
- let merged = vec2<f32>(partialM[0], partialD[0]);
177
- // Trailing barrier so back-to-back calls cannot race a next call's partial
178
- // stores against this call's reads of slot 0.
179
- workgroupBarrier();
180
- return merged;
181
- }
182
- {% endif %}
183
- {% endif %}
184
-
185
 
186
  // An explicit-zero specialization bakes the scale as 0. Otherwise,
187
  // params.scale == 0 encodes an omitted scale and selects 1/sqrt(headDim).
188
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
189
- fn scale_value() -> f32 {
190
  if (params.scale != 0.0) { return params.scale; }
191
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
192
  }
193
-
194
  {% if hasBias %}
195
 
 
 
196
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
197
  let offset = base + d4 * 4u;
198
- return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
199
  }
200
 
201
  {% endif %}
202
  @compute @workgroup_size(WG, 1, 1)
203
  fn main(
204
  @builtin(workgroup_id) wg: vec3<u32>,
205
- @builtin(local_invocation_id) lid: vec3<u32>{% if combineSubgroups %},
206
- @builtin(subgroup_size) sgSize: u32{% endif %}
207
- ) {
208
  let qi = wg.x;
209
  let h = wg.y;
210
  let b = wg.z;
@@ -243,7 +172,11 @@ fn main(
243
  // Causal upper bound: query qi attends only keys 0..qi, so stop after the tile
244
  // containing qi and skip the unattended tail. Non-causal keeps the full kvSeq
245
  // sweep.
 
 
 
246
  var keyBoundV = params.kvSeq;
 
247
  var keyFloor: u32 = 0u;
248
  {% if hasWindow %}
249
  // Sliding window: query qi sits at absolute position p = kvSeq - qSeq
@@ -269,7 +202,7 @@ fn main(
269
  var score = -FLT_MAX;
270
  var m = -FLT_MAX;
271
  var dPart = 0.0;
272
- var keyAllowed = kj < keyBound{% if hasWindow %} && kj >= keyFloor{% endif %};
273
  if (keyAllowed) {
274
  let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
275
  {% if hasMask %}
@@ -277,17 +210,6 @@ fn main(
277
  // [q, k] masks set batch/head strides to 0).
278
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qi * params.maskSeqStride + kj;
279
  {% endif %}
280
- {% if hasMask and maskIsBool %}
281
- if (attn_mask[maskIndex] == 0u) {
282
- keyAllowed = false;
283
- } else {
284
- var acc: f32 = 0.0;
285
- for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
286
- acc = acc + dot(q_shared[d4], vec4<f32>(key[kRowV4 + d4]));
287
- }
288
- score = acc * scale;
289
- }
290
- {% else %}
291
  var acc: f32 = 0.0;
292
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
293
  acc = acc + dot(q_shared[d4], vec4<f32>(key[kRowV4 + d4]));
@@ -295,12 +217,11 @@ fn main(
295
  score = acc * scale;
296
  {% if hasMask %}
297
  score = score + f32(attn_mask[maskIndex]);
298
- {% endif %}
299
  {% endif %}
300
  m = score;
301
  dPart = select(0.0, 1.0, keyAllowed);
302
  }
303
- let tile = combine_partials(m, dPart, tid{% if combineSubgroups %}, sgSize{% endif %});
304
 
305
  // Online merge of the tile into the running state (softmax-online rule).
306
  // Merge one key tile's online-softmax (maximum, denominator) partial into the
@@ -317,7 +238,6 @@ fn main(
317
  probs[tid] = prob;
318
  workgroupBarrier();
319
 
320
-
321
  // running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
322
  // tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
323
  let tileCount = min(WG, keyBound - kjBase);
 
 
1
  enable subgroups;
 
2
  {{ env.wgsl.resourceDeclarations }}
3
 
4
  // Flash-style tiled online-softmax attention for vec4-aligned head dimensions.
 
26
  const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
27
  const Q_HEADS: u32 = {{ qNumHeads }}u;
28
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
29
+ {% set qHeads = "Q_HEADS" %}
30
+ {% set kvHeads = "KV_HEADS" %}
31
  const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
32
  // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
33
  // `m - m` finite so an empty lane / all--inf row contributes the exact
 
48
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
49
  return select(value - maxValue, 0.0, equalFiniteMax);
50
  }
51
+
52
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
53
  return exp(shifted_value(value, maxValue));
54
  }
 
62
  // Both the subgroup and portable barrier-tree engines return the same merged
63
  // pair to every invocation. Repeated merges require a workgroup barrier between
64
  // calls before their shared partial storage is reused.
 
 
65
  // Cross-subgroup merge that assumes nothing about which invocations share a
66
  // subgroup or how many subgroups there are: each subgroup's elected lane
67
  // publishes the subgroup pair in the slot at its OWN invocation index and sets
 
112
  workgroupBarrier();
113
  return combinedMD;
114
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
115
 
116
  // An explicit-zero specialization bakes the scale as 0. Otherwise,
117
  // params.scale == 0 encodes an omitted scale and selects 1/sqrt(headDim).
118
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
119
  if (params.scale != 0.0) { return params.scale; }
120
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
121
  }
 
122
  {% if hasBias %}
123
 
124
+ {% set BW = "" %}
125
+ {% set BC = "" %}
126
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
127
  let offset = base + d4 * 4u;
128
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
129
  }
130
 
131
  {% endif %}
132
  @compute @workgroup_size(WG, 1, 1)
133
  fn main(
134
  @builtin(workgroup_id) wg: vec3<u32>,
135
+ @builtin(local_invocation_id) lid: vec3<u32>,
136
+ @builtin(subgroup_size) sgSize: u32) {
 
137
  let qi = wg.x;
138
  let h = wg.y;
139
  let b = wg.z;
 
172
  // Causal upper bound: query qi attends only keys 0..qi, so stop after the tile
173
  // containing qi and skip the unattended tail. Non-causal keeps the full kvSeq
174
  // sweep.
175
+ {% if hasCausal %}
176
+ var keyBoundV = select(params.kvSeq, min(params.kvSeq, qi + 1u), params.isCausal != 0u);
177
+ {% else %}
178
  var keyBoundV = params.kvSeq;
179
+ {% endif %}
180
  var keyFloor: u32 = 0u;
181
  {% if hasWindow %}
182
  // Sliding window: query qi sits at absolute position p = kvSeq - qSeq
 
202
  var score = -FLT_MAX;
203
  var m = -FLT_MAX;
204
  var dPart = 0.0;
205
+ var keyAllowed = kj < keyBound{% if hasCausal %} && (params.isCausal == 0u || kj <= qi){% endif %}{% if hasWindow %} && kj >= keyFloor{% endif %};
206
  if (keyAllowed) {
207
  let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
208
  {% if hasMask %}
 
210
  // [q, k] masks set batch/head strides to 0).
211
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qi * params.maskSeqStride + kj;
212
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
213
  var acc: f32 = 0.0;
214
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
215
  acc = acc + dot(q_shared[d4], vec4<f32>(key[kRowV4 + d4]));
 
217
  score = acc * scale;
218
  {% if hasMask %}
219
  score = score + f32(attn_mask[maskIndex]);
 
220
  {% endif %}
221
  m = score;
222
  dPart = select(0.0, 1.0, keyAllowed);
223
  }
224
+ let tile = combine_partials(m, dPart, tid, sgSize);
225
 
226
  // Online merge of the tile into the running state (softmax-online rule).
227
  // Merge one key tile's online-softmax (maximum, denominator) partial into the
 
238
  probs[tid] = prob;
239
  workgroupBarrier();
240
 
 
241
  // running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
242
  // tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
243
  let tileCount = min(WG, keyBound - kjBase);
build/webgpu/attn-flash-prefill-cluster.wgsl.jinja CHANGED
@@ -1,22 +1,18 @@
1
- {% set sourceProfile = sourceProfile if sourceProfile is defined else 0 %}
2
- {% set scaling = scaling | default("0.0") %}
3
- {% set qkvStrideV4 = qkvStrideV4 | default(0) %}
4
- {% set QSEQ = "params.seq_len" if sourceProfile == 1 else "params.qSeq" %}
5
- {% set KVSEQ = "(params.past_len + params.seq_len)" if sourceProfile == 1 else "params.kvSeq" %}
6
- {% set IS_CAUSAL = "1u" if sourceProfile == 1 else "params.isCausal" %}
7
- {% set Q_STRIDE = "QKV_STRIDE_V4" if sourceProfile == 1 else "Q_HIDDEN_V4" %}
8
- {% set QUERY = "qkv" if sourceProfile == 1 else "query" %}
9
- {% set KEY = "cache_keys" if sourceProfile == 1 else "key" %}
10
- {% set VALUE = "cache_values" if sourceProfile == 1 else "value" %}
11
- {% set OUTPUT = "attn_out" if sourceProfile == 1 else "output" %}
12
- {% if sourceProfile == 1 %}{% set ATTN_SCALE_OVERRIDE = scaling %}{% endif %}
13
  {% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
14
  {% if batchNoSgReduction is not defined %}{% set batchNoSgReduction = false %}{% endif %}
 
15
  {% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
16
- {% if maskIsKeyKeep is not defined %}{% set maskIsKeyKeep = false %}{% endif %}
17
- {% if maskIsBool is not defined %}{% set maskIsBool = false %}{% endif %}
18
- {% if hasSoftcap is not defined %}{% set hasSoftcap = false %}{% endif %}
19
  {% if hasHeadSink is not defined %}{% set hasHeadSink = false %}{% endif %}
 
20
  {% if quantCacheFormat is not defined %}{% set quantCacheFormat = "" %}{% endif %}
21
  {% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
22
  // A windowed cache binds a fixed CAPACITY but keeps only the most recent
@@ -24,7 +20,14 @@
24
  // count, which is still the right batch stride but the wrong attention bound, so
25
  // the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
26
  // in both roles.
27
- {% set KVA = "kvActive" if useSeqlens else KVSEQ %}
 
 
 
 
 
 
 
28
  {% macro score_expr(part) %}{% if hasSoftcap %}params.softcap * tanh(clamp(({{ part }} * SCALE) / params.softcap, -30.0, 30.0)){% else %}{{ part }} * SCALE{% endif %}{% endmacro %}
29
  {% if useSubgroups %}
30
  enable subgroups;
@@ -40,11 +43,30 @@ enable subgroups;
40
  {% set MASK_TILE_TYPE = "u32" if MASK_IS_INT else "f32" %}
41
  {% set MASK_TILE_LOAD = "attn_mask[maskIndex]" if MASK_IS_INT else "f32(attn_mask[maskIndex])" %}
42
  {% set MASK_TILE_ZERO = "0u" if MASK_IS_INT else "0.0" %}
43
- {% set MASK_ELEMENT = "maskValue" if STAGE_MASK else "attn_mask[maskIndex]" %}
44
  {% set MASK_ADDITIVE = "maskValue" if STAGE_MASK else "f32(attn_mask[maskIndex])" %}
45
  {% set SLICE_COUNT = ((headDimV4 / LPQ) | int) %}
46
  {% set QPL = (2 if (not usesF16 and SLICE_COUNT <= 4 and TILE_K <= 16 and TILE_Q % 2 == 0) else 1) if useSubgroups else 1 %}
47
  {% macro qn(name, qi) %}{{ name }}{% if qi > 0 %}_q{{ qi }}{% endif %}{% endmacro %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
48
  {% set ROPE_LANE_XOR = ((LPQ / 2) | int) %}
49
  {% set QL = qLayout if qLayout is defined else layout %}
50
  {% set KL = kvLayout if kvLayout is defined else layout %}
@@ -69,8 +91,7 @@ const HEAD_DIM: u32 = {{ headDim }}u;
69
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
70
  {% if QL != "bhsd" %}
71
  const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
72
- {% if sourceProfile == 1 %}const QKV_STRIDE_V4: u32 = {{ qkvStrideV4 }}u; // packed [Q; K; V] input row stride
73
- {% endif %}{% endif %}
74
  {% if KL != "bhsd" %}
75
  const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
76
  {% endif %}
@@ -91,16 +112,7 @@ const WG: u32 = (TILE_Q / QPL) * LPQ;
91
  {% else %}
92
  const WG: u32 = TILE_Q * LPQ;
93
  {% endif %}
94
- {% if MASK_IS_INT %}
95
- // Key-keep masks in contrib attention use a finite low logit for a rejected
96
- // key. Logical ONNX bool masks use the exclusion sentinel instead, so a fully
97
- // masked row has zero mass. Keeping one declaration shape for both mask modes
98
- // lets both mask modes share the score loop below.
99
- const NEG_INF: f32 = -3.4028234663852886e38;
100
- const MASK_NEG: f32 = {{ "-1e38" if maskIsKeyKeep else "-3.4028234663852886e38" }};
101
- {% else %}
102
  const NEG_INF: f32 = -3.4028234663852886e38;
103
- {% endif %}
104
 
105
  var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
106
  var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
@@ -121,16 +133,10 @@ var<workgroup> red: array<f32, WG>;
121
  {% endif %}
122
  {% endif %}
123
 
124
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
125
- fn scale_value() -> f32 {
126
- {% if ATTN_SCALE_OVERRIDE is defined %}
127
- return {{ ATTN_SCALE_OVERRIDE }};
128
- {% else %}
129
  if (params.scale != 0.0) { return params.scale; }
130
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
131
- {% endif %}
132
  }
133
-
134
  {% if quantCacheFormat %}
135
  // A quantized cache is dequantized once per key into the staged tile, then read
136
  // by all TILE_Q queries in the workgroup. The unpack cost is amortized over the
@@ -147,14 +153,13 @@ fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
147
  {{ scaleBuffer }}[base + 2u],
148
  {{ scaleBuffer }}[base + 3u]
149
  );
150
- }
151
- {%- endmacro %}
152
- {%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
153
  {{ emit_quant_scale4(kind, scaleBuffer) }}
154
  fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
155
- {%- if format == "int8" %}
156
  return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
157
- {%- else %}
158
  // Two elements cover this vec4: each carries two +8-biased nibbles, low first.
159
  let rowBase = indexV4 - d4;
160
  let lo = {{ buffer }}[rowBase + d4 * 2u];
@@ -165,14 +170,21 @@ fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
165
  );
166
  let signed = nibbles - vec4<i32>(8);
167
  return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
168
- {%- endif %}
169
- }
170
- {%- endmacro %}
171
-
172
  {{ emit_quant_load4(quantCacheFormat, "key", KEY, "k_scale") }}
173
  {{ emit_quant_load4(quantCacheFormat, "value", VALUE, "v_scale") }}
174
  {% endif %}
175
 
 
 
 
 
 
 
 
 
 
176
  @compute @workgroup_size(WG, 1, 1)
177
  fn main(
178
  @builtin(workgroup_id) wg: vec3<u32>,
@@ -211,6 +223,9 @@ fn main(
211
  {% endif %}
212
  {% for c in range(SLICE_COUNT) %}
213
  var {{ qn("qr" ~ c, qi) }} = vec4<f32>({{ QUERY }}[{{ qn("qBase4", qi) }} + {{ c }}u]);
 
 
 
214
  var {{ qn("o" ~ c, qi) }} = vec4<f32>(0.0);
215
  {% endfor %}
216
  {% endfor %}
@@ -390,9 +405,13 @@ fn main(
390
  // TILE_K is a small shader constant; the loop updates the named q/o slices in place.
391
  {% if not useSubgroups and batchNoSgReduction %}
392
  // First publish every key's partial dot without intervening barriers.
 
393
  var s: array<f32, TILE_K>;
 
394
  for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
 
395
  s[kk] = NEG_INF;
 
396
  var part: f32 = 0.0;
397
  let kb = kk * HEAD_DIM_V4 + lane8 * SLICE;
398
  {% for c in range(SLICE_COUNT) %}
@@ -422,26 +441,41 @@ fn main(
422
  workgroupBarrier();
423
  // The barrier at the start of the next K tile protects this scratch before
424
  // it is overwritten.
 
 
 
 
 
 
425
  for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
426
  let kj = kStart + kk;
427
  let part = red[kk * WG + tid - lane8];
 
 
 
428
  if (kj >= minKj && kj < maxKj) {
429
  {% if hasMask %}
430
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
431
- {% if maskIsBool %}
432
- if (attn_mask[maskIndex] != 0u) {
433
- s[kk] = {{ score_expr("part") }};
434
- } else {
435
- s[kk] = MASK_NEG;
436
- }
437
- {% else %}
438
  s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
439
- {% endif %}
440
  {% else %}
441
  s[kk] = {{ score_expr("part") }};
442
  {% endif %}
443
  }
444
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
445
  {% elif QPL > 1 %}
446
  {% for qi in range(QPL) %}
447
  var {{ qn("s", qi) }}: array<f32, TILE_K>;
@@ -481,21 +515,7 @@ fn main(
481
  // in-bounds for padding queries in the last tile (their output is dropped).
482
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
483
  {% endif %}
484
- {% if maskIsKeyKeep %}
485
- // A broadcast key mask uses 1 for a retained key and 0 for padding.
486
- {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
487
- {% elif maskIsBool %}
488
- // Logical bool: a rejected key contributes no softmax mass. Leaving
489
- // the initialized NEG_INF sentinel in place makes a fully masked row
490
- // land on the zero-denominator output guard below.
491
- if ({{ MASK_ELEMENT }} != 0u) {
492
- {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
493
- } else {
494
- {{ qn("s", qi) }}[kk] = MASK_NEG;
495
- }
496
- {% else %}
497
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
498
- {% endif %}
499
  {% else %}
500
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
501
  {% endif %}
@@ -541,21 +561,7 @@ fn main(
541
  // in-bounds for padding queries in the last tile (their output is dropped).
542
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
543
  {% endif %}
544
- {% if maskIsKeyKeep %}
545
- // A broadcast key mask uses 1 for a retained key and 0 for padding.
546
- s[kk] = {{ score_expr("part") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
547
- {% elif maskIsBool %}
548
- // Logical bool: a rejected key contributes no softmax mass. Leaving
549
- // the initialized NEG_INF sentinel in place makes a fully masked row
550
- // land on the zero-denominator output guard below.
551
- if ({{ MASK_ELEMENT }} != 0u) {
552
- s[kk] = {{ score_expr("part") }};
553
- } else {
554
- s[kk] = MASK_NEG;
555
- }
556
- {% else %}
557
  s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
558
- {% endif %}
559
  {% else %}
560
  s[kk] = {{ score_expr("part") }};
561
  {% endif %}
@@ -563,30 +569,36 @@ fn main(
563
  }
564
  {% endif %}
565
 
566
- // Per-thread online softmax over the tile. s[kk] is reused to hold the
567
- // exponentiated probabilities for the PV accumulation below.
568
- {% for qi in range(QPL) %}
569
- var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
570
- for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
571
- {{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
572
- }
573
- let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
574
- let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
575
- var {{ qn("pSum", qi) }}: f32 = 0.0;
576
- for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
577
- let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
578
- {{ qn("s", qi) }}[kk] = pk;
579
- {{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
580
- }
581
- {{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
582
- {{ qn("m", qi) }} = {{ qn("newMax", qi) }};
583
- {% endfor %}
584
  // A boundary tile can address V rows outside a query's attended range, and
585
  // a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
586
  // NaN, the guarded loop selects the V operand away for range-excluded keys.
587
  // Interior tiles use the unguarded FMA chain. Mask exclusion applies to
588
  // materialized V rows and does not require this range guard.
589
- {% if QPL > 1 %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
590
  let tileInterior = {% for qi in range(QPL) %}(kStart >= {{ qn("minKj", qi) }} && kStart + TILE_K <= {{ qn("maxKj", qi) }}){% if not loop.last %} && {% endif %}{% endfor %};
591
  {% for c in range(SLICE_COUNT) %}
592
  {
@@ -635,7 +647,7 @@ fn main(
635
  kStart = kStart + TILE_K;
636
  }
637
 
638
- {% macro attention_value(c, qi) %}{{ qn("o" ~ c, qi) }} * inv{% endmacro %}
639
  {% for qi in range(QPL) %}
640
  if ({{ qn("qValid", qi) }}) {
641
  {% if QL == "bhsd" %}
 
1
+ {% set QSEQ = "params.qSeq" %}
2
+ {% set KVSEQ = "params.kvSeq" %}
3
+ {% set IS_CAUSAL = "params.isCausal" %}
4
+ {% set Q_STRIDE = "Q_HIDDEN_V4" %}
5
+ {% set QUERY = "query" %}
6
+ {% set KEY = "key" %}
7
+ {% set VALUE = "value" %}
8
+ {% set OUTPUT = "output" %}
 
 
 
 
9
  {% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
10
  {% if batchNoSgReduction is not defined %}{% set batchNoSgReduction = false %}{% endif %}
11
+ {% set SHARE_PROB = shareNoSgProb is defined and shareNoSgProb %}
12
  {% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
13
+ {% set maskIsKeyKeep = false %}{% set maskIsBool = false %}{% if hasSoftcap is not defined %}{% set hasSoftcap = false %}{% endif %}
 
 
14
  {% if hasHeadSink is not defined %}{% set hasHeadSink = false %}{% endif %}
15
+ {% set Q_HIDDEN = qHidden | default(0) %}
16
  {% if quantCacheFormat is not defined %}{% set quantCacheFormat = "" %}{% endif %}
17
  {% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
18
  // A windowed cache binds a fixed CAPACITY but keeps only the most recent
 
20
  // count, which is still the right batch stride but the wrong attention bound, so
21
  // the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
22
  // in both roles.
23
+ {% if hasKeyLimit is not defined %}{% set hasKeyLimit = false %}{% endif %}
24
+ {% if useSeqlens %}
25
+ {% set KVA = "kvActive" %}
26
+ {% elif hasKeyLimit %}
27
+ {% set KVA = "select(params.kvSeq, min(params.keyLimit, params.kvSeq), params.keyLimit > 0u)" %}
28
+ {% else %}
29
+ {% set KVA = KVSEQ %}
30
+ {% endif %}
31
  {% macro score_expr(part) %}{% if hasSoftcap %}params.softcap * tanh(clamp(({{ part }} * SCALE) / params.softcap, -30.0, 30.0)){% else %}{{ part }} * SCALE{% endif %}{% endmacro %}
32
  {% if useSubgroups %}
33
  enable subgroups;
 
43
  {% set MASK_TILE_TYPE = "u32" if MASK_IS_INT else "f32" %}
44
  {% set MASK_TILE_LOAD = "attn_mask[maskIndex]" if MASK_IS_INT else "f32(attn_mask[maskIndex])" %}
45
  {% set MASK_TILE_ZERO = "0u" if MASK_IS_INT else "0.0" %}
 
46
  {% set MASK_ADDITIVE = "maskValue" if STAGE_MASK else "f32(attn_mask[maskIndex])" %}
47
  {% set SLICE_COUNT = ((headDimV4 / LPQ) | int) %}
48
  {% set QPL = (2 if (not usesF16 and SLICE_COUNT <= 4 and TILE_K <= 16 and TILE_Q % 2 == 0) else 1) if useSubgroups else 1 %}
49
  {% macro qn(name, qi) %}{{ name }}{% if qi > 0 %}_q{{ qi }}{% endif %}{% endmacro %}
50
+ {% macro emit_tile_softmax() %}
51
+ // Per-thread online softmax over the tile. s[kk] is reused to hold the
52
+ // exponentiated probabilities for the PV accumulation below.
53
+ {% for qi in range(QPL) %}
54
+ var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
55
+ for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
56
+ {{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
57
+ }
58
+ let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
59
+ let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
60
+ var {{ qn("pSum", qi) }}: f32 = 0.0;
61
+ for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
62
+ let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
63
+ {{ qn("s", qi) }}[kk] = pk;
64
+ {{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
65
+ }
66
+ {{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
67
+ {{ qn("m", qi) }} = {{ qn("newMax", qi) }};
68
+ {% endfor %}
69
+ {% endmacro %}
70
  {% set ROPE_LANE_XOR = ((LPQ / 2) | int) %}
71
  {% set QL = qLayout if qLayout is defined else layout %}
72
  {% set KL = kvLayout if kvLayout is defined else layout %}
 
91
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
92
  {% if QL != "bhsd" %}
93
  const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
94
+ {% endif %}
 
95
  {% if KL != "bhsd" %}
96
  const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
97
  {% endif %}
 
112
  {% else %}
113
  const WG: u32 = TILE_Q * LPQ;
114
  {% endif %}
 
 
 
 
 
 
 
 
115
  const NEG_INF: f32 = -3.4028234663852886e38;
 
116
 
117
  var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
118
  var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
 
133
  {% endif %}
134
  {% endif %}
135
 
136
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
 
 
 
137
  if (params.scale != 0.0) { return params.scale; }
138
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
 
139
  }
 
140
  {% if quantCacheFormat %}
141
  // A quantized cache is dequantized once per key into the staged tile, then read
142
  // by all TILE_Q queries in the workgroup. The unpack cost is amortized over the
 
153
  {{ scaleBuffer }}[base + 2u],
154
  {{ scaleBuffer }}[base + 3u]
155
  );
156
+ }{% endmacro %}
157
+ {% macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
 
158
  {{ emit_quant_scale4(kind, scaleBuffer) }}
159
  fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
160
+ {% if format == "int8" %}
161
  return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
162
+ {% else %}
163
  // Two elements cover this vec4: each carries two +8-biased nibbles, low first.
164
  let rowBase = indexV4 - d4;
165
  let lo = {{ buffer }}[rowBase + d4 * 2u];
 
170
  );
171
  let signed = nibbles - vec4<i32>(8);
172
  return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
173
+ {% endif %}
174
+ }{% endmacro %}
 
 
175
  {{ emit_quant_load4(quantCacheFormat, "key", KEY, "k_scale") }}
176
  {{ emit_quant_load4(quantCacheFormat, "value", VALUE, "v_scale") }}
177
  {% endif %}
178
 
179
+ {% if hasBias %}
180
+ {% set BW = "" %}
181
+ {% set BC = "" %}
182
+ fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
183
+ let offset = base + d4 * 4u;
184
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
185
+ }
186
+
187
+ {% endif %}
188
  @compute @workgroup_size(WG, 1, 1)
189
  fn main(
190
  @builtin(workgroup_id) wg: vec3<u32>,
 
223
  {% endif %}
224
  {% for c in range(SLICE_COUNT) %}
225
  var {{ qn("qr" ~ c, qi) }} = vec4<f32>({{ QUERY }}[{{ qn("qBase4", qi) }} + {{ c }}u]);
226
+ {% if hasBias %}
227
+ {{ qn("qr" ~ c, qi) }} = {{ qn("qr" ~ c, qi) }} + load_bias4(h * HEAD_DIM, lane8 * SLICE + {{ c }}u);
228
+ {% endif %}
229
  var {{ qn("o" ~ c, qi) }} = vec4<f32>(0.0);
230
  {% endfor %}
231
  {% endfor %}
 
405
  // TILE_K is a small shader constant; the loop updates the named q/o slices in place.
406
  {% if not useSubgroups and batchNoSgReduction %}
407
  // First publish every key's partial dot without intervening barriers.
408
+ {% if not SHARE_PROB %}
409
  var s: array<f32, TILE_K>;
410
+ {% endif %}
411
  for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
412
+ {% if not SHARE_PROB %}
413
  s[kk] = NEG_INF;
414
+ {% endif %}
415
  var part: f32 = 0.0;
416
  let kb = kk * HEAD_DIM_V4 + lane8 * SLICE;
417
  {% for c in range(SLICE_COUNT) %}
 
441
  workgroupBarrier();
442
  // The barrier at the start of the next K tile protects this scratch before
443
  // it is overwritten.
444
+ {% if SHARE_PROB %}
445
+ // One lane computes the query softmax, reusing the reduced-score slots for
446
+ // probabilities. The discarded partial lanes hold the three online states.
447
+ if (lane8 == 0u) {
448
+ var s: array<f32, TILE_K>;
449
+ {% endif %}
450
  for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
451
  let kj = kStart + kk;
452
  let part = red[kk * WG + tid - lane8];
453
+ {% if SHARE_PROB %}
454
+ s[kk] = NEG_INF;
455
+ {% endif %}
456
  if (kj >= minKj && kj < maxKj) {
457
  {% if hasMask %}
458
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
 
 
 
 
 
 
 
459
  s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
 
460
  {% else %}
461
  s[kk] = {{ score_expr("part") }};
462
  {% endif %}
463
  }
464
  }
465
+ {% if SHARE_PROB %}
466
+ {{ emit_tile_softmax() }}
467
+ for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
468
+ red[kk * WG + cbase] = s[kk];
469
+ }
470
+ red[cbase + 1u] = corr;
471
+ red[cbase + 2u] = m;
472
+ red[cbase + 3u] = l;
473
+ }
474
+ workgroupBarrier();
475
+ let corr = red[cbase + 1u];
476
+ m = red[cbase + 2u];
477
+ l = red[cbase + 3u];
478
+ {% endif %}
479
  {% elif QPL > 1 %}
480
  {% for qi in range(QPL) %}
481
  var {{ qn("s", qi) }}: array<f32, TILE_K>;
 
515
  // in-bounds for padding queries in the last tile (their output is dropped).
516
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
517
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
518
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
 
519
  {% else %}
520
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
521
  {% endif %}
 
561
  // in-bounds for padding queries in the last tile (their output is dropped).
562
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
563
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
564
  s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
 
565
  {% else %}
566
  s[kk] = {{ score_expr("part") }};
567
  {% endif %}
 
569
  }
570
  {% endif %}
571
 
572
+ {% if not SHARE_PROB %}
573
+ {{ emit_tile_softmax() }}
574
+ {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
575
  // A boundary tile can address V rows outside a query's attended range, and
576
  // a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
577
  // NaN, the guarded loop selects the V operand away for range-excluded keys.
578
  // Interior tiles use the unguarded FMA chain. Mask exclusion applies to
579
  // materialized V rows and does not require this range guard.
580
+ {% if SHARE_PROB %}
581
+ {% macro shared_pv(masked) %}
582
+ for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
583
+ let weight = red[kk * WG + cbase];
584
+ let vb = kk * HEAD_DIM_V4 + lane8 * SLICE;
585
+ {% for c in range(SLICE_COUNT) %}
586
+ o{{ c }} = o{{ c }} + weight * {% if masked %}select(vec4<f32>(), {% endif %}vec4<f32>(v_tile[vb + {{ c }}u]){% if masked %}, weight != 0.0){% endif %};
587
+ {% endfor %}
588
+ }
589
+ {% endmacro %}
590
+ // Consume each shared probability once across all output slices. This
591
+ // retains the key order of every output accumulator without a private tile.
592
+ {% for c in range(SLICE_COUNT) %}
593
+ o{{ c }} = o{{ c }} * corr;
594
+ {% endfor %}
595
+ let tileInterior = kStart >= minKj && kStart + TILE_K <= maxKj;
596
+ if (tileInterior) {
597
+ {{ shared_pv(false) }}
598
+ } else {
599
+ {{ shared_pv(true) }}
600
+ }
601
+ {% elif QPL > 1 %}
602
  let tileInterior = {% for qi in range(QPL) %}(kStart >= {{ qn("minKj", qi) }} && kStart + TILE_K <= {{ qn("maxKj", qi) }}){% if not loop.last %} && {% endif %}{% endfor %};
603
  {% for c in range(SLICE_COUNT) %}
604
  {
 
647
  kStart = kStart + TILE_K;
648
  }
649
 
650
+ {% macro attention_value(c, qi) %}{{ qn("o" ~ c, qi) }} * inv{% if hasBias %} + load_bias4(2u * {{ Q_HIDDEN }}u + h * HEAD_DIM, lane8 * SLICE + {{ c }}u){% endif %}{% endmacro %}
651
  {% for qi in range(QPL) %}
652
  if ({{ qn("qValid", qi) }}) {
653
  {% if QL == "bhsd" %}
build/webgpu/attn-flash-q32-broadcast.wgsl.jinja CHANGED
@@ -1,4 +1,5 @@
1
- // Register-resident flash prefill uses one 32-lane subgroup per workgroup. Each
 
2
  // lane owns a query row and keeps its Q slice and f32 output in registers. Lanes
3
  // cooperatively load K/V and broadcast them with subgroupShuffle, so each query
4
  // computes q·k and p·v without cross-lane reductions or workgroup storage.
@@ -10,10 +11,15 @@
10
  {% set ST = "f16" if usesF16 else "f32" %}
11
  {% set components = ["x", "y", "z", "w"] %}
12
  {% set USE_SUBGROUPS = useSubgroups if useSubgroups is defined else true %}
 
13
  {% set Q_STEP = qStep if qStep is defined else 32 %}
 
14
  {% if USE_SUBGROUPS %}
15
  enable subgroups;
16
  {% endif %}
 
 
 
17
  {{ env.wgsl.resourceDeclarations }}
18
 
19
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
@@ -43,8 +49,7 @@ fn scale_value() -> f32 {
43
  }
44
 
45
 
46
-
47
- @compute @workgroup_size({{ Q_STEP }}, 1, 1)
48
  fn main(
49
  @builtin(workgroup_id) wg: vec3<u32>,
50
  @builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
@@ -85,8 +90,19 @@ fn main(
85
  // Causal key ceiling. The key-loop bound (kvEnd) uses the workgroup's LAST query
86
  // so every lane shares a uniform trip count (subgroup ops stay reconverged);
87
  // each lane masks its own keys past myMaxKj to NEG_INF.
 
 
 
 
 
 
 
 
 
 
88
  let kvEnd = params.kvSeq;
89
  let myMaxKj = params.kvSeq;
 
90
  let kvBase = b * params.kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
91
  let kvKeyStride = KV_HIDDEN_V4;
92
 
@@ -174,6 +190,13 @@ fn main(
174
  {% endfor %}
175
  previous_max = new_max;
176
  previous_denom = denom;
 
 
 
 
 
 
 
177
 
178
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
179
  {% if USE_SUBGROUPS %}
@@ -190,21 +213,37 @@ fn main(
190
  }
191
  {% endif %}
192
  {% endif %}
 
 
 
 
193
  var acc: vec4<f32> = vec4<f32>(0.0);
 
194
  {% for g in range(qkGroups) %}
195
  {% for lane in range(4) %}
196
- {% if USE_SUBGROUPS %}
197
- {% if g < 8 %}
198
- acc = acc + vec4<f32>(subgroupShuffle(v_local0, {{ g * 4 + lane }}u)) * qk{{ g }}.{{ components[lane] }};
 
199
  {% else %}
200
- acc = acc + vec4<f32>(subgroupShuffle(v_local1, {{ (g - 8) * 4 + lane }}u)) * qk{{ g }}.{{ components[lane] }};
201
  {% endif %}
 
 
202
  {% else %}
203
- acc = acc + vec4<f32>(valueTile[d4 * K_STEP + {{ g * 4 + lane }}u]) * qk{{ g }}.{{ components[lane] }};
204
  {% endif %}
205
  {% endfor %}
 
 
 
 
206
  {% endfor %}
 
 
 
207
  o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
 
208
  }
209
  {% if not USE_SUBGROUPS %}
210
  workgroupBarrier();
 
1
+ // Register-resident flash prefill uses one 32-lane subgroup per workgroup (fixed
2
+ // by the adapter, or pinned where the adapter can compile exactly 32 lanes). Each
3
  // lane owns a query row and keeps its Q slice and f32 output in registers. Lanes
4
  // cooperatively load K/V and broadcast them with subgroupShuffle, so each query
5
  // computes q·k and p·v without cross-lane reductions or workgroup storage.
 
11
  {% set ST = "f16" if usesF16 else "f32" %}
12
  {% set components = ["x", "y", "z", "w"] %}
13
  {% set USE_SUBGROUPS = useSubgroups if useSubgroups is defined else true %}
14
+ {% set PIN_SUBGROUP_32 = USE_SUBGROUPS and pinSubgroupSize32 is defined and pinSubgroupSize32 %}
15
  {% set Q_STEP = qStep if qStep is defined else 32 %}
16
+ {% set CAUSAL = false if (hasCausal is defined and not hasCausal) else true %}
17
  {% if USE_SUBGROUPS %}
18
  enable subgroups;
19
  {% endif %}
20
+ {% if PIN_SUBGROUP_32 %}
21
+ enable subgroup_size_control;
22
+ {% endif %}
23
  {{ env.wgsl.resourceDeclarations }}
24
 
25
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
 
49
  }
50
 
51
 
52
+ @compute @workgroup_size({{ Q_STEP }}, 1, 1){{ " @subgroup_size(32)" if PIN_SUBGROUP_32 else "" }}
 
53
  fn main(
54
  @builtin(workgroup_id) wg: vec3<u32>,
55
  @builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
 
90
  // Causal key ceiling. The key-loop bound (kvEnd) uses the workgroup's LAST query
91
  // so every lane shares a uniform trip count (subgroup ops stay reconverged);
92
  // each lane masks its own keys past myMaxKj to NEG_INF.
93
+ {% if CAUSAL %}
94
+ {% set FIXED_CAUSAL = fixedCausal is defined and fixedCausal %}
95
+ {% macro key_ceiling(query) %}
96
+ {% if FIXED_CAUSAL %}
97
+ min({{ query }} + 1u, params.kvSeq){% else %}
98
+ select(params.kvSeq, min({{ query }} + 1u, params.kvSeq), params.isCausal != 0u){% endif %}{% endmacro %}
99
+ let lastQ = min(wg.x * Q_STEP + Q_STEP - 1u, params.qSeq - 1u);
100
+ let kvEnd = {{ key_ceiling("lastQ") }};
101
+ let myMaxKj = {{ key_ceiling("qi") }};
102
+ {% else %}
103
  let kvEnd = params.kvSeq;
104
  let myMaxKj = params.kvSeq;
105
+ {% endif %}
106
  let kvBase = b * params.kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
107
  let kvKeyStride = KV_HIDDEN_V4;
108
 
 
190
  {% endfor %}
191
  previous_max = new_max;
192
  previous_denom = denom;
193
+ {% set PV_HALF = ST == "f16" and not (precisePv | default(false)) %}
194
+ {% set PV_FLUSH_GROUPS = 4 %}
195
+ {% if PV_HALF %}
196
+ {% for g in range(qkGroups) %}
197
+ let p{{ g }} = vec4<f16>(qk{{ g }});
198
+ {% endfor %}
199
+ {% endif %}
200
 
201
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
202
  {% if USE_SUBGROUPS %}
 
213
  }
214
  {% endif %}
215
  {% endif %}
216
+ {% if PV_HALF %}
217
+ var acc_f32: vec4<f32> = vec4<f32>(0.0);
218
+ var acc: vec4<f16> = vec4<f16>(0.0);
219
+ {% else %}
220
  var acc: vec4<f32> = vec4<f32>(0.0);
221
+ {% endif %}
222
  {% for g in range(qkGroups) %}
223
  {% for lane in range(4) %}
224
+ {% if not USE_SUBGROUPS %}
225
+ {% set pvValue = "valueTile[d4 * K_STEP + " ~ (g * 4 + lane) ~ "u]" %}
226
+ {% elif g < 8 %}
227
+ {% set pvValue = "subgroupShuffle(v_local0, " ~ (g * 4 + lane) ~ "u)" %}
228
  {% else %}
229
+ {% set pvValue = "subgroupShuffle(v_local1, " ~ ((g - 8) * 4 + lane) ~ "u)" %}
230
  {% endif %}
231
+ {% if PV_HALF %}
232
+ acc = fma({{ pvValue }}, vec4<f16>(p{{ g }}.{{ components[lane] }}), acc);
233
  {% else %}
234
+ acc = acc + vec4<f32>({{ pvValue }}) * qk{{ g }}.{{ components[lane] }};
235
  {% endif %}
236
  {% endfor %}
237
+ {% if PV_HALF and (g + 1) % PV_FLUSH_GROUPS == 0 %}
238
+ acc_f32 = acc_f32 + vec4<f32>(acc);
239
+ acc = vec4<f16>(0.0);
240
+ {% endif %}
241
  {% endfor %}
242
+ {% if PV_HALF %}
243
+ o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc_f32;
244
+ {% else %}
245
  o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
246
+ {% endif %}
247
  }
248
  {% if not USE_SUBGROUPS %}
249
  workgroupBarrier();
build/webgpu/attn-materialized-rowstats-combine-f32.wgsl.jinja CHANGED
@@ -12,6 +12,12 @@
12
  // `maxOnly` means the producer published only a row max per slot because
13
  // computing the denominator there would double the exp count. Fold maxima and
14
  // leave the denominator to the apply pass, which sees every row element anyway.
 
 
 
 
 
 
15
  {{ env.wgsl.resourceDeclarations }}
16
 
17
  const SLOTS: u32 = {{ statSlots }}u;
@@ -37,6 +43,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
37
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
38
  return select(value - maxValue, 0.0, equalFiniteMax);
39
  }
 
40
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
41
  return exp(shifted_value(value, maxValue));
42
  }
@@ -45,8 +52,7 @@ fn exp_shift(value: f32, maxValue: f32) -> f32 {
45
  fn main(
46
  @builtin(global_invocation_id) gid: vec3<u32>
47
  ) {
48
- let row = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
49
- if (row >= params.rows) { return; }
50
 
51
  // `row` already runs over (batch, head, query) together, and the partial
52
  // layout puts that same product one axis out from the slot, so the stride
 
12
  // `maxOnly` means the producer published only a row max per slot because
13
  // computing the denominator there would double the exp count. Fold maxima and
14
  // leave the denominator to the apply pass, which sees every row element anyway.
15
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
16
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
17
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
18
+ // per-axis workgroup fold width.
19
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
20
+ if ({{ name }} >= {{ bound }}) { return; }{% endmacro %}
21
  {{ env.wgsl.resourceDeclarations }}
22
 
23
  const SLOTS: u32 = {{ statSlots }}u;
 
43
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
44
  return select(value - maxValue, 0.0, equalFiniteMax);
45
  }
46
+
47
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
48
  return exp(shifted_value(value, maxValue));
49
  }
 
52
  fn main(
53
  @builtin(global_invocation_id) gid: vec3<u32>
54
  ) {
55
+ {{ flat_index_2d("WG", "row", "params.rows", guardInline=true) }}
 
56
 
57
  // `row` already runs over (batch, head, query) together, and the partial
58
  // layout puts that same product one axis out from the slot, so the stride
build/webgpu/attn-materialized-sgmat-f32.wgsl.jinja CHANGED
@@ -1,7 +1,5 @@
1
  {% set MT = "f16" if (operandF16 is defined and operandF16) else "f32" %}
2
- {% if MT == "f16" %}
3
- enable f16;
4
- {% endif %}
5
  enable subgroups;
6
  {% if pinSubgroupSize32 %}
7
  enable subgroup_size_control;
@@ -9,18 +7,16 @@ enable subgroup_size_control;
9
  enable chromium_experimental_subgroup_matrix;
10
  diagnostic(off, chromium.subgroup_matrix_uniformity);
11
 
12
-
13
  {{ env.wgsl.resourceDeclarations }}
14
 
15
- {% set layout = layout | default("bsh") %}
16
- {% set headMajor = layout == "bhsd" %}
17
  {% set kvHeadMajor = (kvLayout | default(layout)) == "bhsd" %}
18
- {% set CAUSAL_UPPER_LEFT = causalUpperLeft is defined and causalUpperLeft %}
19
  {% set CAUSAL = (causalRightAlign is defined and causalRightAlign) or CAUSAL_UPPER_LEFT %}
20
- {% macro q_index(row, d) %}{% if headMajor %}((b * HEADS + h) * params.qSeq + {{ row }}) * HEAD_DIM + {{ d }}{% else %}(b * params.qSeq + {{ row }}) * HIDDEN + h * HEAD_DIM + {{ d }}{% endif %}{% endmacro %}
21
  {% macro kv_index(seq, d) %}{% if kvHeadMajor %}((b * KV_HEADS + h_kv) * params.kvSeq + {{ seq }}) * HEAD_DIM + {{ d }}{% else %}(b * params.kvSeq + {{ seq }}) * KV_HIDDEN + h_kv * HEAD_DIM + {{ d }}{% endif %}{% endmacro %}
22
- {% set OUT_ROW_STRIDE = "HEAD_DIM" if headMajor else "HIDDEN" %}
23
  {% set KV_ROW_STRIDE = "HEAD_DIM" if kvHeadMajor else "KV_HIDDEN" %}
 
24
  {% set scorePhase = phase == "score" %}
25
  {% set SCORE_BIAS = scorePhase and scoreBias is defined and scoreBias %}
26
  {% set SCORE_WINDOW = CAUSAL and scoreWindow is defined and scoreWindow %}
@@ -28,45 +24,38 @@ diagnostic(off, chromium.subgroup_matrix_uniformity);
28
  {% set FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
29
  {% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
30
  {% macro score_value(index, guard) %}
31
- {% if FUSED_SOFTMAX %}
32
- {% if PRIVATE_ROW_STATS %}
33
- select(0.0, exp_shift(scores[{{ index }}], private_softmax_m) / private_softmax_d, {{ guard[1] }})
34
- {%- else %}
35
- select(0.0, exp_shift(scores[{{ index }}], softmax_m[{{ guard[0] }}]) / softmax_d[{{ guard[0] }}], {{ guard[1] }})
36
- {%- endif %}
37
- {% else %}
38
- select(0.0, scores[{{ index }}], {{ guard[1] }})
39
- {%- endif %}
40
- {% endmacro %}
41
  {% set TILE_M_VALUE = materializedSgmatQueryTile %}
42
  {% set TILE_N_VALUE = materializedSgmatKeyTile %}
43
  {% set TILE_K_VALUE = materializedSgmatInnerTile %}
44
- {% set SUB_ROWS_VALUE = materializedSgmatSubgroupTileRows if materializedSgmatSubgroupTileRows is defined else 16 %}
45
- {% set SUB_COLS_VALUE = materializedSgmatSubgroupTileCols if materializedSgmatSubgroupTileCols is defined else 32 %}
46
- {% set ROW_BLOCKS = (SUB_ROWS_VALUE / 8)|int %}
47
- {% set COL_BLOCKS = (SUB_COLS_VALUE / 8)|int %}
48
  {% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
49
  {% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
50
  {% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
51
  {% set WORKGROUP_THREADS = SUBGROUP_COUNT * 32 %}
52
  {% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
53
  {% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
54
- {% set DIRECT_SCORE_STORE = materializedSgmatDirectScoreStore and not EMIT_ROW_STATS %}
55
- {% set DIRECT_APPLY_STORE = materializedSgmatDirectApplyStore and not hasBias and MT == "f32" %}
56
  {% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
57
  {% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
58
  and materializedSgmatRuntimeDirectStore
59
- and (scorePhase or (not hasBias and MT == "f32"))
60
  and not EMIT_ROW_STATS %}
61
  {% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
62
- {% macro q_tile_value(index) %}{% if hasBias %}(query[{{ index }}] + bias[h * HEAD_DIM + k]){% else %}query[{{ index }}]{% endif %}{% endmacro %}
 
 
63
 
64
  const HEADS: u32 = {{ qNumHeads }}u;
65
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
66
  const HEAD_DIM: u32 = {{ headDim }}u;
67
- {% if layout != "bhsd" or hasBias %}
68
  const HIDDEN: u32 = {{ qHidden }}u;
69
- {% endif %}
70
  {% if hasBias %}
71
  /* Packed [Q; K; V] bias. Q folds into the query tile before the GEMM. K is
72
  * omitted: expanding (q + bq).(k + bk) leaves a (q + bq).bk term that is
@@ -108,14 +97,12 @@ var<workgroup> softmax_m: array<f32, {{ TILE_M_VALUE }}>;
108
  var<workgroup> softmax_d: array<f32, {{ TILE_M_VALUE }}>;
109
  {% endif %}
110
  {% if FUSED_SOFTMAX or EMIT_ROW_STATS %}
111
- {% set stableUsage = stableHelperUsage if stableHelperUsage is defined else "all" %}
112
  // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
113
  // `m - m` finite so an empty lane / all--inf row contributes the exact
114
  // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
115
  // a zero final denominator according to their public semantics. Using -inf
116
  // here changes +inf-row behavior.
117
  const FLT_MAX: f32 = 3.4028234663852886e38;
118
- {% if stableUsage != "constant" %}
119
 
120
  fn is_finite_f32(value: f32) -> bool {
121
  return select(false, value <= FLT_MAX, value >= -FLT_MAX);
@@ -129,14 +116,10 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
129
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
130
  return select(value - maxValue, 0.0, equalFiniteMax);
131
  }
132
- {%- endif %}
133
- {% if stableUsage == "all" %}
134
 
135
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
136
  return exp(shifted_value(value, maxValue));
137
  }
138
- {%- endif %}
139
-
140
  {% endif %}
141
 
142
  @compute @workgroup_size({{ WORKGROUP_THREADS }}, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
@@ -177,7 +160,7 @@ fn main(
177
  let stat_row = m_base + idx / {{ SUBGROUP_COLS }}u;
178
  if (stat_row < params.qSeq) {
179
  let slot = wg.x * {{ SUBGROUP_COLS }}u + idx % {{ SUBGROUP_COLS }}u;
180
- let out_index = (((b * HEADS + h) * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
181
  scorePartials[out_index] = -FLT_MAX;
182
  scorePartials[out_index + 1u] = 0.0;
183
  }
@@ -211,30 +194,18 @@ fn main(
211
  {% endif %}
212
  {% endif %}
213
  {% if FUSED_SOFTMAX %}
214
- {% if PRIVATE_ROW_STATS %}
215
- // In the admitted BM64/BN64/BK32/WG256 loader, four adjacent lanes own the
216
- // same query row for every reduction tile. Keep that row's constants private:
217
- // this removes both 512 bytes of workgroup storage and the initialization
218
- // barrier while preserving the exact exp/divide sequence of the shared-memory
219
- // row-stats arm.
220
- let private_stat_row =
221
- (b * HEADS + h) * params.qSeq + min(m_base + li / 4u, params.qSeq - 1u);
222
- let private_softmax_m = rowStats[private_stat_row * 2u];
223
- let private_softmax_d = rowStats[private_stat_row * 2u + 1u];
224
- {% else %}
225
  // One row-stats pair per query row of the tile. A query tail clamps to the last
226
  // real row rather than reading past the buffer; those lanes are discarded by the
227
  // staging guard anyway, and the clamp keeps the denominator non-zero.
228
  for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
229
- let stat_row = (b * HEADS + h) * params.qSeq + min(m_base + r, params.qSeq - 1u);
230
  softmax_m[r] = rowStats[stat_row * 2u];
231
  softmax_d[r] = rowStats[stat_row * 2u + 1u];
232
  }
233
  workgroupBarrier();
234
- {% endif %}
235
  {% endif %}
236
  for (var k_base = {% if SCORE_WINDOW and not scorePhase %}inner_start{% else %}0u{% endif %}; k_base < {% if CAUSAL and not scorePhase %}inner_bound{% else %}inner{% endif %}; k_base += TILE_K) {
237
- {% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" %}
238
  // Full interior PV tiles can be loaded directly from storage. Query,
239
  // reduction, and output-dimension tails use the guarded shared path below.
240
  if (
@@ -244,7 +215,7 @@ fn main(
244
  ) {
245
  for (var step = 0u; step < TILE_K; step += 8u) {
246
  {% for row_block in range(ROW_BLOCKS) %}
247
- let score_offset{{ row_block }} = (b * HEADS + h) * params.qSeq * params.kvSeq
248
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
249
  var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
250
  subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
@@ -294,7 +265,7 @@ fn main(
294
  tile_A[a_row * TILE_K + a_col + i] = loaded;
295
  {% endif %}
296
  {% else %}
297
- let score_base = (b * HEADS + h) * params.qSeq * params.kvSeq;
298
  tile_A[a_row * TILE_K + a_col + i] =
299
  {{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["a_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
300
  {% endif %}
@@ -309,20 +280,20 @@ fn main(
309
  {% if headDim % 32 == 0 %}
310
  tile_B[b_row * TILE_K + b_col + i] = select(
311
  {{ "0.0h" if MT == "f16" else "0.0" }},
312
- key[{{ kv_index("col", "k") }}],
313
  col < params.kvSeq
314
  );
315
  {% else %}
316
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
317
  if (col < params.kvSeq && k < HEAD_DIM) {
318
- loaded = key[{{ kv_index("col", "k") }}];
319
  }
320
  tile_B[b_row * TILE_K + b_col + i] = loaded;
321
  {% endif %}
322
  {% else %}
323
  tile_B[b_row * TILE_K + b_col + i] = select(
324
  {{ "0.0h" if MT == "f16" else "0.0" }},
325
- value[{{ kv_index("k", "col") }}],
326
  k < params.kvSeq && col < HEAD_DIM
327
  );
328
  {% endif %}
@@ -341,13 +312,13 @@ fn main(
341
  loaded = {{ q_tile_value(q_index("row", "k")) }};
342
  }
343
  {% elif FUSED_SOFTMAX %}
344
- let score_base = (b * HEADS + h) * params.qSeq * params.kvSeq;
345
  let loaded =
346
  {{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["tile_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
347
  {% else %}
348
  var loaded = 0.0;
349
  if (row < params.qSeq && k < params.kvSeq) {
350
- let score_base = (b * HEADS + h) * params.qSeq * params.kvSeq;
351
  loaded = scores[score_base + row * params.kvSeq + k];
352
  }
353
  {% endif %}
@@ -362,12 +333,12 @@ fn main(
362
  {% if scorePhase %}
363
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
364
  if (col < params.kvSeq && k < HEAD_DIM) {
365
- loaded = key[{{ kv_index("col", "k") }}];
366
  }
367
  {% else %}
368
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
369
  if (k < params.kvSeq && col < HEAD_DIM) {
370
- loaded = value[{{ kv_index("k", "col") }}];
371
  }
372
  {% endif %}
373
  tile_B[idx] = loaded;
@@ -410,7 +381,7 @@ fn main(
410
  {% for col_block in range(COL_BLOCKS) %}
411
  {% if scorePhase %}
412
  let output_offset{{ row_block }}{{ col_block }} =
413
- (b * HEADS + h) * params.qSeq * params.kvSeq
414
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
415
  + n_base + base_B + {{ col_block * 8 }}u;
416
  subgroupMatrixStore<row_major>(
@@ -467,7 +438,9 @@ fn main(
467
  + row_in_block * 8u + col_in_block + pair
468
  ];
469
  {% if scorePhase %}
 
470
  let scale = {{ attentionScaleExpression }};
 
471
  {% if CAUSAL %}
472
  var scored = result * scale;
473
  {% if SCORE_BIAS %}
@@ -485,10 +458,10 @@ fn main(
485
  if (i32(col) + i32(params.windowSize) <= kv_causal_off + i32(row)) { scored = -FLT_MAX; }
486
  {% endif %}
487
  {% else %}
488
- let scored = result * scale;
489
  {% endif %}
490
  scores[
491
- (b * HEADS + h) * params.qSeq * params.kvSeq + row * params.kvSeq + col
492
  ] = scored;
493
  {% if EMIT_ROW_STATS %}
494
  // Softmax sees the STORED value, so the statistics have to be taken on it
@@ -502,7 +475,7 @@ fn main(
502
  {% if hasBias %}
503
  // V bias row base: skip the packed Q and K blocks, then index this head.
504
  {% endif %}
505
- output[{{ q_index("row", "col") }}] = {{ "f16(" if MT == "f16" else "" }}result{{ ")" if MT == "f16" else "" }}{% if hasBias %} + bias[2u * HIDDEN + h * HEAD_DIM + col]{% endif %};
506
  {% endif %}
507
  }
508
  }
@@ -530,7 +503,7 @@ fn main(
530
  // eight consecutive pairs, and the combine pass reads a slot's whole column
531
  // of rows contiguously.
532
  let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
533
- let out_index = (((b * HEADS + h) * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
534
  scorePartials[out_index] = stat_m{{ row_block }};
535
  scorePartials[out_index + 1u] = stat_d{{ row_block }};
536
  }
 
1
  {% set MT = "f16" if (operandF16 is defined and operandF16) else "f32" %}
2
+ {% set ST = "f16" if (storageF16 is defined and storageF16) else MT %}
 
 
3
  enable subgroups;
4
  {% if pinSubgroupSize32 %}
5
  enable subgroup_size_control;
 
7
  enable chromium_experimental_subgroup_matrix;
8
  diagnostic(off, chromium.subgroup_matrix_uniformity);
9
 
 
10
  {{ env.wgsl.resourceDeclarations }}
11
 
 
 
12
  {% set kvHeadMajor = (kvLayout | default(layout)) == "bhsd" %}
13
+ {% set CAUSAL_UPPER_LEFT = false %}
14
  {% set CAUSAL = (causalRightAlign is defined and causalRightAlign) or CAUSAL_UPPER_LEFT %}
15
+ {% macro q_index(row, d) %}(b * params.qSeq + {{ row }}) * HIDDEN + h * HEAD_DIM + {{ d }}{% endmacro %}
16
  {% macro kv_index(seq, d) %}{% if kvHeadMajor %}((b * KV_HEADS + h_kv) * params.kvSeq + {{ seq }}) * HEAD_DIM + {{ d }}{% else %}(b * params.kvSeq + {{ seq }}) * KV_HIDDEN + h_kv * HEAD_DIM + {{ d }}{% endif %}{% endmacro %}
17
+ {% set OUT_ROW_STRIDE = "HIDDEN" %}
18
  {% set KV_ROW_STRIDE = "HEAD_DIM" if kvHeadMajor else "KV_HIDDEN" %}
19
+ {% set SCRATCH_HEAD = "(b * HEADS + h)" %}
20
  {% set scorePhase = phase == "score" %}
21
  {% set SCORE_BIAS = scorePhase and scoreBias is defined and scoreBias %}
22
  {% set SCORE_WINDOW = CAUSAL and scoreWindow is defined and scoreWindow %}
 
24
  {% set FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
25
  {% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
26
  {% macro score_value(index, guard) %}
27
+ {% set rowMax = "private_softmax_m" if PRIVATE_ROW_STATS else "softmax_m[" ~ guard[0] ~ "]" %}
28
+ {% set rowDenom = "private_softmax_d" if PRIVATE_ROW_STATS else "softmax_d[" ~ guard[0] ~ "]" %}
29
+ select(0.0, {{ "exp_shift(scores[" ~ index ~ "], " ~ rowMax ~ ") / " ~ rowDenom if FUSED_SOFTMAX else "scores[" ~ index ~ "]" }}, {{ guard[1] }}){% endmacro %}
 
 
 
 
 
 
 
30
  {% set TILE_M_VALUE = materializedSgmatQueryTile %}
31
  {% set TILE_N_VALUE = materializedSgmatKeyTile %}
32
  {% set TILE_K_VALUE = materializedSgmatInnerTile %}
33
+ {% set SUB_ROWS_VALUE = 16 %}
34
+ {% set SUB_COLS_VALUE = 32 %}
35
+ {% set ROW_BLOCKS = 2 %}
36
+ {% set COL_BLOCKS = 4 %}
37
  {% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
38
  {% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
39
  {% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
40
  {% set WORKGROUP_THREADS = SUBGROUP_COUNT * 32 %}
41
  {% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
42
  {% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
43
+ {% set DIRECT_SCORE_STORE = false %}
44
+ {% set DIRECT_APPLY_STORE = false %}
45
  {% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
46
  {% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
47
  and materializedSgmatRuntimeDirectStore
48
+ and (scorePhase or (not hasBias and MT == "f32" and ST == "f32"))
49
  and not EMIT_ROW_STATS %}
50
  {% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
51
+ {% set SCALE_IN_Q = DIRECT_SCORE_STORE or (RUNTIME_DIRECT_STORE and scorePhase) %}
52
+ {% macro operand_load(name, index) %}{% if ST != MT %}{{ MT }}({% endif %}{{ name }}[{{ index }}]{% if ST != MT %}){% endif %}{% endmacro %}
53
+ {% macro q_tile_value(index) %}{% if hasBias %}({{ operand_load("query", index) }} + bias[h * HEAD_DIM + k]){% else %}{{ operand_load("query", index) }}{% endif %}{% if SCALE_IN_Q %} * score_scale{% endif %}{% endmacro %}
54
 
55
  const HEADS: u32 = {{ qNumHeads }}u;
56
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
57
  const HEAD_DIM: u32 = {{ headDim }}u;
 
58
  const HIDDEN: u32 = {{ qHidden }}u;
 
59
  {% if hasBias %}
60
  /* Packed [Q; K; V] bias. Q folds into the query tile before the GEMM. K is
61
  * omitted: expanding (q + bq).(k + bk) leaves a (q + bq).bk term that is
 
97
  var<workgroup> softmax_d: array<f32, {{ TILE_M_VALUE }}>;
98
  {% endif %}
99
  {% if FUSED_SOFTMAX or EMIT_ROW_STATS %}
 
100
  // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
101
  // `m - m` finite so an empty lane / all--inf row contributes the exact
102
  // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
103
  // a zero final denominator according to their public semantics. Using -inf
104
  // here changes +inf-row behavior.
105
  const FLT_MAX: f32 = 3.4028234663852886e38;
 
106
 
107
  fn is_finite_f32(value: f32) -> bool {
108
  return select(false, value <= FLT_MAX, value >= -FLT_MAX);
 
116
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
117
  return select(value - maxValue, 0.0, equalFiniteMax);
118
  }
 
 
119
 
120
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
121
  return exp(shifted_value(value, maxValue));
122
  }
 
 
123
  {% endif %}
124
 
125
  @compute @workgroup_size({{ WORKGROUP_THREADS }}, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
 
160
  let stat_row = m_base + idx / {{ SUBGROUP_COLS }}u;
161
  if (stat_row < params.qSeq) {
162
  let slot = wg.x * {{ SUBGROUP_COLS }}u + idx % {{ SUBGROUP_COLS }}u;
163
+ let out_index = (({{ SCRATCH_HEAD }} * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
164
  scorePartials[out_index] = -FLT_MAX;
165
  scorePartials[out_index + 1u] = 0.0;
166
  }
 
194
  {% endif %}
195
  {% endif %}
196
  {% if FUSED_SOFTMAX %}
 
 
 
 
 
 
 
 
 
 
 
197
  // One row-stats pair per query row of the tile. A query tail clamps to the last
198
  // real row rather than reading past the buffer; those lanes are discarded by the
199
  // staging guard anyway, and the clamp keeps the denominator non-zero.
200
  for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
201
+ let stat_row = {{ SCRATCH_HEAD }} * params.qSeq + min(m_base + r, params.qSeq - 1u);
202
  softmax_m[r] = rowStats[stat_row * 2u];
203
  softmax_d[r] = rowStats[stat_row * 2u + 1u];
204
  }
205
  workgroupBarrier();
 
206
  {% endif %}
207
  for (var k_base = {% if SCORE_WINDOW and not scorePhase %}inner_start{% else %}0u{% endif %}; k_base < {% if CAUSAL and not scorePhase %}inner_bound{% else %}inner{% endif %}; k_base += TILE_K) {
208
+ {% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" and ST == "f32" %}
209
  // Full interior PV tiles can be loaded directly from storage. Query,
210
  // reduction, and output-dimension tails use the guarded shared path below.
211
  if (
 
215
  ) {
216
  for (var step = 0u; step < TILE_K; step += 8u) {
217
  {% for row_block in range(ROW_BLOCKS) %}
218
+ let score_offset{{ row_block }} = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
219
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
220
  var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
221
  subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
 
265
  tile_A[a_row * TILE_K + a_col + i] = loaded;
266
  {% endif %}
267
  {% else %}
268
+ let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
269
  tile_A[a_row * TILE_K + a_col + i] =
270
  {{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["a_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
271
  {% endif %}
 
280
  {% if headDim % 32 == 0 %}
281
  tile_B[b_row * TILE_K + b_col + i] = select(
282
  {{ "0.0h" if MT == "f16" else "0.0" }},
283
+ {{ operand_load("key", kv_index("col", "k")) }},
284
  col < params.kvSeq
285
  );
286
  {% else %}
287
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
288
  if (col < params.kvSeq && k < HEAD_DIM) {
289
+ loaded = {{ operand_load("key", kv_index("col", "k")) }};
290
  }
291
  tile_B[b_row * TILE_K + b_col + i] = loaded;
292
  {% endif %}
293
  {% else %}
294
  tile_B[b_row * TILE_K + b_col + i] = select(
295
  {{ "0.0h" if MT == "f16" else "0.0" }},
296
+ {{ operand_load("value", kv_index("k", "col")) }},
297
  k < params.kvSeq && col < HEAD_DIM
298
  );
299
  {% endif %}
 
312
  loaded = {{ q_tile_value(q_index("row", "k")) }};
313
  }
314
  {% elif FUSED_SOFTMAX %}
315
+ let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
316
  let loaded =
317
  {{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["tile_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
318
  {% else %}
319
  var loaded = 0.0;
320
  if (row < params.qSeq && k < params.kvSeq) {
321
+ let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
322
  loaded = scores[score_base + row * params.kvSeq + k];
323
  }
324
  {% endif %}
 
333
  {% if scorePhase %}
334
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
335
  if (col < params.kvSeq && k < HEAD_DIM) {
336
+ loaded = {{ operand_load("key", kv_index("col", "k")) }};
337
  }
338
  {% else %}
339
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
340
  if (k < params.kvSeq && col < HEAD_DIM) {
341
+ loaded = {{ operand_load("value", kv_index("k", "col")) }};
342
  }
343
  {% endif %}
344
  tile_B[idx] = loaded;
 
381
  {% for col_block in range(COL_BLOCKS) %}
382
  {% if scorePhase %}
383
  let output_offset{{ row_block }}{{ col_block }} =
384
+ {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
385
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
386
  + n_base + base_B + {{ col_block * 8 }}u;
387
  subgroupMatrixStore<row_major>(
 
438
  + row_in_block * 8u + col_in_block + pair
439
  ];
440
  {% if scorePhase %}
441
+ {% if not SCALE_IN_Q %}
442
  let scale = {{ attentionScaleExpression }};
443
+ {% endif %}
444
  {% if CAUSAL %}
445
  var scored = result * scale;
446
  {% if SCORE_BIAS %}
 
458
  if (i32(col) + i32(params.windowSize) <= kv_causal_off + i32(row)) { scored = -FLT_MAX; }
459
  {% endif %}
460
  {% else %}
461
+ let scored = result{% if not SCALE_IN_Q %} * scale{% endif %};
462
  {% endif %}
463
  scores[
464
+ {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq + row * params.kvSeq + col
465
  ] = scored;
466
  {% if EMIT_ROW_STATS %}
467
  // Softmax sees the STORED value, so the statistics have to be taken on it
 
475
  {% if hasBias %}
476
  // V bias row base: skip the packed Q and K blocks, then index this head.
477
  {% endif %}
478
+ output[{{ q_index("row", "col") }}] = {{ "f16(" if ST == "f16" else "" }}result{{ ")" if ST == "f16" else "" }}{% if hasBias %} + bias[2u * HIDDEN + h * HEAD_DIM + col]{% endif %};
479
  {% endif %}
480
  }
481
  }
 
503
  // eight consecutive pairs, and the combine pass reads a slot's whole column
504
  // of rows contiguously.
505
  let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
506
+ let out_index = (({{ SCRATCH_HEAD }} * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
507
  scorePartials[out_index] = stat_m{{ row_block }};
508
  scorePartials[out_index + 1u] = stat_d{{ row_block }};
509
  }
build/webgpu/attn-online-scalar.wgsl.jinja CHANGED
@@ -1,14 +1,17 @@
1
  {{ env.wgsl.resourceDeclarations }}
2
- {% set MASK_BATCH = "batch * params.maskBatchStride + " if layout == "bsh" else "" %}
3
 
4
  // Online-softmax attention fallback with no feature requirements: one
5
  // workgroup per (batch, head, query token) walks the keys serially; the
6
- // workgroup cooperates on each q·k dot (tree reduction) and on the running
7
- // V accumulator, with the online rescale applied per key. This path requires
8
- // no subgroup or subgroup-matrix features.
9
  // Layout: rank-3 token-major [batch, seq, heads * headDim].
10
  // An omitted scale uses 1/sqrt(headDim); explicit zero remains zero through
11
  // the scaleIsExplicitZero specialization.
 
 
 
12
  {% if hasMask %}
13
  // Additive per-score attention bias with broadcast strides: a stride of 0
14
  // collapses that axis (batch and/or head broadcast).
@@ -18,18 +21,39 @@ const Q_HIDDEN: u32 = {{ qHidden }}u;
18
  const KV_HIDDEN: u32 = {{ kvHidden }}u;
19
  const Q_HEADS: u32 = {{ qNumHeads }}u;
20
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
21
- {% set qHeads = "params.qHeads" if headsFromParams else "Q_HEADS" %}
22
- {% set kvHeads = "params.kvHeads" if headsFromParams else "KV_HEADS" %}
23
- {% set scale = scale | default("0.0") %}
24
  const WG: u32 = {{ workgroupSize }}u;
25
 
26
- var<workgroup> partial: array<f32, WG>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
27
  var<workgroup> running_max: f32;
28
  var<workgroup> running_denom: f32;
29
  var<workgroup> running_out: array<f32, HEAD_DIM>;
30
  var<workgroup> previous_scale: f32;
31
- {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
32
- fn {{ name }}(value: f32, tid: u32) -> f32 {
33
  {{ buffer }}[tid] = value;
34
  workgroupBarrier();
35
  // Ceil-halving keeps every lane when the workgroup size is not a power of
@@ -39,11 +63,7 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
39
  loop {
40
  let half = (n + 1u) / 2u;
41
  if (tid < n - half) {
42
- {% if mode == "max" %}
43
- {{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
44
- {% else %}
45
- {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
46
- {% endif %}
47
  }
48
  workgroupBarrier();
49
  n = half;
@@ -55,22 +75,16 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
55
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
56
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
57
  let reduced = {{ buffer }}[0];
58
- {% if trailingBarrier %}
59
  workgroupBarrier();
60
- {% endif %}
61
  return reduced;
62
  }
63
  {% endmacro %}
64
-
65
- {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
66
-
67
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
68
- fn scale_value() -> f32 {
69
  if (params.scale != 0.0) { return params.scale; }
70
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
71
  }
72
 
73
-
74
  @compute @workgroup_size(WG, 1, 1)
75
  fn main(
76
  @builtin(workgroup_id) wg: vec3<u32>,
@@ -110,6 +124,9 @@ fn main(
110
 
111
  var maxKj = params.kvSeq;
112
  var minKj: u32 = 0u;
 
 
 
113
  {% if hasWindow %}
114
  // local_window_size: the query at relative index
115
  // query_token sits at absolute position p = kvSeq - qSeq + query_token, so it
@@ -124,21 +141,26 @@ fn main(
124
 
125
  for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
126
  let kRow = kvBase + key_token * kvTokenStride;
127
- var partial_dot = 0.0;
128
  for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
129
  var q_value = f32(query[qBase + d]);
130
  var k_value = f32(key[kRow + d]);
131
- partial_dot = partial_dot + q_value * k_value;
 
 
 
 
 
132
  }
133
 
134
- // reduce_sum returns the same partial[0] to every lane, so `score` is already
135
  // workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
136
  // through shared memory by the barrier below).
137
  {% if hasMask %}
138
  let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
139
- let score = reduce_sum(partial_dot, tid) * scale_value() + f32(attn_mask[maskIndex]);
140
  {% else %}
141
- let score = reduce_sum(partial_dot, tid) * scale_value();
142
  {% endif %}
143
  if (tid == 0u) {
144
  let next_max = max(running_max, score);
@@ -153,6 +175,9 @@ fn main(
153
  let probability_numerator = exp(score - running_max);
154
  for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
155
  var v_value = f32(value[kRow + d]);
 
 
 
156
  running_out[d] = running_out[d] * previous_scale + probability_numerator * v_value;
157
  }
158
  workgroupBarrier();
 
1
  {{ env.wgsl.resourceDeclarations }}
2
+ {% set MASK_BATCH = "batch * params.maskBatchStride + " %}
3
 
4
  // Online-softmax attention fallback with no feature requirements: one
5
  // workgroup per (batch, head, query token) walks the keys serially; the
6
+ // workgroup cooperates on each q·k dot (compensated tree reduction) and on the
7
+ // running V accumulator, with the online rescale applied per key. This path
8
+ // requires no subgroup or subgroup-matrix features.
9
  // Layout: rank-3 token-major [batch, seq, heads * headDim].
10
  // An omitted scale uses 1/sqrt(headDim); explicit zero remains zero through
11
  // the scaleIsExplicitZero specialization.
12
+ {% if hasBias %}
13
+ // Packed [Q; K; V] bias rows are applied during the serial key walk.
14
+ {% endif %}
15
  {% if hasMask %}
16
  // Additive per-score attention bias with broadcast strides: a stride of 0
17
  // collapses that axis (batch and/or head broadcast).
 
21
  const KV_HIDDEN: u32 = {{ kvHidden }}u;
22
  const Q_HEADS: u32 = {{ qNumHeads }}u;
23
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
24
+ {% set qHeads = "Q_HEADS" %}
25
+ {% set kvHeads = "KV_HEADS" %}
 
26
  const WG: u32 = {{ workgroupSize }}u;
27
 
28
+ {% set dotType = "f32" %}
29
+ // Retain product and addition residuals across a dot product. Explicit fma
30
+ // boundaries preserve the addition error transform under reassociation.
31
+ struct DotAccumulator {
32
+ hi: {{ dotType }},
33
+ lo: {{ dotType }},
34
+ }
35
+ fn dot_accumulate(acc: DotAccumulator, a: {{ dotType }}, b: {{ dotType }}) -> DotAccumulator {
36
+ let product = fma(a, b, {{ dotType }}(0.0));
37
+ let productError = fma(a, b, -product);
38
+ let sum = fma(acc.hi, {{ dotType }}(1.0), product);
39
+ let bv = fma({{ dotType }}(-1.0), acc.hi, sum);
40
+ let av = fma({{ dotType }}(-1.0), bv, sum);
41
+ let ae = fma({{ dotType }}(-1.0), av, acc.hi);
42
+ let be = fma({{ dotType }}(-1.0), bv, product);
43
+ let error = fma(fma(fma(ae, {{ dotType }}(1.0), be), {{ dotType }}(1.0), acc.lo), {{ dotType }}(1.0), productError);
44
+ let hi = fma(sum, {{ dotType }}(1.0), error);
45
+ return DotAccumulator(hi, fma({{ dotType }}(-1.0), fma({{ dotType }}(-1.0), sum, hi), error));
46
+ }
47
+ fn dot_value(acc: DotAccumulator) -> {{ dotType }} {
48
+ return fma(acc.hi, {{ dotType }}(1.0), acc.lo);
49
+ }
50
+ var<workgroup> partial: array<DotAccumulator, WG>;
51
  var<workgroup> running_max: f32;
52
  var<workgroup> running_denom: f32;
53
  var<workgroup> running_out: array<f32, HEAD_DIM>;
54
  var<workgroup> previous_scale: f32;
55
+ {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
56
+ fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
57
  {{ buffer }}[tid] = value;
58
  workgroupBarrier();
59
  // Ceil-halving keeps every lane when the workgroup size is not a power of
 
63
  loop {
64
  let half = (n + 1u) / 2u;
65
  if (tid < n - half) {
66
+ {{ buffer }}[tid] = dot_accumulate(dot_accumulate({{ buffer }}[tid], {{ buffer }}[tid + half].hi, 1.0), {{ buffer }}[tid + half].lo, 1.0);
 
 
 
 
67
  }
68
  workgroupBarrier();
69
  n = half;
 
75
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
76
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
77
  let reduced = {{ buffer }}[0];
 
78
  workgroupBarrier();
 
79
  return reduced;
80
  }
81
  {% endmacro %}
82
+ {{ wgsl_tree_reduce_f32("reduce_dot", "compensated", "partial", "WG", valueType="DotAccumulator") }}
83
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
 
 
84
  if (params.scale != 0.0) { return params.scale; }
85
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
86
  }
87
 
 
88
  @compute @workgroup_size(WG, 1, 1)
89
  fn main(
90
  @builtin(workgroup_id) wg: vec3<u32>,
 
124
 
125
  var maxKj = params.kvSeq;
126
  var minKj: u32 = 0u;
127
+ {% if hasCausal %}
128
+ maxKj = min(maxKj, select(params.kvSeq, query_token + 1u, params.isCausal != 0u));
129
+ {% endif %}
130
  {% if hasWindow %}
131
  // local_window_size: the query at relative index
132
  // query_token sits at absolute position p = kvSeq - qSeq + query_token, so it
 
141
 
142
  for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
143
  let kRow = kvBase + key_token * kvTokenStride;
144
+ var partial_dot = DotAccumulator(0.0, 0.0);
145
  for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
146
  var q_value = f32(query[qBase + d]);
147
  var k_value = f32(key[kRow + d]);
148
+ {% if hasBias %}
149
+ let channel = h * HEAD_DIM + d;
150
+ q_value = q_value + f32(bias[channel]);
151
+ k_value = k_value + f32(bias[Q_HIDDEN + channel]);
152
+ {% endif %}
153
+ partial_dot = dot_accumulate(partial_dot, q_value, k_value);
154
  }
155
 
156
+ // reduce_dot returns the same partial[0] to every lane, so `score` is already
157
  // workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
158
  // through shared memory by the barrier below).
159
  {% if hasMask %}
160
  let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
161
+ let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value() + f32(attn_mask[maskIndex]);
162
  {% else %}
163
+ let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value();
164
  {% endif %}
165
  if (tid == 0u) {
166
  let next_max = max(running_max, score);
 
175
  let probability_numerator = exp(score - running_max);
176
  for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
177
  var v_value = f32(value[kRow + d]);
178
+ {% if hasBias %}
179
+ v_value = v_value + f32(bias[2u * Q_HIDDEN + h * HEAD_DIM + d]);
180
+ {% endif %}
181
  running_out[d] = running_out[d] * previous_scale + probability_numerator * v_value;
182
  }
183
  workgroupBarrier();
build/webgpu/bench.json CHANGED
@@ -710,7 +710,9 @@
710
  {
711
  "name": "quant-int8-prefill-h32kv8-d128-s512-pathology",
712
  "preset": "stress",
713
- "provenance": { "notes": "A production-sized prefill exercises the int8 KV-cache attention contract." },
 
 
714
  "attrs": {
715
  "num_heads": 32,
716
  "kv_num_heads": 8,
@@ -765,7 +767,9 @@
765
  {
766
  "name": "quant-int4-prefill-h32kv8-d128-s512-pathology",
767
  "preset": "stress",
768
- "provenance": { "notes": "A production-sized prefill exercises the packed int4 KV-cache attention contract." },
 
 
769
  "attrs": {
770
  "num_heads": 32,
771
  "kv_num_heads": 8,
@@ -1043,7 +1047,7 @@
1043
  "name": "headsink-prefill-h32kv8-d128-s512-generic-pathology",
1044
  "preset": "stress",
1045
  "provenance": {
1046
- "notes": "A Llama-sized GQA prefill exercises smooth-softmax head sinks on a production attention shape."
1047
  },
1048
  "attrs": { "num_heads": 32, "kv_num_heads": 8 },
1049
  "inputs": {
@@ -1086,7 +1090,7 @@
1086
  "name": "bias-headsink-prefill-h32kv8-d128-s512-generic-pathology",
1087
  "preset": "stress",
1088
  "provenance": {
1089
- "notes": "A production prefill combining additive bias with smooth-softmax head sinks exercises the generic attention route."
1090
  },
1091
  "attrs": { "num_heads": 32, "kv_num_heads": 8 },
1092
  "inputs": {
@@ -1133,7 +1137,9 @@
1133
  {
1134
  "name": "softcap-prefill-h32kv8-d128-s512-generic-pathology",
1135
  "preset": "stress",
1136
- "provenance": { "notes": "A production prefill exercises model-used score soft-capping on the past-KV contract." },
 
 
1137
  "attrs": { "num_heads": 32, "kv_num_heads": 8, "softcap": 30 },
1138
  "inputs": {
1139
  "queryT": {
@@ -1460,6 +1466,74 @@
1460
  ]
1461
  }
1462
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1463
  {
1464
  "name": "sharedkv-chunk-append-32h8kv-d128-cap1024-q128",
1465
  "preset": "smoke",
@@ -1551,7 +1625,97 @@
1551
  },
1552
  "provenance": {
1553
  "source": "register-geometry gate asymmetry",
1554
- "notes": "Cached f32 prefill at headDim 256, above the shared-memory cluster's register-geometry boundary. This guards consistent route admission between no-past and cached prefill families."
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1555
  }
1556
  },
1557
  {
@@ -1576,6 +1740,439 @@
1576
  { "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
1577
  ]
1578
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1579
  }
1580
  ]
1581
  }
 
710
  {
711
  "name": "quant-int8-prefill-h32kv8-d128-s512-pathology",
712
  "preset": "stress",
713
+ "provenance": {
714
+ "notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with an int8 quantized key/value cache."
715
+ },
716
  "attrs": {
717
  "num_heads": 32,
718
  "kv_num_heads": 8,
 
767
  {
768
  "name": "quant-int4-prefill-h32kv8-d128-s512-pathology",
769
  "preset": "stress",
770
+ "provenance": {
771
+ "notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with a packed int4 quantized key/value cache."
772
+ },
773
  "attrs": {
774
  "num_heads": 32,
775
  "kv_num_heads": 8,
 
1047
  "name": "headsink-prefill-h32kv8-d128-s512-generic-pathology",
1048
  "preset": "stress",
1049
  "provenance": {
1050
+ "notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with per-head attention sink logits."
1051
  },
1052
  "attrs": { "num_heads": 32, "kv_num_heads": 8 },
1053
  "inputs": {
 
1090
  "name": "bias-headsink-prefill-h32kv8-d128-s512-generic-pathology",
1091
  "preset": "stress",
1092
  "provenance": {
1093
+ "notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with an additive attention bias and per-head attention sink logits."
1094
  },
1095
  "attrs": { "num_heads": 32, "kv_num_heads": 8 },
1096
  "inputs": {
 
1137
  {
1138
  "name": "softcap-prefill-h32kv8-d128-s512-generic-pathology",
1139
  "preset": "stress",
1140
+ "provenance": {
1141
+ "notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with a logit softcap of 30."
1142
+ },
1143
  "attrs": { "num_heads": 32, "kv_num_heads": 8, "softcap": 30 },
1144
  "inputs": {
1145
  "queryT": {
 
1466
  ]
1467
  }
1468
  },
1469
+ {
1470
+ "name": "window-headsink-rotary-decode-gptoss-64h8kv-d64-cap384-w128",
1471
+ "preset": "smoke",
1472
+ "attrs": {
1473
+ "num_heads": 64,
1474
+ "kv_num_heads": 8,
1475
+ "do_rotary": 1,
1476
+ "sliding_window_cache": 1,
1477
+ "local_window_size": 128
1478
+ },
1479
+ "inputs": {
1480
+ "queryT": {
1481
+ "dtype": "float32",
1482
+ "shape": [1, 1, 4096],
1483
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
1484
+ },
1485
+ "keyT": {
1486
+ "dtype": "float32",
1487
+ "shape": [1, 1, 512],
1488
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
1489
+ },
1490
+ "valueT": {
1491
+ "dtype": "float32",
1492
+ "shape": [1, 1, 512],
1493
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
1494
+ },
1495
+ "pastKeyT": {
1496
+ "dtype": "float32",
1497
+ "shape": [1, 8, 384, 64],
1498
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
1499
+ },
1500
+ "pastValueT": {
1501
+ "dtype": "float32",
1502
+ "shape": [1, 8, 384, 64],
1503
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
1504
+ },
1505
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [4095] } },
1506
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [4096] } },
1507
+ "cosCacheT": {
1508
+ "dtype": "float32",
1509
+ "shape": [4096, 32],
1510
+ "data": { "kind": "rotaryCos", "thetaStart": 0.0, "thetaStep": 0.017 }
1511
+ },
1512
+ "sinCacheT": {
1513
+ "dtype": "float32",
1514
+ "shape": [4096, 32],
1515
+ "data": { "kind": "rotarySin", "thetaStart": 0.0, "thetaStep": 0.017 }
1516
+ },
1517
+ "headSinkT": {
1518
+ "dtype": "float32",
1519
+ "shape": [64],
1520
+ "data": { "kind": "fillFloat32", "scale": 0.5, "sinStep": 0.041, "cosStep": 0.011 }
1521
+ }
1522
+ },
1523
+ "outputs": {
1524
+ "outputT": { "dtype": "float32", "shape": [1, 1, 4096] },
1525
+ "presentKeyT": { "dtype": "float32", "shape": [1, 8, 384, 64] },
1526
+ "presentValueT": { "dtype": "float32", "shape": [1, 8, 384, 64] }
1527
+ },
1528
+ "bench": {
1529
+ "metrics": [
1530
+ {
1531
+ "type": "bandwidth",
1532
+ "value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
1533
+ }
1534
+ ]
1535
+ }
1536
+ },
1537
  {
1538
  "name": "sharedkv-chunk-append-32h8kv-d128-cap1024-q128",
1539
  "preset": "smoke",
 
1625
  },
1626
  "provenance": {
1627
  "source": "register-geometry gate asymmetry",
1628
+ "notes": "Measures cached float32 prefill that appends 128 tokens after 896 cached ones, filling a 1024-slot key/value buffer, with 8 query heads over 2 key/value heads of width 256."
1629
+ }
1630
+ },
1631
+ {
1632
+ "name": "bidirectional-share-append-32h8kv-d128-cap1024-q128",
1633
+ "preset": "smoke",
1634
+ "vars": { "batch": 1, "qSeq": 128, "kvSeq": 1024, "heads": 32, "headDim": 128 },
1635
+ "attrs": { "num_heads": 32, "kv_num_heads": 8, "causal": 0 },
1636
+ "inputs": {
1637
+ "queryT": {
1638
+ "dtype": "float32",
1639
+ "shape": [1, 128, 4096],
1640
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
1641
+ },
1642
+ "keyT": {
1643
+ "dtype": "float32",
1644
+ "shape": [1, 128, 1024],
1645
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
1646
+ },
1647
+ "valueT": {
1648
+ "dtype": "float32",
1649
+ "shape": [1, 128, 1024],
1650
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
1651
+ },
1652
+ "pastKeyT": {
1653
+ "dtype": "float32",
1654
+ "shape": [1, 8, 1024, 128],
1655
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
1656
+ },
1657
+ "pastValueT": {
1658
+ "dtype": "float32",
1659
+ "shape": [1, 8, 1024, 128],
1660
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
1661
+ },
1662
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [1023] } },
1663
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [1024] } }
1664
+ },
1665
+ "outputs": {
1666
+ "outputT": { "dtype": "float32", "shape": [1, 128, 4096] },
1667
+ "presentKeyT": { "dtype": "float32", "shape": [1, 8, 1024, 128] },
1668
+ "presentValueT": { "dtype": "float32", "shape": [1, 8, 1024, 128] }
1669
+ },
1670
+ "bench": {
1671
+ "metrics": [
1672
+ { "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
1673
+ ]
1674
+ }
1675
+ },
1676
+ {
1677
+ "name": "bidirectional-share-append-f16-32h8kv-d128-cap512-q64",
1678
+ "preset": "smoke",
1679
+ "vars": { "batch": 1, "qSeq": 64, "kvSeq": 512, "heads": 32, "headDim": 128 },
1680
+ "attrs": { "num_heads": 32, "kv_num_heads": 8, "causal": 0 },
1681
+ "inputs": {
1682
+ "queryT": {
1683
+ "dtype": "float16",
1684
+ "shape": [1, 64, 4096],
1685
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
1686
+ },
1687
+ "keyT": {
1688
+ "dtype": "float16",
1689
+ "shape": [1, 64, 1024],
1690
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
1691
+ },
1692
+ "valueT": {
1693
+ "dtype": "float16",
1694
+ "shape": [1, 64, 1024],
1695
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
1696
+ },
1697
+ "pastKeyT": {
1698
+ "dtype": "float16",
1699
+ "shape": [1, 8, 512, 128],
1700
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
1701
+ },
1702
+ "pastValueT": {
1703
+ "dtype": "float16",
1704
+ "shape": [1, 8, 512, 128],
1705
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
1706
+ },
1707
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [511] } },
1708
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [512] } }
1709
+ },
1710
+ "outputs": {
1711
+ "outputT": { "dtype": "float16", "shape": [1, 64, 4096] },
1712
+ "presentKeyT": { "dtype": "float16", "shape": [1, 8, 512, 128] },
1713
+ "presentValueT": { "dtype": "float16", "shape": [1, 8, 512, 128] }
1714
+ },
1715
+ "bench": {
1716
+ "metrics": [
1717
+ { "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
1718
+ ]
1719
  }
1720
  },
1721
  {
 
1740
  { "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
1741
  ]
1742
  }
1743
+ },
1744
+ {
1745
+ "name": "causal-prefill-f16-h4kv2-d64-q64",
1746
+ "provenance": {
1747
+ "source": "synthetic",
1748
+ "test": "Causal half-precision prompt prefill",
1749
+ "notes": "Each query attends only itself and earlier key positions. Verify the complete attention output and exact retention of both key/value inputs in the present caches."
1750
+ },
1751
+ "attrs": { "num_heads": 4, "kv_num_heads": 2 },
1752
+ "inputs": {
1753
+ "queryT": {
1754
+ "dtype": "float16",
1755
+ "shape": [1, 64, 256],
1756
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.0625 }
1757
+ },
1758
+ "keyT": {
1759
+ "dtype": "float16",
1760
+ "shape": [1, 64, 128],
1761
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.0625 }
1762
+ },
1763
+ "valueT": {
1764
+ "dtype": "float16",
1765
+ "shape": [1, 64, 128],
1766
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.031, "scale": 0.0625 }
1767
+ },
1768
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [63] } },
1769
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [64] } }
1770
+ },
1771
+ "outputs": {
1772
+ "outputT": { "dtype": "float16", "shape": [1, 64, 256] },
1773
+ "presentKeyT": { "dtype": "float16", "shape": [1, 2, 64, 64] },
1774
+ "presentValueT": { "dtype": "float16", "shape": [1, 2, 64, 64] }
1775
+ },
1776
+ "preset": "model",
1777
+ "bench": {
1778
+ "metrics": [
1779
+ {
1780
+ "type": "gflops",
1781
+ "value": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * (dim(shapes.queryT, 1) + 1) * dim(shapes.queryT, 2)"
1782
+ }
1783
+ ]
1784
+ }
1785
+ },
1786
+ {
1787
+ "name": "causal-prefill-f16-h4kv2-d128-q65",
1788
+ "provenance": {
1789
+ "source": "synthetic",
1790
+ "test": "Causal half-precision prompt prefill",
1791
+ "notes": "Each query attends only itself and earlier key positions. Verify the complete attention output and exact retention of both key/value inputs in the present caches."
1792
+ },
1793
+ "attrs": { "num_heads": 4, "kv_num_heads": 2 },
1794
+ "inputs": {
1795
+ "queryT": {
1796
+ "dtype": "float16",
1797
+ "shape": [1, 65, 512],
1798
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.0625 }
1799
+ },
1800
+ "keyT": {
1801
+ "dtype": "float16",
1802
+ "shape": [1, 65, 256],
1803
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.0625 }
1804
+ },
1805
+ "valueT": {
1806
+ "dtype": "float16",
1807
+ "shape": [1, 65, 256],
1808
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.031, "scale": 0.0625 }
1809
+ },
1810
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [64] } },
1811
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65] } }
1812
+ },
1813
+ "outputs": {
1814
+ "outputT": { "dtype": "float16", "shape": [1, 65, 512] },
1815
+ "presentKeyT": { "dtype": "float16", "shape": [1, 2, 65, 128] },
1816
+ "presentValueT": { "dtype": "float16", "shape": [1, 2, 65, 128] }
1817
+ },
1818
+ "preset": "model",
1819
+ "bench": {
1820
+ "metrics": [
1821
+ {
1822
+ "type": "gflops",
1823
+ "value": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * (dim(shapes.queryT, 1) + 1) * dim(shapes.queryT, 2)"
1824
+ }
1825
+ ]
1826
+ }
1827
+ },
1828
+ {
1829
+ "name": "causal-prefill-f16-h8kv2-d64-q512",
1830
+ "provenance": {
1831
+ "source": "synthetic",
1832
+ "test": "Causal half-precision prompt prefill",
1833
+ "notes": "Each query attends only itself and earlier key positions. Verify the complete attention output and exact retention of both key/value inputs in the present caches."
1834
+ },
1835
+ "attrs": { "num_heads": 8, "kv_num_heads": 2 },
1836
+ "inputs": {
1837
+ "queryT": {
1838
+ "dtype": "float16",
1839
+ "shape": [1, 512, 512],
1840
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.0625 }
1841
+ },
1842
+ "keyT": {
1843
+ "dtype": "float16",
1844
+ "shape": [1, 512, 128],
1845
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.0625 }
1846
+ },
1847
+ "valueT": {
1848
+ "dtype": "float16",
1849
+ "shape": [1, 512, 128],
1850
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.031, "scale": 0.0625 }
1851
+ },
1852
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [511] } },
1853
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [512] } }
1854
+ },
1855
+ "outputs": {
1856
+ "outputT": { "dtype": "float16", "shape": [1, 512, 512] },
1857
+ "presentKeyT": { "dtype": "float16", "shape": [1, 2, 512, 64] },
1858
+ "presentValueT": { "dtype": "float16", "shape": [1, 2, 512, 64] }
1859
+ },
1860
+ "preset": "model",
1861
+ "bench": {
1862
+ "metrics": [
1863
+ {
1864
+ "type": "gflops",
1865
+ "value": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * (dim(shapes.queryT, 1) + 1) * dim(shapes.queryT, 2)"
1866
+ }
1867
+ ]
1868
+ }
1869
+ },
1870
+ {
1871
+ "name": "prefill-f32-h8kv1-d32-q33",
1872
+ "preset": "model",
1873
+ "provenance": { "notes": "Small bidirectional prefill with query tails and varying head widths." },
1874
+ "vars": { "batch": 1, "qSeq": 33, "kvSeq": 33, "heads": 8, "kvHeads": 1, "headDim": 32 },
1875
+ "attrs": { "num_heads": 8, "kv_num_heads": 1, "scale": 0.125, "causal": 0 },
1876
+ "inputs": {
1877
+ "queryT": { "shape": [1, 33, 256], "dtype": "float32", "dist": "normal", "seed": 1075, "scale": 0.2 },
1878
+ "keyT": { "shape": [1, 33, 32], "dtype": "float32", "dist": "normal", "seed": 1076, "scale": 0.2 },
1879
+ "valueT": { "shape": [1, 33, 32], "dtype": "float32", "dist": "normal", "seed": 1077, "scale": 0.2 },
1880
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [32] } },
1881
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [33] } }
1882
+ },
1883
+ "outputs": {
1884
+ "outputT": { "shape": [1, 33, 256], "dtype": "float32" },
1885
+ "presentKeyT": { "shape": [1, 1, 33, 32], "dtype": "float32" },
1886
+ "presentValueT": { "shape": [1, 1, 33, 32], "dtype": "float32" }
1887
+ },
1888
+ "bench": {
1889
+ "metrics": [
1890
+ { "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
1891
+ ]
1892
+ }
1893
+ },
1894
+ {
1895
+ "name": "prefill-f32-h8kv1-d128-q33",
1896
+ "preset": "model",
1897
+ "provenance": { "notes": "Small bidirectional prefill with query tails and varying head widths." },
1898
+ "vars": { "batch": 1, "qSeq": 33, "kvSeq": 33, "heads": 8, "kvHeads": 1, "headDim": 128 },
1899
+ "attrs": { "num_heads": 8, "kv_num_heads": 1, "scale": 0.125, "causal": 0 },
1900
+ "inputs": {
1901
+ "queryT": { "shape": [1, 33, 1024], "dtype": "float32", "dist": "normal", "seed": 1075, "scale": 0.2 },
1902
+ "keyT": { "shape": [1, 33, 128], "dtype": "float32", "dist": "normal", "seed": 1076, "scale": 0.2 },
1903
+ "valueT": { "shape": [1, 33, 128], "dtype": "float32", "dist": "normal", "seed": 1077, "scale": 0.2 },
1904
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [32] } },
1905
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [33] } }
1906
+ },
1907
+ "outputs": {
1908
+ "outputT": { "shape": [1, 33, 1024], "dtype": "float32" },
1909
+ "presentKeyT": { "shape": [1, 1, 33, 128], "dtype": "float32" },
1910
+ "presentValueT": { "shape": [1, 1, 33, 128], "dtype": "float32" }
1911
+ },
1912
+ "bench": {
1913
+ "metrics": [
1914
+ { "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
1915
+ ]
1916
+ }
1917
+ },
1918
+ {
1919
+ "name": "prefill-f32-h8kv1-d256-q32",
1920
+ "preset": "model",
1921
+ "provenance": { "notes": "Small bidirectional prefill with query tails and varying head widths." },
1922
+ "vars": { "batch": 1, "qSeq": 32, "kvSeq": 32, "heads": 8, "kvHeads": 1, "headDim": 256 },
1923
+ "attrs": { "num_heads": 8, "kv_num_heads": 1, "scale": 0.125, "causal": 0 },
1924
+ "inputs": {
1925
+ "queryT": { "shape": [1, 32, 2048], "dtype": "float32", "dist": "normal", "seed": 1075, "scale": 0.2 },
1926
+ "keyT": { "shape": [1, 32, 256], "dtype": "float32", "dist": "normal", "seed": 1076, "scale": 0.2 },
1927
+ "valueT": { "shape": [1, 32, 256], "dtype": "float32", "dist": "normal", "seed": 1077, "scale": 0.2 },
1928
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [31] } },
1929
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [32] } }
1930
+ },
1931
+ "outputs": {
1932
+ "outputT": { "shape": [1, 32, 2048], "dtype": "float32" },
1933
+ "presentKeyT": { "shape": [1, 1, 32, 256], "dtype": "float32" },
1934
+ "presentValueT": { "shape": [1, 1, 32, 256], "dtype": "float32" }
1935
+ },
1936
+ "bench": {
1937
+ "metrics": [
1938
+ { "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
1939
+ ]
1940
+ }
1941
+ },
1942
+ {
1943
+ "name": "cooperative-window-h4kv2-d122-cap256",
1944
+ "preset": "smoke",
1945
+ "attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
1946
+ "inputs": {
1947
+ "queryT": {
1948
+ "dtype": "float32",
1949
+ "shape": [1, 1, 488],
1950
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
1951
+ },
1952
+ "keyT": {
1953
+ "dtype": "float32",
1954
+ "shape": [1, 1, 244],
1955
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
1956
+ },
1957
+ "valueT": {
1958
+ "dtype": "float32",
1959
+ "shape": [1, 1, 244],
1960
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
1961
+ },
1962
+ "pastKeyT": {
1963
+ "dtype": "float32",
1964
+ "shape": [1, 2, 256, 122],
1965
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
1966
+ },
1967
+ "pastValueT": {
1968
+ "dtype": "float32",
1969
+ "shape": [1, 2, 256, 122],
1970
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
1971
+ },
1972
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
1973
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
1974
+ },
1975
+ "outputs": {
1976
+ "outputT": { "dtype": "float32", "shape": [1, 1, 488] },
1977
+ "presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 122] },
1978
+ "presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 122] }
1979
+ },
1980
+ "bench": {
1981
+ "metrics": [
1982
+ {
1983
+ "type": "bandwidth",
1984
+ "value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
1985
+ }
1986
+ ]
1987
+ }
1988
+ },
1989
+ {
1990
+ "name": "cooperative-window-h4kv2-d124-cap256",
1991
+ "preset": "smoke",
1992
+ "attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
1993
+ "inputs": {
1994
+ "queryT": {
1995
+ "dtype": "float32",
1996
+ "shape": [1, 1, 496],
1997
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
1998
+ },
1999
+ "keyT": {
2000
+ "dtype": "float32",
2001
+ "shape": [1, 1, 248],
2002
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
2003
+ },
2004
+ "valueT": {
2005
+ "dtype": "float32",
2006
+ "shape": [1, 1, 248],
2007
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
2008
+ },
2009
+ "pastKeyT": {
2010
+ "dtype": "float32",
2011
+ "shape": [1, 2, 256, 124],
2012
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
2013
+ },
2014
+ "pastValueT": {
2015
+ "dtype": "float32",
2016
+ "shape": [1, 2, 256, 124],
2017
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
2018
+ },
2019
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
2020
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
2021
+ },
2022
+ "outputs": {
2023
+ "outputT": { "dtype": "float32", "shape": [1, 1, 496] },
2024
+ "presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 124] },
2025
+ "presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 124] }
2026
+ },
2027
+ "bench": {
2028
+ "metrics": [
2029
+ {
2030
+ "type": "bandwidth",
2031
+ "value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
2032
+ }
2033
+ ]
2034
+ }
2035
+ },
2036
+ {
2037
+ "name": "cooperative-window-h4kv2-d246-cap256",
2038
+ "preset": "smoke",
2039
+ "attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
2040
+ "inputs": {
2041
+ "queryT": {
2042
+ "dtype": "float32",
2043
+ "shape": [1, 1, 984],
2044
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
2045
+ },
2046
+ "keyT": {
2047
+ "dtype": "float32",
2048
+ "shape": [1, 1, 492],
2049
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
2050
+ },
2051
+ "valueT": {
2052
+ "dtype": "float32",
2053
+ "shape": [1, 1, 492],
2054
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
2055
+ },
2056
+ "pastKeyT": {
2057
+ "dtype": "float32",
2058
+ "shape": [1, 2, 256, 246],
2059
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
2060
+ },
2061
+ "pastValueT": {
2062
+ "dtype": "float32",
2063
+ "shape": [1, 2, 256, 246],
2064
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
2065
+ },
2066
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
2067
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
2068
+ },
2069
+ "outputs": {
2070
+ "outputT": { "dtype": "float32", "shape": [1, 1, 984] },
2071
+ "presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 246] },
2072
+ "presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 246] }
2073
+ },
2074
+ "bench": {
2075
+ "metrics": [
2076
+ {
2077
+ "type": "bandwidth",
2078
+ "value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
2079
+ }
2080
+ ]
2081
+ }
2082
+ },
2083
+ {
2084
+ "name": "cooperative-window-h4kv2-d248-cap256",
2085
+ "preset": "smoke",
2086
+ "attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
2087
+ "inputs": {
2088
+ "queryT": {
2089
+ "dtype": "float32",
2090
+ "shape": [1, 1, 992],
2091
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
2092
+ },
2093
+ "keyT": {
2094
+ "dtype": "float32",
2095
+ "shape": [1, 1, 496],
2096
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
2097
+ },
2098
+ "valueT": {
2099
+ "dtype": "float32",
2100
+ "shape": [1, 1, 496],
2101
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
2102
+ },
2103
+ "pastKeyT": {
2104
+ "dtype": "float32",
2105
+ "shape": [1, 2, 256, 248],
2106
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
2107
+ },
2108
+ "pastValueT": {
2109
+ "dtype": "float32",
2110
+ "shape": [1, 2, 256, 248],
2111
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
2112
+ },
2113
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
2114
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
2115
+ },
2116
+ "outputs": {
2117
+ "outputT": { "dtype": "float32", "shape": [1, 1, 992] },
2118
+ "presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 248] },
2119
+ "presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 248] }
2120
+ },
2121
+ "bench": {
2122
+ "metrics": [
2123
+ {
2124
+ "type": "bandwidth",
2125
+ "value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
2126
+ }
2127
+ ]
2128
+ }
2129
+ },
2130
+ {
2131
+ "name": "cooperative-window-h4kv2-d256-cap256",
2132
+ "preset": "smoke",
2133
+ "attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
2134
+ "inputs": {
2135
+ "queryT": {
2136
+ "dtype": "float32",
2137
+ "shape": [1, 1, 1024],
2138
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
2139
+ },
2140
+ "keyT": {
2141
+ "dtype": "float32",
2142
+ "shape": [1, 1, 512],
2143
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
2144
+ },
2145
+ "valueT": {
2146
+ "dtype": "float32",
2147
+ "shape": [1, 1, 512],
2148
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
2149
+ },
2150
+ "pastKeyT": {
2151
+ "dtype": "float32",
2152
+ "shape": [1, 2, 256, 256],
2153
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
2154
+ },
2155
+ "pastValueT": {
2156
+ "dtype": "float32",
2157
+ "shape": [1, 2, 256, 256],
2158
+ "data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
2159
+ },
2160
+ "seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
2161
+ "totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
2162
+ },
2163
+ "outputs": {
2164
+ "outputT": { "dtype": "float32", "shape": [1, 1, 1024] },
2165
+ "presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 256] },
2166
+ "presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 256] }
2167
+ },
2168
+ "bench": {
2169
+ "metrics": [
2170
+ {
2171
+ "type": "bandwidth",
2172
+ "value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
2173
+ }
2174
+ ]
2175
+ }
2176
  }
2177
  ]
2178
  }
build/webgpu/gqa-attention.wgsl.jinja CHANGED
@@ -2,8 +2,12 @@
2
  {% set quantized = quantized | default(false) %}
3
  {% set bits = bits | default(0) %}
4
  {% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
5
- {% if usesF16 %}enable f16;
6
- {% endif %}{{ env.wgsl.resourceDeclarations }}
 
 
 
 
7
  // f16 inputs widen to f32 on load; computation and shared memory remain f32,
8
  // and results narrow only on store. The casts are identities for f32.
9
  {% set IO = "f16" if usesF16 else "f32" %}
@@ -29,15 +33,13 @@ const Q_HIDDEN: u32 = {{ qHidden }}u;
29
  const GROUP: u32 = {{ qHeads }}u / {{ kvHeads }}u;
30
  const PACKED: u32 = {{ packed }}u;
31
  {% if hasQNorm %}const QK_EPS: f32 = {{ qkEps }};{% endif %}
32
- {% if cooperative %}const WG: u32 = {{ tunables.COOPERATIVE_WORKGROUP_SIZE }}u;{% else %}const WG: u32 = {{ scalarWorkgroupSize }}u;{% endif %}
33
  const NEG_INF: f32 = -3.4028234663852886e38;
34
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
35
- fn scale_value() -> f32 {
36
  if (params.scale != 0.0) { return params.scale; }
37
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
38
  }
39
 
40
-
41
  {% if quantized %}
42
  fn kscale(d: u32, hk: u32) -> f32 { return k_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
43
  fn vscale(d: u32, hk: u32) -> f32 { return v_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
@@ -67,10 +69,15 @@ fn read_v(base: u32, d: u32{% if quantized %}, hk: u32{% endif %}) -> f32 {
67
  }
68
 
69
  {% if cooperative %}
70
- var<workgroup> qsh: array<f32, HEAD_DIM>; // query head-vector (post norm/rotary)
71
- var<workgroup> acc_sh: array<f32, HEAD_DIM * 32u>; // per-lane V accumulator [lane*HEAD_DIM + d]
72
- var<workgroup> m_sh: array<f32, 32u>;
73
- var<workgroup> l_sh: array<f32, 32u>;
 
 
 
 
 
74
 
75
  @compute @workgroup_size(WG)
76
  fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
@@ -92,42 +99,47 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
92
  {% else %}
93
  let activeEnd = totalSeq;
94
  {% endif %}
95
- let absPos = activeEnd - params.qSeq + s;
 
 
 
 
 
96
 
97
  // Cooperative query-vector prep (post norm/rotary) into shared memory.
98
  let qBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
99
- for (var d = lane; d < HEAD_DIM; d = d + WG) { qsh[d] = f32(query[qBase + d]); }
100
  workgroupBarrier();
101
  {% if hasQNorm %}
102
  var part = 0.0;
103
- for (var d = lane; d < HEAD_DIM; d = d + WG) { part = part + qsh[d] * qsh[d]; }
104
- m_sh[lane] = part;
105
  workgroupBarrier();
106
  var ms = 0.0;
107
- for (var L = 0u; L < WG; L = L + 1u) { ms = ms + m_sh[L]; }
108
  let invRms = inverseSqrt(ms / f32(HEAD_DIM) + QK_EPS);
109
- for (var d = lane; d < HEAD_DIM; d = d + WG) { qsh[d] = qsh[d] * invRms * f32(q_norm_weight[d]); }
110
  workgroupBarrier();
111
  {% endif %}
112
  {% if hasRotary %}
113
  // Each lane owns disjoint pairs (d, d+HALF), so the in-place rotate is safe.
114
  for (var d = lane; d < HALF; d = d + WG) {
115
- let cs = f32(cos_cache[absPos * HALF + d]);
116
- let sn = f32(sin_cache[absPos * HALF + d]);
117
- let x0 = qsh[d];
118
- let x1 = qsh[d + HALF];
119
- qsh[d] = x0 * cs - x1 * sn;
120
- qsh[d + HALF] = x1 * cs + x0 * sn;
121
  }
122
  workgroupBarrier();
123
  {% endif %}
124
  let scale = scale_value();
125
- let maxKj = absPos + 1u;
126
  var minKj = 0u;
127
  if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
128
 
129
  let accBase = lane * HEAD_DIM;
130
- for (var d = 0u; d < HEAD_DIM; d = d + 1u) { acc_sh[accBase + d] = 0.0; }
131
  var m = NEG_INF;
132
  var l = 0.0;
133
 
@@ -135,34 +147,34 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
135
  for (var j = minKj + lane; j < maxKj; j = j + WG) {
136
  let base = ((b * KV_HEADS + hk) * totalSeq + j) * PACKED;
137
  var dot = 0.0;
138
- for (var d = 0u; d < HEAD_DIM; d = d + 1u) { dot = dot + qsh[d] * read_k(base, d{% if quantized %}, hk{% endif %}); }
139
  var score = dot * scale;
140
  if (params.softcap > 0.0) { score = params.softcap * tanh(clamp(score / params.softcap, -30.0, 30.0)); }
141
  {% if hasBias %}
142
  let bb = select(b, 0u, params.biasBatch == 1u);
143
  let bh = select(h, 0u, params.biasHeads == 1u);
144
- score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) * totalSeq + j]);
145
  {% endif %}
146
  let newM = max(m, score);
147
  let corr = exp(m - newM);
148
  let p = exp(score - newM);
149
  l = l * corr + p;
150
- for (var d = 0u; d < HEAD_DIM; d = d + 1u) { acc_sh[accBase + d] = acc_sh[accBase + d] * corr + p * read_v(base, d{% if quantized %}, hk{% endif %}); }
151
  m = newM;
152
  }
153
 
154
  // Flash merge across the WG lanes: global max, rescale, summed denom.
155
- m_sh[lane] = m;
156
  workgroupBarrier();
157
  var gm = NEG_INF;
158
- for (var L = 0u; L < WG; L = L + 1u) { gm = max(gm, m_sh[L]); }
159
  let fctr = exp(m - gm);
160
  l = l * fctr;
161
- for (var d = 0u; d < HEAD_DIM; d = d + 1u) { acc_sh[accBase + d] = acc_sh[accBase + d] * fctr; }
162
- l_sh[lane] = l;
163
  workgroupBarrier();
164
  var glsum = 0.0;
165
- for (var L = 0u; L < WG; L = L + 1u) { glsum = glsum + l_sh[L]; }
166
 
167
  var sink = 0.0;
168
  {% if hasHeadSink %}
@@ -187,7 +199,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
187
  let oBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
188
  for (var d = lane; d < HEAD_DIM; d = d + WG) {
189
  var sumacc = 0.0;
190
- for (var L = 0u; L < WG; L = L + 1u) { sumacc = sumacc + acc_sh[L * HEAD_DIM + d]; }
191
  output[oBase + d] = {{ IO }}(sumacc * invDenom);
192
  }
193
  }
@@ -215,7 +227,12 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
215
  {% else %}
216
  let activeEnd = totalSeq;
217
  {% endif %}
218
- let absPos = activeEnd - params.qSeq + s;
 
 
 
 
 
219
 
220
  // Private query-vector prep (post norm/rotary).
221
  var q: array<f32, HEAD_DIM>;
@@ -229,8 +246,8 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
229
  {% endif %}
230
  {% if hasRotary %}
231
  for (var d = 0u; d < HALF; d = d + 1u) {
232
- let cs = f32(cos_cache[absPos * HALF + d]);
233
- let sn = f32(sin_cache[absPos * HALF + d]);
234
  let x0 = q[d];
235
  let x1 = q[d + HALF];
236
  q[d] = x0 * cs - x1 * sn;
@@ -238,7 +255,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
238
  }
239
  {% endif %}
240
  let scale = scale_value();
241
- let maxKj = absPos + 1u;
242
  var minKj = 0u;
243
  if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
244
 
@@ -256,7 +273,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
256
  {% if hasBias %}
257
  let bb = select(b, 0u, params.biasBatch == 1u);
258
  let bh = select(h, 0u, params.biasHeads == 1u);
259
- score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) * totalSeq + j]);
260
  {% endif %}
261
  let newM = max(m, score);
262
  let corr = exp(m - newM);
 
2
  {% set quantized = quantized | default(false) %}
3
  {% set bits = bits | default(0) %}
4
  {% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
5
+ {% if windowOrigin is not defined %}{% set windowOrigin = false %}{% endif %}
6
+ {% set needsOrigin = windowOrigin and (hasRotary or hasBias) %}
7
+ {% set ropePos = "(absPos + origin)" if needsOrigin else "absPos" %}
8
+ {% set biasStride = "params.biasCols" if needsOrigin else "totalSeq" %}
9
+ {% set biasCol = "(origin + j)" if needsOrigin else "j" %}
10
+ {{ env.wgsl.resourceDeclarations }}
11
  // f16 inputs widen to f32 on load; computation and shared memory remain f32,
12
  // and results narrow only on store. The casts are identities for f32.
13
  {% set IO = "f16" if usesF16 else "f32" %}
 
33
  const GROUP: u32 = {{ qHeads }}u / {{ kvHeads }}u;
34
  const PACKED: u32 = {{ packed }}u;
35
  {% if hasQNorm %}const QK_EPS: f32 = {{ qkEps }};{% endif %}
36
+ {% if cooperative %}const WG: u32 = {{ gqaCooperativeWorkgroupSize }}u;{% else %}const WG: u32 = {{ scalarWorkgroupSize }}u;{% endif %}
37
  const NEG_INF: f32 = -3.4028234663852886e38;
38
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
39
  if (params.scale != 0.0) { return params.scale; }
40
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
41
  }
42
 
 
43
  {% if quantized %}
44
  fn kscale(d: u32, hk: u32) -> f32 { return k_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
45
  fn vscale(d: u32, hk: u32) -> f32 { return v_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
 
69
  }
70
 
71
  {% if cooperative %}
72
+ // One allocation avoids the per-variable 16-byte padding between arrays.
73
+ // The manifest budgets the structure's complete size, rounded to 16 bytes.
74
+ struct CooperativeState {
75
+ q: array<f32, HEAD_DIM>,
76
+ acc: array<f32, HEAD_DIM * WG>,
77
+ m: array<f32, WG>,
78
+ l: array<f32, WG>,
79
+ }
80
+ var<workgroup> state: CooperativeState;
81
 
82
  @compute @workgroup_size(WG)
83
  fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
 
99
  {% else %}
100
  let activeEnd = totalSeq;
101
  {% endif %}
102
+ {% if not bidirectional or hasRotary %} let absPos = activeEnd - params.qSeq + s;
103
+ {% endif %}{% if needsOrigin %}
104
+ // Absolute position of cache row 0 (windowed cache): the query end is the
105
+ // unclamped total while the cache holds only its last `activeEnd` rows.
106
+ let origin = max(params.qSeq, u32(seqlens_k[b]) + 1u) - activeEnd;
107
+ {% endif %}
108
 
109
  // Cooperative query-vector prep (post norm/rotary) into shared memory.
110
  let qBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
111
+ for (var d = lane; d < HEAD_DIM; d = d + WG) { state.q[d] = f32(query[qBase + d]); }
112
  workgroupBarrier();
113
  {% if hasQNorm %}
114
  var part = 0.0;
115
+ for (var d = lane; d < HEAD_DIM; d = d + WG) { part = part + state.q[d] * state.q[d]; }
116
+ state.m[lane] = part;
117
  workgroupBarrier();
118
  var ms = 0.0;
119
+ for (var L = 0u; L < WG; L = L + 1u) { ms = ms + state.m[L]; }
120
  let invRms = inverseSqrt(ms / f32(HEAD_DIM) + QK_EPS);
121
+ for (var d = lane; d < HEAD_DIM; d = d + WG) { state.q[d] = state.q[d] * invRms * f32(q_norm_weight[d]); }
122
  workgroupBarrier();
123
  {% endif %}
124
  {% if hasRotary %}
125
  // Each lane owns disjoint pairs (d, d+HALF), so the in-place rotate is safe.
126
  for (var d = lane; d < HALF; d = d + WG) {
127
+ let cs = f32(cos_cache[{{ ropePos }} * HALF + d]);
128
+ let sn = f32(sin_cache[{{ ropePos }} * HALF + d]);
129
+ let x0 = state.q[d];
130
+ let x1 = state.q[d + HALF];
131
+ state.q[d] = x0 * cs - x1 * sn;
132
+ state.q[d + HALF] = x1 * cs + x0 * sn;
133
  }
134
  workgroupBarrier();
135
  {% endif %}
136
  let scale = scale_value();
137
+ let maxKj = {% if bidirectional %}activeEnd{% else %}absPos + 1u{% endif %};
138
  var minKj = 0u;
139
  if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
140
 
141
  let accBase = lane * HEAD_DIM;
142
+ for (var d = 0u; d < HEAD_DIM; d = d + 1u) { state.acc[accBase + d] = 0.0; }
143
  var m = NEG_INF;
144
  var l = 0.0;
145
 
 
147
  for (var j = minKj + lane; j < maxKj; j = j + WG) {
148
  let base = ((b * KV_HEADS + hk) * totalSeq + j) * PACKED;
149
  var dot = 0.0;
150
+ for (var d = 0u; d < HEAD_DIM; d = d + 1u) { dot = dot + state.q[d] * read_k(base, d{% if quantized %}, hk{% endif %}); }
151
  var score = dot * scale;
152
  if (params.softcap > 0.0) { score = params.softcap * tanh(clamp(score / params.softcap, -30.0, 30.0)); }
153
  {% if hasBias %}
154
  let bb = select(b, 0u, params.biasBatch == 1u);
155
  let bh = select(h, 0u, params.biasHeads == 1u);
156
+ score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) * {{ biasStride }} + {{ biasCol }}]);
157
  {% endif %}
158
  let newM = max(m, score);
159
  let corr = exp(m - newM);
160
  let p = exp(score - newM);
161
  l = l * corr + p;
162
+ for (var d = 0u; d < HEAD_DIM; d = d + 1u) { state.acc[accBase + d] = state.acc[accBase + d] * corr + p * read_v(base, d{% if quantized %}, hk{% endif %}); }
163
  m = newM;
164
  }
165
 
166
  // Flash merge across the WG lanes: global max, rescale, summed denom.
167
+ state.m[lane] = m;
168
  workgroupBarrier();
169
  var gm = NEG_INF;
170
+ for (var L = 0u; L < WG; L = L + 1u) { gm = max(gm, state.m[L]); }
171
  let fctr = exp(m - gm);
172
  l = l * fctr;
173
+ for (var d = 0u; d < HEAD_DIM; d = d + 1u) { state.acc[accBase + d] = state.acc[accBase + d] * fctr; }
174
+ state.l[lane] = l;
175
  workgroupBarrier();
176
  var glsum = 0.0;
177
+ for (var L = 0u; L < WG; L = L + 1u) { glsum = glsum + state.l[L]; }
178
 
179
  var sink = 0.0;
180
  {% if hasHeadSink %}
 
199
  let oBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
200
  for (var d = lane; d < HEAD_DIM; d = d + WG) {
201
  var sumacc = 0.0;
202
+ for (var L = 0u; L < WG; L = L + 1u) { sumacc = sumacc + state.acc[L * HEAD_DIM + d]; }
203
  output[oBase + d] = {{ IO }}(sumacc * invDenom);
204
  }
205
  }
 
227
  {% else %}
228
  let activeEnd = totalSeq;
229
  {% endif %}
230
+ {% if not bidirectional or hasRotary %} let absPos = activeEnd - params.qSeq + s;
231
+ {% endif %}{% if needsOrigin %}
232
+ // Absolute position of cache row 0 (windowed cache): the query end is the
233
+ // unclamped total while the cache holds only its last `activeEnd` rows.
234
+ let origin = max(params.qSeq, u32(seqlens_k[b]) + 1u) - activeEnd;
235
+ {% endif %}
236
 
237
  // Private query-vector prep (post norm/rotary).
238
  var q: array<f32, HEAD_DIM>;
 
246
  {% endif %}
247
  {% if hasRotary %}
248
  for (var d = 0u; d < HALF; d = d + 1u) {
249
+ let cs = f32(cos_cache[{{ ropePos }} * HALF + d]);
250
+ let sn = f32(sin_cache[{{ ropePos }} * HALF + d]);
251
  let x0 = q[d];
252
  let x1 = q[d + HALF];
253
  q[d] = x0 * cs - x1 * sn;
 
255
  }
256
  {% endif %}
257
  let scale = scale_value();
258
+ let maxKj = {% if bidirectional %}activeEnd{% else %}absPos + 1u{% endif %};
259
  var minKj = 0u;
260
  if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
261
 
 
273
  {% if hasBias %}
274
  let bb = select(b, 0u, params.biasBatch == 1u);
275
  let bh = select(h, 0u, params.biasHeads == 1u);
276
+ score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) * {{ biasStride }} + {{ biasCol }}]);
277
  {% endif %}
278
  let newM = max(m, score);
279
  let corr = exp(m - newM);
build/webgpu/gqa-present.wgsl.jinja CHANGED
@@ -9,8 +9,15 @@
9
  {% set bits = bits | default(0) %}
10
  {% set qmax = qmax | default(0) %}
11
  {% set qmin = qmin | default(0) %}
12
- {% if usesF16 is defined and usesF16 %}enable f16;
13
- {% endif %}{{ env.wgsl.resourceDeclarations }}
 
 
 
 
 
 
 
14
  {% if mode == "transpose" %}
15
 
16
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
@@ -31,9 +38,7 @@ const KV_HIDDEN: u32 = {{ kvHidden }}u;
31
  fn main(
32
  @builtin(global_invocation_id) gid: vec3<u32>
33
  ) {
34
- // The dispatch folds oversized one-dimensional grids into x/y. Rebuild the
35
- // flat invocation index; this reduces to gid.x when no fold is needed.
36
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
37
  {% if presentVec4 %}
38
  let total = params.batchSize * KV_HEADS * params.kvSeq * HEAD_DIM_V4;
39
  if (index >= total) {
@@ -60,9 +65,11 @@ fn main(
60
  present_value[index] = {{ presentScalar }}(value[packed_index]);
61
  {% endif %}
62
  }
63
- {%- else %}
64
 
65
  {% set shareAppend = mode == "merge_share" and shareRegion == "append" %}
 
 
66
  {% set cooperativeCopy = device.adapterInfo.vendor != "apple" %}
67
  {% set cooperativeMerge = cooperativeCopy or shareAppend %}
68
  {% set cooperativeMode = cooperativeMerge and (mode == "copy" or mode == "merge" or mode == "merge_share") %}
@@ -95,7 +102,7 @@ const HEAD_DIM: u32 = {{ headDim }}u;
95
  {% if mode != "copy" %}const KV_HEADS: u32 = {{ kvHeads }}u;
96
  {% endif %}
97
  const WG: u32 = {{ tunables.COPY_WORKGROUP_SIZE }}u;
98
- {% if mode != "copy" and (mode != "merge_share" or shareAppend) %}const KV_HIDDEN: u32 = {{ kvHeads }}u * {{ headDim }}u;
99
  {% endif %}
100
  {% if mode == "build_quant" or mode == "append_quant" %}const PACKED: u32 = {{ packed }}u;
101
  {% endif %}
@@ -153,8 +160,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
153
  {% endmacro %}
154
  {% macro append_per_row_walk() %}
155
  // Per-row append walk: one thread streams its row's contiguous bytes.
156
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
157
- if (i >= params.count) { return; }
158
  let j = i % params.keySeq;
159
  let tmp = i / params.keySeq;
160
  let hk = tmp % KV_HEADS;
@@ -263,8 +269,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
263
  {% endif %}
264
  {% elif mode == "copy" %}
265
  // Thread i streams one contiguous (batch, kvHead, token) row.
266
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
267
- if (i >= params.count) { return; }
268
  let base = i * HEAD_DIM;
269
  for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
270
  present_key[base + d] = src_k[base + d];
@@ -272,8 +277,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
272
  }
273
  {% elif mode == "merge" %}
274
  // Per-row merge walk.
275
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
276
- if (i >= params.count) { return; }
277
  let t = i % params.seq;
278
  let tmp = i / params.seq;
279
  let hk = tmp % KV_HEADS;
@@ -296,8 +300,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
296
  {% elif mode == "merge_share" %}
297
  // Past and present share the full-capacity stride, so outside-window rows
298
  // copy at the same index.
299
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
300
- if (i >= params.count) { return; }
301
  let t = i % params.seq;
302
  let tmp = i / params.seq;
303
  let hk = tmp % KV_HEADS;
@@ -316,28 +319,51 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
316
  }
317
  }
318
  {% elif mode == "window_shift" %}
319
- // Windowed cache: `params.seq` is a fixed capacity C, and the invariant is that
320
- // the L = min(T, C) most recent tokens live contiguously at rows [0, L).
321
- //
322
- // residentBefore = min(T - S, C) E = max(0, residentBefore + S - C)
323
- // rows [0, appendStart) <- past rows shifted down by E
324
- // rows [appendStart, append+S) <- the new tokens
325
- // rows beyond that <- outside the window, cleared
326
- //
327
- // The shift reads past and writes present, which must be distinct buffers: an
328
- // in-place compaction would have one invocation overwrite row t while another
329
- // still needs it as the source for row t-E, with no ordering between them.
330
- //
331
- // Attention needs no change for this layout. It derives its key range from
332
- // `min(capacity, seqlens_k[b] + 1)`, which is exactly L, and both the causal
333
- // and local-window masks depend only on the query/key distance:
334
- // q_abs - k_abs = (T - qSeq + s) - (origin + t) = (L - qSeq + s) - t
335
- // so scoring a windowed cache as if it were a full L-length one is the same
336
- // arithmetic. Only RoPE would need the true absolute position, which is why
337
- // `windowShiftOk` refuses a rotary request outright rather than silently
338
- // rotating at the cache row.
339
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
340
- if (i >= params.count) { return; }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
341
  let t = i % params.seq;
342
  let tmp = i / params.seq;
343
  let hk = tmp % KV_HEADS;
@@ -358,6 +384,9 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
358
  present_key[dstBase + d] = past_k[srcBase + d];
359
  present_value[dstBase + d] = past_v[srcBase + d];
360
  }
 
 
 
361
  } else if (t < appendStart + params.keySeq) {
362
  let nSrc = (b * params.keySeq + (t - appendStart)) * KV_HIDDEN + hk * HEAD_DIM;
363
  for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
@@ -365,6 +394,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
365
  present_value[dstBase + d] = new_v[nSrc + d];
366
  }
367
  } else {
 
368
  // Outside the resident window. Cleared rather than left stale so the present
369
  // buffer does not expose stale contents from its distinct allocation.
370
  for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
@@ -372,11 +402,9 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
372
  present_value[dstBase + d] = {{ zeroScalar }}(0.0);
373
  }
374
  }
 
375
  {% else %}
376
- // The dispatch folds oversized one-dimensional grids into x/y. Rebuild the
377
- // flat invocation index; this reduces to gid.x when no fold is needed.
378
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
379
- if (i >= params.count) { return; }
380
  let t = i % params.seq;
381
  let tmp = i / params.seq;
382
  let hk = tmp % KV_HEADS;
@@ -435,4 +463,4 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
435
  {% endif %}
436
  {% endif %}
437
  }
438
- {%- endif %}
 
9
  {% set bits = bits | default(0) %}
10
  {% set qmax = qmax | default(0) %}
11
  {% set qmin = qmin | default(0) %}
12
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
13
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
14
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
15
+ // per-axis workgroup fold width.
16
+ {% if bound == "" %}
17
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% else %}
18
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
19
+ if ({{ name }} >= {{ bound }}) { return; }{% endif %}{% endmacro %}
20
+ {{ env.wgsl.resourceDeclarations }}
21
  {% if mode == "transpose" %}
22
 
23
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
 
38
  fn main(
39
  @builtin(global_invocation_id) gid: vec3<u32>
40
  ) {
41
+ {{ flat_index_2d("WG", "index", "") }}
 
 
42
  {% if presentVec4 %}
43
  let total = params.batchSize * KV_HEADS * params.kvSeq * HEAD_DIM_V4;
44
  if (index >= total) {
 
65
  present_value[index] = {{ presentScalar }}(value[packed_index]);
66
  {% endif %}
67
  }
68
+ {% else %}
69
 
70
  {% set shareAppend = mode == "merge_share" and shareRegion == "append" %}
71
+ {% set shiftRegion = shiftRegion | default("all") %}
72
+ {% set shiftRetainOnly = mode == "window_shift" and shiftRegion == "shift" %}
73
  {% set cooperativeCopy = device.adapterInfo.vendor != "apple" %}
74
  {% set cooperativeMerge = cooperativeCopy or shareAppend %}
75
  {% set cooperativeMode = cooperativeMerge and (mode == "copy" or mode == "merge" or mode == "merge_share") %}
 
102
  {% if mode != "copy" %}const KV_HEADS: u32 = {{ kvHeads }}u;
103
  {% endif %}
104
  const WG: u32 = {{ tunables.COPY_WORKGROUP_SIZE }}u;
105
+ {% if mode != "copy" and (mode != "merge_share" or shareAppend) and not shiftRetainOnly %}const KV_HIDDEN: u32 = {{ kvHeads }}u * {{ headDim }}u;
106
  {% endif %}
107
  {% if mode == "build_quant" or mode == "append_quant" %}const PACKED: u32 = {{ packed }}u;
108
  {% endif %}
 
160
  {% endmacro %}
161
  {% macro append_per_row_walk() %}
162
  // Per-row append walk: one thread streams its row's contiguous bytes.
163
+ {{ flat_index_2d("WG", guardInline=true) }}
 
164
  let j = i % params.keySeq;
165
  let tmp = i / params.keySeq;
166
  let hk = tmp % KV_HEADS;
 
269
  {% endif %}
270
  {% elif mode == "copy" %}
271
  // Thread i streams one contiguous (batch, kvHead, token) row.
272
+ {{ flat_index_2d("WG", guardInline=true) }}
 
273
  let base = i * HEAD_DIM;
274
  for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
275
  present_key[base + d] = src_k[base + d];
 
277
  }
278
  {% elif mode == "merge" %}
279
  // Per-row merge walk.
280
+ {{ flat_index_2d("WG", guardInline=true) }}
 
281
  let t = i % params.seq;
282
  let tmp = i / params.seq;
283
  let hk = tmp % KV_HEADS;
 
300
  {% elif mode == "merge_share" %}
301
  // Past and present share the full-capacity stride, so outside-window rows
302
  // copy at the same index.
303
+ {{ flat_index_2d("WG", guardInline=true) }}
 
304
  let t = i % params.seq;
305
  let tmp = i / params.seq;
306
  let hk = tmp % KV_HEADS;
 
319
  }
320
  }
321
  {% elif mode == "window_shift" %}
322
+ // Keep the most recent min(total, capacity) tokens in contiguous cache rows.
323
+ // Read past and write a distinct present buffer so compaction cannot race.
324
+ // Rotary append uses the absolute token position; retained keys are already rotated.
325
+ {% if shiftRegion == "append" %}
326
+ // Append-only pass: the thread range spans just the appended window (batch x
327
+ // kvHead x keySeq rows), while `params.seq` stays the cache capacity. The shift
328
+ // pass owns every other row and the two ranges are disjoint, so their order is
329
+ // immaterial.
330
+ {{ flat_index_2d("WG", guardInline=true) }}
331
+ let j = i % params.keySeq;
332
+ let tmp = i / params.keySeq;
333
+ let hk = tmp % KV_HEADS;
334
+ let b = tmp / KV_HEADS;
335
+ let absTotal = max(params.keySeq, u32(seqlens_k[b]) + 1u);
336
+ let residentBefore = min(absTotal - params.keySeq, params.seq);
337
+ // `max(0u, a + b - c)` does not clamp in u32: the subtraction wraps first, so a
338
+ // step that evicts nothing reads back ~2^32 instead of 0. Compare before
339
+ // subtracting.
340
+ let filled = residentBefore + params.keySeq;
341
+ let evicted = select(0u, filled - params.seq, filled > params.seq);
342
+ let appendStart = residentBefore - evicted;
343
+ let t = appendStart + j;
344
+ let dstBase = ((b * KV_HEADS + hk) * params.seq + t) * HEAD_DIM;
345
+ let nSrc = (b * params.keySeq + j) * KV_HIDDEN + hk * HEAD_DIM;
346
+ // Row appendStart + j carries absolute position absTotal - keySeq + j: the cache
347
+ // holds only the last L positions, so the row index is short of the absolute one
348
+ // by the window origin. The reference rotates the same rows at
349
+ // `appendStart + cacheOrigin`, which is the same number.
350
+ let pos = absTotal - params.keySeq + j;
351
+ var k: array<f32, HEAD_DIM>;
352
+ for (var d = 0u; d < HEAD_DIM; d = d + 1u) { k[d] = new_k[nSrc + d]; }
353
+ for (var d = 0u; d < HALF; d = d + 1u) {
354
+ let cs = cos_cache[pos * HALF + d];
355
+ let sn = sin_cache[pos * HALF + d];
356
+ let x0 = k[d];
357
+ let x1 = k[d + HALF];
358
+ k[d] = x0 * cs - x1 * sn;
359
+ k[d + HALF] = x1 * cs + x0 * sn;
360
+ }
361
+ for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
362
+ present_key[dstBase + d] = k[d];
363
+ present_value[dstBase + d] = new_v[nSrc + d];
364
+ }
365
+ {% else %}
366
+ {{ flat_index_2d("WG", guardInline=true) }}
367
  let t = i % params.seq;
368
  let tmp = i / params.seq;
369
  let hk = tmp % KV_HEADS;
 
384
  present_key[dstBase + d] = past_k[srcBase + d];
385
  present_value[dstBase + d] = past_v[srcBase + d];
386
  }
387
+ {% if shiftRegion == "shift" %}
388
+ } else if (t >= appendStart + params.keySeq) {
389
+ {% else %}
390
  } else if (t < appendStart + params.keySeq) {
391
  let nSrc = (b * params.keySeq + (t - appendStart)) * KV_HIDDEN + hk * HEAD_DIM;
392
  for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
 
394
  present_value[dstBase + d] = new_v[nSrc + d];
395
  }
396
  } else {
397
+ {% endif %}
398
  // Outside the resident window. Cleared rather than left stale so the present
399
  // buffer does not expose stale contents from its distinct allocation.
400
  for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
 
402
  present_value[dstBase + d] = {{ zeroScalar }}(0.0);
403
  }
404
  }
405
+ {% endif %}
406
  {% else %}
407
+ {{ flat_index_2d("WG", guardInline=true) }}
 
 
 
408
  let t = i % params.seq;
409
  let tmp = i / params.seq;
410
  let hk = tmp % KV_HEADS;
 
463
  {% endif %}
464
  {% endif %}
465
  }
466
+ {% endif %}
build/webgpu/gqa-qprep.wgsl.jinja CHANGED
@@ -1,5 +1,9 @@
1
- {% if usesF16 is defined and usesF16 %}enable f16;
2
- {% endif %}{{ env.wgsl.resourceDeclarations }}
 
 
 
 
3
 
4
  // One invocation per (batch, query head, query token) applies optional per-head
5
  // RMS normalization followed by NeoX half-split rotary embedding at the query's
@@ -21,9 +25,7 @@ const QK_EPS: f32 = {{ qkEps }};
21
 
22
  @compute @workgroup_size(WG)
23
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
24
- // 2D-folded flat index: gid.y carries the high bits past the per-axis dispatch fold width.
25
- // Reduces to gid.x when the dispatch does not fold.
26
- let qi = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
27
  let total = params.batch * Q_HEADS * params.qSeq;
28
  if (qi >= total) { return; }
29
  let s = qi % params.qSeq;
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
6
+ {{ env.wgsl.resourceDeclarations }}
7
 
8
  // One invocation per (batch, query head, query token) applies optional per-head
9
  // RMS normalization followed by NeoX half-split rotary embedding at the query's
 
25
 
26
  @compute @workgroup_size(WG)
27
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
28
+ {{ flat_index_2d("WG", "qi", "") }}
 
 
29
  let total = params.batch * Q_HEADS * params.qSeq;
30
  if (qi >= total) { return; }
31
  let s = qi % params.qSeq;
build/webgpu/manifest.json CHANGED
The diff for this file is too large to render. See raw diff
 
build/webgpu/metadata.json CHANGED
@@ -1,32 +1,32 @@
1
  {
2
  "name": "com.microsoft.GroupQueryAttention",
3
- "id": "_com_microsoft_groupqueryattention_webgpu_1ea0022",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "attention-rank4-tiled.wgsl.jinja": "gYl58ecIVTxJJkvnK/6lNcbPKaRT24cfz5Yx0MxxxxE=",
11
- "attn-flash-decode-splitk-merge.wgsl.jinja": "y3ubiijNd5rGJa2KVdluBdaoKYAGeqrPVd6KSppJPHw=",
12
- "attn-flash-decode-splitk.wgsl.jinja": "xRRcDB3IGuE25xq3sG1qVPHRsESDmfA15KFVKlq5l4w=",
13
- "attn-flash-online.wgsl.jinja": "ontnzJw9RHN7DklqqyavsAiRuKrsnWuEiU+vllbKfGE=",
14
- "attn-flash-prefill-cluster.wgsl.jinja": "UHb5IdDeGHB/0gsQdAHcfg9FI73nyzk3zoZiIhaWbgI=",
15
- "attn-flash-q32-broadcast.wgsl.jinja": "U/O9TNr3pXvFP+VZ2EAJLuXaWSz62De0TQEmEiU7TLc=",
16
- "attn-materialized-rowstats-combine-f32.wgsl.jinja": "zcR02XUVlPyYeLeWH7gLOX65BcozZ1qje9hSIZSOaBA=",
17
- "attn-materialized-sgmat-f32.wgsl.jinja": "rNruz2BmIcsQeu5eUnQZh41ftfDt/PjqnZiy2XT+9bk=",
18
- "attn-online-scalar.wgsl.jinja": "AJGUYoMkrTHnwCPxu1gVp2CIrdHlca5PMgr9mgqukz4=",
19
- "bench.json": "bUGLPxdrsT20K/2ptNbqVs+28i4d2W/9vOyPld+joKs=",
20
- "gqa-attention.wgsl.jinja": "9JBHgllUq4dekL6Y1uXkyi2NoEYKjHTkHqnhktbgdzY=",
21
- "gqa-present.wgsl.jinja": "/63mkauF/TKCsH4dxgIVlTZnREzSd/j64WGkvu6i2ls=",
22
- "gqa-qprep.wgsl.jinja": "kRFlYxr5SdQJlbjmcF/Ii+0w56q0xhF8jyAOKPNx47I=",
23
- "manifest.json": "U9Xgwr8xQCIMv/BNUNjiaMDyUjEVXkawKtl5Bsnn4UA=",
24
- "test.json": "GGjXGSZThloU2mCpXKJf9lJVgddpxc8ouDX+JRD5zTs="
25
  }
26
  },
27
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
28
  "webgpu": {
29
- "manifestSpec": "2.0",
30
  "variants": {
31
  "qkv_present_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
32
  "past_kv_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
@@ -40,9 +40,12 @@
40
  "new_kv_share_append_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
41
  "new_kv_share_append_headsink_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
42
  "new_kv_share_append_rotary_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
43
  "qkv_present_tiled_nosg": ["attention-rank4-tiled.wgsl.jinja", "gqa-present.wgsl.jinja"],
44
  "qkv_present_flash": ["attn-flash-online.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
45
  "qkv_present": ["attn-online-scalar.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
46
  "quant_int8": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
47
  "quant_int4": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
48
  "quant_int8_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
@@ -50,17 +53,9 @@
50
  "qkv_present_flash_cluster": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
51
  "past_kv_bias_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
52
  "past_kv_qnorm_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja", "gqa-qprep.wgsl.jinja"],
53
- "quant_int8_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
54
- "qkv_present_flash_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
55
- "qkv_present_flash_cluster_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
56
- "past_kv_bias_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
57
- "past_kv_qnorm_rotary_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja", "gqa-qprep.wgsl.jinja"],
58
  "past_kv_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
59
  "new_kv_past_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
60
  "window_shift_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
61
- "past_kv_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
62
- "new_kv_past_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
63
- "window_shift_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
64
  "past_kv_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
65
  "new_kv_past_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
66
  "past_kv_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
@@ -68,32 +63,27 @@
68
  "past_kv_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
69
  "past_kv_bias_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
70
  "window_shift_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
71
- "past_kv_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
72
- "new_kv_past_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
73
- "past_kv_rotary_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
74
- "past_kv_softcap_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
75
- "past_kv_headsink_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
76
- "past_kv_bias_headsink_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
77
- "window_shift_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
78
  "qkv_present_flash_q32_broadcast": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
79
  "qkv_present_flash_q32_shared": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
80
  "past_kv": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
81
  "past_kv_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
82
  "past_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
 
 
83
  "new_kv_past": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
84
  "window_shift_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
 
85
  "new_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
86
  "past_kv_bias": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
87
  "past_kv_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
88
  "past_kv_bias_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
89
  "quant_int8_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
90
  "quant_int4_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
91
- "quant_int8_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
92
- "quant_int4_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
93
  "share_append_split_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
94
- "share_append_split_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
95
  "share_append_split_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
96
- "share_append_split_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"]
97
  }
98
  }
99
  }
 
1
  {
2
  "name": "com.microsoft.GroupQueryAttention",
3
+ "id": "_com_microsoft_groupqueryattention_webgpu_0c84e3e",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "attention-rank4-tiled.wgsl.jinja": "084x5bbpxeFR95/Q4nGzLrLnk9JJrd9XA1KJHxwU9CY=",
11
+ "attn-flash-decode-splitk-merge.wgsl.jinja": "nKMUcaMP7mcKykLHoGsn/KnsIaJiJRHDsceE7dKVcwo=",
12
+ "attn-flash-decode-splitk.wgsl.jinja": "ZD6C2fn5I7xZvvpTlsHdLtai8JBK2Ar+xc0SrLn3SAE=",
13
+ "attn-flash-online.wgsl.jinja": "QGaUiZyIJuQ9xS00I8cuZWYQI8bGldV+UJ0Mp/5kj6w=",
14
+ "attn-flash-prefill-cluster.wgsl.jinja": "Q7DUmulv66HOVBMkAannNVhEtOWkD+uhDn5muWo4p+E=",
15
+ "attn-flash-q32-broadcast.wgsl.jinja": "vogUJggIxK8tRg2tDORhIXZLheesY3Y221iVFFyHVY4=",
16
+ "attn-materialized-rowstats-combine-f32.wgsl.jinja": "SG+V1o5KhJ1LvkUQL0jcq62C+ETRGpy0XV9tvvOQivk=",
17
+ "attn-materialized-sgmat-f32.wgsl.jinja": "R0+zNbbnfRKpA+R7lcgY5FPlkmLNzNBJwJe1PmaSOvE=",
18
+ "attn-online-scalar.wgsl.jinja": "ysj3CuuKeu+MPlnaTwnNvCu3sSf8egjHE/uhE/Essmo=",
19
+ "bench.json": "v06hV3rz9UPzelx7YbIHmyOWgbj/0Y9ci7iz9mQS6m0=",
20
+ "gqa-attention.wgsl.jinja": "Wy1oR/7g/EzxQZizt/zGtOh4qiD60SVN4pfLSjmfsDo=",
21
+ "gqa-present.wgsl.jinja": "r0ELXXscqaVYqTYzA9IYr/xl9jWAe20ZFPScyiw34UQ=",
22
+ "gqa-qprep.wgsl.jinja": "VI/IHqwEMTpvrhFza3kKRDuLO3Eau91C4aMZRW0nAjg=",
23
+ "manifest.json": "8m8RL1h3mBW9P8voV+gr/KpJiMheynqdWWhmcrOoi8E=",
24
+ "test.json": "I771sUbwb/9MXIc27bUlvNl3UsiiFxP3H0/wRMt/B9M="
25
  }
26
  },
27
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
28
  "webgpu": {
29
+ "manifestSpec": "2.1",
30
  "variants": {
31
  "qkv_present_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
32
  "past_kv_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
40
  "new_kv_share_append_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
41
  "new_kv_share_append_headsink_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
42
  "new_kv_share_append_rotary_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
43
+ "new_kv_share_append_bidirectional_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
44
  "qkv_present_tiled_nosg": ["attention-rank4-tiled.wgsl.jinja", "gqa-present.wgsl.jinja"],
45
  "qkv_present_flash": ["attn-flash-online.wgsl.jinja", "gqa-present.wgsl.jinja"],
46
+ "qkv_present_flash_causal": ["attn-flash-online.wgsl.jinja", "gqa-present.wgsl.jinja"],
47
  "qkv_present": ["attn-online-scalar.wgsl.jinja", "gqa-present.wgsl.jinja"],
48
+ "qkv_present_causal": ["attn-online-scalar.wgsl.jinja", "gqa-present.wgsl.jinja"],
49
  "quant_int8": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
50
  "quant_int4": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
51
  "quant_int8_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
53
  "qkv_present_flash_cluster": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
54
  "past_kv_bias_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
55
  "past_kv_qnorm_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja", "gqa-qprep.wgsl.jinja"],
 
 
 
 
 
56
  "past_kv_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
57
  "new_kv_past_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
58
  "window_shift_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
 
 
59
  "past_kv_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
60
  "new_kv_past_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
61
  "past_kv_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
63
  "past_kv_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
64
  "past_kv_bias_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
65
  "window_shift_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
 
 
 
 
 
 
66
  "qkv_present_flash_q32_broadcast": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
67
  "qkv_present_flash_q32_shared": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
68
  "past_kv": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
69
  "past_kv_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
70
  "past_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
71
+ "window_shift_rotary_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
72
+ "window_shift_rotary_headsink_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
73
+ "window_shift_rotary_bias_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
74
  "new_kv_past": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
75
  "window_shift_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
76
+ "window_shift_headsink_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
77
+ "window_shift_bias_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
78
  "new_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
79
  "past_kv_bias": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
80
  "past_kv_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
81
  "past_kv_bias_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
82
  "quant_int8_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
83
  "quant_int4_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
 
84
  "share_append_split_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
 
85
  "share_append_split_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
86
+ "share_append_bidirectional_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"]
87
  }
88
  }
89
  }
build/webgpu/test.json CHANGED
The diff for this file is too large to render. See raw diff