Xenova HF Staff commited on
Commit
43f449f
·
verified ·
1 Parent(s): 2157e54

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -52,6 +52,8 @@ Attributes and default values (overridable per request):
52
 
53
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
54
 
 
 
55
  - `qkv_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences with bias and optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
56
  - `qkv_no_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences without bias and with optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
57
  - `qkv_no_bias_small_seq` — Whole-head attention for short bidirectional float32 requests without bias. One workgroup stages a head's complete key and value planes, and each participating invocation owns one query row through scoring, normalization, and context accumulation.
@@ -100,7 +102,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
100
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
101
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
102
  - [`test.json`](build/webgpu/test.json) — correctness cases
103
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
104
  - [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
105
  - [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
106
  - [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
@@ -114,13 +116,14 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
114
  - [`attn-materialized-softmax-f32.wgsl.jinja`](build/webgpu/attn-materialized-softmax-f32.wgsl.jinja)
115
  - [`attn-online-scalar.wgsl.jinja`](build/webgpu/attn-online-scalar.wgsl.jinja)
116
  - [`attn-small-head-parallel.wgsl.jinja`](build/webgpu/attn-small-head-parallel.wgsl.jinja)
 
117
  - [`mha-small-seq-blocked.wgsl.jinja`](build/webgpu/mha-small-seq-blocked.wgsl.jinja)
118
  - [`mha-small-seq.wgsl.jinja`](build/webgpu/mha-small-seq.wgsl.jinja)
119
 
120
  ## Use with `@huggingface/kernels`
121
 
122
  ```sh
123
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
124
  ```
125
 
126
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
52
 
53
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
54
 
55
+ - `qkv_no_bias_small_head_value_subgroups` — Cooperatively reduces values across key lanes and reuses keys/values for two queries when device workgroup storage fits. Uses portable-width subgroup collectives with f32 scores, probabilities and accumulation. Short key sequences use the serial value path.
56
+ - `qkv_no_bias_small_head_value_tree` — Cooperatively reduces values across key lanes and reuses keys/values for two queries when device workgroup storage fits. Uses shared-memory tree reductions with f32 scores, probabilities and accumulation. Short key sequences use the serial value path.
57
  - `qkv_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences with bias and optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
58
  - `qkv_no_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences without bias and with optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
59
  - `qkv_no_bias_small_seq` — Whole-head attention for short bidirectional float32 requests without bias. One workgroup stages a head's complete key and value planes, and each participating invocation owns one query row through scoring, normalization, and context accumulation.
 
102
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
103
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
104
  - [`test.json`](build/webgpu/test.json) — correctness cases
105
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
106
  - [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
107
  - [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
108
  - [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
 
116
  - [`attn-materialized-softmax-f32.wgsl.jinja`](build/webgpu/attn-materialized-softmax-f32.wgsl.jinja)
117
  - [`attn-online-scalar.wgsl.jinja`](build/webgpu/attn-online-scalar.wgsl.jinja)
118
  - [`attn-small-head-parallel.wgsl.jinja`](build/webgpu/attn-small-head-parallel.wgsl.jinja)
119
+ - [`attn-small-head-value.wgsl.jinja`](build/webgpu/attn-small-head-value.wgsl.jinja)
120
  - [`mha-small-seq-blocked.wgsl.jinja`](build/webgpu/mha-small-seq-blocked.wgsl.jinja)
121
  - [`mha-small-seq.wgsl.jinja`](build/webgpu/mha-small-seq.wgsl.jinja)
122
 
123
  ## Use with `@huggingface/kernels`
124
 
125
  ```sh
126
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
127
  ```
128
 
129
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
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
@@ -49,6 +46,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
49
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
50
  return select(value - maxValue, 0.0, equalFiniteMax);
51
  }
 
52
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
53
  return exp(shifted_value(value, maxValue));
54
  }
@@ -104,11 +102,6 @@ fn main(
104
  // V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
105
  let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
106
  outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
107
- {% endif %}
108
- {% if hasGate %}
109
- // The gated route multiplies the normalized attention output elementwise by its gate.
110
- let gateV = vec4<f32>(gate[qBaseV4 + d4]);
111
- outValue = outValue * (vec4<f32>(1.0) / (vec4<f32>(1.0) + exp(-gateV)));
112
  {% endif %}
113
  output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
114
  }
 
 
 
 
1
  {% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
 
2
  {% set qSeq = qSeq | default(0) %}
3
+ {% set qHidden = qHidden | default(0) %}
4
  {{ env.wgsl.resourceDeclarations }}
5
 
6
  // Split-K flash decode, pass 2 of 2. Combines the per-split un-normalized online
 
46
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
47
  return select(value - maxValue, 0.0, equalFiniteMax);
48
  }
49
+
50
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
51
  return exp(shifted_value(value, maxValue));
52
  }
 
102
  // V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
103
  let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
104
  outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
 
 
 
 
 
105
  {% endif %}
106
  output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
107
  }
build/webgpu/attn-flash-decode-splitk.wgsl.jinja CHANGED
@@ -1,11 +1,6 @@
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 scale = scale | default("0.0") %}
8
- {% if hasRotary is not defined %}{% set hasRotary = false %}{% endif %}
9
  {% set qSeq = qSeq | default(0) %}
10
  {% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
11
  {% if useSubgroups %}
@@ -13,7 +8,8 @@ enable subgroups;
13
  {% endif %}
14
  {{ env.wgsl.resourceDeclarations }}
15
 
16
- // Split-K flash attention, pass 1 of 2; the merge pass follows. This geometry
 
17
  // handles decode and short-query, long-context prefill inputs.
18
  //
19
  // The non-split flash decode launches only `batch * numHeads` workgroups, each
@@ -56,6 +52,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
56
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
57
  return select(value - maxValue, 0.0, equalFiniteMax);
58
  }
 
59
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
60
  return exp(shifted_value(value, maxValue));
61
  }
@@ -63,7 +60,7 @@ fn exp_shift(value: f32, maxValue: f32) -> f32 {
63
  var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
64
  var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
65
  var<workgroup> probs: array<f32, WG>;
66
- {% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) %}
67
  {% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
68
  {% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
69
  {% if coopQk %}
@@ -132,41 +129,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
132
  return combinedMD;
133
  }
134
  {% else %}
135
- {% set mdStreamed = mdStreams is defined %}
136
- {% set mdStreams = mdStreams if mdStreams is defined else 1 %}
137
- {% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
138
  var<workgroup> partialM: array<f32, {{ mdExtent }}>;
139
  var<workgroup> partialD: array<f32, {{ mdExtent }}>;
140
- {% if mdStreamed %}
141
-
142
- // In-place fold of {{ mdStreams }} streams. Input partials occupy
143
- // partialM/partialD; stream s returns its merged pair in slot s * WG.
144
- fn combine_partials_streams(lidx: u32) {
145
- workgroupBarrier();
146
- var stride = WG / 2u;
147
- loop {
148
- if (stride == 0u) {
149
- break;
150
- }
151
- if (lidx < stride) {
152
- {% for s in range(mdStreams) %}
153
- {
154
- let slot = {{ s }}u * WG + lidx;
155
- let m1 = partialM[slot];
156
- let d1 = partialD[slot];
157
- let m2 = partialM[slot + stride];
158
- let d2 = partialD[slot + stride];
159
- let mNew = max(m1, m2);
160
- partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
161
- partialM[slot] = mNew;
162
- }
163
- {% endfor %}
164
- }
165
- workgroupBarrier();
166
- stride = stride / 2u;
167
- }
168
- }
169
- {% else %}
170
 
171
  fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
172
  partialM[lidx] = m;
@@ -196,54 +161,12 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
196
  return merged;
197
  }
198
  {% endif %}
199
- {% endif %}
200
-
201
 
202
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
203
- fn scale_value() -> f32 {
204
  if (params.scale != 0.0) { return params.scale; }
205
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
206
  }
207
 
208
-
209
- {% if quantizedCache %}
210
- {% macro emit_quant_scale4(kind, scaleBuffer) %}
211
- fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
212
- if (params.perChannel == 0u) {
213
- return vec4<f32>({{ scaleBuffer }}[0]);
214
- }
215
- let base = hk * HEAD_DIM + d4 * 4u;
216
- return vec4<f32>(
217
- {{ scaleBuffer }}[base],
218
- {{ scaleBuffer }}[base + 1u],
219
- {{ scaleBuffer }}[base + 2u],
220
- {{ scaleBuffer }}[base + 3u]
221
- );
222
- }
223
- {%- endmacro %}
224
- {%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
225
- {{ emit_quant_scale4(kind, scaleBuffer) }}
226
- fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
227
- {%- if format == "int8" %}
228
- return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
229
- {%- else %}
230
- // Two elements cover this vec4: each carries two +8-biased nibbles, low first.
231
- let rowBase = indexV4 - d4;
232
- let lo = {{ buffer }}[rowBase + d4 * 2u];
233
- let hi = {{ buffer }}[rowBase + d4 * 2u + 1u];
234
- let nibbles = vec4<i32>(
235
- i32(lo & 0xFu), i32((lo >> 4u) & 0xFu),
236
- i32(hi & 0xFu), i32((hi >> 4u) & 0xFu)
237
- );
238
- let signed = nibbles - vec4<i32>(8);
239
- return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
240
- {%- endif %}
241
- }
242
- {%- endmacro %}
243
-
244
- {{ emit_quant_load4("int8", "key", "key", "k_scale") }}
245
- {{ emit_quant_load4("int8", "value", "value", "v_scale") }}
246
- {% else %}
247
  fn load_key4(indexV4: u32) -> vec4<f32> {
248
  return vec4<f32>(key[indexV4]);
249
  }
@@ -251,16 +174,17 @@ fn load_key4(indexV4: u32) -> vec4<f32> {
251
  fn load_value4(indexV4: u32) -> vec4<f32> {
252
  return vec4<f32>(value[indexV4]);
253
  }
254
- {% endif %}
255
 
256
  {% if hasBias %}
257
  // Packed [Q; K; V] bias rows (token-independent). The Q bias folds into the
258
  // query row before the Q.K dots; the K bias adds a constant to every key score
259
  // that softmax cancels, so it is skipped; the V bias is token-independent and
260
  // is applied once in the merge pass after the final normalize.
 
 
261
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
262
  let offset = base + d4 * 4u;
263
- return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
264
  }
265
 
266
  {% endif %}
@@ -290,13 +214,7 @@ fn main(
290
  let tid = lid.x;
291
  let hKv = h / (Q_HEADS / KV_HEADS);
292
  let cacheSeq = params.kvSeq;
293
- {% if cacheSeqlens %}
294
- // Buffer-sharing caches retain their capacity in the physical BNSH stride;
295
- // seqlens_k supplies the active end independently for each batch.
296
- let kvSeq = min(cacheSeq, u32(seqlens_k[b]) + 1u);
297
- {% else %}
298
  let kvSeq = cacheSeq;
299
- {% endif %}
300
 
301
  // Query row (decode uses token zero; short-query prefill folds the token into wg.x).
302
  {% if splitQueries %}
@@ -317,7 +235,6 @@ fn main(
317
  splitEnd = kvSeq;
318
  }
319
 
320
- {% set hasBias = hasBias is defined and hasBias %}
321
  for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
322
  var qv = vec4<f32>(query[qBaseV4 + d4]);
323
  {% if hasBias %}
@@ -357,7 +274,7 @@ fn main(
357
  if (j < tileCount) {
358
  let kRowV4 = kvBaseV4 + (kjBase + j) * kvTokenStrideV4;
359
  for (var d4: u32 = lane; d4 < HEAD_DIM_V4; d4 = d4 + sgSize) {
360
- accS = accS + dot(q_shared[d4], load_key4(kRowV4 + d4{% if quantizedCache %}, d4, hKv{% endif %}));
361
  }
362
  }
363
  let sj = subgroupAdd(accS);
@@ -376,7 +293,7 @@ fn main(
376
  let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
377
  var acc: f32 = 0.0;
378
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
379
- acc = acc + dot(q_shared[d4], load_key4(kRowV4 + d4{% if quantizedCache %}, d4, hKv{% endif %}));
380
  }
381
  score = acc * scale;
382
  m = score;
@@ -391,21 +308,10 @@ fn main(
391
  let maskQuery = 0u;
392
  {% endif %}
393
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
394
- {% if maskIsBool %}
395
- // A rejected bool-mask key contributes no probability mass. The merge
396
- // pass already maps a zero global denominator to an all-zero output row.
397
- if (attn_mask[maskIndex] == 0u) {
398
- keyAllowed = false;
399
- score = -FLT_MAX;
400
- dPart = 0.0;
401
- }
402
- {% else %}
403
  score = score + f32(attn_mask[maskIndex]);
404
- {% endif %}
405
  m = score;
406
  }
407
- {% endif %}
408
- let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
409
 
410
  // Merge one key tile's online-softmax (maximum, denominator) partial into the
411
  // running state, then store the per-key probabilities consumed by V accumulation.
@@ -421,7 +327,6 @@ fn main(
421
  probs[tid] = prob;
422
  workgroupBarrier();
423
 
424
-
425
  {% if jSplitV %}
426
  // j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
427
  // J_GROUPS for dim block d4v into a register, then the groups combine
@@ -434,10 +339,7 @@ fn main(
434
  loop {
435
  if (jj >= tileCount) { break; }
436
  vacc = vacc + probs[jj] * load_value4(
437
- kvBaseV4 + (kjBase + jj) * kvTokenStrideV4 + d4v{% if quantizedCache %},
438
- d4v,
439
- hKv{% endif %}
440
- );
441
  jj = jj + J_GROUPS;
442
  }
443
  vacc_sh[tid] = vacc;
@@ -455,10 +357,7 @@ fn main(
455
  var vSum = vec4<f32>(0.0);
456
  for (var i: u32 = 0u; i < tileCount; i = i + 1u) {
457
  vSum = vSum + probs[i] * load_value4(
458
- kvBaseV4 + (kjBase + i) * kvTokenStrideV4 + d4{% if quantizedCache %},
459
- d4,
460
- hKv{% endif %}
461
- );
462
  }
463
  running_out[d4] = running_out[d4] * correction + vSum;
464
  }
 
1
  {% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
2
  {% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
 
 
3
  {% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
 
 
 
4
  {% set qSeq = qSeq | default(0) %}
5
  {% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
6
  {% if useSubgroups %}
 
8
  {% endif %}
9
  {{ env.wgsl.resourceDeclarations }}
10
 
11
+ // Split-K flash attention. Single-partition direct output skips the merge pass.
12
+ // Otherwise this is the first of two passes. This geometry
13
  // handles decode and short-query, long-context prefill inputs.
14
  //
15
  // The non-split flash decode launches only `batch * numHeads` workgroups, each
 
52
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
53
  return select(value - maxValue, 0.0, equalFiniteMax);
54
  }
55
+
56
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
57
  return exp(shifted_value(value, maxValue));
58
  }
 
60
  var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
61
  var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
62
  var<workgroup> probs: array<f32, WG>;
63
+ {% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) and (allowCooperativeQk if allowCooperativeQk is defined else true) %}
64
  {% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
65
  {% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
66
  {% if coopQk %}
 
129
  return combinedMD;
130
  }
131
  {% else %}
132
+ {% set mdExtent = "WG" %}
 
 
133
  var<workgroup> partialM: array<f32, {{ mdExtent }}>;
134
  var<workgroup> partialD: array<f32, {{ mdExtent }}>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
135
 
136
  fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
137
  partialM[lidx] = m;
 
161
  return merged;
162
  }
163
  {% endif %}
 
 
164
 
165
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
166
  if (params.scale != 0.0) { return params.scale; }
167
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
168
  }
169
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
170
  fn load_key4(indexV4: u32) -> vec4<f32> {
171
  return vec4<f32>(key[indexV4]);
172
  }
 
174
  fn load_value4(indexV4: u32) -> vec4<f32> {
175
  return vec4<f32>(value[indexV4]);
176
  }
 
177
 
178
  {% if hasBias %}
179
  // Packed [Q; K; V] bias rows (token-independent). The Q bias folds into the
180
  // query row before the Q.K dots; the K bias adds a constant to every key score
181
  // that softmax cancels, so it is skipped; the V bias is token-independent and
182
  // is applied once in the merge pass after the final normalize.
183
+ {% set BW = "" %}
184
+ {% set BC = "" %}
185
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
186
  let offset = base + d4 * 4u;
187
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
188
  }
189
 
190
  {% endif %}
 
214
  let tid = lid.x;
215
  let hKv = h / (Q_HEADS / KV_HEADS);
216
  let cacheSeq = params.kvSeq;
 
 
 
 
 
217
  let kvSeq = cacheSeq;
 
218
 
219
  // Query row (decode uses token zero; short-query prefill folds the token into wg.x).
220
  {% if splitQueries %}
 
235
  splitEnd = kvSeq;
236
  }
237
 
 
238
  for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
239
  var qv = vec4<f32>(query[qBaseV4 + d4]);
240
  {% if hasBias %}
 
274
  if (j < tileCount) {
275
  let kRowV4 = kvBaseV4 + (kjBase + j) * kvTokenStrideV4;
276
  for (var d4: u32 = lane; d4 < HEAD_DIM_V4; d4 = d4 + sgSize) {
277
+ accS = accS + dot(q_shared[d4], load_key4(kRowV4 + d4));
278
  }
279
  }
280
  let sj = subgroupAdd(accS);
 
293
  let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
294
  var acc: f32 = 0.0;
295
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
296
+ acc = acc + dot(q_shared[d4], load_key4(kRowV4 + d4));
297
  }
298
  score = acc * scale;
299
  m = score;
 
308
  let maskQuery = 0u;
309
  {% endif %}
310
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
 
 
 
 
 
 
 
 
 
311
  score = score + f32(attn_mask[maskIndex]);
 
312
  m = score;
313
  }
314
+ {% endif %} let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
 
315
 
316
  // Merge one key tile's online-softmax (maximum, denominator) partial into the
317
  // running state, then store the per-key probabilities consumed by V accumulation.
 
327
  probs[tid] = prob;
328
  workgroupBarrier();
329
 
 
330
  {% if jSplitV %}
331
  // j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
332
  // J_GROUPS for dim block d4v into a register, then the groups combine
 
339
  loop {
340
  if (jj >= tileCount) { break; }
341
  vacc = vacc + probs[jj] * load_value4(
342
+ kvBaseV4 + (kjBase + jj) * kvTokenStrideV4 + d4v );
 
 
 
343
  jj = jj + J_GROUPS;
344
  }
345
  vacc_sh[tid] = vacc;
 
357
  var vSum = vec4<f32>(0.0);
358
  for (var i: u32 = 0u; i < tileCount; i = i + 1u) {
359
  vSum = vSum + probs[i] * load_value4(
360
+ kvBaseV4 + (kjBase + i) * kvTokenStrideV4 + d4 );
 
 
 
361
  }
362
  running_out[d4] = running_out[d4] * correction + vSum;
363
  }
build/webgpu/attn-flash-online.wgsl.jinja CHANGED
@@ -28,8 +28,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 +50,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,7 +64,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
@@ -116,41 +116,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
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;
@@ -180,22 +148,20 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
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 %}
@@ -221,7 +187,6 @@ fn main(
221
  let kvTokenStrideV4 = KV_HIDDEN_V4;
222
 
223
  // Cooperative vec4 Q-row load; init the output accumulator.
224
- {% set hasBias = hasBias is defined and hasBias %}
225
  for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
226
  var qv = vec4<f32>(query[qBaseV4 + d4]);
227
  {% if hasBias %}
@@ -297,7 +262,6 @@ fn main(
297
  probs[tid] = prob;
298
  workgroupBarrier();
299
 
300
-
301
  // running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
302
  // tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
303
  let tileCount = min(WG, keyBound - kjBase);
 
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 = "Q_HEADS" %}
32
+ {% set kvHeads = "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
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
51
  return select(value - maxValue, 0.0, equalFiniteMax);
52
  }
53
+
54
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
55
  return exp(shifted_value(value, maxValue));
56
  }
 
64
  // Both the subgroup and portable barrier-tree engines return the same merged
65
  // pair to every invocation. Repeated merges require a workgroup barrier between
66
  // calls before their shared partial storage is reused.
 
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
 
116
  return combinedMD;
117
  }
118
  {% else %}
119
+ {% set mdExtent = "WG" %}
 
 
120
  var<workgroup> partialM: array<f32, {{ mdExtent }}>;
121
  var<workgroup> partialD: array<f32, {{ mdExtent }}>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
122
 
123
  fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
124
  partialM[lidx] = m;
 
148
  return merged;
149
  }
150
  {% endif %}
 
 
151
 
152
  // An explicit-zero specialization bakes the scale as 0. Otherwise,
153
  // params.scale == 0 encodes an omitted scale and selects 1/sqrt(headDim).
154
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
155
  if (params.scale != 0.0) { return params.scale; }
156
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
157
  }
 
158
  {% if hasBias %}
159
 
160
+ {% set BW = "" %}
161
+ {% set BC = "" %}
162
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
163
  let offset = base + d4 * 4u;
164
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
165
  }
166
 
167
  {% endif %}
 
187
  let kvTokenStrideV4 = KV_HIDDEN_V4;
188
 
189
  // Cooperative vec4 Q-row load; init the output accumulator.
 
190
  for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
191
  var qv = vec4<f32>(query[qBaseV4 + d4]);
192
  {% if hasBias %}
 
262
  probs[tid] = prob;
263
  workgroupBarrier();
264
 
 
265
  // running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
266
  // tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
267
  let tileCount = min(WG, keyBound - kjBase);
build/webgpu/attn-flash-prefill-cluster.wgsl.jinja CHANGED
@@ -1,32 +1,27 @@
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
  {% set Q_HIDDEN = qHidden | default(0) %}
21
- {% if quantCacheFormat is not defined %}{% set quantCacheFormat = "" %}{% endif %}
22
- {% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
23
  // A windowed cache binds a fixed CAPACITY but keeps only the most recent
24
  // min(total, capacity) rows resident. params.kvSeq then names the physical row
25
  // count, which is still the right batch stride but the wrong attention bound, so
26
  // the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
27
  // in both roles.
28
- {% set KVA = "kvActive" if useSeqlens else KVSEQ %}
29
- {% macro score_expr(part) %}{% if hasSoftcap %}params.softcap * tanh(clamp(({{ part }} * SCALE) / params.softcap, -30.0, 30.0)){% else %}{{ part }} * SCALE{% endif %}{% endmacro %}
 
 
 
 
 
30
  {% if useSubgroups %}
31
  enable subgroups;
32
  {% endif %}
@@ -36,15 +31,30 @@ enable subgroups;
36
  // workgroup storage; values widen to f32 when read. Other specializations stage
37
  // them as f32. Score and weighted-value accumulation remain in f32 throughout.
38
  {% set STAGE_MASK = hasMask and useSubgroups and stageMask is defined and stageMask %}
39
- {% set MASK_IS_INT = hasMask and (maskIsKeyKeep or maskIsBool) %}
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
 
49
  // Tiled flash prefill attention with configurable-width query clusters for
50
  // token-major [batch, seq, heads*headDim] attention ops. Each workgroup covers
@@ -64,8 +74,7 @@ enable subgroups;
64
  const HEAD_DIM: u32 = {{ headDim }}u;
65
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
66
  const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
67
- {% if sourceProfile == 1 %}const QKV_STRIDE_V4: u32 = {{ qkvStrideV4 }}u; // packed [Q; K; V] input row stride
68
- {% endif %}const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
69
  const Q_HEADS: u32 = {{ qNumHeads }}u;
70
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
71
  const TILE_Q: u32 = {{ TILE_Q }}u;
@@ -79,24 +88,10 @@ const WG: u32 = (TILE_Q / QPL) * LPQ;
79
  {% else %}
80
  const WG: u32 = TILE_Q * LPQ;
81
  {% endif %}
82
- {% if MASK_IS_INT %}
83
- // Key-keep masks in contrib attention use a finite low logit for a rejected
84
- // key. Logical ONNX bool masks use the exclusion sentinel instead, so a fully
85
- // masked row has zero mass. Keeping one declaration shape for both mask modes
86
- // lets both mask modes share the score loop below.
87
  const NEG_INF: f32 = -3.4028234663852886e38;
88
- const MASK_NEG: f32 = {{ "-1e38" if maskIsKeyKeep else "-3.4028234663852886e38" }};
89
- {% else %}
90
- const NEG_INF: f32 = -3.4028234663852886e38;
91
- {% endif %}
92
 
93
  var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
94
  var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
95
- {% if STAGE_MASK %}
96
- // Each LPQ cluster consumes one mask value per (query,key), so stage the
97
- // TILE_Q x TILE_K mask tile once instead of issuing LPQ duplicate global loads.
98
- var<workgroup> mask_tile: array<{{ MASK_TILE_TYPE }}, TILE_Q * TILE_K>;
99
- {% endif %}
100
  {% if not useSubgroups %}
101
  {% if batchNoSgReduction %}
102
  // No-subgroups cluster reduction scratch for a whole K tile. Staging every
@@ -109,62 +104,17 @@ var<workgroup> red: array<f32, WG>;
109
  {% endif %}
110
  {% endif %}
111
 
112
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
113
- fn scale_value() -> f32 {
114
- {% if ATTN_SCALE_OVERRIDE is defined %}
115
- return {{ ATTN_SCALE_OVERRIDE }};
116
- {% else %}
117
  if (params.scale != 0.0) { return params.scale; }
118
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
119
- {% endif %}
120
- }
121
-
122
- {% if quantCacheFormat %}
123
- // A quantized cache is dequantized once per key into the staged tile, then read
124
- // by all TILE_Q queries in the workgroup. The unpack cost is amortized over the
125
- // tile height instead of paid once per (query, key).
126
- {% macro emit_quant_scale4(kind, scaleBuffer) %}
127
- fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
128
- if (params.perChannel == 0u) {
129
- return vec4<f32>({{ scaleBuffer }}[0]);
130
- }
131
- let base = hk * HEAD_DIM + d4 * 4u;
132
- return vec4<f32>(
133
- {{ scaleBuffer }}[base],
134
- {{ scaleBuffer }}[base + 1u],
135
- {{ scaleBuffer }}[base + 2u],
136
- {{ scaleBuffer }}[base + 3u]
137
- );
138
- }
139
- {%- endmacro %}
140
- {%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
141
- {{ emit_quant_scale4(kind, scaleBuffer) }}
142
- fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
143
- {%- if format == "int8" %}
144
- return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
145
- {%- else %}
146
- // Two elements cover this vec4: each carries two +8-biased nibbles, low first.
147
- let rowBase = indexV4 - d4;
148
- let lo = {{ buffer }}[rowBase + d4 * 2u];
149
- let hi = {{ buffer }}[rowBase + d4 * 2u + 1u];
150
- let nibbles = vec4<i32>(
151
- i32(lo & 0xFu), i32((lo >> 4u) & 0xFu),
152
- i32(hi & 0xFu), i32((hi >> 4u) & 0xFu)
153
- );
154
- let signed = nibbles - vec4<i32>(8);
155
- return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
156
- {%- endif %}
157
  }
158
- {%- endmacro %}
159
-
160
- {{ emit_quant_load4(quantCacheFormat, "key", KEY, "k_scale") }}
161
- {{ emit_quant_load4(quantCacheFormat, "value", VALUE, "v_scale") }}
162
- {% endif %}
163
 
164
  {% if hasBias %}
 
 
165
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
166
  let offset = base + d4 * 4u;
167
- return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
168
  }
169
 
170
  {% endif %}
@@ -212,17 +162,6 @@ fn main(
212
  var {{ qn("m", qi) }}: f32 = NEG_INF;
213
  var {{ qn("l", qi) }}: f32 = 0.0;
214
  {% endfor %}
215
- {% if useSeqlens %}
216
-
217
- // Resident rows of a windowed cache: the survivors were shifted down to [0, kvActive),
218
- // so query/key DISTANCE is unchanged and every bound below reads as if the cache were
219
- // exactly kvActive long. Rotary is excluded from this path (it would need the absolute
220
- // position, not the cache-relative one), so pastLenForRope keeps the physical length.
221
- // Clamping below by the query count matches the cache-update passes and the scalar
222
- // path: a right-padded batch (seqlens_k[b]+1 < qSeq) still appends its whole chunk,
223
- // so its queries score against all of it.
224
- let kvActive = min({{ KVSEQ }}, max({{ QSEQ }}, u32(seqlens_k[b]) + 1u));
225
- {% endif %}
226
  // Causal ceiling per query; the key loop runs over the workgroup's union range
227
  // (uniform trip count), masking out-of-range (query, key) pairs.
228
  // Upper-left causal and/or sliding-window bounds. Query qIdx sits at
@@ -260,33 +199,13 @@ fn main(
260
  let kj = kStart + slot;
261
  if (kj < wgEnd) {
262
  let base4 = kvBatch4 + kj * KV_HIDDEN_V4 + d4;
263
- {% if quantCacheFormat %}
264
- k_tile[i] = vec4<{{ ST }}>(load_key4(base4, d4, hKv));
265
- v_tile[i] = vec4<{{ ST }}>(load_value4(base4, d4, hKv));
266
- {% else %}
267
  k_tile[i] = vec4<{{ ST }}>({{ KEY }}[base4]);
268
  v_tile[i] = vec4<{{ ST }}>({{ VALUE }}[base4]);
269
- {% endif %}
270
  } else {
271
  k_tile[i] = vec4<{{ ST }}>(0.0);
272
  v_tile[i] = vec4<{{ ST }}>(0.0);
273
  }
274
  }
275
- {% if STAGE_MASK %}
276
- // The K/V load barrier also publishes this compact mask tile.
277
- for (var i: u32 = tid; i < TILE_Q * TILE_K; i = i + WG) {
278
- let qSlot = i / TILE_K;
279
- let kSlot = i % TILE_K;
280
- let maskQ = min(wg.x * TILE_Q + qSlot, {{ QSEQ }} - 1u);
281
- let maskK = kStart + kSlot;
282
- if (maskK < wgEnd) {
283
- let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQ * params.maskSeqStride + maskK;
284
- mask_tile[i] = {{ MASK_TILE_LOAD }};
285
- } else {
286
- mask_tile[i] = {{ MASK_TILE_ZERO }};
287
- }
288
- }
289
- {% endif %}
290
  workgroupBarrier();
291
  // TILE_K is a small shader constant; the loop updates the named q/o slices in place.
292
  {% if not useSubgroups and batchNoSgReduction %}
@@ -329,15 +248,7 @@ fn main(
329
  if (kj >= minKj && kj < maxKj) {
330
  {% if hasMask %}
331
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
332
- {% if maskIsBool %}
333
- if (attn_mask[maskIndex] != 0u) {
334
- s[kk] = {{ score_expr("part") }};
335
- } else {
336
- s[kk] = MASK_NEG;
337
- }
338
- {% else %}
339
  s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
340
- {% endif %}
341
  {% else %}
342
  s[kk] = {{ score_expr("part") }};
343
  {% endif %}
@@ -374,29 +285,11 @@ fn main(
374
  let sc = partv[{{ qi }}];
375
  if (kj >= {{ qn("minKj", qi) }} && kj < {{ qn("maxKj", qi) }}) {
376
  {% if hasMask %}
377
- {% if STAGE_MASK %}
378
- let maskValue = mask_tile[(qSub + {{ qi }}u) * TILE_K + kk];
379
- {% else %}
380
  // Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
381
  // [q, k] masks set batch/head strides to 0). {{ qn("qClamped", qi) }} keeps the seq index
382
  // in-bounds for padding queries in the last tile (their output is dropped).
383
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
384
- {% endif %}
385
- {% if maskIsKeyKeep %}
386
- // A broadcast key mask uses 1 for a retained key and 0 for padding.
387
- {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
388
- {% elif maskIsBool %}
389
- // Logical bool: a rejected key contributes no softmax mass. Leaving
390
- // the initialized NEG_INF sentinel in place makes a fully masked row
391
- // land on the zero-denominator output guard below.
392
- if ({{ MASK_ELEMENT }} != 0u) {
393
- {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
394
- } else {
395
- {{ qn("s", qi) }}[kk] = MASK_NEG;
396
- }
397
- {% else %}
398
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
399
- {% endif %}
400
  {% else %}
401
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
402
  {% endif %}
@@ -434,29 +327,11 @@ fn main(
434
  {% endif %}
435
  if (kj >= minKj && kj < maxKj) {
436
  {% if hasMask %}
437
- {% if STAGE_MASK %}
438
- let maskValue = mask_tile[qSub * TILE_K + kk];
439
- {% else %}
440
  // Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
441
  // [q, k] masks set batch/head strides to 0). qClamped keeps the seq index
442
  // in-bounds for padding queries in the last tile (their output is dropped).
443
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
444
- {% endif %}
445
- {% if maskIsKeyKeep %}
446
- // A broadcast key mask uses 1 for a retained key and 0 for padding.
447
- s[kk] = {{ score_expr("part") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
448
- {% elif maskIsBool %}
449
- // Logical bool: a rejected key contributes no softmax mass. Leaving
450
- // the initialized NEG_INF sentinel in place makes a fully masked row
451
- // land on the zero-denominator output guard below.
452
- if ({{ MASK_ELEMENT }} != 0u) {
453
- s[kk] = {{ score_expr("part") }};
454
- } else {
455
- s[kk] = MASK_NEG;
456
- }
457
- {% else %}
458
  s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
459
- {% endif %}
460
  {% else %}
461
  s[kk] = {{ score_expr("part") }};
462
  {% endif %}
@@ -464,24 +339,7 @@ fn main(
464
  }
465
  {% endif %}
466
 
467
- // Per-thread online softmax over the tile. s[kk] is reused to hold the
468
- // exponentiated probabilities for the PV accumulation below.
469
- {% for qi in range(QPL) %}
470
- var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
471
- for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
472
- {{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
473
- }
474
- let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
475
- let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
476
- var {{ qn("pSum", qi) }}: f32 = 0.0;
477
- for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
478
- let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
479
- {{ qn("s", qi) }}[kk] = pk;
480
- {{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
481
- }
482
- {{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
483
- {{ qn("m", qi) }} = {{ qn("newMax", qi) }};
484
- {% endfor %}
485
  // A boundary tile can address V rows outside a query's attended range, and
486
  // a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
487
  // NaN, the guarded loop selects the V operand away for range-excluded keys.
@@ -540,22 +398,11 @@ fn main(
540
  {% for qi in range(QPL) %}
541
  if ({{ qn("qValid", qi) }}) {
542
  let outBase4 = (b * {{ QSEQ }} + {{ qn("qIdx", qi) }}) * Q_HIDDEN_V4 + h * HEAD_DIM_V4 + lane8 * SLICE;
543
- {% if hasHeadSink %}
544
- // The head sink is a learned logit that competes with the keys but carries
545
- // no value, so it enters the denominator only and the weighted sum above is
546
- // untouched. Renormalizing against max({{ qn("m", qi) }}, sink) keeps the exponentials in
547
- // range when the sink dominates a fully-masked row.
548
- let sink = f32(head_sink[h]);
549
- let finalM = max({{ qn("m", qi) }}, sink);
550
- let accScale = exp({{ qn("m", qi) }} - finalM);
551
- let inv = accScale / (exp(sink - finalM) + {{ qn("l", qi) }} * accScale);
552
- {% else %}
553
  // {{ qn("l", qi) }} == 0 means this query had no probability-bearing key: either its
554
  // causal/window range is empty or its logical bool mask rejects every key.
555
  // Emit 0 rather than 0/0. Other attention paths enforce the same empty-row
556
  // contract by selecting on positive global mass.
557
  let inv = select(0.0, 1.0 / {{ qn("l", qi) }}, {{ qn("l", qi) }} > 0.0);
558
- {% endif %}
559
  {% for c in range(SLICE_COUNT) %}
560
  {{ OUTPUT }}[outBase4 + {{ c }}u] = vec4<{{ scalar }}>({{ attention_value(c, qi) }});
561
  {% endfor %}
 
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
  {% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
 
 
 
 
12
  {% set Q_HIDDEN = qHidden | default(0) %}
 
 
13
  // A windowed cache binds a fixed CAPACITY but keeps only the most recent
14
  // min(total, capacity) rows resident. params.kvSeq then names the physical row
15
  // count, which is still the right batch stride but the wrong attention bound, so
16
  // the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
17
  // in both roles.
18
+ {% if hasKeyLimit is not defined %}{% set hasKeyLimit = false %}{% endif %}
19
+ {% if hasKeyLimit %}
20
+ {% set KVA = "select(params.kvSeq, min(params.keyLimit, params.kvSeq), params.keyLimit > 0u)" %}
21
+ {% else %}
22
+ {% set KVA = KVSEQ %}
23
+ {% endif %}
24
+ {% macro score_expr(part) %}{{ part }} * SCALE{% endmacro %}
25
  {% if useSubgroups %}
26
  enable subgroups;
27
  {% endif %}
 
31
  // workgroup storage; values widen to f32 when read. Other specializations stage
32
  // them as f32. Score and weighted-value accumulation remain in f32 throughout.
33
  {% set STAGE_MASK = hasMask and useSubgroups and stageMask is defined and stageMask %}
 
 
 
 
 
34
  {% set MASK_ADDITIVE = "maskValue" if STAGE_MASK else "f32(attn_mask[maskIndex])" %}
35
  {% set SLICE_COUNT = ((headDimV4 / LPQ) | int) %}
36
  {% 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 %}
37
  {% macro qn(name, qi) %}{{ name }}{% if qi > 0 %}_q{{ qi }}{% endif %}{% endmacro %}
38
+ {% macro emit_tile_softmax() %}
39
+ // Per-thread online softmax over the tile. s[kk] is reused to hold the
40
+ // exponentiated probabilities for the PV accumulation below.
41
+ {% for qi in range(QPL) %}
42
+ var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
43
+ for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
44
+ {{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
45
+ }
46
+ let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
47
+ let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
48
+ var {{ qn("pSum", qi) }}: f32 = 0.0;
49
+ for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
50
+ let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
51
+ {{ qn("s", qi) }}[kk] = pk;
52
+ {{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
53
+ }
54
+ {{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
55
+ {{ qn("m", qi) }} = {{ qn("newMax", qi) }};
56
+ {% endfor %}
57
+ {% endmacro %}
58
 
59
  // Tiled flash prefill attention with configurable-width query clusters for
60
  // token-major [batch, seq, heads*headDim] attention ops. Each workgroup covers
 
74
  const HEAD_DIM: u32 = {{ headDim }}u;
75
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
76
  const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
77
+ const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
 
78
  const Q_HEADS: u32 = {{ qNumHeads }}u;
79
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
80
  const TILE_Q: u32 = {{ TILE_Q }}u;
 
88
  {% else %}
89
  const WG: u32 = TILE_Q * LPQ;
90
  {% endif %}
 
 
 
 
 
91
  const NEG_INF: f32 = -3.4028234663852886e38;
 
 
 
 
92
 
93
  var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
94
  var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
 
 
 
 
 
95
  {% if not useSubgroups %}
96
  {% if batchNoSgReduction %}
97
  // No-subgroups cluster reduction scratch for a whole K tile. Staging every
 
104
  {% endif %}
105
  {% endif %}
106
 
107
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
 
 
 
108
  if (params.scale != 0.0) { return params.scale; }
109
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
110
  }
 
 
 
 
 
111
 
112
  {% if hasBias %}
113
+ {% set BW = "" %}
114
+ {% set BC = "" %}
115
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
116
  let offset = base + d4 * 4u;
117
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
118
  }
119
 
120
  {% endif %}
 
162
  var {{ qn("m", qi) }}: f32 = NEG_INF;
163
  var {{ qn("l", qi) }}: f32 = 0.0;
164
  {% endfor %}
 
 
 
 
 
 
 
 
 
 
 
165
  // Causal ceiling per query; the key loop runs over the workgroup's union range
166
  // (uniform trip count), masking out-of-range (query, key) pairs.
167
  // Upper-left causal and/or sliding-window bounds. Query qIdx sits at
 
199
  let kj = kStart + slot;
200
  if (kj < wgEnd) {
201
  let base4 = kvBatch4 + kj * KV_HIDDEN_V4 + d4;
 
 
 
 
202
  k_tile[i] = vec4<{{ ST }}>({{ KEY }}[base4]);
203
  v_tile[i] = vec4<{{ ST }}>({{ VALUE }}[base4]);
 
204
  } else {
205
  k_tile[i] = vec4<{{ ST }}>(0.0);
206
  v_tile[i] = vec4<{{ ST }}>(0.0);
207
  }
208
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
209
  workgroupBarrier();
210
  // TILE_K is a small shader constant; the loop updates the named q/o slices in place.
211
  {% if not useSubgroups and batchNoSgReduction %}
 
248
  if (kj >= minKj && kj < maxKj) {
249
  {% if hasMask %}
250
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
 
 
 
 
 
 
 
251
  s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
 
252
  {% else %}
253
  s[kk] = {{ score_expr("part") }};
254
  {% endif %}
 
285
  let sc = partv[{{ qi }}];
286
  if (kj >= {{ qn("minKj", qi) }} && kj < {{ qn("maxKj", qi) }}) {
287
  {% if hasMask %}
 
 
 
288
  // Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
289
  // [q, k] masks set batch/head strides to 0). {{ qn("qClamped", qi) }} keeps the seq index
290
  // in-bounds for padding queries in the last tile (their output is dropped).
291
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
292
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
 
293
  {% else %}
294
  {{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
295
  {% endif %}
 
327
  {% endif %}
328
  if (kj >= minKj && kj < maxKj) {
329
  {% if hasMask %}
 
 
 
330
  // Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
331
  // [q, k] masks set batch/head strides to 0). qClamped keeps the seq index
332
  // in-bounds for padding queries in the last tile (their output is dropped).
333
  let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
334
  s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
 
335
  {% else %}
336
  s[kk] = {{ score_expr("part") }};
337
  {% endif %}
 
339
  }
340
  {% endif %}
341
 
342
+ {{ emit_tile_softmax() }}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
343
  // A boundary tile can address V rows outside a query's attended range, and
344
  // a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
345
  // NaN, the guarded loop selects the V operand away for range-excluded keys.
 
398
  {% for qi in range(QPL) %}
399
  if ({{ qn("qValid", qi) }}) {
400
  let outBase4 = (b * {{ QSEQ }} + {{ qn("qIdx", qi) }}) * Q_HIDDEN_V4 + h * HEAD_DIM_V4 + lane8 * SLICE;
 
 
 
 
 
 
 
 
 
 
401
  // {{ qn("l", qi) }} == 0 means this query had no probability-bearing key: either its
402
  // causal/window range is empty or its logical bool mask rejects every key.
403
  // Emit 0 rather than 0/0. Other attention paths enforce the same empty-row
404
  // contract by selecting on positive global mass.
405
  let inv = select(0.0, 1.0 / {{ qn("l", qi) }}, {{ qn("l", qi) }} > 0.0);
 
406
  {% for c in range(SLICE_COUNT) %}
407
  {{ OUTPUT }}[outBase4 + {{ c }}u] = vec4<{{ scalar }}>({{ attention_value(c, qi) }});
408
  {% endfor %}
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,11 +11,16 @@
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
  {% set CAUSAL = false if (hasCausal is defined and not hasCausal) else true %}
 
15
  {% if USE_SUBGROUPS %}
16
  enable subgroups;
17
  {% endif %}
 
 
 
18
  {{ env.wgsl.resourceDeclarations }}
19
 
20
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
@@ -46,15 +52,16 @@ fn scale_value() -> f32 {
46
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
47
  }
48
 
49
-
50
  {% if hasBias %}
 
 
51
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
52
  let offset = base + d4 * 4u;
53
- return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
54
  }
55
  {% endif %}
56
 
57
- @compute @workgroup_size({{ Q_STEP }}, 1, 1)
58
  fn main(
59
  @builtin(workgroup_id) wg: vec3<u32>,
60
  @builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
@@ -99,9 +106,14 @@ fn main(
99
  // so every lane shares a uniform trip count (subgroup ops stay reconverged);
100
  // each lane masks its own keys past myMaxKj to NEG_INF.
101
  {% if CAUSAL %}
 
 
 
 
 
102
  let lastQ = min(wg.x * Q_STEP + Q_STEP - 1u, params.qSeq - 1u);
103
- let kvEnd = select(params.kvSeq, min(lastQ + 1u, params.kvSeq), params.isCausal != 0u);
104
- let myMaxKj = select(params.kvSeq, min(qi + 1u, params.kvSeq), params.isCausal != 0u);
105
  {% else %}
106
  let kvEnd = params.kvSeq;
107
  let myMaxKj = params.kvSeq;
@@ -193,6 +205,13 @@ fn main(
193
  {% endfor %}
194
  previous_max = new_max;
195
  previous_denom = denom;
 
 
 
 
 
 
 
196
 
197
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
198
  {% if USE_SUBGROUPS %}
@@ -209,21 +228,37 @@ fn main(
209
  }
210
  {% endif %}
211
  {% endif %}
 
 
 
 
212
  var acc: vec4<f32> = vec4<f32>(0.0);
 
213
  {% for g in range(qkGroups) %}
214
  {% for lane in range(4) %}
215
- {% if USE_SUBGROUPS %}
216
- {% if g < 8 %}
217
- acc = acc + vec4<f32>(subgroupShuffle(v_local0, {{ g * 4 + lane }}u)) * qk{{ g }}.{{ components[lane] }};
 
218
  {% else %}
219
- acc = acc + vec4<f32>(subgroupShuffle(v_local1, {{ (g - 8) * 4 + lane }}u)) * qk{{ g }}.{{ components[lane] }};
220
  {% endif %}
 
 
221
  {% else %}
222
- acc = acc + vec4<f32>(valueTile[d4 * K_STEP + {{ g * 4 + lane }}u]) * qk{{ g }}.{{ components[lane] }};
223
  {% endif %}
224
  {% endfor %}
 
 
 
 
225
  {% endfor %}
 
 
 
226
  o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
 
227
  }
228
  {% if not USE_SUBGROUPS %}
229
  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
+ {% set qHidden = qHidden | default(0) %}
18
  {% if USE_SUBGROUPS %}
19
  enable subgroups;
20
  {% endif %}
21
+ {% if PIN_SUBGROUP_32 %}
22
+ enable subgroup_size_control;
23
+ {% endif %}
24
  {{ env.wgsl.resourceDeclarations }}
25
 
26
  const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
 
52
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
53
  }
54
 
 
55
  {% if hasBias %}
56
+ {% set BW = "" %}
57
+ {% set BC = "" %}
58
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
59
  let offset = base + d4 * 4u;
60
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
61
  }
62
  {% endif %}
63
 
64
+ @compute @workgroup_size({{ Q_STEP }}, 1, 1){{ " @subgroup_size(32)" if PIN_SUBGROUP_32 else "" }}
65
  fn main(
66
  @builtin(workgroup_id) wg: vec3<u32>,
67
  @builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
 
106
  // so every lane shares a uniform trip count (subgroup ops stay reconverged);
107
  // each lane masks its own keys past myMaxKj to NEG_INF.
108
  {% if CAUSAL %}
109
+ {% set FIXED_CAUSAL = false %}
110
+ {% macro key_ceiling(query) %}
111
+ {% if FIXED_CAUSAL %}
112
+ min({{ query }} + 1u, params.kvSeq){% else %}
113
+ select(params.kvSeq, min({{ query }} + 1u, params.kvSeq), params.isCausal != 0u){% endif %}{% endmacro %}
114
  let lastQ = min(wg.x * Q_STEP + Q_STEP - 1u, params.qSeq - 1u);
115
+ let kvEnd = {{ key_ceiling("lastQ") }};
116
+ let myMaxKj = {{ key_ceiling("qi") }};
117
  {% else %}
118
  let kvEnd = params.kvSeq;
119
  let myMaxKj = params.kvSeq;
 
205
  {% endfor %}
206
  previous_max = new_max;
207
  previous_denom = denom;
208
+ {% set PV_HALF = ST == "f16" and not (precisePv | default(false)) %}
209
+ {% set PV_FLUSH_GROUPS = 4 %}
210
+ {% if PV_HALF %}
211
+ {% for g in range(qkGroups) %}
212
+ let p{{ g }} = vec4<f16>(qk{{ g }});
213
+ {% endfor %}
214
+ {% endif %}
215
 
216
  for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
217
  {% if USE_SUBGROUPS %}
 
228
  }
229
  {% endif %}
230
  {% endif %}
231
+ {% if PV_HALF %}
232
+ var acc_f32: vec4<f32> = vec4<f32>(0.0);
233
+ var acc: vec4<f16> = vec4<f16>(0.0);
234
+ {% else %}
235
  var acc: vec4<f32> = vec4<f32>(0.0);
236
+ {% endif %}
237
  {% for g in range(qkGroups) %}
238
  {% for lane in range(4) %}
239
+ {% if not USE_SUBGROUPS %}
240
+ {% set pvValue = "valueTile[d4 * K_STEP + " ~ (g * 4 + lane) ~ "u]" %}
241
+ {% elif g < 8 %}
242
+ {% set pvValue = "subgroupShuffle(v_local0, " ~ (g * 4 + lane) ~ "u)" %}
243
  {% else %}
244
+ {% set pvValue = "subgroupShuffle(v_local1, " ~ ((g - 8) * 4 + lane) ~ "u)" %}
245
  {% endif %}
246
+ {% if PV_HALF %}
247
+ acc = fma({{ pvValue }}, vec4<f16>(p{{ g }}.{{ components[lane] }}), acc);
248
  {% else %}
249
+ acc = acc + vec4<f32>({{ pvValue }}) * qk{{ g }}.{{ components[lane] }};
250
  {% endif %}
251
  {% endfor %}
252
+ {% if PV_HALF and (g + 1) % PV_FLUSH_GROUPS == 0 %}
253
+ acc_f32 = acc_f32 + vec4<f32>(acc);
254
+ acc = vec4<f16>(0.0);
255
+ {% endif %}
256
  {% endfor %}
257
+ {% if PV_HALF %}
258
+ o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc_f32;
259
+ {% else %}
260
  o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
261
+ {% endif %}
262
  }
263
  {% if not USE_SUBGROUPS %}
264
  workgroupBarrier();
build/webgpu/attn-materialized-apply-f32.wgsl.jinja CHANGED
@@ -1,3 +1,4 @@
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // Register-blocked normalized-scores @ V GEMM with BHSD and BSH specializations.
@@ -39,7 +40,7 @@ var<workgroup> rowMax: array<f32, BM>;
39
 
40
  {% endif %}
41
  fn v_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
42
- return value[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d];
43
  }
44
 
45
  fn store_y(b: u32, h: u32, q: u32, d: u32, result: f32) {
@@ -48,7 +49,7 @@ fn store_y(b: u32, h: u32, q: u32, d: u32, result: f32) {
48
  {% else %}
49
  let biased = result;
50
  {% endif %}
51
- output[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d] = biased;
52
  }
53
 
54
  @compute @workgroup_size(WG_DIM, WG_DIM, 1)
 
1
+ {% set STORAGE_F16 = false %}
2
  {{ env.wgsl.resourceDeclarations }}
3
 
4
  // Register-blocked normalized-scores @ V GEMM with BHSD and BSH specializations.
 
40
 
41
  {% endif %}
42
  fn v_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
43
+ return {{ "f32(" if STORAGE_F16 else "" }}value[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d]{{ ")" if STORAGE_F16 else "" }};
44
  }
45
 
46
  fn store_y(b: u32, h: u32, q: u32, d: u32, result: f32) {
 
49
  {% else %}
50
  let biased = result;
51
  {% endif %}
52
+ output[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d] = {{ "f16(" if STORAGE_F16 else "" }}biased{{ ")" if STORAGE_F16 else "" }};
53
  }
54
 
55
  @compute @workgroup_size(WG_DIM, WG_DIM, 1)
build/webgpu/attn-materialized-rowstats-combine-f32.wgsl.jinja CHANGED
@@ -13,6 +13,12 @@
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
  {% set MAX_ONLY = maxOnly is defined and maxOnly %}
 
 
 
 
 
 
16
  {{ env.wgsl.resourceDeclarations }}
17
 
18
  const SLOTS: u32 = {{ statSlots }}u;
@@ -41,21 +47,19 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
41
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
42
  return select(value - maxValue, 0.0, equalFiniteMax);
43
  }
44
- {%- endif %}
45
  {% if stableUsage == "all" %}
46
 
47
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
48
  return exp(shifted_value(value, maxValue));
49
  }
50
- {%- endif %}
51
-
52
 
53
  @compute @workgroup_size(WG, 1, 1)
54
  fn main(
55
  @builtin(global_invocation_id) gid: vec3<u32>
56
  ) {
57
- let row = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
58
- if (row >= params.rows) { return; }
59
 
60
  // `row` already runs over (batch, head, query) together, and the partial
61
  // layout puts that same product one axis out from the slot, so the stride
 
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
  {% set MAX_ONLY = maxOnly is defined and maxOnly %}
16
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
17
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
18
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
19
+ // per-axis workgroup fold width.
20
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
21
+ if ({{ name }} >= {{ bound }}) { return; }{% endmacro %}
22
  {{ env.wgsl.resourceDeclarations }}
23
 
24
  const SLOTS: u32 = {{ statSlots }}u;
 
47
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
48
  return select(value - maxValue, 0.0, equalFiniteMax);
49
  }
50
+ {% endif %}
51
  {% if stableUsage == "all" %}
52
 
53
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
54
  return exp(shifted_value(value, maxValue));
55
  }
56
+ {% endif %}
 
57
 
58
  @compute @workgroup_size(WG, 1, 1)
59
  fn main(
60
  @builtin(global_invocation_id) gid: vec3<u32>
61
  ) {
62
+ {{ flat_index_2d("WG", "row", "params.rows", guardInline=true) }}
 
63
 
64
  // `row` already runs over (batch, head, query) together, and the partial
65
  // layout puts that same product one axis out from the slot, so the stride
build/webgpu/attn-materialized-score-f32.wgsl.jinja CHANGED
@@ -1,3 +1,4 @@
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // f32 prefill score GEMM with BHSD and packed BSH specializations. It computes
@@ -8,22 +9,14 @@
8
  // for TM * TN * 4 fused multiply-adds. K keeps a one-word shared stride
9
  // padding, which changes the transposed workgroup-memory stride and reduces
10
  // bank-conflict risk on banked implementations.
11
- {% if scalarAccumulators is not defined %}{% set scalarAccumulators = false %}{% endif %}
12
- {% if maskIsKeyKeep is not defined %}{% set maskIsKeyKeep = false %}{% endif %}
13
- {% if scalarAccumulators %}
14
- {% set TM_JINJA = (materializedQueryTile / materializedWorkgroupDim)|int %}
15
- {% set TN_JINJA = (materializedKeyTile / materializedWorkgroupDim)|int %}
16
- {% endif %}
17
  {% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
 
18
 
19
  const Q_HEADS: u32 = {{ qNumHeads }}u;
20
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
21
  const HEAD_DIM: u32 = {{ headDim }}u;
22
  const Q_HIDDEN: u32 = Q_HEADS * HEAD_DIM;
23
  const KV_HIDDEN: u32 = KV_HEADS * HEAD_DIM;
24
- {% if maskIsKeyKeep %}
25
- const NEG_INF: f32 = -3.4028234663852886e38;
26
- {% endif %}
27
 
28
  const BK: u32 = {{ materializedInnerTile }}u;
29
  const BM: u32 = {{ materializedQueryTile }}u;
@@ -45,46 +38,20 @@ const STAT_SLOTS: u32 = {{ statSlots }}u;
45
  var<workgroup> statScratch: array<f32, WG_THREADS>;
46
  {% endif %}
47
 
48
- {% if EMIT_ROW_STATS %}{% set stableHelperUsage = "constant" %}
49
- {% set stableUsage = stableHelperUsage if stableHelperUsage is defined else "all" %}
50
- // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
51
  // `m - m` finite so an empty lane / all--inf row contributes the exact
52
  // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
53
  // a zero final denominator according to their public semantics. Using -inf
54
  // here changes +inf-row behavior.
55
  const FLT_MAX: f32 = 3.4028234663852886e38;
56
- {% if stableUsage != "constant" %}
57
-
58
- fn is_finite_f32(value: f32) -> bool {
59
- return select(false, value <= FLT_MAX, value >= -FLT_MAX);
60
- }
61
-
62
- // x - m that is exactly 0 when x equals a finite m, so exp(shifted) == 1
63
- // exactly at the row max. `x - x` on an infinite max is a legal fast-math
64
- // fold to 0, which would silently turn +inf rows finite — the explicit
65
- // equality test keeps the NaN propagation of the serial kernels.
66
- fn shifted_value(value: f32, maxValue: f32) -> f32 {
67
- let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
68
- return select(value - maxValue, 0.0, equalFiniteMax);
69
- }
70
- {%- endif %}
71
- {% if stableUsage == "all" %}
72
-
73
- fn exp_shift(value: f32, maxValue: f32) -> f32 {
74
- return exp(shifted_value(value, maxValue));
75
- }
76
- {%- endif %}
77
-
78
  {% endif %}
79
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
80
- fn scale_value() -> f32 {
81
  if (params.scale != 0.0) { return params.scale; }
82
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
83
  }
84
 
85
-
86
  fn q_at(b: u32, h: u32, q: u32, d: u32) -> f32 {
87
- var value = query[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d];
88
  {% if hasBias %}
89
  value = value + bias[h * HEAD_DIM + d];
90
  {% endif %}
@@ -92,7 +59,7 @@ fn q_at(b: u32, h: u32, q: u32, d: u32) -> f32 {
92
  }
93
 
94
  fn k_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
95
- return key[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d];
96
  }
97
 
98
  @compute @workgroup_size(WG_DIM, WG_DIM, 1)
@@ -107,17 +74,8 @@ fn main(
107
  let kBase = wg.x * BN;
108
  let li = lid.y * WG_DIM + lid.x;
109
 
110
- {% if scalarAccumulators %}
111
- // Keep every register-tile cell statically addressable.
112
- {% for i in range(TM_JINJA) %}
113
- {% for j in range(TN_JINJA) %}
114
- var acc{{ i }}_{{ j }}: f32 = 0.0;
115
- {% endfor %}
116
- {% endfor %}
117
- {% else %}
118
  var acc: array<f32, TM * TN>;
119
  for (var i = 0u; i < TM * TN; i = i + 1u) { acc[i] = 0.0; }
120
- {% endif %}
121
 
122
  for (var dBase = 0u; dBase < HEAD_DIM; dBase = dBase + BK) {
123
  // Both source tiles are loaded in their physical row-major direction.
@@ -152,19 +110,6 @@ fn main(
152
  let qr0 = lid.y * TM;
153
  let kc0 = lid.x * TN;
154
  for (var dv = 0u; dv < D_VECS; dv = dv + 1u) {
155
- {% if scalarAccumulators %}
156
- {% for i in range(TM_JINJA) %}
157
- let qv{{ i }} = tileQ[qr0 + {{ i }}u][dv];
158
- {% endfor %}
159
- {% for j in range(TN_JINJA) %}
160
- let kv{{ j }} = tileK[kc0 + {{ j }}u][dv];
161
- {% endfor %}
162
- {% for i in range(TM_JINJA) %}
163
- {% for j in range(TN_JINJA) %}
164
- acc{{ i }}_{{ j }} = acc{{ i }}_{{ j }} + dot(qv{{ i }}, kv{{ j }});
165
- {% endfor %}
166
- {% endfor %}
167
- {% else %}
168
  var qv: array<vec4<f32>, TM>;
169
  var kv: array<vec4<f32>, TN>;
170
  for (var i = 0u; i < TM; i = i + 1u) { qv[i] = tileQ[qr0 + i][dv]; }
@@ -174,7 +119,6 @@ fn main(
174
  acc[i * TN + j] = acc[i * TN + j] + dot(qv[i], kv[j]);
175
  }
176
  }
177
- {% endif %}
178
  }
179
  workgroupBarrier();
180
  }
@@ -189,28 +133,6 @@ fn main(
189
  var rowM: array<f32, TM>;
190
  for (var i = 0u; i < TM; i = i + 1u) { rowM[i] = -FLT_MAX; }
191
  {% endif %}
192
- {% if scalarAccumulators %}
193
- {% for i in range(TM_JINJA) %}
194
- {
195
- let qi = qi0 + {{ i }}u;
196
- if (qi < params.qSeq) {
197
- {% for j in range(TN_JINJA) %}
198
- let ki{{ j }} = ki0 + {{ j }}u;
199
- if (ki{{ j }} < params.kvSeq) {
200
- var score = acc{{ i }}_{{ j }} * scale;
201
- {% if maskIsKeyKeep %}
202
- score = score + (1.0 - f32(attn_mask[b * params.kvSeq + ki{{ j }}])) * NEG_INF;
203
- {% endif %}
204
- scores[scoreBase + qi * params.kvSeq + ki{{ j }}] = score;
205
- {% if EMIT_ROW_STATS %}
206
- rowM[{{ i }}u] = max(rowM[{{ i }}u], score);
207
- {% endif %}
208
- }
209
- {% endfor %}
210
- }
211
- }
212
- {% endfor %}
213
- {% else %}
214
  for (var i = 0u; i < TM; i = i + 1u) {
215
  let qi = qi0 + i;
216
  if (qi >= params.qSeq) { continue; }
@@ -225,7 +147,6 @@ fn main(
225
  }
226
  }
227
  }
228
- {% endif %}
229
 
230
  {% if EMIT_ROW_STATS %}
231
  // Fold each row's partial (m, d) across the WG_DIM threads that own its columns. Every lane
 
1
+ {% set STORAGE_F16 = false %}
2
  {{ env.wgsl.resourceDeclarations }}
3
 
4
  // f32 prefill score GEMM with BHSD and packed BSH specializations. It computes
 
9
  // for TM * TN * 4 fused multiply-adds. K keeps a one-word shared stride
10
  // padding, which changes the transposed workgroup-memory stride and reduces
11
  // bank-conflict risk on banked implementations.
 
 
 
 
 
 
12
  {% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
13
+ {% set statSlots = statSlots | default(0) %}
14
 
15
  const Q_HEADS: u32 = {{ qNumHeads }}u;
16
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
17
  const HEAD_DIM: u32 = {{ headDim }}u;
18
  const Q_HIDDEN: u32 = Q_HEADS * HEAD_DIM;
19
  const KV_HIDDEN: u32 = KV_HEADS * HEAD_DIM;
 
 
 
20
 
21
  const BK: u32 = {{ materializedInnerTile }}u;
22
  const BM: u32 = {{ materializedQueryTile }}u;
 
38
  var<workgroup> statScratch: array<f32, WG_THREADS>;
39
  {% endif %}
40
 
41
+ {% if EMIT_ROW_STATS %}// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
 
 
42
  // `m - m` finite so an empty lane / all--inf row contributes the exact
43
  // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
44
  // a zero final denominator according to their public semantics. Using -inf
45
  // here changes +inf-row behavior.
46
  const FLT_MAX: f32 = 3.4028234663852886e38;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
47
  {% endif %}
48
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
49
  if (params.scale != 0.0) { return params.scale; }
50
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
51
  }
52
 
 
53
  fn q_at(b: u32, h: u32, q: u32, d: u32) -> f32 {
54
+ var value = {{ "f32(" if STORAGE_F16 else "" }}query[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d]{{ ")" if STORAGE_F16 else "" }};
55
  {% if hasBias %}
56
  value = value + bias[h * HEAD_DIM + d];
57
  {% endif %}
 
59
  }
60
 
61
  fn k_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
62
+ return {{ "f32(" if STORAGE_F16 else "" }}key[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d]{{ ")" if STORAGE_F16 else "" }};
63
  }
64
 
65
  @compute @workgroup_size(WG_DIM, WG_DIM, 1)
 
74
  let kBase = wg.x * BN;
75
  let li = lid.y * WG_DIM + lid.x;
76
 
 
 
 
 
 
 
 
 
77
  var acc: array<f32, TM * TN>;
78
  for (var i = 0u; i < TM * TN; i = i + 1u) { acc[i] = 0.0; }
 
79
 
80
  for (var dBase = 0u; dBase < HEAD_DIM; dBase = dBase + BK) {
81
  // Both source tiles are loaded in their physical row-major direction.
 
110
  let qr0 = lid.y * TM;
111
  let kc0 = lid.x * TN;
112
  for (var dv = 0u; dv < D_VECS; dv = dv + 1u) {
 
 
 
 
 
 
 
 
 
 
 
 
 
113
  var qv: array<vec4<f32>, TM>;
114
  var kv: array<vec4<f32>, TN>;
115
  for (var i = 0u; i < TM; i = i + 1u) { qv[i] = tileQ[qr0 + i][dv]; }
 
119
  acc[i * TN + j] = acc[i * TN + j] + dot(qv[i], kv[j]);
120
  }
121
  }
 
122
  }
123
  workgroupBarrier();
124
  }
 
133
  var rowM: array<f32, TM>;
134
  for (var i = 0u; i < TM; i = i + 1u) { rowM[i] = -FLT_MAX; }
135
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
136
  for (var i = 0u; i < TM; i = i + 1u) {
137
  let qi = qi0 + i;
138
  if (qi >= params.qSeq) { continue; }
 
147
  }
148
  }
149
  }
 
150
 
151
  {% if EMIT_ROW_STATS %}
152
  // Fold each row's partial (m, d) across the WG_DIM threads that own its columns. Every lane
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,7 +7,6 @@ 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") %}
@@ -19,27 +16,21 @@ diagnostic(off, chromium.subgroup_matrix_uniformity);
19
  {% 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 %}
20
  {% set OUT_ROW_STRIDE = "HEAD_DIM" if headMajor else "HIDDEN" %}
21
  {% set KV_ROW_STRIDE = "HEAD_DIM" if kvHeadMajor else "KV_HIDDEN" %}
 
22
  {% set scorePhase = phase == "score" %}
23
  {% set FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
24
  {% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
25
  {% macro score_value(index, guard) %}
26
- {% if FUSED_SOFTMAX %}
27
- {% if PRIVATE_ROW_STATS %}
28
- select(0.0, exp_shift(scores[{{ index }}], private_softmax_m) / private_softmax_d, {{ guard[1] }})
29
- {%- else %}
30
- select(0.0, exp_shift(scores[{{ index }}], softmax_m[{{ guard[0] }}]) / softmax_d[{{ guard[0] }}], {{ guard[1] }})
31
- {%- endif %}
32
- {% else %}
33
- select(0.0, scores[{{ index }}], {{ guard[1] }})
34
- {%- endif %}
35
- {% endmacro %}
36
  {% set TILE_M_VALUE = materializedSgmatQueryTile %}
37
  {% set TILE_N_VALUE = materializedSgmatKeyTile %}
38
  {% set TILE_K_VALUE = materializedSgmatInnerTile %}
39
- {% set SUB_ROWS_VALUE = materializedSgmatSubgroupTileRows if materializedSgmatSubgroupTileRows is defined else 16 %}
40
- {% set SUB_COLS_VALUE = materializedSgmatSubgroupTileCols if materializedSgmatSubgroupTileCols is defined else 32 %}
41
- {% set ROW_BLOCKS = (SUB_ROWS_VALUE / 8)|int %}
42
- {% set COL_BLOCKS = (SUB_COLS_VALUE / 8)|int %}
43
  {% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
44
  {% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
45
  {% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
@@ -47,15 +38,16 @@ select(0.0, scores[{{ index }}], {{ guard[1] }})
47
  {% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
48
  {% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
49
  {% set DIRECT_SCORE_STORE = materializedSgmatDirectScoreStore and not EMIT_ROW_STATS %}
50
- {% set DIRECT_APPLY_STORE = materializedSgmatDirectApplyStore and not hasBias and MT == "f32" %}
51
  {% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
52
  {% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
53
  and materializedSgmatRuntimeDirectStore
54
- and (scorePhase or (not hasBias and MT == "f32"))
55
  and not EMIT_ROW_STATS %}
56
  {% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
57
  {% set SCALE_IN_Q = DIRECT_SCORE_STORE or (RUNTIME_DIRECT_STORE and scorePhase) %}
58
- {% macro q_tile_value(index) %}{% if hasBias %}(query[{{ index }}] + bias[h * HEAD_DIM + k]){% else %}query[{{ index }}]{% endif %}{% if SCALE_IN_Q %} * score_scale{% endif %}{% endmacro %}
 
59
 
60
  const HEADS: u32 = {{ qNumHeads }}u;
61
  {% if kvNumHeads is defined %}
@@ -108,14 +100,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 +119,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 "" }}
@@ -173,30 +159,18 @@ fn main(
173
  let inner = params.kvSeq;
174
  {% endif %}
175
  {% if FUSED_SOFTMAX %}
176
- {% if PRIVATE_ROW_STATS %}
177
- // In the admitted BM64/BN64/BK32/WG256 loader, four adjacent lanes own the
178
- // same query row for every reduction tile. Keep that row's constants private:
179
- // this removes both 512 bytes of workgroup storage and the initialization
180
- // barrier while preserving the exact exp/divide sequence of the shared-memory
181
- // row-stats arm.
182
- let private_stat_row =
183
- (b * HEADS + h) * params.qSeq + min(m_base + li / 4u, params.qSeq - 1u);
184
- let private_softmax_m = rowStats[private_stat_row * 2u];
185
- let private_softmax_d = rowStats[private_stat_row * 2u + 1u];
186
- {% else %}
187
  // One row-stats pair per query row of the tile. A query tail clamps to the last
188
  // real row rather than reading past the buffer; those lanes are discarded by the
189
  // staging guard anyway, and the clamp keeps the denominator non-zero.
190
  for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
191
- let stat_row = (b * HEADS + h) * params.qSeq + min(m_base + r, params.qSeq - 1u);
192
  softmax_m[r] = rowStats[stat_row * 2u];
193
  softmax_d[r] = rowStats[stat_row * 2u + 1u];
194
  }
195
  workgroupBarrier();
196
- {% endif %}
197
  {% endif %}
198
  for (var k_base = 0u; k_base < inner; k_base += TILE_K) {
199
- {% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" %}
200
  // Full interior PV tiles can be loaded directly from storage. Query,
201
  // reduction, and output-dimension tails use the guarded shared path below.
202
  if (
@@ -206,7 +180,7 @@ fn main(
206
  ) {
207
  for (var step = 0u; step < TILE_K; step += 8u) {
208
  {% for row_block in range(ROW_BLOCKS) %}
209
- let score_offset{{ row_block }} = (b * HEADS + h) * params.qSeq * params.kvSeq
210
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
211
  var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
212
  subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
@@ -256,7 +230,7 @@ fn main(
256
  tile_A[a_row * TILE_K + a_col + i] = loaded;
257
  {% endif %}
258
  {% else %}
259
- let score_base = (b * HEADS + h) * params.qSeq * params.kvSeq;
260
  tile_A[a_row * TILE_K + a_col + i] =
261
  {{ "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 "" }};
262
  {% endif %}
@@ -271,20 +245,20 @@ fn main(
271
  {% if headDim % 32 == 0 %}
272
  tile_B[b_row * TILE_K + b_col + i] = select(
273
  {{ "0.0h" if MT == "f16" else "0.0" }},
274
- key[{{ kv_index("col", "k") }}],
275
  col < params.kvSeq
276
  );
277
  {% else %}
278
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
279
  if (col < params.kvSeq && k < HEAD_DIM) {
280
- loaded = key[{{ kv_index("col", "k") }}];
281
  }
282
  tile_B[b_row * TILE_K + b_col + i] = loaded;
283
  {% endif %}
284
  {% else %}
285
  tile_B[b_row * TILE_K + b_col + i] = select(
286
  {{ "0.0h" if MT == "f16" else "0.0" }},
287
- value[{{ kv_index("k", "col") }}],
288
  k < params.kvSeq && col < HEAD_DIM
289
  );
290
  {% endif %}
@@ -303,13 +277,13 @@ fn main(
303
  loaded = {{ q_tile_value(q_index("row", "k")) }};
304
  }
305
  {% elif FUSED_SOFTMAX %}
306
- let score_base = (b * HEADS + h) * params.qSeq * params.kvSeq;
307
  let loaded =
308
  {{ "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 "" }};
309
  {% else %}
310
  var loaded = 0.0;
311
  if (row < params.qSeq && k < params.kvSeq) {
312
- let score_base = (b * HEADS + h) * params.qSeq * params.kvSeq;
313
  loaded = scores[score_base + row * params.kvSeq + k];
314
  }
315
  {% endif %}
@@ -324,12 +298,12 @@ fn main(
324
  {% if scorePhase %}
325
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
326
  if (col < params.kvSeq && k < HEAD_DIM) {
327
- loaded = key[{{ kv_index("col", "k") }}];
328
  }
329
  {% else %}
330
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
331
  if (k < params.kvSeq && col < HEAD_DIM) {
332
- loaded = value[{{ kv_index("k", "col") }}];
333
  }
334
  {% endif %}
335
  tile_B[idx] = loaded;
@@ -382,7 +356,7 @@ fn main(
382
  {% for col_block in range(COL_BLOCKS) %}
383
  {% if scorePhase %}
384
  let output_offset{{ row_block }}{{ col_block }} =
385
- (b * HEADS + h) * params.qSeq * params.kvSeq
386
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
387
  + n_base + base_B + {{ col_block * 8 }}u;
388
  subgroupMatrixStore<row_major>(
@@ -448,7 +422,7 @@ fn main(
448
  {% endif %}
449
  let scored = result{% if not SCALE_IN_Q %} * scale{% endif %};
450
  scores[
451
- (b * HEADS + h) * params.qSeq * params.kvSeq + row * params.kvSeq + col
452
  ] = scored;
453
  {% if EMIT_ROW_STATS %}
454
  // Softmax sees the STORED value, so the statistics have to be taken on it
@@ -462,7 +436,7 @@ fn main(
462
  {% if hasBias %}
463
  // V bias row base: skip the packed Q and K blocks, then index this head.
464
  {% endif %}
465
- 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 %};
466
  {% endif %}
467
  }
468
  }
@@ -490,7 +464,7 @@ fn main(
490
  // eight consecutive pairs, and the combine pass reads a slot's whole column
491
  // of rows contiguously.
492
  let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
493
- let out_index = (((b * HEADS + h) * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
494
  scorePartials[out_index] = stat_m{{ row_block }};
495
  scorePartials[out_index + 1u] = stat_d{{ row_block }};
496
  }
 
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 layout = layout | default("bsh") %}
 
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 = "HEAD_DIM" if headMajor else "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 FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
22
  {% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
23
  {% macro score_value(index, guard) %}
24
+ {% set rowMax = "private_softmax_m" if PRIVATE_ROW_STATS else "softmax_m[" ~ guard[0] ~ "]" %}
25
+ {% set rowDenom = "private_softmax_d" if PRIVATE_ROW_STATS else "softmax_d[" ~ guard[0] ~ "]" %}
26
+ select(0.0, {{ "exp_shift(scores[" ~ index ~ "], " ~ rowMax ~ ") / " ~ rowDenom if FUSED_SOFTMAX else "scores[" ~ index ~ "]" }}, {{ guard[1] }}){% endmacro %}
 
 
 
 
 
 
 
27
  {% set TILE_M_VALUE = materializedSgmatQueryTile %}
28
  {% set TILE_N_VALUE = materializedSgmatKeyTile %}
29
  {% set TILE_K_VALUE = materializedSgmatInnerTile %}
30
+ {% set SUB_ROWS_VALUE = 16 %}
31
+ {% set SUB_COLS_VALUE = 32 %}
32
+ {% set ROW_BLOCKS = 2 %}
33
+ {% set COL_BLOCKS = 4 %}
34
  {% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
35
  {% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
36
  {% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
 
38
  {% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
39
  {% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
40
  {% set DIRECT_SCORE_STORE = materializedSgmatDirectScoreStore and not EMIT_ROW_STATS %}
41
+ {% set DIRECT_APPLY_STORE = materializedSgmatDirectApplyStore and not hasBias and MT == "f32" and ST == "f32" %}
42
  {% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
43
  {% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
44
  and materializedSgmatRuntimeDirectStore
45
+ and (scorePhase or (not hasBias and MT == "f32" and ST == "f32"))
46
  and not EMIT_ROW_STATS %}
47
  {% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
48
  {% set SCALE_IN_Q = DIRECT_SCORE_STORE or (RUNTIME_DIRECT_STORE and scorePhase) %}
49
+ {% macro operand_load(name, index) %}{% if ST != MT %}{{ MT }}({% endif %}{{ name }}[{{ index }}]{% if ST != MT %}){% endif %}{% endmacro %}
50
+ {% 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 %}
51
 
52
  const HEADS: u32 = {{ qNumHeads }}u;
53
  {% if kvNumHeads is defined %}
 
100
  var<workgroup> softmax_d: array<f32, {{ TILE_M_VALUE }}>;
101
  {% endif %}
102
  {% if FUSED_SOFTMAX or EMIT_ROW_STATS %}
 
103
  // FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
104
  // `m - m` finite so an empty lane / all--inf row contributes the exact
105
  // accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
106
  // a zero final denominator according to their public semantics. Using -inf
107
  // here changes +inf-row behavior.
108
  const FLT_MAX: f32 = 3.4028234663852886e38;
 
109
 
110
  fn is_finite_f32(value: f32) -> bool {
111
  return select(false, value <= FLT_MAX, value >= -FLT_MAX);
 
119
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
120
  return select(value - maxValue, 0.0, equalFiniteMax);
121
  }
 
 
122
 
123
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
124
  return exp(shifted_value(value, maxValue));
125
  }
 
 
126
  {% endif %}
127
 
128
  @compute @workgroup_size({{ WORKGROUP_THREADS }}, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
 
159
  let inner = params.kvSeq;
160
  {% endif %}
161
  {% if FUSED_SOFTMAX %}
 
 
 
 
 
 
 
 
 
 
 
162
  // One row-stats pair per query row of the tile. A query tail clamps to the last
163
  // real row rather than reading past the buffer; those lanes are discarded by the
164
  // staging guard anyway, and the clamp keeps the denominator non-zero.
165
  for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
166
+ let stat_row = {{ SCRATCH_HEAD }} * params.qSeq + min(m_base + r, params.qSeq - 1u);
167
  softmax_m[r] = rowStats[stat_row * 2u];
168
  softmax_d[r] = rowStats[stat_row * 2u + 1u];
169
  }
170
  workgroupBarrier();
 
171
  {% endif %}
172
  for (var k_base = 0u; k_base < inner; k_base += TILE_K) {
173
+ {% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" and ST == "f32" %}
174
  // Full interior PV tiles can be loaded directly from storage. Query,
175
  // reduction, and output-dimension tails use the guarded shared path below.
176
  if (
 
180
  ) {
181
  for (var step = 0u; step < TILE_K; step += 8u) {
182
  {% for row_block in range(ROW_BLOCKS) %}
183
+ let score_offset{{ row_block }} = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
184
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
185
  var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
186
  subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
 
230
  tile_A[a_row * TILE_K + a_col + i] = loaded;
231
  {% endif %}
232
  {% else %}
233
+ let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
234
  tile_A[a_row * TILE_K + a_col + i] =
235
  {{ "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 "" }};
236
  {% endif %}
 
245
  {% if headDim % 32 == 0 %}
246
  tile_B[b_row * TILE_K + b_col + i] = select(
247
  {{ "0.0h" if MT == "f16" else "0.0" }},
248
+ {{ operand_load("key", kv_index("col", "k")) }},
249
  col < params.kvSeq
250
  );
251
  {% else %}
252
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
253
  if (col < params.kvSeq && k < HEAD_DIM) {
254
+ loaded = {{ operand_load("key", kv_index("col", "k")) }};
255
  }
256
  tile_B[b_row * TILE_K + b_col + i] = loaded;
257
  {% endif %}
258
  {% else %}
259
  tile_B[b_row * TILE_K + b_col + i] = select(
260
  {{ "0.0h" if MT == "f16" else "0.0" }},
261
+ {{ operand_load("value", kv_index("k", "col")) }},
262
  k < params.kvSeq && col < HEAD_DIM
263
  );
264
  {% endif %}
 
277
  loaded = {{ q_tile_value(q_index("row", "k")) }};
278
  }
279
  {% elif FUSED_SOFTMAX %}
280
+ let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
281
  let loaded =
282
  {{ "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 "" }};
283
  {% else %}
284
  var loaded = 0.0;
285
  if (row < params.qSeq && k < params.kvSeq) {
286
+ let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
287
  loaded = scores[score_base + row * params.kvSeq + k];
288
  }
289
  {% endif %}
 
298
  {% if scorePhase %}
299
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
300
  if (col < params.kvSeq && k < HEAD_DIM) {
301
+ loaded = {{ operand_load("key", kv_index("col", "k")) }};
302
  }
303
  {% else %}
304
  var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
305
  if (k < params.kvSeq && col < HEAD_DIM) {
306
+ loaded = {{ operand_load("value", kv_index("k", "col")) }};
307
  }
308
  {% endif %}
309
  tile_B[idx] = loaded;
 
356
  {% for col_block in range(COL_BLOCKS) %}
357
  {% if scorePhase %}
358
  let output_offset{{ row_block }}{{ col_block }} =
359
+ {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
360
  + (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
361
  + n_base + base_B + {{ col_block * 8 }}u;
362
  subgroupMatrixStore<row_major>(
 
422
  {% endif %}
423
  let scored = result{% if not SCALE_IN_Q %} * scale{% endif %};
424
  scores[
425
+ {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq + row * params.kvSeq + col
426
  ] = scored;
427
  {% if EMIT_ROW_STATS %}
428
  // Softmax sees the STORED value, so the statistics have to be taken on it
 
436
  {% if hasBias %}
437
  // V bias row base: skip the packed Q and K blocks, then index this head.
438
  {% endif %}
439
+ 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 %};
440
  {% endif %}
441
  }
442
  }
 
464
  // eight consecutive pairs, and the combine pass reads a slot's whole column
465
  // of rows contiguously.
466
  let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
467
+ let out_index = (({{ SCRATCH_HEAD }} * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
468
  scorePartials[out_index] = stat_m{{ row_block }};
469
  scorePartials[out_index + 1u] = stat_d{{ row_block }};
470
  }
build/webgpu/attn-materialized-softmax-f32.wgsl.jinja CHANGED
@@ -1,9 +1,7 @@
1
  {% set CACHE_VEC4 = cacheVec4 if cacheVec4 is defined else false %}
2
- {% set CAUSAL_ROWS = causalRows is defined and causalRows %}
3
- {% set SCALE_IN = scaleInSoftmax is defined and scaleInSoftmax %}
4
- {% set RUNTIME_COLS = colsFromParams is defined and colsFromParams %}
5
- {% set COLS_EXPR = "params.keyLen" if RUNTIME_COLS else "COLS" %}
6
- {% set ROWS_EXPR = rowsExpr if rowsExpr is defined else "params.rows" %}
7
  {% if useSubgroups %}
8
  enable subgroups;
9
  {% endif %}
@@ -12,7 +10,7 @@ enable subgroups;
12
  // In-place row softmax for materialized f32 attention. The default path reads
13
  // each scalar twice; the register-cached vec4 path eliminates the second read.
14
  const WG: u32 = {{ materializedSoftmaxWg }}u;
15
- {% set SCALED_SCORE = "scores[base + c] * SCALE" if SCALE_IN else "scores[base + c]" %}
16
  {% if CACHE_VEC4 %}
17
  const COLS4: u32 = {{ materializedSoftmaxCols4 }}u;
18
  const CACHE_VECS: u32 = (COLS4 + WG - 1u) / WG;
@@ -38,6 +36,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
38
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
39
  return select(value - maxValue, 0.0, equalFiniteMax);
40
  }
 
41
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
42
  return exp(shifted_value(value, maxValue));
43
  }
@@ -101,41 +100,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
101
  return combinedMD;
102
  }
103
  {% else %}
104
- {% set mdStreamed = mdStreams is defined %}
105
- {% set mdStreams = mdStreams if mdStreams is defined else 1 %}
106
- {% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
107
  var<workgroup> partialM: array<f32, {{ mdExtent }}>;
108
  var<workgroup> partialD: array<f32, {{ mdExtent }}>;
109
- {% if mdStreamed %}
110
-
111
- // In-place fold of {{ mdStreams }} streams. Input partials occupy
112
- // partialM/partialD; stream s returns its merged pair in slot s * WG.
113
- fn combine_partials_streams(lidx: u32) {
114
- workgroupBarrier();
115
- var stride = WG / 2u;
116
- loop {
117
- if (stride == 0u) {
118
- break;
119
- }
120
- if (lidx < stride) {
121
- {% for s in range(mdStreams) %}
122
- {
123
- let slot = {{ s }}u * WG + lidx;
124
- let m1 = partialM[slot];
125
- let d1 = partialD[slot];
126
- let m2 = partialM[slot + stride];
127
- let d2 = partialD[slot + stride];
128
- let mNew = max(m1, m2);
129
- partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
130
- partialM[slot] = mNew;
131
- }
132
- {% endfor %}
133
- }
134
- workgroupBarrier();
135
- stride = stride / 2u;
136
- }
137
- }
138
- {% else %}
139
 
140
  fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
141
  partialM[lidx] = m;
@@ -165,8 +132,6 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
165
  return merged;
166
  }
167
  {% endif %}
168
- {% endif %}
169
-
170
 
171
  @compute @workgroup_size(WG, 1, 1)
172
  fn main(
 
1
  {% set CACHE_VEC4 = cacheVec4 if cacheVec4 is defined else false %}
2
+ {% set CAUSAL_ROWS = false %}
3
+ {% set COLS_EXPR = "COLS" %}
4
+ {% set ROWS_EXPR = "params.rows" %}
 
 
5
  {% if useSubgroups %}
6
  enable subgroups;
7
  {% endif %}
 
10
  // In-place row softmax for materialized f32 attention. The default path reads
11
  // each scalar twice; the register-cached vec4 path eliminates the second read.
12
  const WG: u32 = {{ materializedSoftmaxWg }}u;
13
+ {% set SCALED_SCORE = "scores[base + c]" %}
14
  {% if CACHE_VEC4 %}
15
  const COLS4: u32 = {{ materializedSoftmaxCols4 }}u;
16
  const CACHE_VECS: u32 = (COLS4 + WG - 1u) / WG;
 
36
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
37
  return select(value - maxValue, 0.0, equalFiniteMax);
38
  }
39
+
40
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
41
  return exp(shifted_value(value, maxValue));
42
  }
 
100
  return combinedMD;
101
  }
102
  {% else %}
103
+ {% set mdExtent = "WG" %}
 
 
104
  var<workgroup> partialM: array<f32, {{ mdExtent }}>;
105
  var<workgroup> partialD: array<f32, {{ mdExtent }}>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
106
 
107
  fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
108
  partialM[lidx] = m;
 
132
  return merged;
133
  }
134
  {% endif %}
 
 
135
 
136
  @compute @workgroup_size(WG, 1, 1)
137
  fn main(
build/webgpu/attn-online-scalar.wgsl.jinja CHANGED
@@ -1,11 +1,11 @@
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.
@@ -21,18 +21,39 @@ const Q_HIDDEN: u32 = {{ qHidden }}u;
21
  const KV_HIDDEN: u32 = {{ kvHidden }}u;
22
  const Q_HEADS: u32 = {{ qNumHeads }}u;
23
  const KV_HEADS: u32 = {{ kvNumHeads }}u;
24
- {% set qHeads = "params.qHeads" if headsFromParams else "Q_HEADS" %}
25
- {% set kvHeads = "params.kvHeads" if headsFromParams else "KV_HEADS" %}
26
- {% set scale = scale | default("0.0") %}
27
  const WG: u32 = {{ workgroupSize }}u;
28
 
29
- var<workgroup> partial: array<f32, WG>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30
  var<workgroup> running_max: f32;
31
  var<workgroup> running_denom: f32;
32
  var<workgroup> running_out: array<f32, HEAD_DIM>;
33
  var<workgroup> previous_scale: f32;
34
- {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
35
- fn {{ name }}(value: f32, tid: u32) -> f32 {
36
  {{ buffer }}[tid] = value;
37
  workgroupBarrier();
38
  // Ceil-halving keeps every lane when the workgroup size is not a power of
@@ -42,11 +63,7 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
42
  loop {
43
  let half = (n + 1u) / 2u;
44
  if (tid < n - half) {
45
- {% if mode == "max" %}
46
- {{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
47
- {% else %}
48
- {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
49
- {% endif %}
50
  }
51
  workgroupBarrier();
52
  n = half;
@@ -58,22 +75,16 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
58
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
59
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
60
  let reduced = {{ buffer }}[0];
61
- {% if trailingBarrier %}
62
  workgroupBarrier();
63
- {% endif %}
64
  return reduced;
65
  }
66
  {% endmacro %}
67
-
68
- {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
69
-
70
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
71
- fn scale_value() -> f32 {
72
  if (params.scale != 0.0) { return params.scale; }
73
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
74
  }
75
 
76
-
77
  @compute @workgroup_size(WG, 1, 1)
78
  fn main(
79
  @builtin(workgroup_id) wg: vec3<u32>,
@@ -117,7 +128,7 @@ fn main(
117
 
118
  for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
119
  let kRow = kvBase + key_token * kvTokenStride;
120
- var partial_dot = 0.0;
121
  for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
122
  var q_value = f32(query[qBase + d]);
123
  var k_value = f32(key[kRow + d]);
@@ -126,17 +137,17 @@ fn main(
126
  q_value = q_value + f32(bias[channel]);
127
  k_value = k_value + f32(bias[Q_HIDDEN + channel]);
128
  {% endif %}
129
- partial_dot = partial_dot + q_value * k_value;
130
  }
131
 
132
- // reduce_sum returns the same partial[0] to every lane, so `score` is already
133
  // workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
134
  // through shared memory by the barrier below).
135
  {% if hasMask %}
136
  let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
137
- let score = reduce_sum(partial_dot, tid) * scale_value() + f32(attn_mask[maskIndex]);
138
  {% else %}
139
- let score = reduce_sum(partial_dot, tid) * scale_value();
140
  {% endif %}
141
  if (tid == 0u) {
142
  let next_max = max(running_max, score);
 
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.
 
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>,
 
128
 
129
  for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
130
  let kRow = kvBase + key_token * kvTokenStride;
131
+ var partial_dot = DotAccumulator(0.0, 0.0);
132
  for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
133
  var q_value = f32(query[qBase + d]);
134
  var k_value = f32(key[kRow + d]);
 
137
  q_value = q_value + f32(bias[channel]);
138
  k_value = k_value + f32(bias[Q_HIDDEN + channel]);
139
  {% endif %}
140
+ partial_dot = dot_accumulate(partial_dot, q_value, k_value);
141
  }
142
 
143
+ // reduce_dot returns the same partial[0] to every lane, so `score` is already
144
  // workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
145
  // through shared memory by the barrier below).
146
  {% if hasMask %}
147
  let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
148
+ let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value() + f32(attn_mask[maskIndex]);
149
  {% else %}
150
+ let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value();
151
  {% endif %}
152
  if (tid == 0u) {
153
  let next_max = max(running_max, score);
build/webgpu/attn-small-head-parallel.wgsl.jinja CHANGED
@@ -9,17 +9,15 @@ const HIDDEN: u32 = {{ qHidden }}u;
9
  const KV_SEQ: u32 = {{ kvSeq }}u;
10
  const WG: u32 = 64u;
11
  const NEG_MAX: f32 = -3.4028234663852886e38;
12
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
13
- fn scale_value() -> f32 {
14
  if (params.scale != 0.0) { return params.scale; }
15
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
16
  }
17
 
18
-
19
  var<workgroup> scores: array<f32, KV_SEQ>;
20
  var<workgroup> partial: array<f32, WG>;
21
- {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
22
- fn {{ name }}(value: f32, tid: u32) -> f32 {
23
  {{ buffer }}[tid] = value;
24
  workgroupBarrier();
25
  // Ceil-halving keeps every lane when the workgroup size is not a power of
@@ -45,16 +43,12 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
45
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
46
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
47
  let reduced = {{ buffer }}[0];
48
- {% if trailingBarrier %}
49
  workgroupBarrier();
50
- {% endif %}
51
  return reduced;
52
  }
53
  {% endmacro %}
54
-
55
  {{ wgsl_tree_reduce_f32("reduce_max", "max", "partial", "WG") }}
56
  {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
57
-
58
  @compute @workgroup_size(WG, 1, 1)
59
  fn main(
60
  @builtin(workgroup_id) wg: vec3<u32>,
 
9
  const KV_SEQ: u32 = {{ kvSeq }}u;
10
  const WG: u32 = 64u;
11
  const NEG_MAX: f32 = -3.4028234663852886e38;
12
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
13
  if (params.scale != 0.0) { return params.scale; }
14
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
15
  }
16
 
 
17
  var<workgroup> scores: array<f32, KV_SEQ>;
18
  var<workgroup> partial: array<f32, WG>;
19
+ {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
20
+ fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
21
  {{ buffer }}[tid] = value;
22
  workgroupBarrier();
23
  // Ceil-halving keeps every lane when the workgroup size is not a power of
 
43
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
44
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
45
  let reduced = {{ buffer }}[0];
 
46
  workgroupBarrier();
 
47
  return reduced;
48
  }
49
  {% endmacro %}
 
50
  {{ wgsl_tree_reduce_f32("reduce_max", "max", "partial", "WG") }}
51
  {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
 
52
  @compute @workgroup_size(WG, 1, 1)
53
  fn main(
54
  @builtin(workgroup_id) wg: vec3<u32>,
build/webgpu/attn-small-head-value.wgsl.jinja ADDED
@@ -0,0 +1,198 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% if useValueSubgroups %}
2
+ enable subgroups;
3
+ {% endif %}
4
+ {{ env.wgsl.resourceDeclarations }}
5
+ {% set rowType = "f32" if queryBlock == 1 else "vec" ~ queryBlock ~ "<f32>" %}
6
+ {% set dotValueType = rowType %}
7
+ {% set dotType = dotValueType | default("f32") %}
8
+ // Retain product and addition residuals across a dot product. Explicit fma
9
+ // boundaries preserve the addition error transform under reassociation.
10
+ struct DotAccumulator {
11
+ hi: {{ dotType }},
12
+ lo: {{ dotType }},
13
+ }
14
+ fn dot_accumulate(acc: DotAccumulator, a: {{ dotType }}, b: {{ dotType }}) -> DotAccumulator {
15
+ let product = fma(a, b, {{ dotType }}(0.0));
16
+ let productError = fma(a, b, -product);
17
+ let sum = fma(acc.hi, {{ dotType }}(1.0), product);
18
+ let bv = fma({{ dotType }}(-1.0), acc.hi, sum);
19
+ let av = fma({{ dotType }}(-1.0), bv, sum);
20
+ let ae = fma({{ dotType }}(-1.0), av, acc.hi);
21
+ let be = fma({{ dotType }}(-1.0), bv, product);
22
+ let error = fma(fma(fma(ae, {{ dotType }}(1.0), be), {{ dotType }}(1.0), acc.lo), {{ dotType }}(1.0), productError);
23
+ let hi = fma(sum, {{ dotType }}(1.0), error);
24
+ return DotAccumulator(hi, fma({{ dotType }}(-1.0), fma({{ dotType }}(-1.0), sum, hi), error));
25
+ }
26
+ fn dot_value(acc: DotAccumulator) -> {{ dotType }} {
27
+ return fma(acc.hi, {{ dotType }}(1.0), acc.lo);
28
+ }
29
+ const HEAD_DIM: u32 = {{ headDim }}u;
30
+ const HIDDEN: u32 = {{ qHidden }}u;
31
+ const KV_SEQ: u32 = {{ kvSeq }}u;
32
+ const WG: u32 = {{ valueWorkgroupSize }}u;
33
+ const QUERY_BLOCK: u32 = {{ queryBlock }}u;
34
+ const NEG_MAX: f32 = -3.4028234663852886e38;
35
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
36
+ if (params.scale != 0.0) { return params.scale; }
37
+ return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
38
+ }
39
+ var<workgroup> scores: array<{{ rowType }}, {{ scoreStorageElements }}u>;
40
+ var<workgroup> partial: array<{{ rowType }}, WG>;
41
+
42
+ {% if useValueSubgroups %}
43
+ {% macro reduce(name, collective, mode, identity) %}
44
+ fn {{ name }}(value: {{ rowType }}, sgLane: u32, sgId: u32, numSg: u32) -> {{ rowType }} {
45
+ let sub = {{ collective }}(value);
46
+ if (sgLane == 0u) { partial[sgId] = sub; }
47
+ workgroupBarrier();
48
+ var total = {{ rowType }}({{ identity }});
49
+ for (var i = 0u; i < numSg; i = i + 1u) {
50
+ {% if mode == "max" %}
51
+ total = max(total, partial[i]);
52
+ {% else %}
53
+ total = total + partial[i];
54
+ {% endif %}
55
+ }
56
+ workgroupBarrier();
57
+ return total;
58
+ }
59
+ {% endmacro %}
60
+ {{ reduce("reduce_max", "subgroupMax", "max", "NEG_MAX") }}
61
+ {{ reduce("reduce_sum", "subgroupAdd", "add", "0.0") }}
62
+ {% else %}
63
+ {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
64
+ fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
65
+ {{ buffer }}[tid] = value;
66
+ workgroupBarrier();
67
+ // Ceil-halving keeps every lane when the workgroup size is not a power of
68
+ // two. For even n this matches the power-of-two tree order; for odd n, lanes
69
+ // [0, n-half) fold the upper tail while the middle lane carries forward.
70
+ var n: u32 = {{ wg }};
71
+ loop {
72
+ let half = (n + 1u) / 2u;
73
+ if (tid < n - half) {
74
+ {% if mode == "max" %}
75
+ {{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
76
+ {% else %}
77
+ {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
78
+ {% endif %}
79
+ }
80
+ workgroupBarrier();
81
+ n = half;
82
+ if (n == 1u) {
83
+ break;
84
+ }
85
+ }
86
+ // The default trailing barrier makes this helper safe for back-to-back calls: every lane reads
87
+ // slot 0 here, so the next call's first store must not run until all lanes have read it.
88
+ // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
89
+ let reduced = {{ buffer }}[0];
90
+ workgroupBarrier();
91
+ return reduced;
92
+ }
93
+ {% endmacro %}
94
+ {{ wgsl_tree_reduce_f32("reduce_max", "max", valueType=rowType) }}
95
+ {{ wgsl_tree_reduce_f32("reduce_sum", "add", valueType=rowType) }}
96
+ {% endif %}
97
+ @compute @workgroup_size(WG, 1, 1)
98
+ fn main(
99
+ @builtin(workgroup_id) wg: vec3<u32>,
100
+ @builtin(num_workgroups) nwg: vec3<u32>,
101
+ @builtin(local_invocation_id) lid: vec3<u32>
102
+ {% if useValueSubgroups %}
103
+ , @builtin(subgroup_invocation_id) sgLane: u32
104
+ , @builtin(subgroup_id) sgId: u32
105
+ , @builtin(num_subgroups) numSg: u32
106
+ {% endif %}
107
+ ) {
108
+ let h = wg.y;
109
+ let batch = wg.z;
110
+ let tid = lid.x;
111
+ let kvBase = batch * KV_SEQ * HIDDEN + h * HEAD_DIM;
112
+ for (var query_token = wg.x * QUERY_BLOCK; query_token < params.qSeq; query_token = query_token + nwg.x * QUERY_BLOCK) {
113
+ {% for row in range(queryBlock) %}
114
+ let qBase_{{ row }} = (batch * params.qSeq + {% if queryTail and row > 0 %}min(query_token + {{ row }}u, params.qSeq - 1u){% else %}query_token{% if row > 0 %} + {{ row }}u{% endif %}{% endif %}) * HIDDEN + h * HEAD_DIM;
115
+ {% endfor %}
116
+ // Query components stay private throughout the key scan. Their lifetime
117
+ // ends before the value accumulators need the same registers.
118
+ {
119
+ {% for d in range(headDim) %}
120
+ let q_{{ d }} = {% if queryBlock > 1 %}{{ rowType }}({% endif %}{% for row in range(queryBlock) %}{% if not loop.first %}, {% endif %}f32(query[qBase_{{ row }} + {{ d }}u]){% endfor %}{% if queryBlock > 1 %}){% endif %};
121
+ {% endfor %}
122
+ for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
123
+ let kBase = kvBase + key_token * HIDDEN;
124
+ var score = DotAccumulator({{ rowType }}(0.0), {{ rowType }}(0.0));
125
+ {% for d in range(headDim) %}
126
+ score = dot_accumulate(score, q_{{ d }}, {{ rowType }}(f32(key[kBase + {{ d }}u])));
127
+ {% endfor %}
128
+ scores[key_token] = dot_value(score) * scale_value();
129
+ }
130
+ }
131
+ workgroupBarrier();
132
+ var laneMax = {{ rowType }}(NEG_MAX);
133
+ for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
134
+ laneMax = max(laneMax, scores[key_token]);
135
+ }
136
+ let rowMax = reduce_max(laneMax, {% if useValueSubgroups %}sgLane, sgId, numSg{% else %}tid{% endif %});
137
+ var laneSum = {{ rowType }}(0.0);
138
+ for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
139
+ let probability = exp(scores[key_token] - rowMax);
140
+ scores[key_token] = probability;
141
+ laneSum = laneSum + probability;
142
+ }
143
+ let invDenom = {{ rowType }}(1.0) / reduce_sum(laneSum, {% if useValueSubgroups %}sgLane, sgId, numSg{% else %}tid{% endif %});
144
+ {% for d in range(headDim) %}
145
+ var value_acc_{{ d }} = {{ rowType }}(0.0);
146
+ {% endfor %}
147
+ for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
148
+ let probability = scores[key_token];
149
+ let vBase = kvBase + key_token * HIDDEN;
150
+ {% for d in range(headDim) %}
151
+ value_acc_{{ d }} = value_acc_{{ d }} + probability * f32(value[vBase + {{ d }}u]);
152
+ {% endfor %}
153
+ }
154
+ {% if useValueSubgroups %}
155
+ {% for d in range(headDim) %}
156
+ let value_sum_{{ d }} = subgroupAdd(value_acc_{{ d }});
157
+ {% endfor %}
158
+ workgroupBarrier();
159
+ if (sgLane == 0u) {
160
+ {% for d in range(headDim) %}
161
+ scores[sgId * HEAD_DIM + {{ d }}u] = value_sum_{{ d }};
162
+ {% endfor %}
163
+ }
164
+ workgroupBarrier();
165
+ if (tid < HEAD_DIM) {
166
+ var total = {{ rowType }}(0.0);
167
+ for (var i = 0u; i < numSg; i = i + 1u) { total = total + scores[i * HEAD_DIM + tid]; }
168
+ let result = total * invDenom;
169
+ {% else %}
170
+ workgroupBarrier();
171
+ {% for d in range(headDim) %}
172
+ scores[{{ d }}u * WG + tid] = value_acc_{{ d }};
173
+ {% endfor %}
174
+ workgroupBarrier();
175
+ for (var stride = WG / 2u; stride > 0u; stride = stride / 2u) {
176
+ if (tid < stride) {
177
+ {% for d in range(headDim) %}
178
+ scores[{{ d }}u * WG + tid] = scores[{{ d }}u * WG + tid] + scores[{{ d }}u * WG + tid + stride];
179
+ {% endfor %}
180
+ }
181
+ workgroupBarrier();
182
+ }
183
+ if (tid < HEAD_DIM) {
184
+ let result = scores[tid * WG] * invDenom;
185
+ {% endif %}
186
+ {% for row in range(queryBlock) %}
187
+ {% if queryTail and row > 0 %}
188
+ if (query_token + {{ row }}u < params.qSeq) {
189
+ {% endif %}
190
+ output[qBase_{{ row }} + tid] = {{ outputScalar }}(result{% if queryBlock > 1 %}[{{ row }}]{% endif %});
191
+ {% if queryTail and row > 0 %}
192
+ }
193
+ {% endif %}
194
+ {% endfor %}
195
+ }
196
+ workgroupBarrier();
197
+ }
198
+ }
build/webgpu/bench.json CHANGED
The diff for this file is too large to render. See raw diff
 
build/webgpu/manifest.json CHANGED
@@ -57,7 +57,8 @@
57
  "canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32",
58
  "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
59
  "wave32Effective": "wave32Adapter or pinSubgroupSize32",
60
- "headDimPlan": "dim(shapes.queryT, 2) / attrs.num_heads if (ranks.queryT == 3 and attrs.num_heads > 0) else 0",
 
61
  "qkvDtypesOk": "tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT",
62
  "floatDtypeOk": "(tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(tensorDtypes.queryT)",
63
  "qkvShapeOk": "ranks.queryT == 3 and ranks.keyT == 3 and ranks.valueT == 3 and ranks.outputT == 3 and attrs.num_heads > 0 and dim(shapes.queryT, 2) % attrs.num_heads == 0 and dim(shapes.keyT, 2) == dim(shapes.queryT, 2) and dim(shapes.valueT, 2) == dim(shapes.queryT, 2) and dim(shapes.keyT, 1) == dim(shapes.valueT, 1) and dim(shapes.queryT, 0) == dim(shapes.keyT, 0) and dim(shapes.queryT, 0) == dim(shapes.valueT, 0) and dim(shapes.outputT, 0) == dim(shapes.queryT, 0) and dim(shapes.outputT, 1) == dim(shapes.queryT, 1) and dim(shapes.outputT, 2) == dim(shapes.valueT, 2)",
@@ -66,40 +67,42 @@
66
  "qkvContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and noAttnBias",
67
  "qkvMaskContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and attnBiasOk",
68
  "biasOk": "present.biasT and ranks.biasT == 1 and tensorDtypes.biasT == tensorDtypes.queryT and dim(shapes.biasT, 0) == 3 * dim(shapes.queryT, 2)",
69
- "q32BroadcastSubgroupLanes": "device.adapterInfo.subgroupMinSize if subgroupsWave32 else 0",
70
- "q32BroadcastF32HeadVectors": "headDimPlan / 4 if headDimPlan % 4 == 0 else 0",
71
- "q32BroadcastF32RegisterGeometry": "subgroupsWave32 and q32BroadcastF32HeadVectors == q32BroadcastSubgroupLanes",
 
 
72
  "subgroupCluster4": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 4 and device.adapterInfo.subgroupMinSize % 4 == 0 and device.adapterInfo.subgroupMaxSize % 4 == 0",
73
  "subgroupCluster8": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 8 and device.adapterInfo.subgroupMinSize % 8 == 0 and device.adapterInfo.subgroupMaxSize % 8 == 0",
74
  "attentionDispatchFits": "dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 1) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
75
- "flashHeadOk": "headDimPlan % 4 == 0 and headDimPlan >= 32 and headDimPlan <= 256",
76
  "flashSizeOk": "flashHeadOk and (dim(shapes.queryT, 1) * attrs.num_heads >= tunables.FLASH_MIN_QUERY_HEADS or (dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512) or (dim(shapes.queryT, 1) > 1 and dim(shapes.keyT, 1) >= 2048)) and attentionDispatchFits",
77
  "flashShapeOk": "qkvContractOk and flashSizeOk",
78
  "flashMaskShapeOk": "qkvMaskContractOk and flashSizeOk",
79
  "noBiasSplitKCount": "min(tunables.DECODE_MAX_SPLITS if dim(shapes.queryT, 1) == 1 else max(1, ceilDiv(tunables.SPLITK_TARGET_WORKGROUPS, dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads)), ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
80
  "biasSplitKCount": "min(tunables.DECODE_MAX_SPLITS, ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
81
- "noBiasPartialOutBytes": "dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * headDimPlan * 4",
82
  "noBiasStatsBytes": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * 4",
83
  "noBiasSplitScratchFits": "noBiasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and noBiasPartialOutBytes <= device.limits.maxBufferSize and noBiasStatsBytes <= device.limits.maxStorageBufferBindingSize and noBiasStatsBytes <= device.limits.maxBufferSize",
84
  "noBiasSplitDispatchFits": "dim(shapes.queryT, 1) * noBiasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
85
- "biasPartialOutBytes": "dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * headDimPlan * 4",
86
  "biasStatsBytes": "2 * dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * 4",
87
  "biasSplitScratchFits": "biasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and biasPartialOutBytes <= device.limits.maxBufferSize and biasStatsBytes <= device.limits.maxStorageBufferBindingSize and biasStatsBytes <= device.limits.maxBufferSize",
88
  "biasSplitDispatchFits": "biasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
89
  "decodeSplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
90
  "shortQuerySplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) >= 2 and dim(shapes.queryT, 1) <= 16 and dim(shapes.keyT, 1) >= 2048 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
91
  "decodeSplitKBiasOk": "biasOk and qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and biasSplitDispatchFits and biasSplitScratchFits",
92
- "decodeSplitKPortablePreferred": "tensorDtypes.queryT == \"float32\" and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and headDimPlan / 4 < device.adapterInfo.subgroupMinSize",
93
  "materializedScoreBytes": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1) * 4",
94
  "materializedScoreFits": "materializedScoreBytes <= device.limits.maxStorageBufferBindingSize and materializedScoreBytes <= device.limits.maxBufferSize",
95
  "materializedWorkgroupSize": "tunables.MATERIALIZED_WORKGROUP_DIM * tunables.MATERIALIZED_WORKGROUP_DIM",
96
  "materializedScoreStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_KEY_TILE * (tunables.MATERIALIZED_INNER_TILE + 4)) * 4",
97
- "materializedApplyTileN": "tunables.MATERIALIZED_VALUE_TILE_D128 if headDimPlan == 128 else tunables.MATERIALIZED_VALUE_TILE",
98
  "materializedApplyStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) * 4",
99
  "materializedTileGeometryOk": "tunables.MATERIALIZED_INNER_TILE % 4 == 0 and (materializedApplyTileN / tunables.MATERIALIZED_WORKGROUP_DIM) % 4 == 0 and tunables.MATERIALIZED_INNER_TILE > 0 and tunables.MATERIALIZED_WORKGROUP_DIM > 0 and tunables.MATERIALIZED_QUERY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_KEY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and materializedApplyTileN % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN >= materializedWorkgroupSize and (tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) % materializedWorkgroupSize == 0",
100
  "materializedDeviceOk": "materializedTileGeometryOk and materializedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeX and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeY and materializedScoreStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and materializedApplyStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
101
  "materializedWideSimdOk": "device.features.has(\"subgroups\") or (has(device.adapterInfo, \"subgroupMinSize\") and device.adapterInfo.subgroupMinSize >= 16)",
102
- "materializedF32CoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDimPlan >= 64 and headDimPlan <= 128 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and materializedScoreFits and materializedDeviceOk",
103
  "materializedSoftmaxStorageBytes": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE * 8 + 8",
104
  "materializedSoftmaxResourcesFit": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE > 0 and pow2ceil(tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE) == tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE and tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE <= deviceWorkgroupCap and materializedSoftmaxStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
105
  "materializedSgmatQueryTile": "tunables.MATERIALIZED_SGMAT_QUERY_TILE",
@@ -123,9 +126,9 @@
123
  "materializedSgmatResourcesFit": "materializedSgmatGeometryOk and materializedSgmatWorkgroupSize <= deviceWorkgroupCap and materializedSgmatCompactStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
124
  "materializedSgmatDispatchFits": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
125
  "materializedSgmatDirectScoreStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and dim(shapes.keyT, 1) % materializedSgmatKeyTile == 0",
126
- "materializedSgmatDirectApplyStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and headDimPlan % materializedSgmatKeyTile == 0",
127
  "materializedSgmatRuntimeDirectStore": "dim(shapes.queryT, 1) >= 2 * materializedSgmatQueryTile and dim(shapes.keyT, 1) >= 2 * materializedSgmatKeyTile",
128
- "materializedSgmatCoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDimPlan >= 32 and headDimPlan <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits",
129
  "materializedCachedSoftmaxWg": "min(tunables.MATERIALIZED_CACHED_SOFTMAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
130
  "materializedCachedSoftmaxVecsPerLane": "ceilDiv(ceilDiv(dim(shapes.keyT, 1), 4), max(1, materializedCachedSoftmaxWg))",
131
  "materializedCachedSoftmaxStorageBytes": "materializedCachedSoftmaxWg * 8 + 8",
@@ -135,41 +138,47 @@
135
  "materializedSgmatOk": "materializedSgmatCoreOk and materializedAdaptiveSoftmaxOk",
136
  "materializedSgmatFusedOk": "materializedSgmatCoreOk",
137
  "materializedF32Ok": "materializedF32CoreOk and materializedAdaptiveSoftmaxOk",
138
- "clusterTileKWg64": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDimPlan * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_SMALL * 4))))",
139
- "clusterTileKWg128": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDimPlan * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_LARGE * 4))))",
140
- "smallHeadParallelOk": "qkvContractOk and headDimPlan < 32 and dim(shapes.keyT, 1) >= 64 and dim(shapes.keyT, 1) <= 2048 and attrs.unidirectional == 0 and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
141
- "prefillTiledStorageBytes": "headDimPlan * tunables.PREFILL_QUERY_TILE * 4",
 
142
  "prefillTiledDeviceOk": "tunables.PREFILL_QUERY_TILE <= deviceWorkgroupCap and prefillTiledStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
143
- "portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, pow2ceil(max(1, headDimPlan)))",
144
- "portableWorkgroupStorageBytes": "portableWorkgroupSize * 4 + max(1, headDimPlan) * 4 + 16",
145
  "portableWorkgroupOk": "tunables.WORKGROUP_SIZE > 0 and pow2ceil(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and portableWorkgroupSize <= deviceWorkgroupCap and portableWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
146
  "fallbackShapeOk": "qkvContractOk and portableWorkgroupOk",
147
  "fallbackMaskShapeOk": "qkvMaskContractOk and portableWorkgroupOk",
148
- "smallSeqShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and dim(shapes.queryT, 1) >= 1 and dim(shapes.queryT, 1) <= tunables.SMALL_SEQ_MAX and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_MAX and headDimPlan >= 1 and attrs.unidirectional == 0",
149
- "smallSeqPrivateFloats": "dim(shapes.keyT, 1) + headDimPlan",
150
  "smallSeqWorkgroupSize": "max(32, pow2ceil(dim(shapes.queryT, 1)))",
151
- "smallSeqSharedBytes": "dim(shapes.keyT, 1) * headDimPlan * 8",
152
  "smallSeqResourcesFit": "smallSeqPrivateFloats <= tunables.SMALL_SEQ_MAX_PRIVATE_FLOATS and smallSeqWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqSharedBytes <= device.limits.maxComputeWorkgroupStorageSize",
153
- "smallSeqBlockedKvBytes": "dim(shapes.keyT, 1) * headDimPlan * 8",
154
- "smallSeqBlockedLaneBytes": "8 + headDimPlan * 4",
155
  "smallSeqBlockedKeyLanes": "min(pow2ceil(dim(shapes.keyT, 1)), 16 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 16 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else (8 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 8 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else 4))",
156
  "smallSeqBlockedWorkgroupSize": "tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * smallSeqBlockedKeyLanes",
157
  "smallSeqBlockedSharedBytes": "smallSeqBlockedKvBytes + smallSeqBlockedWorkgroupSize * smallSeqBlockedLaneBytes",
158
- "smallSeqBlockedShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and headDimPlan % 4 == 0 and headDimPlan >= 4 and headDimPlan <= tunables.SMALL_SEQ_BLOCKED_MAX_HEAD_DIM and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_BLOCKED_MAX_KV and dim(shapes.queryT, 1) >= 1",
159
  "smallSeqBlockedFits": "smallSeqBlockedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqBlockedWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqBlockedSharedBytes <= device.limits.maxComputeWorkgroupStorageSize and ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
160
  "smallSeqDispatchFits": "attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
161
- "materializedSgmatCoreF16Ok": "qkvContractOk and tensorDtypes.queryT == \"float16\" and tensorDtypes.keyT == \"float16\" and tensorDtypes.valueT == \"float16\" and device.features.has(\"shader-f16\") and attrs.unidirectional == 0 and headDimPlan >= 32 and headDimPlan <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits",
162
  "materializedSgmatFusedF16Ok": "materializedSgmatCoreF16Ok",
163
- "attentionScaleExpression": "\"select(inverseSqrt(f32(HEAD_DIM)), params.scale, params.scale != 0.0)\""
 
 
 
 
 
 
164
  },
165
  "bindings": {
166
- "query": { "arg": "queryT", "buffer": "read-only-storage", "elementType": "$inputElement" },
167
- "key": { "arg": "keyT", "buffer": "read-only-storage", "elementType": "$inputElement" },
168
- "value": { "arg": "valueT", "buffer": "read-only-storage", "elementType": "$inputElement" },
169
- "bias": { "arg": "biasT", "buffer": "read-only-storage", "elementType": "$inputScalar" },
170
- "output": { "arg": "outputT", "buffer": "storage", "elementType": "$outputElement" },
171
  "params": {
172
- "buffer": "uniform",
173
  "struct": [
174
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
175
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
@@ -177,77 +186,67 @@
177
  { "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" }
178
  ]
179
  },
180
- "params_2": {
181
  "name": "params",
182
- "buffer": "uniform",
183
  "struct": [
184
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
185
  { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
186
  ]
187
  },
188
- "q": { "arg": "queryT", "buffer": "read-only-storage", "elementType": "$scalar" },
189
- "k": { "arg": "keyT", "buffer": "read-only-storage", "elementType": "$scalar" },
190
- "v": { "arg": "valueT", "buffer": "read-only-storage", "elementType": "$scalar" },
191
- "y": { "arg": "outputT", "buffer": "storage", "elementType": "$scalar" },
192
- "params_4": {
193
  "name": "params",
194
- "buffer": "uniform",
195
  "struct": [
196
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
197
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
198
  { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
199
  ]
200
  },
201
- "query_2": { "arg": "queryT", "name": "query", "buffer": "read-only-storage", "elementType": "f32" },
202
- "key_2": { "arg": "keyT", "name": "key", "buffer": "read-only-storage", "elementType": "f32" },
203
- "scores": { "scratch": "materializedScores", "buffer": "storage", "elementType": "f32" },
204
- "scorePartials": { "scratch": "materializedScorePartials", "buffer": "storage", "elementType": "f32" },
205
- "bias_2": { "arg": "biasT", "name": "bias", "buffer": "read-only-storage", "elementType": "f32" },
206
- "scorePartials_2": {
207
  "scratch": "materializedScorePartials",
208
  "name": "scorePartials",
209
  "buffer": "read-only-storage",
210
  "elementType": "f32"
211
  },
212
- "rowStats": { "scratch": "materializedRowStats", "buffer": "storage", "elementType": "f32" },
213
- "params_6": {
214
  "name": "params",
215
- "buffer": "uniform",
216
  "struct": [
217
  { "name": "rows", "type": "u32", "value": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)" }
218
  ]
219
  },
220
- "scores_2": {
221
  "scratch": "materializedScores",
222
  "name": "scores",
223
  "buffer": "read-only-storage",
224
  "elementType": "f32"
225
  },
226
- "value_2": { "arg": "valueT", "name": "value", "buffer": "read-only-storage", "elementType": "f32" },
227
- "output_2": { "arg": "outputT", "name": "output", "buffer": "storage", "elementType": "f32" },
228
- "rowStats_2": {
229
  "scratch": "materializedRowStats",
230
  "name": "rowStats",
231
  "buffer": "read-only-storage",
232
  "elementType": "f32"
233
  },
234
- "params_7": {
235
  "name": "params",
236
- "buffer": "uniform",
237
  "struct": [
238
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
239
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }
240
  ]
241
  },
242
- "attn_mask_2": {
243
- "arg": "attentionBiasT",
244
- "name": "attn_mask",
245
- "buffer": "read-only-storage",
246
- "elementType": "$maskElement"
247
- },
248
- "params_8": {
249
  "name": "params",
250
- "buffer": "uniform",
251
  "struct": [
252
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
253
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
@@ -266,53 +265,109 @@
266
  { "name": "maskSeqStride", "type": "u32", "value": "dim(shapes.attentionBiasT, 3)" }
267
  ]
268
  },
269
- "query_3": { "arg": "queryT", "name": "query", "buffer": "read-only-storage", "elementType": "$inputVec4" },
270
- "key_3": { "arg": "keyT", "name": "key", "buffer": "read-only-storage", "elementType": "$inputVec4" },
271
- "value_3": { "arg": "valueT", "name": "value", "buffer": "read-only-storage", "elementType": "$inputVec4" },
272
- "partial_out": { "scratch": "partialOut", "buffer": "storage", "elementType": "vec4<f32>" },
273
- "partial_stats": { "scratch": "partialStats", "buffer": "storage", "elementType": "vec2<f32>" },
274
- "params_9": {
275
  "name": "params",
276
- "buffer": "uniform",
277
  "struct": [
278
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
279
  { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
280
  ]
281
  },
282
- "partial_out_2": {
283
  "scratch": "partialOut",
284
  "name": "partial_out",
285
  "buffer": "read-only-storage",
286
  "elementType": "vec4<f32>"
287
  },
288
- "partial_stats_2": {
289
  "scratch": "partialStats",
290
  "name": "partial_stats",
291
  "buffer": "read-only-storage",
292
  "elementType": "vec2<f32>"
293
  },
294
- "output_3": { "arg": "outputT", "name": "output", "buffer": "storage", "elementType": "$inputVec4" },
295
- "scores_3": {
296
- "scratch": "materializedScores",
297
- "name": "scores",
298
- "buffer": "storage",
299
- "elementType": "$softmaxElementType"
300
- }
301
  },
302
  "variants": [
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
303
  {
304
  "id": "qkv_bias_small_seq_blocked",
305
  "priority": 40,
306
  "when": ["biasOk", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
307
  "derive": {
308
- "usesF16": false,
309
- "scalar": "\"f32\"",
310
  "inputElement": "\"vec4<f32>\"",
311
  "outputElement": "\"vec4<f32>\"",
312
  "inputScalar": "\"f32\"",
313
- "hasBias": true,
314
- "headDim": "headDimPlan",
315
- "headDimV4": "headDimPlan / 4",
316
  "hidden": "dim(shapes.queryT, 2)",
317
  "hiddenV4": "dim(shapes.queryT, 2) / 4",
318
  "kvSeq": "dim(shapes.keyT, 1)",
@@ -329,12 +384,6 @@
329
  "x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
330
  "y": "attrs.num_heads",
331
  "z": "dim(shapes.queryT, 0)"
332
- },
333
- "profile": {
334
- "op": "\"multi_head_attention\"",
335
- "variant": "\"qkv_bias_small_seq_blocked\"",
336
- "numHeads": "attrs.num_heads",
337
- "headDim": "headDim"
338
  }
339
  }
340
  ]
@@ -344,14 +393,9 @@
344
  "priority": 40,
345
  "when": ["not present.biasT", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
346
  "derive": {
347
- "usesF16": false,
348
- "scalar": "\"f32\"",
349
  "inputElement": "\"vec4<f32>\"",
350
  "outputElement": "\"vec4<f32>\"",
351
- "inputScalar": "\"f32\"",
352
- "hasBias": false,
353
- "headDim": "headDimPlan",
354
- "headDimV4": "headDimPlan / 4",
355
  "hidden": "dim(shapes.queryT, 2)",
356
  "hiddenV4": "dim(shapes.queryT, 2) / 4",
357
  "kvSeq": "dim(shapes.keyT, 1)",
@@ -368,12 +412,6 @@
368
  "x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
369
  "y": "attrs.num_heads",
370
  "z": "dim(shapes.queryT, 0)"
371
- },
372
- "profile": {
373
- "op": "\"multi_head_attention\"",
374
- "variant": "\"qkv_no_bias_small_seq_blocked\"",
375
- "numHeads": "attrs.num_heads",
376
- "headDim": "headDim"
377
  }
378
  }
379
  ]
@@ -383,7 +421,6 @@
383
  "priority": 60,
384
  "when": ["not present.biasT", "smallSeqShapeOk", "not flashShapeOk", "smallSeqResourcesFit", "smallSeqDispatchFits"],
385
  "derive": {
386
- "usesF16": false,
387
  "inputElement": "\"f32\"",
388
  "outputElement": "\"f32\"",
389
  "headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
@@ -396,25 +433,18 @@
396
  "id": "main",
397
  "name": "MultiHeadAttention",
398
  "shader": "mha-small-seq.wgsl.jinja",
399
- "bindings": ["query", "key", "value", "output", "params_2"],
400
- "dispatch": { "x": "attrs.num_heads", "y": "dim(shapes.queryT, 0)" },
401
- "profile": {
402
- "op": "\"multi_head_attention\"",
403
- "variant": "\"qkv_no_bias_small_seq\"",
404
- "numHeads": "attrs.num_heads",
405
- "headDim": "dim(shapes.queryT, 2) / attrs.num_heads"
406
- }
407
  }
408
  ]
409
  },
410
  {
411
  "id": "qkv_no_bias_tiled_nosg",
412
  "priority": 19,
413
- "when": ["not present.biasT", "qkvContractOk", "not smallHeadParallelOk", "headDimPlan % 4 == 0", "headDimPlan <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"],
414
  "supersededBy": ["qkv_no_bias_flash_cluster_nosg", "qkv_no_bias_flash_cluster_lpq4_nosg"],
415
  "derive": {
416
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
417
- "usesF16": "tensorDtypes.queryT == \"float16\"",
418
  "blockM": "tunables.PREFILL_QUERY_TILE",
419
  "vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
420
  },
@@ -449,8 +479,8 @@
449
  }
450
  ],
451
  "dispatch": {
452
- "x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / 32) * 32), (32)), 65535)",
453
- "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / 32) * 32), (32)), 65535)",
454
  "z": 1
455
  }
456
  }
@@ -459,14 +489,12 @@
459
  {
460
  "id": "qkv_bias_flash_q32_broadcast_f32_d128",
461
  "priority": 30,
462
- "when": ["tensorDtypes.queryT == \"float32\"", "biasOk", "attrs.unidirectional == 0", "flashShapeOk", "q32BroadcastF32RegisterGeometry", "dim(shapes.queryT, 1) >= 31", "ceilDiv(dim(shapes.queryT, 1), 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "subgroupsWave32"],
463
  "requires": { "features": ["subgroups"] },
464
  "derive": {
465
- "hasBias": true,
466
  "hasCausal": false,
467
  "usesF16": false,
468
  "scalar": "\"f32\"",
469
- "inputVec4": "\"vec4<f32>\"",
470
  "inputElement": "\"vec4<f32>\"",
471
  "outputElement": "\"vec4<f32>\"",
472
  "inputScalar": "\"f32\"",
@@ -484,13 +512,12 @@
484
  "name": "MultiHeadAttention.FlashQ32BroadcastF32Bias",
485
  "shader": "attn-flash-q32-broadcast.wgsl.jinja",
486
  "derive": { "layout": "\"bsh\"" },
487
- "bindings": ["query", "key", "value", "bias", "output", "params_4"],
488
  "dispatch": {
489
  "x": "ceilDiv(dim(shapes.queryT, 1), 32)",
490
  "y": "attrs.num_heads",
491
  "z": "dim(shapes.queryT, 0)"
492
- },
493
- "subgroupCollectivesWidth": 32
494
  }
495
  ]
496
  },
@@ -499,8 +526,6 @@
499
  "priority": 10,
500
  "when": ["not present.biasT", "smallHeadParallelOk"],
501
  "derive": {
502
- "usesF16": "tensorDtypes.queryT == \"float16\"",
503
- "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
504
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
505
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
506
  "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -511,9 +536,9 @@
511
  "passes": [
512
  {
513
  "id": "main",
514
- "name": "MultiHeadAttentionSmallHeadParallel",
515
  "shader": "attn-small-head-parallel.wgsl.jinja",
516
- "bindings": ["query", "key", "value", "output", "params_2"],
517
  "dispatch": {
518
  "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
519
  "y": "attrs.num_heads",
@@ -525,12 +550,9 @@
525
  {
526
  "id": "qkv_no_bias_tiled_attn_bias_nosg",
527
  "priority": 17,
528
- "when": ["not present.biasT", "qkvMaskContractOk", "not smallHeadParallelOk", "headDimPlan % 4 == 0", "headDimPlan <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"],
529
  "derive": {
530
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
531
- "usesF16": "tensorDtypes.queryT == \"float16\"",
532
- "hasMask": true,
533
- "maskIsBool": false,
534
  "blockM": "tunables.PREFILL_QUERY_TILE",
535
  "vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
536
  },
@@ -577,8 +599,8 @@
577
  }
578
  ],
579
  "dispatch": {
580
- "x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / 32) * 32), (32)), 65535)",
581
- "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / 32) * 32), (32)), 65535)",
582
  "z": 1
583
  }
584
  }
@@ -594,10 +616,7 @@
594
  },
595
  "derive": {
596
  "qNumHeads": "attrs.num_heads",
597
- "headDim": "headDimPlan",
598
  "qHidden": "dim(shapes.queryT, 2)",
599
- "hasBias": false,
600
- "useSubgroups": true,
601
  "statSlots": "materializedSgmatStatSlots",
602
  "statQuerySeq": "dim(shapes.queryT, 1)"
603
  },
@@ -616,19 +635,18 @@
616
  "name": "MultiHeadAttention.MaterializedScoresSgmat",
617
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
618
  "derive": { "phase": "\"score\"", "emitRowStats": true },
619
- "bindings": ["query_2", "key_2", "scores", "scorePartials", "params_4"],
620
  "dispatch": {
621
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
622
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
623
  "z": "dim(shapes.queryT, 0) * attrs.num_heads"
624
- },
625
- "subgroupCollectivesWidth": 32
626
  },
627
  {
628
  "id": "rowstats",
629
  "name": "MultiHeadAttention.MaterializedRowStatsCombine",
630
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
631
- "bindings": ["scorePartials_2", "rowStats", "params_6"],
632
  "dispatch": {
633
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
634
  "y": 1,
@@ -640,7 +658,7 @@
640
  "name": "MultiHeadAttention.MaterializedApplySgmat",
641
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
642
  "derive": { "phase": "\"apply\"", "fusedSoftmax": true },
643
- "bindings": ["scores_2", "value_2", "output_2", "rowStats_2", "params_7"],
644
  "dispatch": {
645
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
646
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
@@ -659,10 +677,7 @@
659
  },
660
  "derive": {
661
  "qNumHeads": "attrs.num_heads",
662
- "headDim": "headDimPlan",
663
  "qHidden": "dim(shapes.queryT, 2)",
664
- "hasBias": true,
665
- "useSubgroups": true,
666
  "statSlots": "materializedSgmatStatSlots",
667
  "statQuerySeq": "dim(shapes.queryT, 1)"
668
  },
@@ -681,19 +696,18 @@
681
  "name": "MultiHeadAttention.MaterializedScoresSgmatBias",
682
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
683
  "derive": { "phase": "\"score\"", "emitRowStats": true },
684
- "bindings": ["query_2", "key_2", "bias_2", "scores", "scorePartials", "params_4"],
685
  "dispatch": {
686
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
687
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
688
  "z": "dim(shapes.queryT, 0) * attrs.num_heads"
689
- },
690
- "subgroupCollectivesWidth": 32
691
  },
692
  {
693
  "id": "rowstats",
694
  "name": "MultiHeadAttention.MaterializedRowStatsCombineBias",
695
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
696
- "bindings": ["scorePartials_2", "rowStats", "params_6"],
697
  "dispatch": {
698
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
699
  "y": 1,
@@ -705,7 +719,7 @@
705
  "name": "MultiHeadAttention.MaterializedApplySgmatBias",
706
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
707
  "derive": { "phase": "\"apply\"", "fusedSoftmax": true },
708
- "bindings": ["scores_2", "value_2", "bias_2", "output_2", "rowStats_2", "params_7"],
709
  "dispatch": {
710
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
711
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
@@ -724,10 +738,7 @@
724
  },
725
  "derive": {
726
  "qNumHeads": "attrs.num_heads",
727
- "headDim": "headDimPlan",
728
  "qHidden": "dim(shapes.queryT, 2)",
729
- "hasBias": false,
730
- "useSubgroups": true,
731
  "statSlots": "materializedSgmatStatSlots",
732
  "statQuerySeq": "dim(shapes.queryT, 1)",
733
  "operandF16": true,
@@ -749,19 +760,18 @@
749
  "name": "MultiHeadAttention.MaterializedScoresSgmatF16",
750
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
751
  "derive": { "phase": "\"score\"", "emitRowStats": true },
752
- "bindings": ["query", "key", "scores", "scorePartials", "params_4"],
753
  "dispatch": {
754
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
755
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
756
  "z": "dim(shapes.queryT, 0) * attrs.num_heads"
757
- },
758
- "subgroupCollectivesWidth": 32
759
  },
760
  {
761
  "id": "rowstats",
762
  "name": "MultiHeadAttention.MaterializedRowStatsCombine",
763
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
764
- "bindings": ["scorePartials_2", "rowStats", "params_6"],
765
  "dispatch": {
766
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
767
  "y": 1,
@@ -773,7 +783,7 @@
773
  "name": "MultiHeadAttention.MaterializedApplySgmatF16",
774
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
775
  "derive": { "phase": "\"apply\"", "fusedSoftmax": true },
776
- "bindings": ["scores_2", "value", "output", "rowStats_2", "params_7"],
777
  "dispatch": {
778
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
779
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
@@ -785,15 +795,12 @@
785
  {
786
  "id": "qkv_no_bias_flash_cluster_lpq4_nosg",
787
  "priority": 20,
788
- "when": ["not present.biasT", "flashShapeOk", "headDimPlan % 16 == 0", "headDimPlan % 32 != 0", "dim(shapes.queryT, 1) >= 31"],
789
  "requires": {},
790
  "derive": {
791
- "hasBias": false,
792
  "hasCausal": true,
793
- "hasWindow": false,
794
  "usesF16": "tensorDtypes.queryT == \"float16\"",
795
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
796
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
797
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
798
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
799
  "qNumHeads": "attrs.num_heads",
@@ -826,15 +833,12 @@
826
  {
827
  "id": "qkv_no_bias_flash_cluster_nosg",
828
  "priority": 20,
829
- "when": ["not present.biasT", "flashShapeOk", "headDimPlan % 32 == 0", "dim(shapes.queryT, 1) >= 31"],
830
  "requires": {},
831
  "derive": {
832
- "hasBias": false,
833
  "hasCausal": true,
834
- "hasWindow": false,
835
  "usesF16": "tensorDtypes.queryT == \"float16\"",
836
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
837
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
838
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
839
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
840
  "qNumHeads": "attrs.num_heads",
@@ -867,15 +871,12 @@
867
  {
868
  "id": "qkv_bias_flash_cluster_nosg",
869
  "priority": 19,
870
- "when": ["biasOk", "flashShapeOk", "headDimPlan % 32 == 0", "dim(shapes.queryT, 1) >= 31"],
871
  "requires": {},
872
  "derive": {
873
- "hasBias": true,
874
  "hasCausal": true,
875
- "hasWindow": false,
876
  "usesF16": "tensorDtypes.queryT == \"float16\"",
877
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
878
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
879
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
880
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
881
  "qNumHeads": "attrs.num_heads",
@@ -910,15 +911,12 @@
910
  {
911
  "id": "qkv_no_bias_flash_cluster_lpq4",
912
  "priority": 22,
913
- "when": ["not present.biasT", "flashShapeOk", "headDimPlan % 16 == 0", "headDimPlan % 32 != 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster4"],
914
  "requires": { "features": ["subgroups"] },
915
  "derive": {
916
- "hasBias": false,
917
  "hasCausal": true,
918
- "hasWindow": false,
919
  "usesF16": "tensorDtypes.queryT == \"float16\"",
920
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
921
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
922
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
923
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
924
  "qNumHeads": "attrs.num_heads",
@@ -942,23 +940,19 @@
942
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
943
  "y": "attrs.num_heads",
944
  "z": "dim(shapes.queryT, 0)"
945
- },
946
- "subgroupCollectivesWidth": "portable"
947
  }
948
  ]
949
  },
950
  {
951
  "id": "qkv_no_bias_flash_cluster",
952
  "priority": 22,
953
- "when": ["not present.biasT", "flashShapeOk", "headDimPlan % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
954
  "requires": { "features": ["subgroups"] },
955
  "derive": {
956
- "hasBias": false,
957
  "hasCausal": true,
958
- "hasWindow": false,
959
  "usesF16": "tensorDtypes.queryT == \"float16\"",
960
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
961
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
962
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
963
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
964
  "qNumHeads": "attrs.num_heads",
@@ -982,23 +976,19 @@
982
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
983
  "y": "attrs.num_heads",
984
  "z": "dim(shapes.queryT, 0)"
985
- },
986
- "subgroupCollectivesWidth": "portable"
987
  }
988
  ]
989
  },
990
  {
991
  "id": "qkv_bias_flash_cluster",
992
  "priority": 21,
993
- "when": ["biasOk", "flashShapeOk", "headDimPlan % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
994
  "requires": { "features": ["subgroups"] },
995
  "derive": {
996
- "hasBias": true,
997
  "hasCausal": true,
998
- "hasWindow": false,
999
  "usesF16": "tensorDtypes.queryT == \"float16\"",
1000
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1001
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1002
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1003
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1004
  "qNumHeads": "attrs.num_heads",
@@ -1024,25 +1014,19 @@
1024
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
1025
  "y": "attrs.num_heads",
1026
  "z": "dim(shapes.queryT, 0)"
1027
- },
1028
- "subgroupCollectivesWidth": "portable"
1029
  }
1030
  ]
1031
  },
1032
  {
1033
  "id": "qkv_no_bias_flash_cluster_attn_bias",
1034
  "priority": 22,
1035
- "when": ["not present.biasT", "flashMaskShapeOk", "headDimPlan % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
1036
  "requires": { "features": ["subgroups"] },
1037
  "derive": {
1038
- "hasBias": false,
1039
  "hasCausal": true,
1040
- "hasWindow": false,
1041
- "hasMask": true,
1042
- "maskIsBool": false,
1043
  "usesF16": "tensorDtypes.queryT == \"float16\"",
1044
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1045
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1046
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1047
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1048
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1062,30 +1046,24 @@
1062
  "name": "MultiHeadAttention.Flash",
1063
  "shader": "attn-flash-prefill-cluster.wgsl.jinja",
1064
  "derive": { "layout": "\"bsh\"" },
1065
- "bindings": ["query", "key", "value", "attn_mask_2", "output", "params_8"],
1066
  "dispatch": {
1067
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
1068
  "y": "attrs.num_heads",
1069
  "z": "dim(shapes.queryT, 0)"
1070
- },
1071
- "subgroupCollectivesWidth": "portable"
1072
  }
1073
  ]
1074
  },
1075
  {
1076
  "id": "qkv_bias_flash_cluster_attn_bias",
1077
  "priority": 21,
1078
- "when": ["biasOk", "flashMaskShapeOk", "headDimPlan % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
1079
  "requires": { "features": ["subgroups"] },
1080
  "derive": {
1081
- "hasBias": true,
1082
  "hasCausal": true,
1083
- "hasWindow": false,
1084
- "hasMask": true,
1085
- "maskIsBool": false,
1086
  "usesF16": "tensorDtypes.queryT == \"float16\"",
1087
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1088
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1089
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1090
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1091
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1107,13 +1085,12 @@
1107
  "name": "MultiHeadAttention.Flash",
1108
  "shader": "attn-flash-prefill-cluster.wgsl.jinja",
1109
  "derive": { "layout": "\"bsh\"" },
1110
- "bindings": ["query", "key", "value", "attn_mask_2", "bias", "output", "params_8"],
1111
  "dispatch": {
1112
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
1113
  "y": "attrs.num_heads",
1114
  "z": "dim(shapes.queryT, 0)"
1115
- },
1116
- "subgroupCollectivesWidth": "portable"
1117
  }
1118
  ]
1119
  },
@@ -1123,9 +1100,7 @@
1123
  "when": ["not present.biasT", "decodeSplitKNoBiasOk or shortQuerySplitKNoBiasOk"],
1124
  "requires": {},
1125
  "derive": {
1126
- "combineSubgroups": false,
1127
  "useSubgroups": false,
1128
- "hasWindow": false,
1129
  "splitQueries": true,
1130
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1131
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
@@ -1157,7 +1132,7 @@
1157
  "name": "MultiHeadAttention.DecodeSplitKNoSg",
1158
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1159
  "derive": { "layout": "\"bsh\"" },
1160
- "bindings": ["query_3", "key_3", "value_3", "partial_out", "partial_stats", "params_9"],
1161
  "dispatch": {
1162
  "x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
1163
  "y": "attrs.num_heads",
@@ -1169,7 +1144,7 @@
1169
  "name": "MultiHeadAttention.DecodeSplitKMergeNoSg",
1170
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1171
  "derive": { "layout": "\"bsh\"" },
1172
- "bindings": ["partial_out_2", "partial_stats_2", "output_3"],
1173
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1174
  }
1175
  ]
@@ -1180,10 +1155,7 @@
1180
  "when": ["biasOk", "decodeSplitKBiasOk"],
1181
  "requires": {},
1182
  "derive": {
1183
- "combineSubgroups": false,
1184
  "useSubgroups": false,
1185
- "hasBias": true,
1186
- "hasWindow": false,
1187
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1188
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1189
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1215,7 +1187,7 @@
1215
  "name": "MultiHeadAttention.DecodeSplitKBiasNoSg",
1216
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1217
  "derive": { "layout": "\"bsh\"" },
1218
- "bindings": ["query_3", "key_3", "value_3", "bias", "partial_out", "partial_stats", "params_9"],
1219
  "dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1220
  },
1221
  {
@@ -1223,7 +1195,7 @@
1223
  "name": "MultiHeadAttention.DecodeSplitKMergeBiasNoSg",
1224
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1225
  "derive": { "layout": "\"bsh\"" },
1226
- "bindings": ["partial_out_2", "partial_stats_2", "bias", "output_3"],
1227
  "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1228
  }
1229
  ]
@@ -1235,8 +1207,6 @@
1235
  "demoteWhen": ["decodeSplitKPortablePreferred"],
1236
  "requires": { "features": ["subgroups"] },
1237
  "derive": {
1238
- "combineSubgroups": true,
1239
- "hasWindow": false,
1240
  "splitQueries": true,
1241
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1242
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
@@ -1268,20 +1238,19 @@
1268
  "name": "MultiHeadAttention.DecodeSplitK",
1269
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1270
  "derive": { "layout": "\"bsh\"" },
1271
- "bindings": ["query_3", "key_3", "value_3", "partial_out", "partial_stats", "params_9"],
1272
  "dispatch": {
1273
  "x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
1274
  "y": "attrs.num_heads",
1275
  "z": "dim(shapes.queryT, 0)"
1276
- },
1277
- "subgroupCollectivesWidth": "portable"
1278
  },
1279
  {
1280
  "id": "merge",
1281
  "name": "MultiHeadAttention.DecodeSplitKMerge",
1282
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1283
  "derive": { "layout": "\"bsh\"" },
1284
- "bindings": ["partial_out_2", "partial_stats_2", "output_3"],
1285
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1286
  }
1287
  ]
@@ -1292,9 +1261,6 @@
1292
  "when": ["biasOk", "decodeSplitKBiasOk"],
1293
  "requires": { "features": ["subgroups"] },
1294
  "derive": {
1295
- "combineSubgroups": true,
1296
- "hasBias": true,
1297
- "hasWindow": false,
1298
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1299
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1300
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1326,16 +1292,15 @@
1326
  "name": "MultiHeadAttention.DecodeSplitKBias",
1327
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1328
  "derive": { "layout": "\"bsh\"" },
1329
- "bindings": ["query_3", "key_3", "value_3", "bias", "partial_out", "partial_stats", "params_9"],
1330
- "dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" },
1331
- "subgroupCollectivesWidth": "portable"
1332
  },
1333
  {
1334
  "id": "merge",
1335
  "name": "MultiHeadAttention.DecodeSplitKMergeBias",
1336
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1337
  "derive": { "layout": "\"bsh\"" },
1338
- "bindings": ["partial_out_2", "partial_stats_2", "bias", "output_3"],
1339
  "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1340
  }
1341
  ]
@@ -1348,7 +1313,6 @@
1348
  "derive": {
1349
  "qNumHeads": "attrs.num_heads",
1350
  "kvNumHeads": "attrs.num_heads",
1351
- "headDim": "headDimPlan",
1352
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1353
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1354
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
@@ -1358,7 +1322,6 @@
1358
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
1359
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
1360
  "applyTileN": "materializedApplyTileN",
1361
- "hasBias": false,
1362
  "useSubgroups": "device.features.has(\"subgroups\")"
1363
  },
1364
  "intermediates": [
@@ -1374,7 +1337,7 @@
1374
  "name": "MultiHeadAttention.MaterializedScores",
1375
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1376
  "derive": { "layout": "\"bsh\"" },
1377
- "bindings": ["query_2", "key_2", "scores", "params_4"],
1378
  "dispatch": {
1379
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1380
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1386,20 +1349,19 @@
1386
  "name": "MultiHeadAttention.MaterializedSoftmax",
1387
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
1388
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
1389
- "bindings": ["scores_3", "params_6"],
1390
  "dispatch": {
1391
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1392
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1393
  "z": 1
1394
- },
1395
- "subgroupCollectivesWidth": "portable"
1396
  },
1397
  {
1398
  "id": "apply",
1399
  "name": "MultiHeadAttention.MaterializedApply",
1400
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1401
  "derive": { "layout": "\"bsh\"" },
1402
- "bindings": ["scores_2", "value_2", "output_2", "params_7"],
1403
  "dispatch": {
1404
  "x": "ceilDiv(headDim, applyTileN)",
1405
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1415,7 +1377,6 @@
1415
  "derive": {
1416
  "qNumHeads": "attrs.num_heads",
1417
  "kvNumHeads": "attrs.num_heads",
1418
- "headDim": "headDimPlan",
1419
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1420
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1421
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
@@ -1425,7 +1386,6 @@
1425
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
1426
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
1427
  "applyTileN": "materializedApplyTileN",
1428
- "hasBias": true,
1429
  "useSubgroups": "device.features.has(\"subgroups\")"
1430
  },
1431
  "intermediates": [
@@ -1441,7 +1401,7 @@
1441
  "name": "MultiHeadAttention.MaterializedScoresBias",
1442
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1443
  "derive": { "layout": "\"bsh\"" },
1444
- "bindings": ["query_2", "key_2", "bias_2", "scores", "params_4"],
1445
  "dispatch": {
1446
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1447
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1453,20 +1413,19 @@
1453
  "name": "MultiHeadAttention.MaterializedSoftmaxBias",
1454
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
1455
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
1456
- "bindings": ["scores_3", "params_6"],
1457
  "dispatch": {
1458
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1459
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1460
  "z": 1
1461
- },
1462
- "subgroupCollectivesWidth": "portable"
1463
  },
1464
  {
1465
  "id": "apply",
1466
  "name": "MultiHeadAttention.MaterializedApplyBias",
1467
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1468
  "derive": { "layout": "\"bsh\"" },
1469
- "bindings": ["scores_2", "value_2", "bias_2", "output_2", "params_7"],
1470
  "dispatch": {
1471
  "x": "ceilDiv(headDim, applyTileN)",
1472
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1483,14 +1442,11 @@
1483
  "derive": {
1484
  "qNumHeads": "attrs.num_heads",
1485
  "kvNumHeads": "attrs.num_heads",
1486
- "headDim": "headDimPlan",
1487
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1488
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1489
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
1490
  "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
1491
  "applyTileN": "materializedApplyTileN",
1492
- "hasBias": false,
1493
- "useSubgroups": "device.features.has(\"subgroups\")",
1494
  "statSlots": "materializedGemmStatSlots",
1495
  "statQuerySeq": "dim(shapes.queryT, 1)"
1496
  },
@@ -1509,7 +1465,7 @@
1509
  "name": "MultiHeadAttention.MaterializedScoresFused",
1510
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1511
  "derive": { "layout": "\"bsh\"", "emitRowStats": true },
1512
- "bindings": ["query_2", "key_2", "scores", "scorePartials", "params_4"],
1513
  "dispatch": {
1514
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1515
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1521,7 +1477,7 @@
1521
  "name": "MultiHeadAttention.MaterializedGemmRowStatsCombine",
1522
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
1523
  "derive": { "maxOnly": true },
1524
- "bindings": ["scorePartials_2", "rowStats", "params_6"],
1525
  "dispatch": {
1526
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1527
  "y": 1,
@@ -1533,7 +1489,7 @@
1533
  "name": "MultiHeadAttention.MaterializedApplyFused",
1534
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1535
  "derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
1536
- "bindings": ["scores_2", "value_2", "output_2", "rowStats_2", "params_7"],
1537
  "dispatch": {
1538
  "x": "ceilDiv(headDim, applyTileN)",
1539
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1549,14 +1505,11 @@
1549
  "derive": {
1550
  "qNumHeads": "attrs.num_heads",
1551
  "kvNumHeads": "attrs.num_heads",
1552
- "headDim": "headDimPlan",
1553
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1554
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1555
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
1556
  "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
1557
  "applyTileN": "materializedApplyTileN",
1558
- "hasBias": true,
1559
- "useSubgroups": "device.features.has(\"subgroups\")",
1560
  "statSlots": "materializedGemmStatSlots",
1561
  "statQuerySeq": "dim(shapes.queryT, 1)"
1562
  },
@@ -1575,7 +1528,7 @@
1575
  "name": "MultiHeadAttention.MaterializedScoresBiasFused",
1576
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1577
  "derive": { "layout": "\"bsh\"", "emitRowStats": true },
1578
- "bindings": ["query_2", "key_2", "bias_2", "scores", "scorePartials", "params_4"],
1579
  "dispatch": {
1580
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1581
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1587,7 +1540,7 @@
1587
  "name": "MultiHeadAttention.MaterializedGemmRowStatsCombineBias",
1588
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
1589
  "derive": { "maxOnly": true },
1590
- "bindings": ["scorePartials_2", "rowStats", "params_6"],
1591
  "dispatch": {
1592
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1593
  "y": 1,
@@ -1599,7 +1552,7 @@
1599
  "name": "MultiHeadAttention.MaterializedApplyBiasFused",
1600
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1601
  "derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
1602
- "bindings": ["scores_2", "value_2", "bias_2", "output_2", "rowStats_2", "params_7"],
1603
  "dispatch": {
1604
  "x": "ceilDiv(headDim, applyTileN)",
1605
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
@@ -1611,12 +1564,11 @@
1611
  {
1612
  "id": "qkv_no_bias_flash_q32_broadcast",
1613
  "priority": 30,
1614
- "when": ["tensorDtypes.queryT == \"float16\"", "not present.biasT", "flashShapeOk", "(dim(shapes.queryT, 2) / attrs.num_heads) % 4 == 0", "(dim(shapes.queryT, 2) / attrs.num_heads) >= 64", "(dim(shapes.queryT, 2) / attrs.num_heads) <= 256", "dim(shapes.queryT, 1) >= 31", "subgroupsWave32"],
1615
  "requires": { "features": ["subgroups", "shader-f16"] },
1616
  "derive": {
1617
  "usesF16": true,
1618
  "scalar": "\"f16\"",
1619
- "inputVec4": "\"vec4<f16>\"",
1620
  "inputElement": "\"vec4<f16>\"",
1621
  "outputElement": "\"vec4<f16>\"",
1622
  "qNumHeads": "attrs.num_heads",
@@ -1637,8 +1589,7 @@
1637
  "x": "ceilDiv(dim(shapes.queryT, 1), 32)",
1638
  "y": "attrs.num_heads",
1639
  "z": "dim(shapes.queryT, 0)"
1640
- },
1641
- "subgroupCollectivesWidth": 32
1642
  }
1643
  ]
1644
  },
@@ -1650,7 +1601,6 @@
1650
  "derive": {
1651
  "usesF16": true,
1652
  "scalar": "\"f16\"",
1653
- "inputVec4": "\"vec4<f16>\"",
1654
  "inputElement": "\"vec4<f16>\"",
1655
  "outputElement": "\"vec4<f16>\"",
1656
  "qNumHeads": "attrs.num_heads",
@@ -1681,14 +1631,8 @@
1681
  "priority": 0,
1682
  "when": ["not present.biasT and fallbackMaskShapeOk"],
1683
  "derive": {
1684
- "headsFromParams": false,
1685
- "hasBias": false,
1686
  "hasCausal": true,
1687
- "hasWindow": false,
1688
  "hasKeyLimit": false,
1689
- "hasMask": true,
1690
- "scaleFallbackRsqrt": true,
1691
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1692
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1693
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1694
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1707,7 +1651,7 @@
1707
  "name": "MultiHeadAttention",
1708
  "shader": "attn-online-scalar.wgsl.jinja",
1709
  "derive": { "layout": "\"bsh\"" },
1710
- "bindings": ["query", "key", "value", "attn_mask_2", "output", "params_8"],
1711
  "dispatch": {
1712
  "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1713
  "y": "attrs.num_heads",
@@ -1721,14 +1665,8 @@
1721
  "priority": 0,
1722
  "when": ["biasOk and fallbackMaskShapeOk"],
1723
  "derive": {
1724
- "headsFromParams": false,
1725
- "hasBias": true,
1726
  "hasCausal": true,
1727
- "hasWindow": false,
1728
  "hasKeyLimit": false,
1729
- "hasMask": true,
1730
- "scaleFallbackRsqrt": true,
1731
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1732
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1733
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1734
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1747,7 +1685,7 @@
1747
  "name": "MultiHeadAttention",
1748
  "shader": "attn-online-scalar.wgsl.jinja",
1749
  "derive": { "layout": "\"bsh\"" },
1750
- "bindings": ["query", "key", "value", "attn_mask_2", "bias", "output", "params_8"],
1751
  "dispatch": {
1752
  "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1753
  "y": "attrs.num_heads",
@@ -1761,13 +1699,8 @@
1761
  "priority": 0,
1762
  "when": ["not present.biasT and fallbackShapeOk and not flashShapeOk"],
1763
  "derive": {
1764
- "headsFromParams": false,
1765
- "hasBias": false,
1766
  "hasCausal": true,
1767
- "hasWindow": false,
1768
  "hasKeyLimit": false,
1769
- "scaleFallbackRsqrt": true,
1770
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1771
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1772
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1773
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1799,13 +1732,8 @@
1799
  "priority": 0,
1800
  "when": ["biasOk and fallbackShapeOk and not flashShapeOk"],
1801
  "derive": {
1802
- "headsFromParams": false,
1803
- "hasBias": true,
1804
  "hasCausal": true,
1805
- "hasWindow": false,
1806
  "hasKeyLimit": false,
1807
- "scaleFallbackRsqrt": true,
1808
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1809
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1810
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1811
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1837,17 +1765,10 @@
1837
  "priority": 20,
1838
  "when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashShapeOk"],
1839
  "derive": {
1840
- "headsFromParams": false,
1841
- "hasBias": false,
1842
  "hasCausal": true,
1843
- "hasWindow": false,
1844
  "combineSubgroups": true,
1845
- "hasMask": false,
1846
- "maskIsBool": false,
1847
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1848
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1849
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1850
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1851
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1852
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1853
  "qNumHeads": "attrs.num_heads",
@@ -1864,8 +1785,7 @@
1864
  "shader": "attn-flash-online.wgsl.jinja",
1865
  "derive": { "layout": "\"bsh\"" },
1866
  "bindings": ["query", "key", "value", "output", "params"],
1867
- "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" },
1868
- "subgroupCollectivesWidth": "portable"
1869
  }
1870
  ]
1871
  },
@@ -1874,17 +1794,10 @@
1874
  "priority": 20,
1875
  "when": ["device.features.has(\"subgroups\")", "biasOk", "flashShapeOk"],
1876
  "derive": {
1877
- "headsFromParams": false,
1878
- "hasBias": true,
1879
  "hasCausal": true,
1880
- "hasWindow": false,
1881
  "combineSubgroups": true,
1882
- "hasMask": false,
1883
- "maskIsBool": false,
1884
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1885
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1886
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1887
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1888
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1889
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1890
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1903,8 +1816,7 @@
1903
  "shader": "attn-flash-online.wgsl.jinja",
1904
  "derive": { "layout": "\"bsh\"" },
1905
  "bindings": ["query", "key", "value", "bias", "output", "params"],
1906
- "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" },
1907
- "subgroupCollectivesWidth": "portable"
1908
  }
1909
  ]
1910
  },
@@ -1913,17 +1825,10 @@
1913
  "priority": 18,
1914
  "when": ["true", "not present.biasT", "flashShapeOk"],
1915
  "derive": {
1916
- "headsFromParams": false,
1917
- "hasBias": false,
1918
  "hasCausal": true,
1919
- "hasWindow": false,
1920
  "combineSubgroups": false,
1921
- "hasMask": false,
1922
- "maskIsBool": false,
1923
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1924
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1925
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1926
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1927
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1928
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1929
  "qNumHeads": "attrs.num_heads",
@@ -1949,17 +1854,10 @@
1949
  "priority": 17,
1950
  "when": ["true", "biasOk", "flashShapeOk"],
1951
  "derive": {
1952
- "headsFromParams": false,
1953
- "hasBias": true,
1954
  "hasCausal": true,
1955
- "hasWindow": false,
1956
  "combineSubgroups": false,
1957
- "hasMask": false,
1958
- "maskIsBool": false,
1959
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1960
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1961
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1962
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1963
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1964
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1965
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -1987,17 +1885,10 @@
1987
  "priority": 20,
1988
  "when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashMaskShapeOk"],
1989
  "derive": {
1990
- "headsFromParams": false,
1991
- "hasBias": false,
1992
  "hasCausal": true,
1993
- "hasWindow": false,
1994
  "combineSubgroups": true,
1995
- "hasMask": true,
1996
- "maskIsBool": false,
1997
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1998
- "usesF16": "tensorDtypes.queryT == \"float16\"",
1999
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
2000
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2001
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2002
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2003
  "qNumHeads": "attrs.num_heads",
@@ -2013,9 +1904,8 @@
2013
  "name": "MultiHeadAttention.Flash",
2014
  "shader": "attn-flash-online.wgsl.jinja",
2015
  "derive": { "layout": "\"bsh\"" },
2016
- "bindings": ["query", "key", "value", "attn_mask_2", "output", "params_8"],
2017
- "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" },
2018
- "subgroupCollectivesWidth": "portable"
2019
  }
2020
  ]
2021
  },
@@ -2024,17 +1914,10 @@
2024
  "priority": 20,
2025
  "when": ["device.features.has(\"subgroups\")", "biasOk", "flashMaskShapeOk"],
2026
  "derive": {
2027
- "headsFromParams": false,
2028
- "hasBias": true,
2029
  "hasCausal": true,
2030
- "hasWindow": false,
2031
  "combineSubgroups": true,
2032
- "hasMask": true,
2033
- "maskIsBool": false,
2034
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
2035
- "usesF16": "tensorDtypes.queryT == \"float16\"",
2036
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
2037
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2038
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2039
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2040
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -2052,9 +1935,8 @@
2052
  "name": "MultiHeadAttention.Flash",
2053
  "shader": "attn-flash-online.wgsl.jinja",
2054
  "derive": { "layout": "\"bsh\"" },
2055
- "bindings": ["query", "key", "value", "attn_mask_2", "bias", "output", "params_8"],
2056
- "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" },
2057
- "subgroupCollectivesWidth": "portable"
2058
  }
2059
  ]
2060
  },
@@ -2063,17 +1945,10 @@
2063
  "priority": 18,
2064
  "when": ["true", "not present.biasT", "flashMaskShapeOk"],
2065
  "derive": {
2066
- "headsFromParams": false,
2067
- "hasBias": false,
2068
  "hasCausal": true,
2069
- "hasWindow": false,
2070
  "combineSubgroups": false,
2071
- "hasMask": true,
2072
- "maskIsBool": false,
2073
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
2074
- "usesF16": "tensorDtypes.queryT == \"float16\"",
2075
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
2076
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2077
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2078
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2079
  "qNumHeads": "attrs.num_heads",
@@ -2089,7 +1964,7 @@
2089
  "name": "MultiHeadAttention.NoBiasOnlineFlashNoSg",
2090
  "shader": "attn-flash-online.wgsl.jinja",
2091
  "derive": { "layout": "\"bsh\"" },
2092
- "bindings": ["query", "key", "value", "attn_mask_2", "output", "params_8"],
2093
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
2094
  }
2095
  ]
@@ -2099,17 +1974,10 @@
2099
  "priority": 17,
2100
  "when": ["true", "biasOk", "flashMaskShapeOk"],
2101
  "derive": {
2102
- "headsFromParams": false,
2103
- "hasBias": true,
2104
  "hasCausal": true,
2105
- "hasWindow": false,
2106
  "combineSubgroups": false,
2107
- "hasMask": true,
2108
- "maskIsBool": false,
2109
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
2110
- "usesF16": "tensorDtypes.queryT == \"float16\"",
2111
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
2112
- "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2113
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2114
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
2115
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
@@ -2127,7 +1995,7 @@
2127
  "name": "MultiHeadAttention.BiasOnlineFlashNoSg",
2128
  "shader": "attn-flash-online.wgsl.jinja",
2129
  "derive": { "layout": "\"bsh\"" },
2130
- "bindings": ["query", "key", "value", "attn_mask_2", "bias", "output", "params_8"],
2131
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
2132
  }
2133
  ]
@@ -2142,13 +2010,11 @@
2142
  },
2143
  "derive": {
2144
  "qNumHeads": "attrs.num_heads",
2145
- "headDim": "headDimPlan",
2146
  "qHidden": "dim(shapes.queryT, 2)",
2147
  "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
2148
  "materializedSoftmaxCols": "dim(shapes.keyT, 1)",
2149
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
2150
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
2151
- "hasBias": false,
2152
  "useSubgroups": true
2153
  },
2154
  "intermediates": [
@@ -2164,7 +2030,7 @@
2164
  "name": "MultiHeadAttention.MaterializedScoresSgmat",
2165
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2166
  "derive": { "phase": "\"score\"" },
2167
- "bindings": ["query_2", "key_2", "scores", "params_4"],
2168
  "dispatch": {
2169
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
2170
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
@@ -2176,20 +2042,19 @@
2176
  "name": "MultiHeadAttention.MaterializedSoftmaxSgmat",
2177
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
2178
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
2179
- "bindings": ["scores_3", "params_6"],
2180
  "dispatch": {
2181
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2182
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2183
  "z": 1
2184
- },
2185
- "subgroupCollectivesWidth": "portable"
2186
  },
2187
  {
2188
  "id": "apply",
2189
  "name": "MultiHeadAttention.MaterializedApplySgmat",
2190
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2191
  "derive": { "phase": "\"apply\"" },
2192
- "bindings": ["scores_2", "value_2", "output_2", "params_7"],
2193
  "dispatch": {
2194
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
2195
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
@@ -2208,13 +2073,11 @@
2208
  },
2209
  "derive": {
2210
  "qNumHeads": "attrs.num_heads",
2211
- "headDim": "headDimPlan",
2212
  "qHidden": "dim(shapes.queryT, 2)",
2213
  "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
2214
  "materializedSoftmaxCols": "dim(shapes.keyT, 1)",
2215
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
2216
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
2217
- "hasBias": true,
2218
  "useSubgroups": true
2219
  },
2220
  "intermediates": [
@@ -2230,7 +2093,7 @@
2230
  "name": "MultiHeadAttention.MaterializedScoresSgmatBias",
2231
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2232
  "derive": { "phase": "\"score\"" },
2233
- "bindings": ["query_2", "key_2", "bias_2", "scores", "params_4"],
2234
  "dispatch": {
2235
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
2236
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
@@ -2242,20 +2105,19 @@
2242
  "name": "MultiHeadAttention.MaterializedSoftmaxSgmatBias",
2243
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
2244
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
2245
- "bindings": ["scores_3", "params_6"],
2246
  "dispatch": {
2247
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2248
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2249
  "z": 1
2250
- },
2251
- "subgroupCollectivesWidth": "portable"
2252
  },
2253
  {
2254
  "id": "apply",
2255
  "name": "MultiHeadAttention.MaterializedApplySgmatBias",
2256
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2257
  "derive": { "phase": "\"apply\"" },
2258
- "bindings": ["scores_2", "value_2", "bias_2", "output_2", "params_7"],
2259
  "dispatch": {
2260
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
2261
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
 
57
  "canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32",
58
  "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
59
  "wave32Effective": "wave32Adapter or pinSubgroupSize32",
60
+ "wave32SubgroupsUsable": "subgroupsWave32 or pinSubgroupSize32",
61
+ "headDim": "dim(shapes.queryT, 2) / attrs.num_heads if (ranks.queryT == 3 and attrs.num_heads > 0) else 0",
62
  "qkvDtypesOk": "tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT",
63
  "floatDtypeOk": "(tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(tensorDtypes.queryT)",
64
  "qkvShapeOk": "ranks.queryT == 3 and ranks.keyT == 3 and ranks.valueT == 3 and ranks.outputT == 3 and attrs.num_heads > 0 and dim(shapes.queryT, 2) % attrs.num_heads == 0 and dim(shapes.keyT, 2) == dim(shapes.queryT, 2) and dim(shapes.valueT, 2) == dim(shapes.queryT, 2) and dim(shapes.keyT, 1) == dim(shapes.valueT, 1) and dim(shapes.queryT, 0) == dim(shapes.keyT, 0) and dim(shapes.queryT, 0) == dim(shapes.valueT, 0) and dim(shapes.outputT, 0) == dim(shapes.queryT, 0) and dim(shapes.outputT, 1) == dim(shapes.queryT, 1) and dim(shapes.outputT, 2) == dim(shapes.valueT, 2)",
 
67
  "qkvContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and noAttnBias",
68
  "qkvMaskContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and attnBiasOk",
69
  "biasOk": "present.biasT and ranks.biasT == 1 and tensorDtypes.biasT == tensorDtypes.queryT and dim(shapes.biasT, 0) == 3 * dim(shapes.queryT, 2)",
70
+ "hasBias": "present.biasT",
71
+ "hasMask": "present.attentionBiasT",
72
+ "q32BroadcastSubgroupLanes": "32 if wave32SubgroupsUsable else 0",
73
+ "q32BroadcastF32HeadVectors": "headDim / 4 if headDim % 4 == 0 else 0",
74
+ "q32BroadcastF32RegisterGeometry": "wave32SubgroupsUsable and q32BroadcastF32HeadVectors == q32BroadcastSubgroupLanes",
75
  "subgroupCluster4": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 4 and device.adapterInfo.subgroupMinSize % 4 == 0 and device.adapterInfo.subgroupMaxSize % 4 == 0",
76
  "subgroupCluster8": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 8 and device.adapterInfo.subgroupMinSize % 8 == 0 and device.adapterInfo.subgroupMaxSize % 8 == 0",
77
  "attentionDispatchFits": "dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 1) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
78
+ "flashHeadOk": "headDim % 4 == 0 and headDim >= 32 and headDim <= 256",
79
  "flashSizeOk": "flashHeadOk and (dim(shapes.queryT, 1) * attrs.num_heads >= tunables.FLASH_MIN_QUERY_HEADS or (dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512) or (dim(shapes.queryT, 1) > 1 and dim(shapes.keyT, 1) >= 2048)) and attentionDispatchFits",
80
  "flashShapeOk": "qkvContractOk and flashSizeOk",
81
  "flashMaskShapeOk": "qkvMaskContractOk and flashSizeOk",
82
  "noBiasSplitKCount": "min(tunables.DECODE_MAX_SPLITS if dim(shapes.queryT, 1) == 1 else max(1, ceilDiv(tunables.SPLITK_TARGET_WORKGROUPS, dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads)), ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
83
  "biasSplitKCount": "min(tunables.DECODE_MAX_SPLITS, ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
84
+ "noBiasPartialOutBytes": "dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * headDim * 4",
85
  "noBiasStatsBytes": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * 4",
86
  "noBiasSplitScratchFits": "noBiasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and noBiasPartialOutBytes <= device.limits.maxBufferSize and noBiasStatsBytes <= device.limits.maxStorageBufferBindingSize and noBiasStatsBytes <= device.limits.maxBufferSize",
87
  "noBiasSplitDispatchFits": "dim(shapes.queryT, 1) * noBiasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
88
+ "biasPartialOutBytes": "dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * headDim * 4",
89
  "biasStatsBytes": "2 * dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * 4",
90
  "biasSplitScratchFits": "biasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and biasPartialOutBytes <= device.limits.maxBufferSize and biasStatsBytes <= device.limits.maxStorageBufferBindingSize and biasStatsBytes <= device.limits.maxBufferSize",
91
  "biasSplitDispatchFits": "biasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
92
  "decodeSplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
93
  "shortQuerySplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) >= 2 and dim(shapes.queryT, 1) <= 16 and dim(shapes.keyT, 1) >= 2048 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
94
  "decodeSplitKBiasOk": "biasOk and qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and biasSplitDispatchFits and biasSplitScratchFits",
95
+ "decodeSplitKPortablePreferred": "tensorDtypes.queryT == \"float32\" and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and headDim / 4 < device.adapterInfo.subgroupMinSize",
96
  "materializedScoreBytes": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1) * 4",
97
  "materializedScoreFits": "materializedScoreBytes <= device.limits.maxStorageBufferBindingSize and materializedScoreBytes <= device.limits.maxBufferSize",
98
  "materializedWorkgroupSize": "tunables.MATERIALIZED_WORKGROUP_DIM * tunables.MATERIALIZED_WORKGROUP_DIM",
99
  "materializedScoreStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_KEY_TILE * (tunables.MATERIALIZED_INNER_TILE + 4)) * 4",
100
+ "materializedApplyTileN": "tunables.MATERIALIZED_VALUE_TILE_D128 if headDim == 128 else tunables.MATERIALIZED_VALUE_TILE",
101
  "materializedApplyStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) * 4",
102
  "materializedTileGeometryOk": "tunables.MATERIALIZED_INNER_TILE % 4 == 0 and (materializedApplyTileN / tunables.MATERIALIZED_WORKGROUP_DIM) % 4 == 0 and tunables.MATERIALIZED_INNER_TILE > 0 and tunables.MATERIALIZED_WORKGROUP_DIM > 0 and tunables.MATERIALIZED_QUERY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_KEY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and materializedApplyTileN % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN >= materializedWorkgroupSize and (tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) % materializedWorkgroupSize == 0",
103
  "materializedDeviceOk": "materializedTileGeometryOk and materializedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeX and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeY and materializedScoreStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and materializedApplyStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
104
  "materializedWideSimdOk": "device.features.has(\"subgroups\") or (has(device.adapterInfo, \"subgroupMinSize\") and device.adapterInfo.subgroupMinSize >= 16)",
105
+ "materializedF32CoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDim >= 64 and headDim <= 128 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and materializedScoreFits and materializedDeviceOk",
106
  "materializedSoftmaxStorageBytes": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE * 8 + 8",
107
  "materializedSoftmaxResourcesFit": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE > 0 and pow2ceil(tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE) == tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE and tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE <= deviceWorkgroupCap and materializedSoftmaxStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
108
  "materializedSgmatQueryTile": "tunables.MATERIALIZED_SGMAT_QUERY_TILE",
 
126
  "materializedSgmatResourcesFit": "materializedSgmatGeometryOk and materializedSgmatWorkgroupSize <= deviceWorkgroupCap and materializedSgmatCompactStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
127
  "materializedSgmatDispatchFits": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
128
  "materializedSgmatDirectScoreStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and dim(shapes.keyT, 1) % materializedSgmatKeyTile == 0",
129
+ "materializedSgmatDirectApplyStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and headDim % materializedSgmatKeyTile == 0",
130
  "materializedSgmatRuntimeDirectStore": "dim(shapes.queryT, 1) >= 2 * materializedSgmatQueryTile and dim(shapes.keyT, 1) >= 2 * materializedSgmatKeyTile",
131
+ "materializedSgmatCoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDim >= 32 and headDim <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits",
132
  "materializedCachedSoftmaxWg": "min(tunables.MATERIALIZED_CACHED_SOFTMAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
133
  "materializedCachedSoftmaxVecsPerLane": "ceilDiv(ceilDiv(dim(shapes.keyT, 1), 4), max(1, materializedCachedSoftmaxWg))",
134
  "materializedCachedSoftmaxStorageBytes": "materializedCachedSoftmaxWg * 8 + 8",
 
138
  "materializedSgmatOk": "materializedSgmatCoreOk and materializedAdaptiveSoftmaxOk",
139
  "materializedSgmatFusedOk": "materializedSgmatCoreOk",
140
  "materializedF32Ok": "materializedF32CoreOk and materializedAdaptiveSoftmaxOk",
141
+ "clusterTileKWg64": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDim * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_SMALL * 4))))",
142
+ "clusterTileKWg128": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDim * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_LARGE * 4))))",
143
+ "smallHeadShapeOk": "qkvContractOk and headDim < 32 and dim(shapes.keyT, 1) >= 64 and attrs.unidirectional == 0 and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
144
+ "smallHeadParallelOk": "smallHeadShapeOk and dim(shapes.keyT, 1) <= 2048",
145
+ "prefillTiledStorageBytes": "headDim * tunables.PREFILL_QUERY_TILE * 4",
146
  "prefillTiledDeviceOk": "tunables.PREFILL_QUERY_TILE <= deviceWorkgroupCap and prefillTiledStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
147
+ "portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, pow2ceil(max(1, headDim)))",
148
+ "portableWorkgroupStorageBytes": "portableWorkgroupSize * 8 + max(1, headDim) * 4 + 16",
149
  "portableWorkgroupOk": "tunables.WORKGROUP_SIZE > 0 and pow2ceil(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and portableWorkgroupSize <= deviceWorkgroupCap and portableWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
150
  "fallbackShapeOk": "qkvContractOk and portableWorkgroupOk",
151
  "fallbackMaskShapeOk": "qkvMaskContractOk and portableWorkgroupOk",
152
+ "smallSeqShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and dim(shapes.queryT, 1) >= 1 and dim(shapes.queryT, 1) <= tunables.SMALL_SEQ_MAX and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_MAX and headDim >= 1 and attrs.unidirectional == 0",
153
+ "smallSeqPrivateFloats": "dim(shapes.keyT, 1) + headDim",
154
  "smallSeqWorkgroupSize": "max(32, pow2ceil(dim(shapes.queryT, 1)))",
155
+ "smallSeqSharedBytes": "dim(shapes.keyT, 1) * headDim * 8",
156
  "smallSeqResourcesFit": "smallSeqPrivateFloats <= tunables.SMALL_SEQ_MAX_PRIVATE_FLOATS and smallSeqWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqSharedBytes <= device.limits.maxComputeWorkgroupStorageSize",
157
+ "smallSeqBlockedKvBytes": "dim(shapes.keyT, 1) * headDim * 8",
158
+ "smallSeqBlockedLaneBytes": "8 + headDim * 4",
159
  "smallSeqBlockedKeyLanes": "min(pow2ceil(dim(shapes.keyT, 1)), 16 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 16 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else (8 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 8 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else 4))",
160
  "smallSeqBlockedWorkgroupSize": "tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * smallSeqBlockedKeyLanes",
161
  "smallSeqBlockedSharedBytes": "smallSeqBlockedKvBytes + smallSeqBlockedWorkgroupSize * smallSeqBlockedLaneBytes",
162
+ "smallSeqBlockedShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and headDim % 4 == 0 and headDim >= 4 and headDim <= tunables.SMALL_SEQ_BLOCKED_MAX_HEAD_DIM and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_BLOCKED_MAX_KV and dim(shapes.queryT, 1) >= 1",
163
  "smallSeqBlockedFits": "smallSeqBlockedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqBlockedWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqBlockedSharedBytes <= device.limits.maxComputeWorkgroupStorageSize and ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
164
  "smallSeqDispatchFits": "attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
165
+ "materializedSgmatCoreF16Ok": "qkvContractOk and tensorDtypes.queryT == \"float16\" and tensorDtypes.keyT == \"float16\" and tensorDtypes.valueT == \"float16\" and device.features.has(\"shader-f16\") and attrs.unidirectional == 0 and headDim >= 32 and headDim <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits",
166
  "materializedSgmatFusedF16Ok": "materializedSgmatCoreF16Ok",
167
+ "attentionScaleExpression": "\"select(inverseSqrt(f32(HEAD_DIM)), params.scale, params.scale != 0.0)\"",
168
+ "smallHeadValueWg": "pow(2, log2ceil(min(64, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX) + 1) - 1)",
169
+ "smallHeadValueStorage": "max(dim(shapes.keyT, 1), smallHeadValueWg * headDim)",
170
+ "smallHeadValueQueryBlock": "2 if dim(shapes.queryT, 1) >= 2 and (smallHeadValueStorage + smallHeadValueWg) * 8 <= device.limits.maxComputeWorkgroupStorageSize else 1",
171
+ "smallHeadValueEstimatedSteps": "headDim * (ceilDiv(dim(shapes.keyT, 1), smallHeadValueWg) + 2 * log2ceil(smallHeadValueWg) * (2 if tensorDtypes.queryT == \"float16\" else 1) + 2)",
172
+ "smallHeadValueSgReductionSteps": "ceilDiv(smallHeadValueWg, max(1, device.adapterInfo.subgroupMinSize)) if has(device.adapterInfo, \"subgroupMinSize\") else log2ceil(smallHeadValueWg)",
173
+ "smallHeadValueSgEstimatedSteps": "headDim * (ceilDiv(dim(shapes.keyT, 1), smallHeadValueWg) + 2 * smallHeadValueSgReductionSteps * (2 if tensorDtypes.queryT == \"float16\" else 1) + 2)"
174
  },
175
  "bindings": {
176
+ "query": { "arg": "queryT", "elementType": "$inputElement" },
177
+ "key": { "arg": "keyT", "elementType": "$inputElement" },
178
+ "value": { "arg": "valueT", "elementType": "$inputElement" },
179
+ "bias": { "arg": "biasT", "elementType": "$inputScalar" },
180
+ "output": { "arg": "outputT", "elementType": "$outputElement" },
181
  "params": {
 
182
  "struct": [
183
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
184
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
 
186
  { "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" }
187
  ]
188
  },
189
+ "params_main": {
190
  "name": "params",
 
191
  "struct": [
192
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
193
  { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
194
  ]
195
  },
196
+ "q": { "arg": "queryT", "elementType": "$scalar" },
197
+ "k": { "arg": "keyT", "elementType": "$scalar" },
198
+ "v": { "arg": "valueT", "elementType": "$scalar" },
199
+ "y": { "arg": "outputT", "elementType": "$scalar" },
200
+ "params_scores": {
201
  "name": "params",
 
202
  "struct": [
203
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
204
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
205
  { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
206
  ]
207
  },
208
+ "query_f32": { "arg": "queryT", "name": "query", "elementType": "f32" },
209
+ "key_f32": { "arg": "keyT", "name": "key", "elementType": "f32" },
210
+ "scores": { "scratch": "materializedScores", "elementType": "f32" },
211
+ "scorePartials": { "scratch": "materializedScorePartials", "elementType": "f32" },
212
+ "bias_f32": { "arg": "biasT", "name": "bias", "elementType": "f32" },
213
+ "scorePartials_f32": {
214
  "scratch": "materializedScorePartials",
215
  "name": "scorePartials",
216
  "buffer": "read-only-storage",
217
  "elementType": "f32"
218
  },
219
+ "rowStats": { "scratch": "materializedRowStats", "elementType": "f32" },
220
+ "params_rows": {
221
  "name": "params",
 
222
  "struct": [
223
  { "name": "rows", "type": "u32", "value": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)" }
224
  ]
225
  },
226
+ "scores_f32": {
227
  "scratch": "materializedScores",
228
  "name": "scores",
229
  "buffer": "read-only-storage",
230
  "elementType": "f32"
231
  },
232
+ "value_f32": { "arg": "valueT", "name": "value", "elementType": "f32" },
233
+ "output_f32": { "arg": "outputT", "name": "output", "elementType": "f32" },
234
+ "rowStats_f32": {
235
  "scratch": "materializedRowStats",
236
  "name": "rowStats",
237
  "buffer": "read-only-storage",
238
  "elementType": "f32"
239
  },
240
+ "params_apply": {
241
  "name": "params",
 
242
  "struct": [
243
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
244
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }
245
  ]
246
  },
247
+ "attn_mask_main": { "arg": "attentionBiasT", "name": "attn_mask", "elementType": "$maskElement" },
248
+ "params__uniform": {
 
 
 
 
 
249
  "name": "params",
 
250
  "struct": [
251
  { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
252
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
 
265
  { "name": "maskSeqStride", "type": "u32", "value": "dim(shapes.attentionBiasT, 3)" }
266
  ]
267
  },
268
+ "query_query_t": { "arg": "queryT", "name": "query", "elementType": "$inputVec4" },
269
+ "key_key_t": { "arg": "keyT", "name": "key", "elementType": "$inputVec4" },
270
+ "value_value_t": { "arg": "valueT", "name": "value", "elementType": "$inputVec4" },
271
+ "partial_out": { "scratch": "partialOut", "elementType": "vec4<f32>" },
272
+ "partial_stats": { "scratch": "partialStats", "elementType": "vec2<f32>" },
273
+ "params_kv_seq_scale": {
274
  "name": "params",
 
275
  "struct": [
276
  { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
277
  { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
278
  ]
279
  },
280
+ "partial_out_merge": {
281
  "scratch": "partialOut",
282
  "name": "partial_out",
283
  "buffer": "read-only-storage",
284
  "elementType": "vec4<f32>"
285
  },
286
+ "partial_stats_merge": {
287
  "scratch": "partialStats",
288
  "name": "partial_stats",
289
  "buffer": "read-only-storage",
290
  "elementType": "vec2<f32>"
291
  },
292
+ "output_merge": { "arg": "outputT", "name": "output", "elementType": "$inputVec4" },
293
+ "scores_softmax": { "scratch": "materializedScores", "name": "scores", "elementType": "$softmaxElementType" }
 
 
 
 
 
294
  },
295
  "variants": [
296
+ {
297
+ "id": "qkv_no_bias_small_head_value_subgroups",
298
+ "priority": 12,
299
+ "when": ["not present.biasT", "smallHeadShapeOk", "headDim > 0", "headDim <= smallHeadValueWg", "(smallHeadValueStorage + smallHeadValueWg) * 4 * smallHeadValueQueryBlock <= device.limits.maxComputeWorkgroupStorageSize", "device.wgslLanguageFeatures.has(\"subgroup_id\")"],
300
+ "derive": {
301
+ "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
302
+ "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
303
+ "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
304
+ "qHidden": "dim(shapes.queryT, 2)",
305
+ "headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
306
+ "kvSeq": "dim(shapes.keyT, 1)",
307
+ "valueWorkgroupSize": "smallHeadValueWg",
308
+ "scoreStorageElements": "smallHeadValueStorage",
309
+ "queryBlock": "smallHeadValueQueryBlock",
310
+ "queryTail": "dim(shapes.queryT, 1) % smallHeadValueQueryBlock != 0",
311
+ "useValueSubgroups": true
312
+ },
313
+ "passes": [
314
+ {
315
+ "id": "main",
316
+ "name": "MultiHeadAttention.SmallHeadValueSubgroups",
317
+ "shader": "attn-small-head-value.wgsl.jinja",
318
+ "bindings": ["query", "key", "value", "output", "params_main"],
319
+ "dispatch": {
320
+ "x": "min(ceilDiv(dim(shapes.queryT, 1), smallHeadValueQueryBlock), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
321
+ "y": "attrs.num_heads",
322
+ "z": "dim(shapes.queryT, 0)"
323
+ }
324
+ }
325
+ ],
326
+ "requires": { "features": ["subgroups"] },
327
+ "demoteWhen": ["dim(shapes.keyT, 1) <= smallHeadValueSgEstimatedSteps"]
328
+ },
329
+ {
330
+ "id": "qkv_no_bias_small_head_value_tree",
331
+ "priority": 11,
332
+ "when": ["not present.biasT", "smallHeadShapeOk", "headDim > 0", "headDim <= smallHeadValueWg", "(smallHeadValueStorage + smallHeadValueWg) * 4 * smallHeadValueQueryBlock <= device.limits.maxComputeWorkgroupStorageSize", "true"],
333
+ "derive": {
334
+ "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
335
+ "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
336
+ "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
337
+ "qHidden": "dim(shapes.queryT, 2)",
338
+ "headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
339
+ "kvSeq": "dim(shapes.keyT, 1)",
340
+ "valueWorkgroupSize": "smallHeadValueWg",
341
+ "scoreStorageElements": "smallHeadValueStorage",
342
+ "queryBlock": "smallHeadValueQueryBlock",
343
+ "queryTail": "dim(shapes.queryT, 1) % smallHeadValueQueryBlock != 0",
344
+ "useValueSubgroups": false
345
+ },
346
+ "passes": [
347
+ {
348
+ "id": "main",
349
+ "name": "MultiHeadAttention.SmallHeadValueTree",
350
+ "shader": "attn-small-head-value.wgsl.jinja",
351
+ "bindings": ["query", "key", "value", "output", "params_main"],
352
+ "dispatch": {
353
+ "x": "min(ceilDiv(dim(shapes.queryT, 1), smallHeadValueQueryBlock), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
354
+ "y": "attrs.num_heads",
355
+ "z": "dim(shapes.queryT, 0)"
356
+ }
357
+ }
358
+ ],
359
+ "requires": { "features": [] },
360
+ "demoteWhen": ["dim(shapes.keyT, 1) <= smallHeadValueEstimatedSteps"]
361
+ },
362
  {
363
  "id": "qkv_bias_small_seq_blocked",
364
  "priority": 40,
365
  "when": ["biasOk", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
366
  "derive": {
 
 
367
  "inputElement": "\"vec4<f32>\"",
368
  "outputElement": "\"vec4<f32>\"",
369
  "inputScalar": "\"f32\"",
370
+ "headDimV4": "headDim / 4",
 
 
371
  "hidden": "dim(shapes.queryT, 2)",
372
  "hiddenV4": "dim(shapes.queryT, 2) / 4",
373
  "kvSeq": "dim(shapes.keyT, 1)",
 
384
  "x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
385
  "y": "attrs.num_heads",
386
  "z": "dim(shapes.queryT, 0)"
 
 
 
 
 
 
387
  }
388
  }
389
  ]
 
393
  "priority": 40,
394
  "when": ["not present.biasT", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
395
  "derive": {
 
 
396
  "inputElement": "\"vec4<f32>\"",
397
  "outputElement": "\"vec4<f32>\"",
398
+ "headDimV4": "headDim / 4",
 
 
 
399
  "hidden": "dim(shapes.queryT, 2)",
400
  "hiddenV4": "dim(shapes.queryT, 2) / 4",
401
  "kvSeq": "dim(shapes.keyT, 1)",
 
412
  "x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
413
  "y": "attrs.num_heads",
414
  "z": "dim(shapes.queryT, 0)"
 
 
 
 
 
 
415
  }
416
  }
417
  ]
 
421
  "priority": 60,
422
  "when": ["not present.biasT", "smallSeqShapeOk", "not flashShapeOk", "smallSeqResourcesFit", "smallSeqDispatchFits"],
423
  "derive": {
 
424
  "inputElement": "\"f32\"",
425
  "outputElement": "\"f32\"",
426
  "headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
 
433
  "id": "main",
434
  "name": "MultiHeadAttention",
435
  "shader": "mha-small-seq.wgsl.jinja",
436
+ "bindings": ["query", "key", "value", "output", "params_main"],
437
+ "dispatch": { "x": "attrs.num_heads", "y": "dim(shapes.queryT, 0)" }
 
 
 
 
 
 
438
  }
439
  ]
440
  },
441
  {
442
  "id": "qkv_no_bias_tiled_nosg",
443
  "priority": 19,
444
+ "when": ["not present.biasT", "qkvContractOk", "not smallHeadParallelOk", "headDim % 4 == 0", "headDim <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"],
445
  "supersededBy": ["qkv_no_bias_flash_cluster_nosg", "qkv_no_bias_flash_cluster_lpq4_nosg"],
446
  "derive": {
447
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
448
  "blockM": "tunables.PREFILL_QUERY_TILE",
449
  "vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
450
  },
 
479
  }
480
  ],
481
  "dispatch": {
482
+ "x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
483
+ "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
484
  "z": 1
485
  }
486
  }
 
489
  {
490
  "id": "qkv_bias_flash_q32_broadcast_f32_d128",
491
  "priority": 30,
492
+ "when": ["tensorDtypes.queryT == \"float32\"", "biasOk", "attrs.unidirectional == 0", "flashShapeOk", "q32BroadcastF32RegisterGeometry", "dim(shapes.queryT, 1) >= 31", "ceilDiv(dim(shapes.queryT, 1), 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32SubgroupsUsable"],
493
  "requires": { "features": ["subgroups"] },
494
  "derive": {
 
495
  "hasCausal": false,
496
  "usesF16": false,
497
  "scalar": "\"f32\"",
 
498
  "inputElement": "\"vec4<f32>\"",
499
  "outputElement": "\"vec4<f32>\"",
500
  "inputScalar": "\"f32\"",
 
512
  "name": "MultiHeadAttention.FlashQ32BroadcastF32Bias",
513
  "shader": "attn-flash-q32-broadcast.wgsl.jinja",
514
  "derive": { "layout": "\"bsh\"" },
515
+ "bindings": ["query", "key", "value", "bias", "output", "params_scores"],
516
  "dispatch": {
517
  "x": "ceilDiv(dim(shapes.queryT, 1), 32)",
518
  "y": "attrs.num_heads",
519
  "z": "dim(shapes.queryT, 0)"
520
+ }
 
521
  }
522
  ]
523
  },
 
526
  "priority": 10,
527
  "when": ["not present.biasT", "smallHeadParallelOk"],
528
  "derive": {
 
 
529
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
530
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
531
  "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
536
  "passes": [
537
  {
538
  "id": "main",
539
+ "name": "MultiHeadAttention.SmallHeadParallel",
540
  "shader": "attn-small-head-parallel.wgsl.jinja",
541
+ "bindings": ["query", "key", "value", "output", "params_main"],
542
  "dispatch": {
543
  "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
544
  "y": "attrs.num_heads",
 
550
  {
551
  "id": "qkv_no_bias_tiled_attn_bias_nosg",
552
  "priority": 17,
553
+ "when": ["not present.biasT", "qkvMaskContractOk", "not smallHeadParallelOk", "headDim % 4 == 0", "headDim <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"],
554
  "derive": {
555
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
 
 
556
  "blockM": "tunables.PREFILL_QUERY_TILE",
557
  "vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
558
  },
 
599
  }
600
  ],
601
  "dispatch": {
602
+ "x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
603
+ "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
604
  "z": 1
605
  }
606
  }
 
616
  },
617
  "derive": {
618
  "qNumHeads": "attrs.num_heads",
 
619
  "qHidden": "dim(shapes.queryT, 2)",
 
 
620
  "statSlots": "materializedSgmatStatSlots",
621
  "statQuerySeq": "dim(shapes.queryT, 1)"
622
  },
 
635
  "name": "MultiHeadAttention.MaterializedScoresSgmat",
636
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
637
  "derive": { "phase": "\"score\"", "emitRowStats": true },
638
+ "bindings": ["query_f32", "key_f32", "scores", "scorePartials", "params_scores"],
639
  "dispatch": {
640
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
641
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
642
  "z": "dim(shapes.queryT, 0) * attrs.num_heads"
643
+ }
 
644
  },
645
  {
646
  "id": "rowstats",
647
  "name": "MultiHeadAttention.MaterializedRowStatsCombine",
648
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
649
+ "bindings": ["scorePartials_f32", "rowStats", "params_rows"],
650
  "dispatch": {
651
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
652
  "y": 1,
 
658
  "name": "MultiHeadAttention.MaterializedApplySgmat",
659
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
660
  "derive": { "phase": "\"apply\"", "fusedSoftmax": true },
661
+ "bindings": ["scores_f32", "value_f32", "output_f32", "rowStats_f32", "params_apply"],
662
  "dispatch": {
663
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
664
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
 
677
  },
678
  "derive": {
679
  "qNumHeads": "attrs.num_heads",
 
680
  "qHidden": "dim(shapes.queryT, 2)",
 
 
681
  "statSlots": "materializedSgmatStatSlots",
682
  "statQuerySeq": "dim(shapes.queryT, 1)"
683
  },
 
696
  "name": "MultiHeadAttention.MaterializedScoresSgmatBias",
697
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
698
  "derive": { "phase": "\"score\"", "emitRowStats": true },
699
+ "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "scorePartials", "params_scores"],
700
  "dispatch": {
701
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
702
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
703
  "z": "dim(shapes.queryT, 0) * attrs.num_heads"
704
+ }
 
705
  },
706
  {
707
  "id": "rowstats",
708
  "name": "MultiHeadAttention.MaterializedRowStatsCombineBias",
709
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
710
+ "bindings": ["scorePartials_f32", "rowStats", "params_rows"],
711
  "dispatch": {
712
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
713
  "y": 1,
 
719
  "name": "MultiHeadAttention.MaterializedApplySgmatBias",
720
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
721
  "derive": { "phase": "\"apply\"", "fusedSoftmax": true },
722
+ "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "rowStats_f32", "params_apply"],
723
  "dispatch": {
724
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
725
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
 
738
  },
739
  "derive": {
740
  "qNumHeads": "attrs.num_heads",
 
741
  "qHidden": "dim(shapes.queryT, 2)",
 
 
742
  "statSlots": "materializedSgmatStatSlots",
743
  "statQuerySeq": "dim(shapes.queryT, 1)",
744
  "operandF16": true,
 
760
  "name": "MultiHeadAttention.MaterializedScoresSgmatF16",
761
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
762
  "derive": { "phase": "\"score\"", "emitRowStats": true },
763
+ "bindings": ["query", "key", "scores", "scorePartials", "params_scores"],
764
  "dispatch": {
765
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
766
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
767
  "z": "dim(shapes.queryT, 0) * attrs.num_heads"
768
+ }
 
769
  },
770
  {
771
  "id": "rowstats",
772
  "name": "MultiHeadAttention.MaterializedRowStatsCombine",
773
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
774
+ "bindings": ["scorePartials_f32", "rowStats", "params_rows"],
775
  "dispatch": {
776
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
777
  "y": 1,
 
783
  "name": "MultiHeadAttention.MaterializedApplySgmatF16",
784
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
785
  "derive": { "phase": "\"apply\"", "fusedSoftmax": true },
786
+ "bindings": ["scores_f32", "value", "output", "rowStats_f32", "params_apply"],
787
  "dispatch": {
788
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
789
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
 
795
  {
796
  "id": "qkv_no_bias_flash_cluster_lpq4_nosg",
797
  "priority": 20,
798
+ "when": ["not present.biasT", "flashShapeOk", "headDim % 16 == 0", "headDim % 32 != 0", "dim(shapes.queryT, 1) >= 31"],
799
  "requires": {},
800
  "derive": {
 
801
  "hasCausal": true,
 
802
  "usesF16": "tensorDtypes.queryT == \"float16\"",
803
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
804
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
805
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
806
  "qNumHeads": "attrs.num_heads",
 
833
  {
834
  "id": "qkv_no_bias_flash_cluster_nosg",
835
  "priority": 20,
836
+ "when": ["not present.biasT", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31"],
837
  "requires": {},
838
  "derive": {
 
839
  "hasCausal": true,
 
840
  "usesF16": "tensorDtypes.queryT == \"float16\"",
841
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
842
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
843
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
844
  "qNumHeads": "attrs.num_heads",
 
871
  {
872
  "id": "qkv_bias_flash_cluster_nosg",
873
  "priority": 19,
874
+ "when": ["biasOk", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31"],
875
  "requires": {},
876
  "derive": {
 
877
  "hasCausal": true,
 
878
  "usesF16": "tensorDtypes.queryT == \"float16\"",
879
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
880
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
881
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
882
  "qNumHeads": "attrs.num_heads",
 
911
  {
912
  "id": "qkv_no_bias_flash_cluster_lpq4",
913
  "priority": 22,
914
+ "when": ["not present.biasT", "flashShapeOk", "headDim % 16 == 0", "headDim % 32 != 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster4"],
915
  "requires": { "features": ["subgroups"] },
916
  "derive": {
 
917
  "hasCausal": true,
 
918
  "usesF16": "tensorDtypes.queryT == \"float16\"",
919
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
920
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
921
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
922
  "qNumHeads": "attrs.num_heads",
 
940
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
941
  "y": "attrs.num_heads",
942
  "z": "dim(shapes.queryT, 0)"
943
+ }
 
944
  }
945
  ]
946
  },
947
  {
948
  "id": "qkv_no_bias_flash_cluster",
949
  "priority": 22,
950
+ "when": ["not present.biasT", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
951
  "requires": { "features": ["subgroups"] },
952
  "derive": {
 
953
  "hasCausal": true,
 
954
  "usesF16": "tensorDtypes.queryT == \"float16\"",
955
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
956
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
957
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
958
  "qNumHeads": "attrs.num_heads",
 
976
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
977
  "y": "attrs.num_heads",
978
  "z": "dim(shapes.queryT, 0)"
979
+ }
 
980
  }
981
  ]
982
  },
983
  {
984
  "id": "qkv_bias_flash_cluster",
985
  "priority": 21,
986
+ "when": ["biasOk", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
987
  "requires": { "features": ["subgroups"] },
988
  "derive": {
 
989
  "hasCausal": true,
 
990
  "usesF16": "tensorDtypes.queryT == \"float16\"",
991
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
992
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
993
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
994
  "qNumHeads": "attrs.num_heads",
 
1014
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
1015
  "y": "attrs.num_heads",
1016
  "z": "dim(shapes.queryT, 0)"
1017
+ }
 
1018
  }
1019
  ]
1020
  },
1021
  {
1022
  "id": "qkv_no_bias_flash_cluster_attn_bias",
1023
  "priority": 22,
1024
+ "when": ["not present.biasT", "flashMaskShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
1025
  "requires": { "features": ["subgroups"] },
1026
  "derive": {
 
1027
  "hasCausal": true,
 
 
 
1028
  "usesF16": "tensorDtypes.queryT == \"float16\"",
1029
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1030
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1031
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1032
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1046
  "name": "MultiHeadAttention.Flash",
1047
  "shader": "attn-flash-prefill-cluster.wgsl.jinja",
1048
  "derive": { "layout": "\"bsh\"" },
1049
+ "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
1050
  "dispatch": {
1051
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
1052
  "y": "attrs.num_heads",
1053
  "z": "dim(shapes.queryT, 0)"
1054
+ }
 
1055
  }
1056
  ]
1057
  },
1058
  {
1059
  "id": "qkv_bias_flash_cluster_attn_bias",
1060
  "priority": 21,
1061
+ "when": ["biasOk", "flashMaskShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
1062
  "requires": { "features": ["subgroups"] },
1063
  "derive": {
 
1064
  "hasCausal": true,
 
 
 
1065
  "usesF16": "tensorDtypes.queryT == \"float16\"",
1066
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1067
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1068
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1069
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1085
  "name": "MultiHeadAttention.Flash",
1086
  "shader": "attn-flash-prefill-cluster.wgsl.jinja",
1087
  "derive": { "layout": "\"bsh\"" },
1088
+ "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
1089
  "dispatch": {
1090
  "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
1091
  "y": "attrs.num_heads",
1092
  "z": "dim(shapes.queryT, 0)"
1093
+ }
 
1094
  }
1095
  ]
1096
  },
 
1100
  "when": ["not present.biasT", "decodeSplitKNoBiasOk or shortQuerySplitKNoBiasOk"],
1101
  "requires": {},
1102
  "derive": {
 
1103
  "useSubgroups": false,
 
1104
  "splitQueries": true,
1105
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1106
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
 
1132
  "name": "MultiHeadAttention.DecodeSplitKNoSg",
1133
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1134
  "derive": { "layout": "\"bsh\"" },
1135
+ "bindings": ["query_query_t", "key_key_t", "value_value_t", "partial_out", "partial_stats", "params_kv_seq_scale"],
1136
  "dispatch": {
1137
  "x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
1138
  "y": "attrs.num_heads",
 
1144
  "name": "MultiHeadAttention.DecodeSplitKMergeNoSg",
1145
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1146
  "derive": { "layout": "\"bsh\"" },
1147
+ "bindings": ["partial_out_merge", "partial_stats_merge", "output_merge"],
1148
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1149
  }
1150
  ]
 
1155
  "when": ["biasOk", "decodeSplitKBiasOk"],
1156
  "requires": {},
1157
  "derive": {
 
1158
  "useSubgroups": false,
 
 
1159
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1160
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1161
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1187
  "name": "MultiHeadAttention.DecodeSplitKBiasNoSg",
1188
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1189
  "derive": { "layout": "\"bsh\"" },
1190
+ "bindings": ["query_query_t", "key_key_t", "value_value_t", "bias", "partial_out", "partial_stats", "params_kv_seq_scale"],
1191
  "dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1192
  },
1193
  {
 
1195
  "name": "MultiHeadAttention.DecodeSplitKMergeBiasNoSg",
1196
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1197
  "derive": { "layout": "\"bsh\"" },
1198
+ "bindings": ["partial_out_merge", "partial_stats_merge", "bias", "output_merge"],
1199
  "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1200
  }
1201
  ]
 
1207
  "demoteWhen": ["decodeSplitKPortablePreferred"],
1208
  "requires": { "features": ["subgroups"] },
1209
  "derive": {
 
 
1210
  "splitQueries": true,
1211
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1212
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
 
1238
  "name": "MultiHeadAttention.DecodeSplitK",
1239
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1240
  "derive": { "layout": "\"bsh\"" },
1241
+ "bindings": ["query_query_t", "key_key_t", "value_value_t", "partial_out", "partial_stats", "params_kv_seq_scale"],
1242
  "dispatch": {
1243
  "x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
1244
  "y": "attrs.num_heads",
1245
  "z": "dim(shapes.queryT, 0)"
1246
+ }
 
1247
  },
1248
  {
1249
  "id": "merge",
1250
  "name": "MultiHeadAttention.DecodeSplitKMerge",
1251
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1252
  "derive": { "layout": "\"bsh\"" },
1253
+ "bindings": ["partial_out_merge", "partial_stats_merge", "output_merge"],
1254
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1255
  }
1256
  ]
 
1261
  "when": ["biasOk", "decodeSplitKBiasOk"],
1262
  "requires": { "features": ["subgroups"] },
1263
  "derive": {
 
 
 
1264
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1265
  "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1266
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1292
  "name": "MultiHeadAttention.DecodeSplitKBias",
1293
  "shader": "attn-flash-decode-splitk.wgsl.jinja",
1294
  "derive": { "layout": "\"bsh\"" },
1295
+ "bindings": ["query_query_t", "key_key_t", "value_value_t", "bias", "partial_out", "partial_stats", "params_kv_seq_scale"],
1296
+ "dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
 
1297
  },
1298
  {
1299
  "id": "merge",
1300
  "name": "MultiHeadAttention.DecodeSplitKMergeBias",
1301
  "shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
1302
  "derive": { "layout": "\"bsh\"" },
1303
+ "bindings": ["partial_out_merge", "partial_stats_merge", "bias", "output_merge"],
1304
  "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1305
  }
1306
  ]
 
1313
  "derive": {
1314
  "qNumHeads": "attrs.num_heads",
1315
  "kvNumHeads": "attrs.num_heads",
 
1316
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1317
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1318
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
 
1322
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
1323
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
1324
  "applyTileN": "materializedApplyTileN",
 
1325
  "useSubgroups": "device.features.has(\"subgroups\")"
1326
  },
1327
  "intermediates": [
 
1337
  "name": "MultiHeadAttention.MaterializedScores",
1338
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1339
  "derive": { "layout": "\"bsh\"" },
1340
+ "bindings": ["query_f32", "key_f32", "scores", "params_scores"],
1341
  "dispatch": {
1342
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1343
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1349
  "name": "MultiHeadAttention.MaterializedSoftmax",
1350
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
1351
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
1352
+ "bindings": ["scores_softmax", "params_rows"],
1353
  "dispatch": {
1354
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1355
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1356
  "z": 1
1357
+ }
 
1358
  },
1359
  {
1360
  "id": "apply",
1361
  "name": "MultiHeadAttention.MaterializedApply",
1362
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1363
  "derive": { "layout": "\"bsh\"" },
1364
+ "bindings": ["scores_f32", "value_f32", "output_f32", "params_apply"],
1365
  "dispatch": {
1366
  "x": "ceilDiv(headDim, applyTileN)",
1367
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1377
  "derive": {
1378
  "qNumHeads": "attrs.num_heads",
1379
  "kvNumHeads": "attrs.num_heads",
 
1380
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1381
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1382
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
 
1386
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
1387
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
1388
  "applyTileN": "materializedApplyTileN",
 
1389
  "useSubgroups": "device.features.has(\"subgroups\")"
1390
  },
1391
  "intermediates": [
 
1401
  "name": "MultiHeadAttention.MaterializedScoresBias",
1402
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1403
  "derive": { "layout": "\"bsh\"" },
1404
+ "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "params_scores"],
1405
  "dispatch": {
1406
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1407
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1413
  "name": "MultiHeadAttention.MaterializedSoftmaxBias",
1414
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
1415
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
1416
+ "bindings": ["scores_softmax", "params_rows"],
1417
  "dispatch": {
1418
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1419
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
1420
  "z": 1
1421
+ }
 
1422
  },
1423
  {
1424
  "id": "apply",
1425
  "name": "MultiHeadAttention.MaterializedApplyBias",
1426
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1427
  "derive": { "layout": "\"bsh\"" },
1428
+ "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "params_apply"],
1429
  "dispatch": {
1430
  "x": "ceilDiv(headDim, applyTileN)",
1431
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1442
  "derive": {
1443
  "qNumHeads": "attrs.num_heads",
1444
  "kvNumHeads": "attrs.num_heads",
 
1445
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1446
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1447
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
1448
  "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
1449
  "applyTileN": "materializedApplyTileN",
 
 
1450
  "statSlots": "materializedGemmStatSlots",
1451
  "statQuerySeq": "dim(shapes.queryT, 1)"
1452
  },
 
1465
  "name": "MultiHeadAttention.MaterializedScoresFused",
1466
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1467
  "derive": { "layout": "\"bsh\"", "emitRowStats": true },
1468
+ "bindings": ["query_f32", "key_f32", "scores", "scorePartials", "params_scores"],
1469
  "dispatch": {
1470
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1471
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1477
  "name": "MultiHeadAttention.MaterializedGemmRowStatsCombine",
1478
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
1479
  "derive": { "maxOnly": true },
1480
+ "bindings": ["scorePartials_f32", "rowStats", "params_rows"],
1481
  "dispatch": {
1482
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1483
  "y": 1,
 
1489
  "name": "MultiHeadAttention.MaterializedApplyFused",
1490
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1491
  "derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
1492
+ "bindings": ["scores_f32", "value_f32", "output_f32", "rowStats_f32", "params_apply"],
1493
  "dispatch": {
1494
  "x": "ceilDiv(headDim, applyTileN)",
1495
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1505
  "derive": {
1506
  "qNumHeads": "attrs.num_heads",
1507
  "kvNumHeads": "attrs.num_heads",
 
1508
  "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
1509
  "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
1510
  "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
1511
  "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
1512
  "applyTileN": "materializedApplyTileN",
 
 
1513
  "statSlots": "materializedGemmStatSlots",
1514
  "statQuerySeq": "dim(shapes.queryT, 1)"
1515
  },
 
1528
  "name": "MultiHeadAttention.MaterializedScoresBiasFused",
1529
  "shader": "attn-materialized-score-f32.wgsl.jinja",
1530
  "derive": { "layout": "\"bsh\"", "emitRowStats": true },
1531
+ "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "scorePartials", "params_scores"],
1532
  "dispatch": {
1533
  "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
1534
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1540
  "name": "MultiHeadAttention.MaterializedGemmRowStatsCombineBias",
1541
  "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
1542
  "derive": { "maxOnly": true },
1543
+ "bindings": ["scorePartials_f32", "rowStats", "params_rows"],
1544
  "dispatch": {
1545
  "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1546
  "y": 1,
 
1552
  "name": "MultiHeadAttention.MaterializedApplyBiasFused",
1553
  "shader": "attn-materialized-apply-f32.wgsl.jinja",
1554
  "derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
1555
+ "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "rowStats_f32", "params_apply"],
1556
  "dispatch": {
1557
  "x": "ceilDiv(headDim, applyTileN)",
1558
  "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
 
1564
  {
1565
  "id": "qkv_no_bias_flash_q32_broadcast",
1566
  "priority": 30,
1567
+ "when": ["tensorDtypes.queryT == \"float16\"", "not present.biasT", "flashShapeOk", "(dim(shapes.queryT, 2) / attrs.num_heads) % 4 == 0", "(dim(shapes.queryT, 2) / attrs.num_heads) >= 64", "(dim(shapes.queryT, 2) / attrs.num_heads) <= 256", "dim(shapes.queryT, 1) >= 31", "wave32SubgroupsUsable"],
1568
  "requires": { "features": ["subgroups", "shader-f16"] },
1569
  "derive": {
1570
  "usesF16": true,
1571
  "scalar": "\"f16\"",
 
1572
  "inputElement": "\"vec4<f16>\"",
1573
  "outputElement": "\"vec4<f16>\"",
1574
  "qNumHeads": "attrs.num_heads",
 
1589
  "x": "ceilDiv(dim(shapes.queryT, 1), 32)",
1590
  "y": "attrs.num_heads",
1591
  "z": "dim(shapes.queryT, 0)"
1592
+ }
 
1593
  }
1594
  ]
1595
  },
 
1601
  "derive": {
1602
  "usesF16": true,
1603
  "scalar": "\"f16\"",
 
1604
  "inputElement": "\"vec4<f16>\"",
1605
  "outputElement": "\"vec4<f16>\"",
1606
  "qNumHeads": "attrs.num_heads",
 
1631
  "priority": 0,
1632
  "when": ["not present.biasT and fallbackMaskShapeOk"],
1633
  "derive": {
 
 
1634
  "hasCausal": true,
 
1635
  "hasKeyLimit": false,
 
 
 
1636
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1637
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1638
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1651
  "name": "MultiHeadAttention",
1652
  "shader": "attn-online-scalar.wgsl.jinja",
1653
  "derive": { "layout": "\"bsh\"" },
1654
+ "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
1655
  "dispatch": {
1656
  "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1657
  "y": "attrs.num_heads",
 
1665
  "priority": 0,
1666
  "when": ["biasOk and fallbackMaskShapeOk"],
1667
  "derive": {
 
 
1668
  "hasCausal": true,
 
1669
  "hasKeyLimit": false,
 
 
 
1670
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1671
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1672
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1685
  "name": "MultiHeadAttention",
1686
  "shader": "attn-online-scalar.wgsl.jinja",
1687
  "derive": { "layout": "\"bsh\"" },
1688
+ "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
1689
  "dispatch": {
1690
  "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
1691
  "y": "attrs.num_heads",
 
1699
  "priority": 0,
1700
  "when": ["not present.biasT and fallbackShapeOk and not flashShapeOk"],
1701
  "derive": {
 
 
1702
  "hasCausal": true,
 
1703
  "hasKeyLimit": false,
 
 
1704
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1705
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1706
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1732
  "priority": 0,
1733
  "when": ["biasOk and fallbackShapeOk and not flashShapeOk"],
1734
  "derive": {
 
 
1735
  "hasCausal": true,
 
1736
  "hasKeyLimit": false,
 
 
1737
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1738
  "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
1739
  "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1765
  "priority": 20,
1766
  "when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashShapeOk"],
1767
  "derive": {
 
 
1768
  "hasCausal": true,
 
1769
  "combineSubgroups": true,
 
 
1770
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1771
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1772
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1773
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1774
  "qNumHeads": "attrs.num_heads",
 
1785
  "shader": "attn-flash-online.wgsl.jinja",
1786
  "derive": { "layout": "\"bsh\"" },
1787
  "bindings": ["query", "key", "value", "output", "params"],
1788
+ "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
 
1789
  }
1790
  ]
1791
  },
 
1794
  "priority": 20,
1795
  "when": ["device.features.has(\"subgroups\")", "biasOk", "flashShapeOk"],
1796
  "derive": {
 
 
1797
  "hasCausal": true,
 
1798
  "combineSubgroups": true,
 
 
1799
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1800
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1801
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1802
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1803
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1816
  "shader": "attn-flash-online.wgsl.jinja",
1817
  "derive": { "layout": "\"bsh\"" },
1818
  "bindings": ["query", "key", "value", "bias", "output", "params"],
1819
+ "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
 
1820
  }
1821
  ]
1822
  },
 
1825
  "priority": 18,
1826
  "when": ["true", "not present.biasT", "flashShapeOk"],
1827
  "derive": {
 
 
1828
  "hasCausal": true,
 
1829
  "combineSubgroups": false,
 
 
1830
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1831
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1832
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1833
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1834
  "qNumHeads": "attrs.num_heads",
 
1854
  "priority": 17,
1855
  "when": ["true", "biasOk", "flashShapeOk"],
1856
  "derive": {
 
 
1857
  "hasCausal": true,
 
1858
  "combineSubgroups": false,
 
 
1859
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1860
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1861
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1862
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1863
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1885
  "priority": 20,
1886
  "when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashMaskShapeOk"],
1887
  "derive": {
 
 
1888
  "hasCausal": true,
 
1889
  "combineSubgroups": true,
 
 
1890
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1891
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1892
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1893
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1894
  "qNumHeads": "attrs.num_heads",
 
1904
  "name": "MultiHeadAttention.Flash",
1905
  "shader": "attn-flash-online.wgsl.jinja",
1906
  "derive": { "layout": "\"bsh\"" },
1907
+ "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
1908
+ "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
 
1909
  }
1910
  ]
1911
  },
 
1914
  "priority": 20,
1915
  "when": ["device.features.has(\"subgroups\")", "biasOk", "flashMaskShapeOk"],
1916
  "derive": {
 
 
1917
  "hasCausal": true,
 
1918
  "combineSubgroups": true,
 
 
1919
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1920
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1921
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1922
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1923
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1935
  "name": "MultiHeadAttention.Flash",
1936
  "shader": "attn-flash-online.wgsl.jinja",
1937
  "derive": { "layout": "\"bsh\"" },
1938
+ "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
1939
+ "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
 
1940
  }
1941
  ]
1942
  },
 
1945
  "priority": 18,
1946
  "when": ["true", "not present.biasT", "flashMaskShapeOk"],
1947
  "derive": {
 
 
1948
  "hasCausal": true,
 
1949
  "combineSubgroups": false,
 
 
1950
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1951
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1952
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1953
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1954
  "qNumHeads": "attrs.num_heads",
 
1964
  "name": "MultiHeadAttention.NoBiasOnlineFlashNoSg",
1965
  "shader": "attn-flash-online.wgsl.jinja",
1966
  "derive": { "layout": "\"bsh\"" },
1967
+ "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
1968
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
1969
  }
1970
  ]
 
1974
  "priority": 17,
1975
  "when": ["true", "biasOk", "flashMaskShapeOk"],
1976
  "derive": {
 
 
1977
  "hasCausal": true,
 
1978
  "combineSubgroups": false,
 
 
1979
  "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1980
  "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1981
  "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1982
  "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
1983
  "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
 
1995
  "name": "MultiHeadAttention.BiasOnlineFlashNoSg",
1996
  "shader": "attn-flash-online.wgsl.jinja",
1997
  "derive": { "layout": "\"bsh\"" },
1998
+ "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
1999
  "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
2000
  }
2001
  ]
 
2010
  },
2011
  "derive": {
2012
  "qNumHeads": "attrs.num_heads",
 
2013
  "qHidden": "dim(shapes.queryT, 2)",
2014
  "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
2015
  "materializedSoftmaxCols": "dim(shapes.keyT, 1)",
2016
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
2017
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
 
2018
  "useSubgroups": true
2019
  },
2020
  "intermediates": [
 
2030
  "name": "MultiHeadAttention.MaterializedScoresSgmat",
2031
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2032
  "derive": { "phase": "\"score\"" },
2033
+ "bindings": ["query_f32", "key_f32", "scores", "params_scores"],
2034
  "dispatch": {
2035
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
2036
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
 
2042
  "name": "MultiHeadAttention.MaterializedSoftmaxSgmat",
2043
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
2044
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
2045
+ "bindings": ["scores_softmax", "params_rows"],
2046
  "dispatch": {
2047
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2048
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2049
  "z": 1
2050
+ }
 
2051
  },
2052
  {
2053
  "id": "apply",
2054
  "name": "MultiHeadAttention.MaterializedApplySgmat",
2055
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2056
  "derive": { "phase": "\"apply\"" },
2057
+ "bindings": ["scores_f32", "value_f32", "output_f32", "params_apply"],
2058
  "dispatch": {
2059
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
2060
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
 
2073
  },
2074
  "derive": {
2075
  "qNumHeads": "attrs.num_heads",
 
2076
  "qHidden": "dim(shapes.queryT, 2)",
2077
  "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
2078
  "materializedSoftmaxCols": "dim(shapes.keyT, 1)",
2079
  "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
2080
  "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
 
2081
  "useSubgroups": true
2082
  },
2083
  "intermediates": [
 
2093
  "name": "MultiHeadAttention.MaterializedScoresSgmatBias",
2094
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2095
  "derive": { "phase": "\"score\"" },
2096
+ "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "params_scores"],
2097
  "dispatch": {
2098
  "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
2099
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
 
2105
  "name": "MultiHeadAttention.MaterializedSoftmaxSgmatBias",
2106
  "shader": "attn-materialized-softmax-f32.wgsl.jinja",
2107
  "derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
2108
+ "bindings": ["scores_softmax", "params_rows"],
2109
  "dispatch": {
2110
  "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2111
  "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
2112
  "z": 1
2113
+ }
 
2114
  },
2115
  {
2116
  "id": "apply",
2117
  "name": "MultiHeadAttention.MaterializedApplySgmatBias",
2118
  "shader": "attn-materialized-sgmat-f32.wgsl.jinja",
2119
  "derive": { "phase": "\"apply\"" },
2120
+ "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "params_apply"],
2121
  "dispatch": {
2122
  "x": "ceilDiv(headDim, materializedSgmatKeyTile)",
2123
  "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
build/webgpu/metadata.json CHANGED
@@ -1,36 +1,39 @@
1
  {
2
  "name": "com.microsoft.MultiHeadAttention",
3
- "id": "_com_microsoft_multiheadattention_webgpu_5144476",
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": "XSAMF/uzQKJg9pZZDeqmViaVXuj47Vd7Pk404Z1J48I=",
12
- "attn-flash-decode-splitk.wgsl.jinja": "dqxX8taw1tpo/QQkFr2zJ7TEEtQtOspTYEYHADXtESs=",
13
- "attn-flash-online.wgsl.jinja": "hCxO/6VGkyixwyoT/kuYeHgcEkFoRwk/ZmiLlnwa0jM=",
14
- "attn-flash-prefill-cluster.wgsl.jinja": "igbsg5xRuFz7alnwmL74hhh1JUqfDB5PfUslsEz5rjA=",
15
- "attn-flash-q32-broadcast.wgsl.jinja": "8qgKXydwiH2shac/fkdHi+FR+c8pOMDFbNp9Zs8KAnc=",
16
- "attn-materialized-apply-f32.wgsl.jinja": "aF7JSJj8rQN8wGTV7IWmxkqMDbXOnGBSI3tsI1CxZNk=",
17
- "attn-materialized-rowstats-combine-f32.wgsl.jinja": "+/5m03sN44zTn9q7cquDea2UNhEO1ujzm1J7yPai4OY=",
18
- "attn-materialized-score-f32.wgsl.jinja": "gCvBC1T1CLEmNxcpo04CLBQdFBkEVMJ9Xhv1wbQxCsI=",
19
- "attn-materialized-sgmat-f32.wgsl.jinja": "Ez9UJggPXMZ6OyzUQu8xLtg8bZd+lTCoEGDwLkNHNL0=",
20
- "attn-materialized-softmax-f32.wgsl.jinja": "DCuz4gFKXDCdR9UMnfD+mYFjBghnv12GxtEsPKH52mE=",
21
- "attn-online-scalar.wgsl.jinja": "6+hU7MzwHz69fjciv94xsPxN+o9rwm42u9Cin2C+AL8=",
22
- "attn-small-head-parallel.wgsl.jinja": "oxFE6hQZKTRugz9hwuS+cob00PskT8iYqP4e+YWtHKA=",
23
- "bench.json": "b/WdaAhEP5NpRXNB5w1mivRNTfALrBDJisDONxTf4Yo=",
24
- "manifest.json": "+/bNCAjkG5OoutFq7O4vQJqb2yMxzICsKvFXsRE50bA=",
25
- "mha-small-seq-blocked.wgsl.jinja": "zznGERrPco5uH9tKiQ2QE+ArO3vho/nEv69wT8IgY1U=",
26
- "mha-small-seq.wgsl.jinja": "UsbUIkvgWL/RgOBSm7lhPGIatA36yTQR4Z2gpHVLqx4=",
27
- "test.json": "NlMA/wlmGSpogAuhSZ/kcQ01dd3M+CvkPF0d/VBmIpM="
 
28
  }
29
  },
30
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
31
  "webgpu": {
32
- "manifestSpec": "2.0",
33
  "variants": {
 
 
34
  "qkv_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
35
  "qkv_no_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
36
  "qkv_no_bias_small_seq": ["mha-small-seq.wgsl.jinja"],
 
1
  {
2
  "name": "com.microsoft.MultiHeadAttention",
3
+ "id": "_com_microsoft_multiheadattention_webgpu_e942a19",
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": "JIhoYhGkuat/Wrvix6ddD19+bmzHnc0yJdwWbW9o+nw=",
12
+ "attn-flash-decode-splitk.wgsl.jinja": "tDj5VlltkM0N99vjpj6z5wtlIv00benlaAboa4e5vWI=",
13
+ "attn-flash-online.wgsl.jinja": "1DGpp4kzd/v4gzAviJKnYGCT7BwDomzXK3He6CIT+1E=",
14
+ "attn-flash-prefill-cluster.wgsl.jinja": "pjDvKLZ3CSluSAPNfb9j3U7yKEy+wVOfOAyZ1jdrjv8=",
15
+ "attn-flash-q32-broadcast.wgsl.jinja": "puOnD1vGGeVg6HmaUi8k4JAZJviJM1dXWJyMZ2xuURM=",
16
+ "attn-materialized-apply-f32.wgsl.jinja": "E2Vas8y/zcjHQb4Ks5TPnC4/lAhVRawXR5S1OwJB1Tc=",
17
+ "attn-materialized-rowstats-combine-f32.wgsl.jinja": "CBltqk0CN9vvscb6D29o1T43vHWZYlU9fFpO9av2Sa8=",
18
+ "attn-materialized-score-f32.wgsl.jinja": "bxW7HdJcF47NTbyXVHcaWrtCO0xCJ8pW5GW9/FTuGGg=",
19
+ "attn-materialized-sgmat-f32.wgsl.jinja": "3e4ZOgVD5XSnYjVldRVq5QL72KiloCkr91OZj/VU0ZU=",
20
+ "attn-materialized-softmax-f32.wgsl.jinja": "bpHIAAuUa+nrniTKxXEScXT5rJr5nObXY/sJ81hby5A=",
21
+ "attn-online-scalar.wgsl.jinja": "AbpGvmzka15nPEj3QZPyMaW48Lo/9TqtndP0SdyqV/U=",
22
+ "attn-small-head-parallel.wgsl.jinja": "Y0otPUGeAa4DzQUi+LnWLRTATc1i/awpIipRqbGDoQ0=",
23
+ "attn-small-head-value.wgsl.jinja": "+6NLq7AJT3lRkAwsW2AvjGE0nhkD7dkBLqh6SpOFjy4=",
24
+ "bench.json": "ateXBEoOVVMHNfiuXPmDW6sKzU1NRVRA9ndzslgsD2A=",
25
+ "manifest.json": "VPWToEdcoP2Mr6k2oQuf/BhwUZdbcPl9J5YLZsMBMaE=",
26
+ "mha-small-seq-blocked.wgsl.jinja": "WudxiAnGnY7dSEnlhHtwvQKvm4AN3ngQ9cbBfI6OUBM=",
27
+ "mha-small-seq.wgsl.jinja": "1MR+JuEbffrS+QkuAZzFzWCLPX4VX9RJ5lmNIOm0/o0=",
28
+ "test.json": "IEMAo5OqrMxRC7xc5Y/8xmjWObXuWw1SJhJeGGzDhwc="
29
  }
30
  },
31
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
32
  "webgpu": {
33
+ "manifestSpec": "2.1",
34
  "variants": {
35
+ "qkv_no_bias_small_head_value_subgroups": ["attn-small-head-value.wgsl.jinja"],
36
+ "qkv_no_bias_small_head_value_tree": ["attn-small-head-value.wgsl.jinja"],
37
  "qkv_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
38
  "qkv_no_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
39
  "qkv_no_bias_small_seq": ["mha-small-seq.wgsl.jinja"],
build/webgpu/mha-small-seq-blocked.wgsl.jinja CHANGED
@@ -37,19 +37,20 @@ 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
  }
43
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
44
- fn scale_value() -> f32 {
45
  if (params.scale != 0.0) { return params.scale; }
46
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
47
  }
48
-
49
  {% if hasBias %}
 
 
50
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
51
  let offset = base + d4 * 4u;
52
- return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
53
  }
54
  {% endif %}
55
 
 
37
  let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
38
  return select(value - maxValue, 0.0, equalFiniteMax);
39
  }
40
+
41
  fn exp_shift(value: f32, maxValue: f32) -> f32 {
42
  return exp(shifted_value(value, maxValue));
43
  }
44
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
45
  if (params.scale != 0.0) { return params.scale; }
46
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
47
  }
 
48
  {% if hasBias %}
49
+ {% set BW = "" %}
50
+ {% set BC = "" %}
51
  fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
52
  let offset = base + d4 * 4u;
53
+ return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
54
  }
55
  {% endif %}
56
 
build/webgpu/mha-small-seq.wgsl.jinja CHANGED
@@ -1,4 +1,26 @@
1
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2
 
3
  // Whole-head attention for tiny bidirectional sequences without bias or mask.
4
  // One workgroup owns one (head, batch) and stages the complete K and V planes in
@@ -9,13 +31,11 @@ const HEAD_DIM: u32 = {{ headDim }}u;
9
  const KV_SEQ: u32 = {{ kvSeq }}u;
10
  const HIDDEN: u32 = {{ hidden }}u;
11
  const WG: u32 = {{ workgroupSize }}u;
12
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
13
- fn scale_value() -> f32 {
14
  if (params.scale != 0.0) { return params.scale; }
15
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
16
  }
17
 
18
-
19
  var<workgroup> kShared: array<f32, KV_SEQ * HEAD_DIM>;
20
  var<workgroup> vShared: array<f32, KV_SEQ * HEAD_DIM>;
21
 
@@ -56,11 +76,11 @@ fn main(
56
  var scores: array<f32, KV_SEQ>;
57
  var maxScore = -3.4028234663852886e38;
58
  for (var k = 0u; k < KV_SEQ; k++) {
59
- var dot = 0.0;
60
  for (var d = 0u; d < HEAD_DIM; d++) {
61
- dot += qRow[d] * kShared[k * HEAD_DIM + d];
62
  }
63
- let s = dot * scale;
64
  scores[k] = s;
65
  maxScore = max(maxScore, s);
66
  }
 
1
  {{ env.wgsl.resourceDeclarations }}
2
+ {% set dotType = "f32" %}
3
+ // Retain product and addition residuals across a dot product. Explicit fma
4
+ // boundaries preserve the addition error transform under reassociation.
5
+ struct DotAccumulator {
6
+ hi: {{ dotType }},
7
+ lo: {{ dotType }},
8
+ }
9
+ fn dot_accumulate(acc: DotAccumulator, a: {{ dotType }}, b: {{ dotType }}) -> DotAccumulator {
10
+ let product = fma(a, b, {{ dotType }}(0.0));
11
+ let productError = fma(a, b, -product);
12
+ let sum = fma(acc.hi, {{ dotType }}(1.0), product);
13
+ let bv = fma({{ dotType }}(-1.0), acc.hi, sum);
14
+ let av = fma({{ dotType }}(-1.0), bv, sum);
15
+ let ae = fma({{ dotType }}(-1.0), av, acc.hi);
16
+ let be = fma({{ dotType }}(-1.0), bv, product);
17
+ let error = fma(fma(fma(ae, {{ dotType }}(1.0), be), {{ dotType }}(1.0), acc.lo), {{ dotType }}(1.0), productError);
18
+ let hi = fma(sum, {{ dotType }}(1.0), error);
19
+ return DotAccumulator(hi, fma({{ dotType }}(-1.0), fma({{ dotType }}(-1.0), sum, hi), error));
20
+ }
21
+ fn dot_value(acc: DotAccumulator) -> {{ dotType }} {
22
+ return fma(acc.hi, {{ dotType }}(1.0), acc.lo);
23
+ }
24
 
25
  // Whole-head attention for tiny bidirectional sequences without bias or mask.
26
  // One workgroup owns one (head, batch) and stages the complete K and V planes in
 
31
  const KV_SEQ: u32 = {{ kvSeq }}u;
32
  const HIDDEN: u32 = {{ hidden }}u;
33
  const WG: u32 = {{ workgroupSize }}u;
34
+ {% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
 
35
  if (params.scale != 0.0) { return params.scale; }
36
  return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
37
  }
38
 
 
39
  var<workgroup> kShared: array<f32, KV_SEQ * HEAD_DIM>;
40
  var<workgroup> vShared: array<f32, KV_SEQ * HEAD_DIM>;
41
 
 
76
  var scores: array<f32, KV_SEQ>;
77
  var maxScore = -3.4028234663852886e38;
78
  for (var k = 0u; k < KV_SEQ; k++) {
79
+ var dot = DotAccumulator(0.0, 0.0);
80
  for (var d = 0u; d < HEAD_DIM; d++) {
81
+ dot = dot_accumulate(dot, qRow[d], kShared[k * HEAD_DIM + d]);
82
  }
83
+ let s = dot_value(dot) * scale;
84
  scores[k] = s;
85
  maxScore = max(maxScore, s);
86
  }
build/webgpu/test.json CHANGED
The diff for this file is too large to render. See raw diff