Xenova HF Staff commited on
Commit
2a3f031
·
verified ·
1 Parent(s): 47389e2

sync 2e7068faf55e

Browse files
README.md CHANGED
@@ -1,3 +1,94 @@
1
  ---
 
2
  license: apache-2.0
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ library_name: kernels
3
  license: apache-2.0
4
+ tags:
5
+ - kernel
6
+ - webgpu
7
+ - wgsl
8
  ---
9
+ # com.microsoft.LinearAttention
10
+
11
+ `com.microsoft` · ONNX Runtime contrib operator · contrib since_version 1
12
+
13
+ ## Description
14
+
15
+ Recurrent linear attention for packed `[B, T, H*D]` decode and prefill. It supports all four update rules, standard and inverse GQA, shared-key heads, and rollback states through `state_window`. Activations and state may independently use float16 or float32; bfloat16 is not implemented. `past_state` is optional for every update rule and defaults to zeros.
16
+
17
+ See the [ONNX Runtime `LinearAttention` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.LinearAttention) for the reference semantics.
18
+
19
+ ## Inputs
20
+
21
+ | Name | Bind key | Logical dtype | Rank | Shape | Description | Presence |
22
+ | --- | --- | --- | --- | --- | --- | --- |
23
+ | `query` | `queryT` | `T` | `3` | — | Query vectors with 3D packed shape `(B, T, H_q * d_k)`; heads are packed into the last dimension. | required |
24
+ | `key` | `keyT` | `T` | `3` | — | Key vectors with 3D packed shape `(B, T, H_k * d_k)`, where positive `H_k` divides `H_kv`; `H_k < H_kv` shares each key head across multiple KV-state heads. Keys should be L2-normalized for `delta`/`gated_delta` modes. | required |
25
+ | `value` | `valueT` | `T` | `3` | — | Value vectors with 3D packed shape `(B, T, H_kv * d_v)`. | required |
26
+ | `past_state` | `pastStateT` | `S` | derived | derived; see description | Recurrent state from the previous step with shape `(B, H_kv, d_k, d_v)`, or `(W, B, H_kv, d_k, d_v)` when `state_window = W > 0`; defaults to zeros if absent. | optional |
27
+ | `decay` | `decayT` | `T` | `3` | — | Exponential decay gate in log-space with shape `(B, T, H_kv * d_k)` or `(B, T, H_kv)`; required for `gated` and `gated_delta` modes. | optional |
28
+ | `beta` | `betaT` | `T` | `3` | — | Update rate (sigmoid output) with shape `(B, T, H_kv)` or `(B, T, 1)`; required for `delta` and `gated_delta` modes. | optional |
29
+
30
+ ## Outputs
31
+
32
+ | Name | Bind key | Logical dtype | Rank | Shape | Description | Presence |
33
+ | --- | --- | --- | --- | --- | --- | --- |
34
+ | `output` | `outputT` | `T` | `3` | derived; see description | Attention output with 3D packed shape `(B, T, max(H_q, H_kv) * d_v)`. | required |
35
+ | `present_state` | `presentStateT` | `S` | derived | derived; see description | Updated recurrent state with shape `(B, H_kv, d_k, d_v)`, or `(W, B, H_kv, d_k, d_v)` when `state_window = W > 0`. | required |
36
+
37
+ ## Attributes
38
+
39
+ Attributes and default values (overridable per request):
40
+
41
+ | Attribute | Default | Description |
42
+ | --- | --- | --- |
43
+ | `chunk_size` | `64` | Accepted for schema compatibility; does not affect the result. |
44
+ | `scale` | `0` | Scale applied to query-key products. Zero selects `1 / sqrt(d_k)`. |
45
+ | `state_window` | `0` | Number of recent recurrent states retained in `present_state`, in the supported range 0 to 8; zero returns only the current state. |
46
+ | `update_rule` | `"gated_delta"` | Recurrent update rule: `linear`, `gated`, `delta`, or `gated_delta`. |
47
+ | `kv_num_heads` | — | Number of key/value heads. |
48
+ | `q_num_heads` | — | Number of query heads. |
49
+
50
+ ## Type constraints
51
+
52
+ | Variable | Allowed dtypes |
53
+ | --- | --- |
54
+ | `T` | `float32`, `float16` |
55
+ | `S` | `float32`, `float16` |
56
+
57
+ ## Files
58
+
59
+ - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, provenance)
60
+ - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
61
+ - [`test.json`](build/webgpu/test.json) — correctness cases
62
+ - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
63
+ - [`chunk-out.wgsl.jinja`](build/webgpu/chunk-out.wgsl.jinja)
64
+ - [`chunk-prep.wgsl.jinja`](build/webgpu/chunk-prep.wgsl.jinja)
65
+ - [`chunk-scan.wgsl.jinja`](build/webgpu/chunk-scan.wgsl.jinja)
66
+ - [`chunk-ut.wgsl.jinja`](build/webgpu/chunk-ut.wgsl.jinja)
67
+ - [`linear-attention.scalar.wgsl.jinja`](build/webgpu/linear-attention.scalar.wgsl.jinja)
68
+ - [`linear-attention.serial.wgsl.jinja`](build/webgpu/linear-attention.serial.wgsl.jinja)
69
+ - [`linear-attention.vec4.wgsl.jinja`](build/webgpu/linear-attention.vec4.wgsl.jinja)
70
+
71
+ ## Use with `@huggingface/kernels`
72
+
73
+ The loader derives every required output's shape and logical dtype from the manifest contract and this call.
74
+ It then allocates the result tensors automatically.
75
+
76
+ The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
77
+
78
+ Replace each `*Data` placeholder with a typed array containing the corresponding input data.
79
+
80
+ ```js
81
+ import { getKernel } from "@huggingface/kernels";
82
+
83
+ const kernel = await getKernel("webgpu-kernels/com.microsoft.LinearAttention", { version: 1 });
84
+ const { outputT, presentStateT } = await kernel({
85
+ queryT: { data: queryTData, shape: [1, 3, 8] },
86
+ keyT: { data: keyTData, shape: [1, 3, 4] },
87
+ valueT: { data: valueTData, shape: [1, 3, 4] },
88
+ pastStateT: { data: pastStateTData, shape: [1, 1, 4, 4] },
89
+ decayT: { data: decayTData, shape: [1, 3, 1] },
90
+ betaT: { data: betaTData, shape: [1, 3, 1] },
91
+ }, {
92
+ attrs: { q_num_heads: 2, kv_num_heads: 1 },
93
+ });
94
+ ```
build/webgpu/bench.json ADDED
@@ -0,0 +1,323 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "op": "com.microsoft.LinearAttention",
3
+ "tunableSpace": {
4
+ "dvGroups": [2, 4, 8],
5
+ "tileV": [4, 8, 16],
6
+ "gatedTileV": [2, 4, 8],
7
+ "chunkSize": [16, 32],
8
+ "chunkTileV": [16, 32]
9
+ },
10
+ "cases": [
11
+ {
12
+ "name": "linear-attention-f32-zero-32x4x16x16",
13
+ "preset": "smoke",
14
+ "vars": { "batch": 1, "seq": 32, "qHeads": 4, "kvHeads": 2, "dk": 16, "dv": 16 },
15
+ "attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
16
+ "inputs": {
17
+ "queryT": { "shape": [1, 32, 64], "dtype": "float32", "dist": "normal", "seed": 201, "scale": 0.2 },
18
+ "keyT": { "shape": [1, 32, 32], "dtype": "float32", "dist": "normal", "seed": 202, "scale": 0.2 },
19
+ "valueT": { "shape": [1, 32, 32], "dtype": "float32", "dist": "normal", "seed": 203, "scale": 0.2 }
20
+ },
21
+ "outputs": {
22
+ "outputT": { "shape": [1, 32, 64], "dtype": "float32" },
23
+ "presentStateT": { "shape": [1, 2, 16, 16], "dtype": "float32" }
24
+ },
25
+ "bench": {
26
+ "primary": true,
27
+ "metrics": [
28
+ { "type": "gflops", "value": "2 * args.batch * args.seq * (args.qHeads + args.kvHeads) * args.dk * args.dv" }
29
+ ]
30
+ }
31
+ },
32
+ {
33
+ "name": "linear-attention-f32-state-32x4x16x16",
34
+ "preset": "smoke",
35
+ "vars": { "batch": 1, "seq": 32, "qHeads": 4, "kvHeads": 2, "dk": 16, "dv": 16 },
36
+ "attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated_delta", "scale": 0.25 },
37
+ "inputs": {
38
+ "queryT": { "shape": [1, 32, 64], "dtype": "float32", "dist": "normal", "seed": 202, "scale": 0.2 },
39
+ "keyT": { "shape": [1, 32, 32], "dtype": "float32", "dist": "normal", "seed": 203, "scale": 0.2 },
40
+ "valueT": { "shape": [1, 32, 32], "dtype": "float32", "dist": "normal", "seed": 204, "scale": 0.2 },
41
+ "pastStateT": { "shape": [1, 2, 16, 16], "dtype": "float32", "dist": "normal", "seed": 205, "scale": 0.1 },
42
+ "decayT": { "shape": [1, 32, 32], "dtype": "float32", "dist": "normal", "seed": 206, "scale": 0.1 },
43
+ "betaT": { "shape": [1, 32, 2], "dtype": "float32", "dist": "normal", "seed": 207, "scale": 0.1 }
44
+ },
45
+ "outputs": {
46
+ "outputT": { "shape": [1, 32, 64], "dtype": "float32" },
47
+ "presentStateT": { "shape": [1, 2, 16, 16], "dtype": "float32" }
48
+ },
49
+ "bench": {
50
+ "primary": true,
51
+ "metrics": [
52
+ {
53
+ "type": "gflops",
54
+ "value": "2 * args.batch * args.seq * (args.qHeads + 2 * args.kvHeads) * args.dk * args.dv"
55
+ }
56
+ ]
57
+ }
58
+ },
59
+ {
60
+ "name": "linear-attention-f32-linear-state-32x4x16x16",
61
+ "preset": "smoke",
62
+ "vars": { "batch": 1, "seq": 32, "qHeads": 4, "kvHeads": 2, "dk": 16, "dv": 16 },
63
+ "attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
64
+ "inputs": {
65
+ "queryT": { "shape": [1, 32, 64], "dtype": "float32", "dist": "normal", "seed": 208, "scale": 0.2 },
66
+ "keyT": { "shape": [1, 32, 32], "dtype": "float32", "dist": "normal", "seed": 209, "scale": 0.2 },
67
+ "valueT": { "shape": [1, 32, 32], "dtype": "float32", "dist": "normal", "seed": 210, "scale": 0.2 },
68
+ "pastStateT": { "shape": [1, 2, 16, 16], "dtype": "float32", "dist": "normal", "seed": 211, "scale": 0.1 }
69
+ },
70
+ "outputs": {
71
+ "outputT": { "shape": [1, 32, 64], "dtype": "float32" },
72
+ "presentStateT": { "shape": [1, 2, 16, 16], "dtype": "float32" }
73
+ },
74
+ "bench": {
75
+ "metrics": [
76
+ { "type": "gflops", "value": "2 * args.batch * args.seq * (args.qHeads + args.kvHeads) * args.dk * args.dv" }
77
+ ]
78
+ }
79
+ },
80
+ {
81
+ "name": "linear-attention-linear-state-scalar-f16-seq1536-pathology",
82
+ "preset": "stress",
83
+ "provenance": {
84
+ "source": "authored for branch coverage",
85
+ "notes": "Long-sequence supplied-state case at head_dim_k 16. It exercises the recurrent small-dk route's serial token recurrence and distinguishes it from the chunked prefill decomposition."
86
+ },
87
+ "vars": { "batch": 4, "seq": 1536, "qHeads": 4, "kvHeads": 2, "dk": 16, "dv": 16 },
88
+ "attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
89
+ "inputs": {
90
+ "queryT": { "shape": [4, 1536, 64], "dtype": "float16", "dist": "normal", "seed": 314, "scale": 0.2 },
91
+ "keyT": { "shape": [4, 1536, 32], "dtype": "float16", "dist": "normal", "seed": 315, "scale": 0.2 },
92
+ "valueT": { "shape": [4, 1536, 32], "dtype": "float16", "dist": "normal", "seed": 316, "scale": 0.2 },
93
+ "pastStateT": { "shape": [4, 2, 16, 16], "dtype": "float16", "dist": "normal", "seed": 317, "scale": 0.1 }
94
+ },
95
+ "outputs": {
96
+ "outputT": { "shape": [4, 1536, 64], "dtype": "float16" },
97
+ "presentStateT": { "shape": [4, 2, 16, 16], "dtype": "float16" }
98
+ },
99
+ "bench": {
100
+ "primary": true,
101
+ "metrics": [
102
+ { "type": "gflops", "value": "2 * args.batch * args.seq * (args.qHeads + args.kvHeads) * args.dk * args.dv" }
103
+ ]
104
+ }
105
+ },
106
+ {
107
+ "name": "linear-attention-gated_delta-scalar-headdimk6-seq1536-stress",
108
+ "preset": "stress",
109
+ "vars": { "batch": 8, "seq": 1536, "qHeads": 4, "kvHeads": 4, "dk": 6, "dv": 12 },
110
+ "attrs": { "q_num_heads": 4, "kv_num_heads": 4, "update_rule": "gated_delta", "scale": 0.25 },
111
+ "inputs": {
112
+ "queryT": { "shape": [8, 1536, 24], "dtype": "float32", "dist": "normal", "seed": 301, "scale": 0.2 },
113
+ "keyT": { "shape": [8, 1536, 24], "dtype": "float32", "dist": "normal", "seed": 302, "scale": 0.2 },
114
+ "valueT": { "shape": [8, 1536, 48], "dtype": "float32", "dist": "normal", "seed": 303, "scale": 0.2 },
115
+ "pastStateT": { "shape": [8, 4, 6, 12], "dtype": "float32", "dist": "normal", "seed": 304, "scale": 0.1 },
116
+ "decayT": { "shape": [8, 1536, 4], "dtype": "float32", "dist": "normal", "seed": 305, "scale": 0.1 },
117
+ "betaT": { "shape": [8, 1536, 4], "dtype": "float32", "dist": "normal", "seed": 306, "scale": 0.1 }
118
+ },
119
+ "outputs": {
120
+ "outputT": { "shape": [8, 1536, 48], "dtype": "float32" },
121
+ "presentStateT": { "shape": [8, 4, 6, 12], "dtype": "float32" }
122
+ },
123
+ "bench": {
124
+ "primary": true,
125
+ "metrics": [
126
+ {
127
+ "type": "gflops",
128
+ "value": "2 * args.batch * args.seq * (args.qHeads + 2 * args.kvHeads) * args.dk * args.dv"
129
+ }
130
+ ]
131
+ }
132
+ },
133
+ {
134
+ "name": "linear-attention-linear-scalar-f16-seq1536-stress",
135
+ "preset": "stress",
136
+ "vars": { "batch": 4, "seq": 1536, "qHeads": 4, "kvHeads": 2, "dk": 16, "dv": 16 },
137
+ "attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
138
+ "inputs": {
139
+ "queryT": { "shape": [4, 1536, 64], "dtype": "float16", "dist": "normal", "seed": 311, "scale": 0.2 },
140
+ "keyT": { "shape": [4, 1536, 32], "dtype": "float16", "dist": "normal", "seed": 312, "scale": 0.2 },
141
+ "valueT": { "shape": [4, 1536, 32], "dtype": "float16", "dist": "normal", "seed": 313, "scale": 0.2 }
142
+ },
143
+ "outputs": {
144
+ "outputT": { "shape": [4, 1536, 64], "dtype": "float16" },
145
+ "presentStateT": { "shape": [4, 2, 16, 16], "dtype": "float16" }
146
+ },
147
+ "bench": {
148
+ "primary": true,
149
+ "metrics": [
150
+ { "type": "gflops", "value": "2 * args.batch * args.seq * (args.qHeads + args.kvHeads) * args.dk * args.dv" }
151
+ ]
152
+ },
153
+ "provenance": {
154
+ "source": "authored for branch coverage",
155
+ "notes": "Zero-state sibling of the supplied-state long-sequence case, covering the same branch and shape without an entry state."
156
+ }
157
+ },
158
+ {
159
+ "name": "linear-attention-gated-delta-f32-bonsai-m16-h48-kv16-dk128-dv128",
160
+ "preset": "stress",
161
+ "vars": { "batch": 1, "seq": 16, "qHeads": 48, "kvHeads": 16, "dk": 128, "dv": 128 },
162
+ "attrs": { "q_num_heads": 48, "kv_num_heads": 16, "update_rule": "gated_delta", "scale": 0.08838834764831845 },
163
+ "inputs": {
164
+ "queryT": { "shape": [1, 16, 6144], "dtype": "float32", "dist": "normal", "seed": 320, "scale": 0.05 },
165
+ "keyT": { "shape": [1, 16, 2048], "dtype": "float32", "dist": "normal", "seed": 321, "scale": 0.05 },
166
+ "valueT": { "shape": [1, 16, 2048], "dtype": "float32", "dist": "normal", "seed": 322, "scale": 0.05 },
167
+ "pastStateT": { "shape": [1, 16, 128, 128], "dtype": "float32", "dist": "normal", "seed": 323, "scale": 0.02 },
168
+ "decayT": { "shape": [1, 16, 16], "dtype": "float32", "dist": "normal", "seed": 324, "scale": 0.08 },
169
+ "betaT": { "shape": [1, 16, 16], "dtype": "float32", "dist": "normal", "seed": 325, "scale": 0.08 }
170
+ },
171
+ "outputs": {
172
+ "outputT": { "shape": [1, 16, 6144], "dtype": "float32", "dist": "empty" },
173
+ "presentStateT": { "shape": [1, 16, 128, 128], "dtype": "float32", "dist": "empty" }
174
+ },
175
+ "bench": {
176
+ "primary": true,
177
+ "metrics": [
178
+ {
179
+ "type": "gflops",
180
+ "value": "2 * args.batch * args.seq * (args.qHeads + 2 * args.kvHeads) * args.dk * args.dv"
181
+ }
182
+ ]
183
+ }
184
+ },
185
+ {
186
+ "name": "linear-attention-gated-delta-f16-bonsai-m16-h48-kv16-dk128-dv128",
187
+ "preset": "stress",
188
+ "vars": { "batch": 1, "seq": 16, "qHeads": 48, "kvHeads": 16, "dk": 128, "dv": 128 },
189
+ "attrs": { "q_num_heads": 48, "kv_num_heads": 16, "update_rule": "gated_delta", "scale": 0.08838834764831845 },
190
+ "inputs": {
191
+ "queryT": { "shape": [1, 16, 6144], "dtype": "float16", "dist": "normal", "seed": 326, "scale": 0.05 },
192
+ "keyT": { "shape": [1, 16, 2048], "dtype": "float16", "dist": "normal", "seed": 327, "scale": 0.05 },
193
+ "valueT": { "shape": [1, 16, 2048], "dtype": "float16", "dist": "normal", "seed": 328, "scale": 0.05 },
194
+ "pastStateT": { "shape": [1, 16, 128, 128], "dtype": "float16", "dist": "normal", "seed": 329, "scale": 0.02 },
195
+ "decayT": { "shape": [1, 16, 16], "dtype": "float16", "dist": "normal", "seed": 330, "scale": 0.08 },
196
+ "betaT": { "shape": [1, 16, 16], "dtype": "float16", "dist": "normal", "seed": 331, "scale": 0.08 }
197
+ },
198
+ "outputs": {
199
+ "outputT": { "shape": [1, 16, 6144], "dtype": "float16", "dist": "empty" },
200
+ "presentStateT": { "shape": [1, 16, 128, 128], "dtype": "float16", "dist": "empty" }
201
+ },
202
+ "bench": {
203
+ "primary": true,
204
+ "metrics": [
205
+ {
206
+ "type": "gflops",
207
+ "value": "2 * args.batch * args.seq * (args.qHeads + 2 * args.kvHeads) * args.dk * args.dv"
208
+ }
209
+ ]
210
+ }
211
+ },
212
+ {
213
+ "name": "linear-attention-qwen3next-decode-s1",
214
+ "preset": "model",
215
+ "provenance": {
216
+ "notes": "Qwen3-Next class defaults (linear_num_value_heads 32, linear_num_key_heads 16, linear_key_head_dim 128, linear_value_head_dim 128) at a decode step, where the recurrence carries the whole cost."
217
+ },
218
+ "vars": { "batch": 1, "seq": 1, "qHeads": 32, "kvHeads": 16, "dk": 128, "dv": 128 },
219
+ "attrs": {
220
+ "q_num_heads": 32,
221
+ "kv_num_heads": 16,
222
+ "update_rule": "gated_delta",
223
+ "scale": 0.08838834764831843,
224
+ "chunk_size": 64
225
+ },
226
+ "inputs": {
227
+ "queryT": { "shape": [1, 1, 4096], "dtype": "float32", "dist": "normal", "seed": 8100, "scale": 0.3 },
228
+ "keyT": { "shape": [1, 1, 2048], "dtype": "float32", "dist": "normal", "seed": 8101, "scale": 0.3 },
229
+ "valueT": { "shape": [1, 1, 2048], "dtype": "float32", "dist": "normal", "seed": 8102, "scale": 0.3 },
230
+ "pastStateT": { "shape": [1, 16, 128, 128], "dtype": "float32", "dist": "normal", "seed": 8103, "scale": 0.1 },
231
+ "decayT": { "shape": [1, 1, 2048], "dtype": "float32", "dist": "uniform", "seed": 8104, "min": 0.9, "max": 1 },
232
+ "betaT": { "shape": [1, 1, 16], "dtype": "float32", "dist": "uniform", "seed": 8105, "min": 0.1, "max": 0.9 }
233
+ },
234
+ "outputs": {
235
+ "outputT": { "shape": [1, 1, 4096], "dtype": "float32" },
236
+ "presentStateT": { "shape": [1, 16, 128, 128], "dtype": "float32" }
237
+ },
238
+ "bench": {
239
+ "metrics": [
240
+ {
241
+ "type": "gflops",
242
+ "value": "2 * args.batch * args.seq * (args.qHeads + 2 * args.kvHeads) * args.dk * args.dv"
243
+ }
244
+ ]
245
+ }
246
+ },
247
+ {
248
+ "name": "linear-attention-qwen3next-prefill-s512",
249
+ "preset": "model",
250
+ "provenance": { "notes": "Qwen3-Next class defaults over a 512-token prefill chunk." },
251
+ "vars": { "batch": 1, "seq": 512, "qHeads": 32, "kvHeads": 16, "dk": 128, "dv": 128 },
252
+ "attrs": {
253
+ "q_num_heads": 32,
254
+ "kv_num_heads": 16,
255
+ "update_rule": "gated_delta",
256
+ "scale": 0.08838834764831843,
257
+ "chunk_size": 64
258
+ },
259
+ "inputs": {
260
+ "queryT": { "shape": [1, 512, 4096], "dtype": "float32", "dist": "normal", "seed": 8200, "scale": 0.3 },
261
+ "keyT": { "shape": [1, 512, 2048], "dtype": "float32", "dist": "normal", "seed": 8201, "scale": 0.3 },
262
+ "valueT": { "shape": [1, 512, 2048], "dtype": "float32", "dist": "normal", "seed": 8202, "scale": 0.3 },
263
+ "pastStateT": { "shape": [1, 16, 128, 128], "dtype": "float32", "dist": "normal", "seed": 8203, "scale": 0.1 },
264
+ "decayT": { "shape": [1, 512, 2048], "dtype": "float32", "dist": "uniform", "seed": 8204, "min": 0.9, "max": 1 },
265
+ "betaT": { "shape": [1, 512, 16], "dtype": "float32", "dist": "uniform", "seed": 8205, "min": 0.1, "max": 0.9 }
266
+ },
267
+ "outputs": {
268
+ "outputT": { "shape": [1, 512, 4096], "dtype": "float32" },
269
+ "presentStateT": { "shape": [1, 16, 128, 128], "dtype": "float32" }
270
+ },
271
+ "bench": {
272
+ "metrics": [
273
+ {
274
+ "type": "gflops",
275
+ "value": "2 * args.batch * args.seq * (args.qHeads + 2 * args.kvHeads) * args.dk * args.dv"
276
+ }
277
+ ]
278
+ }
279
+ },
280
+ {
281
+ "name": "linear-attention-qwen3next-prefill-s2048",
282
+ "preset": "model",
283
+ "provenance": {
284
+ "notes": "Qwen3-Next class defaults over a 2048-token prefill chunk, eight chunk_size 64 blocks per workgroup pass."
285
+ },
286
+ "vars": { "batch": 1, "seq": 2048, "qHeads": 32, "kvHeads": 16, "dk": 128, "dv": 128 },
287
+ "attrs": {
288
+ "q_num_heads": 32,
289
+ "kv_num_heads": 16,
290
+ "update_rule": "gated_delta",
291
+ "scale": 0.08838834764831843,
292
+ "chunk_size": 64
293
+ },
294
+ "inputs": {
295
+ "queryT": { "shape": [1, 2048, 4096], "dtype": "float32", "dist": "normal", "seed": 8300, "scale": 0.3 },
296
+ "keyT": { "shape": [1, 2048, 2048], "dtype": "float32", "dist": "normal", "seed": 8301, "scale": 0.3 },
297
+ "valueT": { "shape": [1, 2048, 2048], "dtype": "float32", "dist": "normal", "seed": 8302, "scale": 0.3 },
298
+ "pastStateT": { "shape": [1, 16, 128, 128], "dtype": "float32", "dist": "normal", "seed": 8303, "scale": 0.1 },
299
+ "decayT": {
300
+ "shape": [1, 2048, 2048],
301
+ "dtype": "float32",
302
+ "dist": "uniform",
303
+ "seed": 8304,
304
+ "min": 0.9,
305
+ "max": 1
306
+ },
307
+ "betaT": { "shape": [1, 2048, 16], "dtype": "float32", "dist": "uniform", "seed": 8305, "min": 0.1, "max": 0.9 }
308
+ },
309
+ "outputs": {
310
+ "outputT": { "shape": [1, 2048, 4096], "dtype": "float32" },
311
+ "presentStateT": { "shape": [1, 16, 128, 128], "dtype": "float32" }
312
+ },
313
+ "bench": {
314
+ "metrics": [
315
+ {
316
+ "type": "gflops",
317
+ "value": "2 * args.batch * args.seq * (args.qHeads + 2 * args.kvHeads) * args.dk * args.dv"
318
+ }
319
+ ]
320
+ }
321
+ }
322
+ ]
323
+ }
build/webgpu/chunk-out.wgsl.jinja ADDED
@@ -0,0 +1,200 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro read_scalar(name, index, dtype) %}
2
+ {% if dtype == "float16" %}
3
+ f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
4
+ {% endmacro %}
5
+ {% macro write_scalar(expr, dtype) %}
6
+ {% if dtype == "float16" %}
7
+ f16({{ expr }}){% else %}{{ expr }}{% endif %}
8
+ {% endmacro -%}
9
+ {% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{%- endmacro %}
10
+
11
+ {% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
12
+ {% if needQuery %}
13
+ // The output scale defaults to 1/sqrt(head_dim_k), matching the recurrent kernels.
14
+ fn out_scale() -> f32 {
15
+ return select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
16
+ }
17
+ {% endif %}
18
+ fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
19
+ let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
20
+ {% if usesDecay %}
21
+ return raw / {{ decay_at("bt", "h", "i") }};
22
+ {% else %}
23
+ return raw;
24
+ {% endif %}
25
+ }
26
+ {% if needKtil %}
27
+
28
+ fn k_til(bt: u32, h: u32, i: u32) -> f32 {
29
+ let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
30
+ {% if usesDecay %}
31
+ return raw * {{ decay_at("bt", "h", "i") }};
32
+ {% else %}
33
+ return raw;
34
+ {% endif %}
35
+ }
36
+ {% endif %}
37
+ {% if needQuery %}
38
+
39
+ // The output scale rides on q_til so both output terms -- q_til * S and P * delta,
40
+ // where P is itself built from q_til -- pick it up without a second pass over y.
41
+ fn q_til(bt: u32, q_head: u32, {% if usesDecay %}h: u32, {% endif %}i: u32, scale: f32) -> f32 {
42
+ let raw = {{ read_scalar("query", "bt * params.qPackedDim + q_head * HEAD_DIM_K + i", queryDtype) }};
43
+ {% if usesDecay %}
44
+ return raw * {{ decay_at("bt", "h", "i") }} * scale;
45
+ {% else %}
46
+ return raw * scale;
47
+ {% endif %}
48
+ }
49
+ {% endif %}
50
+ {%- endmacro %}{% macro q_til_call(bt, q_head, h, i, scale) %}q_til({{ bt }}, {{ q_head }}, {% if usesDecay %}{{ h }}, {% endif %}{{ i }}, {{ scale }}){%- endmacro %}{% macro emit_chunk_head_setup(needHeads=false) %}
51
+ {% if needHeads %}
52
+ let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
53
+ let packed_out = max(params.qNumHeads, params.kvNumHeads) * HEAD_DIM_V;
54
+ {% endif %}
55
+ let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;
56
+ {%- endmacro %}
57
+
58
+ {% if queryDtype == "float16" %}
59
+ enable f16;
60
+ {% endif %}
61
+ {{ env.wgsl.resourceDeclarations }}
62
+
63
+ // com.microsoft.LinearAttention, chunked prefill: outputs, from entry states.
64
+ // y_t = q_til_t . S_c + sum_{s <= t} P[t, s] * delta_s, P[t, s] = dot(k_hat_s, q_til_t)
65
+ // S_c is the state entering this token's chunk, so once the sequential pass has
66
+ // published one state per chunk, every chunk's output block is independent -- this
67
+ // pass carries the largest share of the operator's arithmetic and runs entirely in
68
+ // parallel. P is built here rather than stored: one workgroup owns a whole (output
69
+ // head, chunk) block, so nothing recomputes it.
70
+ //
71
+ // A thread owns one value column and TOKEN_ROWS of the chunk, and the token block is
72
+ // the outer loop, so nothing staged here grows with the chunk length. That is what
73
+ // lets the chunk be long: the entry-state buffer the sequential pass publishes shrinks
74
+ // as 1 / CHUNK, while this pass's re-read of it depends only on TOKEN_ROWS.
75
+ const CHUNK: u32 = {{ chunkSize }}u;
76
+ const HEAD_DIM_K: u32 = {{ headDimK }}u;
77
+ const HEAD_DIM_V: u32 = {{ headDimV }}u;
78
+ const KV_PER_KEY_HEAD: u32 = {{ kvPerKeyHead }}u;
79
+ const TOKEN_ROWS: u32 = {{ chunkOutRows }}u;
80
+ const TK: u32 = {{ chunkTileK }}u;
81
+ const WG: u32 = {{ workgroupSize }}u;
82
+
83
+ {{ emit_chunk_operands(needQuery=true) }}
84
+
85
+ var<workgroup> qblock: array<f32, TOKEN_ROWS * HEAD_DIM_K>;
86
+ var<workgroup> ktile: array<f32, CHUNK * TK>;
87
+ var<workgroup> pmrow: array<f32, TOKEN_ROWS * CHUNK>;
88
+
89
+ @compute @workgroup_size(WG, 1, 1)
90
+ fn main(
91
+ @builtin(workgroup_id) wg: vec3<u32>,
92
+ @builtin(num_workgroups) nwg: vec3<u32>,
93
+ @builtin(local_invocation_id) lid: vec3<u32>,
94
+ ) {
95
+ let tid = lid.x;
96
+ let flat = wg.x + wg.y * nwg.x;
97
+ {{ emit_chunk_head_setup(needHeads=true) }}
98
+ let out_heads = max(params.qNumHeads, params.kvNumHeads);
99
+ let chunk = flat % num_chunks;
100
+ let out_head = (flat / num_chunks) % out_heads;
101
+ let batch = flat / (num_chunks * out_heads);
102
+ if (batch >= params.batchSize) {
103
+ return;
104
+ }
105
+ let head = out_head / heads_per_group;
106
+ let q_head = (head * params.qNumHeads) / params.kvNumHeads + out_head % heads_per_group;
107
+ let base_t = batch * params.seqLength + chunk * CHUNK;
108
+ let live = min(CHUNK, params.seqLength - chunk * CHUNK);
109
+ let scale = out_scale();
110
+ let state_base = ((batch * params.kvNumHeads + head) * num_chunks + chunk) * HEAD_DIM_K * HEAD_DIM_V;
111
+ {% if usesBeta %}
112
+ let chunk_base = ((batch * params.kvNumHeads + head) * num_chunks + chunk) * CHUNK;
113
+ {% endif %}
114
+
115
+ for (var tb = 0u; tb < CHUNK; tb = tb + TOKEN_ROWS) {
116
+ for (var e = tid; e < TOKEN_ROWS * HEAD_DIM_K; e = e + WG) {
117
+ let t = tb + e / HEAD_DIM_K;
118
+ var qv = 0.0;
119
+ if (t < live) {
120
+ qv = {{ q_til_call("base_t + t", "q_head", "head", "e % HEAD_DIM_K", "scale") }};
121
+ }
122
+ qblock[e] = qv;
123
+ }
124
+ for (var e = tid; e < TOKEN_ROWS * CHUNK; e = e + WG) {
125
+ pmrow[e] = 0.0;
126
+ }
127
+ var acc: array<f32, TOKEN_ROWS>;
128
+ for (var m = 0u; m < TOKEN_ROWS; m = m + 1u) {
129
+ acc[m] = 0.0;
130
+ }
131
+ workgroupBarrier();
132
+
133
+ // One sweep of the reduction axis feeds both output terms: this block's rows of P,
134
+ // and its q_til * S contribution. The state column a thread loads is used by all
135
+ // TOKEN_ROWS of its accumulators.
136
+ for (var kb = 0u; kb < HEAD_DIM_K; kb = kb + TK) {
137
+ for (var e = tid; e < CHUNK * TK; e = e + WG) {
138
+ let s = e / TK;
139
+ let i = kb + e % TK;
140
+ var kv = 0.0;
141
+ if (s < live && i < HEAD_DIM_K) {
142
+ kv = k_hat(base_t + s, head, i);
143
+ }
144
+ ktile[e] = kv;
145
+ }
146
+ workgroupBarrier();
147
+ // P is inclusive-lower: token t reads the state after its own update, matching
148
+ // the recurrent kernels.
149
+ for (var e = tid; e < TOKEN_ROWS * CHUNK; e = e + WG) {
150
+ let m = e / CHUNK;
151
+ let s = e % CHUNK;
152
+ if (s <= tb + m) {
153
+ var total = 0.0;
154
+ for (var u = 0u; u < TK; u = u + 1u) {
155
+ if (kb + u < HEAD_DIM_K) {
156
+ total = total + qblock[m * HEAD_DIM_K + kb + u] * ktile[s * TK + u];
157
+ }
158
+ }
159
+ pmrow[e] = pmrow[e] + total;
160
+ }
161
+ }
162
+ if (tid < HEAD_DIM_V) {
163
+ for (var u = 0u; u < TK; u = u + 1u) {
164
+ let i = kb + u;
165
+ if (i < HEAD_DIM_K) {
166
+ let sv = states[state_base + i * HEAD_DIM_V + tid];
167
+ for (var m = 0u; m < TOKEN_ROWS; m = m + 1u) {
168
+ acc[m] = acc[m] + qblock[m * HEAD_DIM_K + i] * sv;
169
+ }
170
+ }
171
+ }
172
+ }
173
+ workgroupBarrier();
174
+ }
175
+
176
+ if (tid < HEAD_DIM_V) {
177
+ for (var s = 0u; s < CHUNK; s = s + 1u) {
178
+ {% if usesBeta %}
179
+ let dv_val = deltas[(chunk_base + s) * HEAD_DIM_V + tid];
180
+ {% else %}
181
+ var dv_val = 0.0;
182
+ if (s < live) {
183
+ dv_val = {{ read_scalar("value", "(base_t + s) * params.vPackedDim + head * HEAD_DIM_V + tid", queryDtype) }};
184
+ }
185
+ {% endif %}
186
+ for (var m = 0u; m < TOKEN_ROWS; m = m + 1u) {
187
+ acc[m] = acc[m] + pmrow[m * CHUNK + s] * dv_val;
188
+ }
189
+ }
190
+ for (var m = 0u; m < TOKEN_ROWS; m = m + 1u) {
191
+ let t = tb + m;
192
+ if (t < live) {
193
+ output[(base_t + t) * packed_out + out_head * HEAD_DIM_V + tid] =
194
+ {{ write_scalar("acc[m]", queryDtype) }};
195
+ }
196
+ }
197
+ }
198
+ workgroupBarrier();
199
+ }
200
+ }
build/webgpu/chunk-prep.wgsl.jinja ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro read_scalar(name, index, dtype) %}
2
+ {% if dtype == "float16" %}
3
+ f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
4
+ {% endmacro %}{% if queryDtype == "float16" %}
5
+ enable f16;
6
+ {% endif %}
7
+ {{ env.wgsl.resourceDeclarations }}
8
+
9
+ // com.microsoft.LinearAttention, chunked prefill: within-chunk decay prefix.
10
+ // gexp[b, t, c] = exp(sum of decay[b, s, c] for s from the chunk start through t)
11
+ // Every later pass reads the recurrence's decay through this one buffer. Writing
12
+ // exp(prefix) rather than the prefix itself keeps the per-element exponential out
13
+ // of the O(chunk^2) inner loops, where it would cost one transcendental per
14
+ // multiply-add instead of one per element.
15
+ const CHUNK: u32 = {{ chunkSize }}u;
16
+ const WG: u32 = {{ workgroupSize }}u;
17
+
18
+ @compute @workgroup_size(WG, 1, 1)
19
+ fn main(
20
+ @builtin(workgroup_id) wg: vec3<u32>,
21
+ @builtin(num_workgroups) nwg: vec3<u32>,
22
+ @builtin(local_invocation_id) lid: vec3<u32>,
23
+ ) {
24
+ // 2D-folded flat (batch * chunk) index: wg.y carries the high bits past the
25
+ // maxComputeWorkgroupsPerDimension dispatch limit.
26
+ let flat = wg.x + wg.y * nwg.x;
27
+ let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;
28
+ let chunk = flat % num_chunks;
29
+ let batch = flat / num_chunks;
30
+ if (batch >= params.batchSize) {
31
+ return;
32
+ }
33
+ let t0 = chunk * CHUNK;
34
+ let t1 = min(t0 + CHUNK, params.seqLength);
35
+ let packed = params.decayPackedDim;
36
+
37
+ // One thread owns a decay column and walks the chunk in order: the prefix is
38
+ // serial in t but independent across columns, so the whole chunk grid runs at once.
39
+ for (var col = lid.x; col < packed; col = col + WG) {
40
+ var prefix = 0.0;
41
+ for (var t = t0; t < t1; t = t + 1u) {
42
+ let idx = (batch * params.seqLength + t) * packed + col;
43
+ prefix = prefix + {{ read_scalar("decay", "idx", queryDtype) }};
44
+ gexp[idx] = exp(prefix);
45
+ }
46
+ }
47
+ }
build/webgpu/chunk-scan.wgsl.jinja ADDED
@@ -0,0 +1,216 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro read_scalar(name, index, dtype) %}
2
+ {% if dtype == "float16" %}
3
+ f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
4
+ {% endmacro %}
5
+ {% macro write_scalar(expr, dtype) %}
6
+ {% if dtype == "float16" %}
7
+ f16({{ expr }}){% else %}{{ expr }}{% endif %}
8
+ {% endmacro -%}
9
+ {% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{%- endmacro %}
10
+
11
+ {% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
12
+ {% if needQuery %}
13
+ // The output scale defaults to 1/sqrt(head_dim_k), matching the recurrent kernels.
14
+ fn out_scale() -> f32 {
15
+ return select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
16
+ }
17
+ {% endif %}
18
+ fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
19
+ let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
20
+ {% if usesDecay %}
21
+ return raw / {{ decay_at("bt", "h", "i") }};
22
+ {% else %}
23
+ return raw;
24
+ {% endif %}
25
+ }
26
+ {% if needKtil %}
27
+
28
+ fn k_til(bt: u32, h: u32, i: u32) -> f32 {
29
+ let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
30
+ {% if usesDecay %}
31
+ return raw * {{ decay_at("bt", "h", "i") }};
32
+ {% else %}
33
+ return raw;
34
+ {% endif %}
35
+ }
36
+ {% endif %}
37
+ {% if needQuery %}
38
+
39
+ // The output scale rides on q_til so both output terms -- q_til * S and P * delta,
40
+ // where P is itself built from q_til -- pick it up without a second pass over y.
41
+ fn q_til(bt: u32, q_head: u32, {% if usesDecay %}h: u32, {% endif %}i: u32, scale: f32) -> f32 {
42
+ let raw = {{ read_scalar("query", "bt * params.qPackedDim + q_head * HEAD_DIM_K + i", queryDtype) }};
43
+ {% if usesDecay %}
44
+ return raw * {{ decay_at("bt", "h", "i") }} * scale;
45
+ {% else %}
46
+ return raw * scale;
47
+ {% endif %}
48
+ }
49
+ {% endif %}
50
+ {%- endmacro %}{% macro emit_chunk_head_setup(needHeads=false) %}
51
+ {% if needHeads %}
52
+ let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
53
+ let packed_out = max(params.qNumHeads, params.kvNumHeads) * HEAD_DIM_V;
54
+ {% endif %}
55
+ let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;
56
+ {%- endmacro %}
57
+
58
+ {% if queryDtype == "float16" or stateDtype == "float16" %}
59
+ enable f16;
60
+ {% endif %}
61
+ {{ env.wgsl.resourceDeclarations }}
62
+
63
+ // com.microsoft.LinearAttention, chunked prefill: the sequential state scan.
64
+ //
65
+ // This is the only pass that must run chunk by chunk, and it is deliberately the
66
+ // smallest: two dense products per chunk rather than one dependent step per token.
67
+ // delta_c = u_c - wk_c S_c (identity for the non-delta rules, where delta = v)
68
+ // S_{c+1} = A_c * (S_c + K_hat_c^T delta_c)
69
+ // It publishes each chunk's entry state and correction so the output pass, which
70
+ // carries most of the arithmetic, can run over every chunk at once.
71
+ const CHUNK: u32 = {{ chunkSize }}u;
72
+ const HEAD_DIM_K: u32 = {{ headDimK }}u;
73
+ const HEAD_DIM_V: u32 = {{ headDimV }}u;
74
+ const KV_PER_KEY_HEAD: u32 = {{ kvPerKeyHead }}u;
75
+ const TILE_V: u32 = {{ chunkTileV }}u;
76
+ const TOKEN_TILE: u32 = {{ chunkScanTokens }}u;
77
+ const WG: u32 = {{ workgroupSize }}u;
78
+ // One thread per (value column, group). The two phases split the workgroup along
79
+ // different axes -- tokens for the correction, reduction rows for the state update --
80
+ // so neither needs a cross-thread fold.
81
+ const GROUPS: u32 = WG / TILE_V;
82
+ {% if usesBeta %}
83
+ const TOKENS_PER_TILE: u32 = TOKEN_TILE / GROUPS;
84
+ {% endif %}
85
+ const ROWS_PER_GROUP: u32 = HEAD_DIM_K / GROUPS;
86
+
87
+ {{ emit_chunk_operands() }}
88
+
89
+ var<workgroup> st: array<f32, HEAD_DIM_K * TILE_V>;
90
+ var<workgroup> dl: array<f32, CHUNK * TILE_V>;
91
+ // Staging for TOKEN_TILE rows of wk, then of k_hat: same shape, disjoint live ranges.
92
+ var<workgroup> stage: array<f32, TOKEN_TILE * HEAD_DIM_K>;
93
+
94
+ @compute @workgroup_size(WG, 1, 1)
95
+ fn main(
96
+ @builtin(workgroup_id) wg: vec3<u32>,
97
+ @builtin(num_workgroups) nwg: vec3<u32>,
98
+ @builtin(local_invocation_id) lid: vec3<u32>,
99
+ ) {
100
+ let tid = lid.x;
101
+ let col = tid % TILE_V;
102
+ let group = tid / TILE_V;
103
+ let flat = wg.x + wg.y * nwg.x;
104
+ {{ emit_chunk_head_setup() }}
105
+ let v_tiles = HEAD_DIM_V / TILE_V;
106
+ let v_tile = flat % v_tiles;
107
+ let head = (flat / v_tiles) % params.kvNumHeads;
108
+ let batch = flat / (v_tiles * params.kvNumHeads);
109
+ if (batch >= params.batchSize) {
110
+ return;
111
+ }
112
+ let dv0 = v_tile * TILE_V;
113
+ let head_base = (batch * params.kvNumHeads + head) * HEAD_DIM_K;
114
+
115
+ for (var e = tid; e < HEAD_DIM_K * TILE_V; e = e + WG) {
116
+ {% if hasPastState %}
117
+ st[e] = {{ read_scalar("past_state", "(head_base + e / TILE_V) * HEAD_DIM_V + dv0 + e % TILE_V", stateDtype) }};
118
+ {% else %}
119
+ st[e] = 0.0;
120
+ {% endif %}
121
+ }
122
+ workgroupBarrier();
123
+
124
+ for (var chunk = 0u; chunk < num_chunks; chunk = chunk + 1u) {
125
+ let base_t = batch * params.seqLength + chunk * CHUNK;
126
+ let live = min(CHUNK, params.seqLength - chunk * CHUNK);
127
+ {% if usesBeta %}
128
+ let chunk_base = ((batch * params.kvNumHeads + head) * num_chunks + chunk) * CHUNK;
129
+ {% endif %}
130
+
131
+ {% if usesBeta %}
132
+ // delta = u - wk S. Staging wk by token tile keeps the reduction loop entirely in
133
+ // workgroup memory without holding the whole chunk.
134
+ for (var tb = 0u; tb < CHUNK; tb = tb + TOKEN_TILE) {
135
+ for (var e = tid; e < TOKEN_TILE * HEAD_DIM_K; e = e + WG) {
136
+ stage[e] = wk[(chunk_base + tb + e / HEAD_DIM_K) * HEAD_DIM_K + e % HEAD_DIM_K];
137
+ }
138
+ workgroupBarrier();
139
+ for (var m = 0u; m < TOKENS_PER_TILE; m = m + 1u) {
140
+ let local_t = group * TOKENS_PER_TILE + m;
141
+ let t = tb + local_t;
142
+ var a = uvec[(chunk_base + t) * HEAD_DIM_V + dv0 + col];
143
+ for (var i = 0u; i < HEAD_DIM_K; i = i + 1u) {
144
+ a = a - stage[local_t * HEAD_DIM_K + i] * st[i * TILE_V + col];
145
+ }
146
+ dl[t * TILE_V + col] = a;
147
+ }
148
+ workgroupBarrier();
149
+ }
150
+ for (var e = tid; e < CHUNK * TILE_V; e = e + WG) {
151
+ deltas[(chunk_base + e / TILE_V) * HEAD_DIM_V + dv0 + e % TILE_V] = dl[e];
152
+ }
153
+ {% else %}
154
+ // The non-delta rules take delta = v with no correction at all.
155
+ for (var e = tid; e < CHUNK * TILE_V; e = e + WG) {
156
+ let t = e / TILE_V;
157
+ var v = 0.0;
158
+ if (t < live) {
159
+ v = {{ read_scalar("value", "(base_t + t) * params.vPackedDim + head * HEAD_DIM_V + dv0 + e % TILE_V", queryDtype) }};
160
+ }
161
+ dl[e] = v;
162
+ }
163
+ {% endif %}
164
+
165
+ // Publish the entry state before the update: it is what the output pass reads.
166
+ let state_base = ((batch * params.kvNumHeads + head) * num_chunks + chunk) * HEAD_DIM_K * HEAD_DIM_V;
167
+ for (var e = tid; e < HEAD_DIM_K * TILE_V; e = e + WG) {
168
+ states[state_base + (e / TILE_V) * HEAD_DIM_V + dv0 + e % TILE_V] = st[e];
169
+ }
170
+ workgroupBarrier();
171
+
172
+ // S = A_c * (S + K_hat^T delta): an outer-product accumulation, so splitting the
173
+ // workgroup along the reduction axis here keeps every thread's rows private. The
174
+ // row accumulators live in registers across the whole chunk, which is what makes
175
+ // each staged k_hat element feed ROWS_PER_GROUP multiply-adds instead of one.
176
+ var acc: array<f32, ROWS_PER_GROUP>;
177
+ for (var m = 0u; m < ROWS_PER_GROUP; m = m + 1u) {
178
+ acc[m] = st[(group * ROWS_PER_GROUP + m) * TILE_V + col];
179
+ }
180
+ for (var tb = 0u; tb < CHUNK; tb = tb + TOKEN_TILE) {
181
+ for (var e = tid; e < TOKEN_TILE * HEAD_DIM_K; e = e + WG) {
182
+ let t = tb + e / HEAD_DIM_K;
183
+ var v = 0.0;
184
+ if (t < live) {
185
+ v = k_hat(base_t + t, head, e % HEAD_DIM_K);
186
+ }
187
+ stage[e] = v;
188
+ }
189
+ workgroupBarrier();
190
+ for (var lt = 0u; lt < TOKEN_TILE; lt = lt + 1u) {
191
+ let dv_val = dl[(tb + lt) * TILE_V + col];
192
+ for (var m = 0u; m < ROWS_PER_GROUP; m = m + 1u) {
193
+ acc[m] = acc[m] + stage[lt * HEAD_DIM_K + group * ROWS_PER_GROUP + m] * dv_val;
194
+ }
195
+ }
196
+ workgroupBarrier();
197
+ }
198
+ {% if usesDecay %}
199
+ let last_t = base_t + live - 1u;
200
+ {% endif %}
201
+ for (var m = 0u; m < ROWS_PER_GROUP; m = m + 1u) {
202
+ let i = group * ROWS_PER_GROUP + m;
203
+ {% if usesDecay %}
204
+ st[i * TILE_V + col] = acc[m] * {{ decay_at("last_t", "head", "i") }};
205
+ {% else %}
206
+ st[i * TILE_V + col] = acc[m];
207
+ {% endif %}
208
+ }
209
+ workgroupBarrier();
210
+ }
211
+
212
+ for (var e = tid; e < HEAD_DIM_K * TILE_V; e = e + WG) {
213
+ present_state[(head_base + e / TILE_V) * HEAD_DIM_V + dv0 + e % TILE_V] =
214
+ {{ write_scalar("st[e]", stateDtype) }};
215
+ }
216
+ }
build/webgpu/chunk-ut.wgsl.jinja ADDED
@@ -0,0 +1,234 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro read_scalar(name, index, dtype) %}
2
+ {% if dtype == "float16" %}
3
+ f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
4
+ {% endmacro %}{% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{%- endmacro %}
5
+
6
+ {% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
7
+ {% if needQuery %}
8
+ // The output scale defaults to 1/sqrt(head_dim_k), matching the recurrent kernels.
9
+ fn out_scale() -> f32 {
10
+ return select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
11
+ }
12
+ {% endif %}
13
+ fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
14
+ let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
15
+ {% if usesDecay %}
16
+ return raw / {{ decay_at("bt", "h", "i") }};
17
+ {% else %}
18
+ return raw;
19
+ {% endif %}
20
+ }
21
+ {% if needKtil %}
22
+
23
+ fn k_til(bt: u32, h: u32, i: u32) -> f32 {
24
+ let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
25
+ {% if usesDecay %}
26
+ return raw * {{ decay_at("bt", "h", "i") }};
27
+ {% else %}
28
+ return raw;
29
+ {% endif %}
30
+ }
31
+ {% endif %}
32
+ {% if needQuery %}
33
+
34
+ // The output scale rides on q_til so both output terms -- q_til * S and P * delta,
35
+ // where P is itself built from q_til -- pick it up without a second pass over y.
36
+ fn q_til(bt: u32, q_head: u32, {% if usesDecay %}h: u32, {% endif %}i: u32, scale: f32) -> f32 {
37
+ let raw = {{ read_scalar("query", "bt * params.qPackedDim + q_head * HEAD_DIM_K + i", queryDtype) }};
38
+ {% if usesDecay %}
39
+ return raw * {{ decay_at("bt", "h", "i") }} * scale;
40
+ {% else %}
41
+ return raw * scale;
42
+ {% endif %}
43
+ }
44
+ {% endif %}
45
+ {%- endmacro %}{% macro emit_chunk_head_setup(needHeads=false) %}
46
+ {% if needHeads %}
47
+ let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
48
+ let packed_out = max(params.qNumHeads, params.kvNumHeads) * HEAD_DIM_V;
49
+ {% endif %}
50
+ let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;
51
+ {%- endmacro %}
52
+
53
+ {% if queryDtype == "float16" %}
54
+ enable f16;
55
+ {% endif %}
56
+ {{ env.wgsl.resourceDeclarations }}
57
+
58
+ // com.microsoft.LinearAttention, chunked prefill: the chunk-local delta transform.
59
+ //
60
+ // The delta rules apply S_t = (I - beta_t k_t k_t^T) diag(a_t) S_{t-1} + beta_t k_t v_t^T,
61
+ // whose per-token correction is what forces the recurrent kernels to serialize. Over
62
+ // one chunk the corrections satisfy a unit lower triangular system
63
+ // (I + W) d = beta * (v - S_0^T k_til), W[t,s] = beta_t * dot(k_til_t, k_hat_s), s < t
64
+ // so d = u - wk * S_0 with u = Tb * V and wk = Tb * K_til, Tb = (I + W)^-1 diag(beta).
65
+ // Both depend only on this chunk's own key, value and beta -- not on the entry state --
66
+ // so every chunk in the sequence computes them at once here, and the sequential pass is
67
+ // left with two dense products per chunk instead of one dependent step per token.
68
+ const CHUNK: u32 = {{ chunkSize }}u;
69
+ const HEAD_DIM_K: u32 = {{ headDimK }}u;
70
+ const HEAD_DIM_V: u32 = {{ headDimV }}u;
71
+ const KV_PER_KEY_HEAD: u32 = {{ kvPerKeyHead }}u;
72
+ const TK: u32 = {{ chunkTileK }}u;
73
+ const WG: u32 = {{ workgroupSize }}u;
74
+ // Entries of the CHUNK x CHUNK system each thread carries across the reduction tiles.
75
+ const ENTRIES: u32 = (CHUNK * CHUNK) / WG;
76
+
77
+ {{ emit_chunk_operands(needKtil=true) }}
78
+
79
+ // `wm` holds W while the system is being built, then the inverse in place: row t of
80
+ // the inverse is produced from row t of W and the already-final rows above it.
81
+ var<workgroup> wm: array<f32, CHUNK * CHUNK>;
82
+ var<workgroup> rowbuf: array<f32, CHUNK>;
83
+ var<workgroup> betas: array<f32, CHUNK>;
84
+ var<workgroup> tile: array<f32, CHUNK * TK>;
85
+ // Both W operands are staged: the inner product below runs CHUNK^2 times per tile,
86
+ // so a global read there would be paid once per pair rather than once per element.
87
+ var<workgroup> tile_til: array<f32, CHUNK * TK>;
88
+
89
+ @compute @workgroup_size(WG, 1, 1)
90
+ fn main(
91
+ @builtin(workgroup_id) wg: vec3<u32>,
92
+ @builtin(num_workgroups) nwg: vec3<u32>,
93
+ @builtin(local_invocation_id) lid: vec3<u32>,
94
+ ) {
95
+ let tid = lid.x;
96
+ let flat = wg.x + wg.y * nwg.x;
97
+ {{ emit_chunk_head_setup() }}
98
+ let chunk = flat % num_chunks;
99
+ let head = (flat / num_chunks) % params.kvNumHeads;
100
+ let batch = flat / (num_chunks * params.kvNumHeads);
101
+ if (batch >= params.batchSize) {
102
+ return;
103
+ }
104
+ let base_t = batch * params.seqLength + chunk * CHUNK;
105
+ // Tokens past the sequence end carry beta 0, which makes their rows of the system
106
+ // the identity and their outputs zero, so the tail chunk needs no separate arm.
107
+ let live = min(CHUNK, params.seqLength - chunk * CHUNK);
108
+
109
+ if (tid < CHUNK) {
110
+ var b = 0.0;
111
+ if (tid < live) {
112
+ let bt = base_t + tid;
113
+ b = {{ read_scalar("beta", "select(bt * params.kvNumHeads + head, bt, params.betaPackedDim == 1u)", queryDtype) }};
114
+ }
115
+ betas[tid] = b;
116
+ }
117
+ workgroupBarrier();
118
+
119
+ // W = beta * tril(K_til K_hat^T, -1), accumulated over reduction-axis tiles so the
120
+ // two operands share one staging buffer pass instead of a full CHUNK x HEAD_DIM_K copy.
121
+ var acc: array<f32, ENTRIES>;
122
+ for (var m = 0u; m < ENTRIES; m = m + 1u) {
123
+ acc[m] = 0.0;
124
+ }
125
+ for (var kb = 0u; kb < HEAD_DIM_K; kb = kb + TK) {
126
+ for (var e = tid; e < CHUNK * TK; e = e + WG) {
127
+ let t = e / TK;
128
+ let i = kb + e % TK;
129
+ var hat = 0.0;
130
+ var til = 0.0;
131
+ if (t < live && i < HEAD_DIM_K) {
132
+ hat = k_hat(base_t + t, head, i);
133
+ til = k_til(base_t + t, head, i);
134
+ }
135
+ tile[e] = hat;
136
+ tile_til[e] = til;
137
+ }
138
+ workgroupBarrier();
139
+ for (var m = 0u; m < ENTRIES; m = m + 1u) {
140
+ let e = tid + m * WG;
141
+ let t = e / CHUNK;
142
+ let s = e % CHUNK;
143
+ if (s < t) {
144
+ var total = 0.0;
145
+ for (var u = 0u; u < TK; u = u + 1u) {
146
+ total = total + tile_til[t * TK + u] * tile[s * TK + u];
147
+ }
148
+ acc[m] = acc[m] + total;
149
+ }
150
+ }
151
+ workgroupBarrier();
152
+ }
153
+ for (var m = 0u; m < ENTRIES; m = m + 1u) {
154
+ let e = tid + m * WG;
155
+ wm[e] = acc[m] * betas[e / CHUNK];
156
+ }
157
+ workgroupBarrier();
158
+
159
+ // Forward substitution for (I + W)^-1, one row per step. W is strictly lower, so
160
+ // row t reads only finished rows; `rowbuf` copies row t out before it is overwritten.
161
+ for (var t = 0u; t < CHUNK; t = t + 1u) {
162
+ if (tid < CHUNK) {
163
+ rowbuf[tid] = wm[t * CHUNK + tid];
164
+ }
165
+ workgroupBarrier();
166
+ if (tid < CHUNK) {
167
+ var a = select(0.0, 1.0, tid == t);
168
+ for (var r = 0u; r < t; r = r + 1u) {
169
+ a = a - rowbuf[r] * wm[r * CHUNK + tid];
170
+ }
171
+ wm[t * CHUNK + tid] = a;
172
+ }
173
+ workgroupBarrier();
174
+ }
175
+ // Tb = T diag(beta): scaling by column completes the transform.
176
+ for (var e = tid; e < CHUNK * CHUNK; e = e + WG) {
177
+ wm[e] = wm[e] * betas[e % CHUNK];
178
+ }
179
+ workgroupBarrier();
180
+
181
+ let chunk_base = ((batch * params.kvNumHeads + head) * num_chunks + chunk) * CHUNK;
182
+
183
+ // wk = Tb K_til, staged one reduction tile at a time.
184
+ for (var kb = 0u; kb < HEAD_DIM_K; kb = kb + TK) {
185
+ for (var e = tid; e < CHUNK * TK; e = e + WG) {
186
+ let s = e / TK;
187
+ let i = kb + e % TK;
188
+ var v = 0.0;
189
+ if (s < live && i < HEAD_DIM_K) {
190
+ v = k_til(base_t + s, head, i);
191
+ }
192
+ tile[e] = v;
193
+ }
194
+ workgroupBarrier();
195
+ for (var e = tid; e < CHUNK * TK; e = e + WG) {
196
+ let t = e / TK;
197
+ let u = e % TK;
198
+ if (kb + u < HEAD_DIM_K) {
199
+ var total = 0.0;
200
+ for (var s = 0u; s <= t; s = s + 1u) {
201
+ total = total + wm[t * CHUNK + s] * tile[s * TK + u];
202
+ }
203
+ wk[(chunk_base + t) * HEAD_DIM_K + kb + u] = total;
204
+ }
205
+ }
206
+ workgroupBarrier();
207
+ }
208
+
209
+ // u = Tb V, over the value axis.
210
+ for (var vb = 0u; vb < HEAD_DIM_V; vb = vb + TK) {
211
+ for (var e = tid; e < CHUNK * TK; e = e + WG) {
212
+ let s = e / TK;
213
+ let j = vb + e % TK;
214
+ var v = 0.0;
215
+ if (s < live && j < HEAD_DIM_V) {
216
+ v = {{ read_scalar("value", "(base_t + s) * params.vPackedDim + head * HEAD_DIM_V + j", queryDtype) }};
217
+ }
218
+ tile[e] = v;
219
+ }
220
+ workgroupBarrier();
221
+ for (var e = tid; e < CHUNK * TK; e = e + WG) {
222
+ let t = e / TK;
223
+ let u = e % TK;
224
+ if (vb + u < HEAD_DIM_V) {
225
+ var total = 0.0;
226
+ for (var s = 0u; s <= t; s = s + 1u) {
227
+ total = total + wm[t * CHUNK + s] * tile[s * TK + u];
228
+ }
229
+ uvec[(chunk_base + t) * HEAD_DIM_V + vb + u] = total;
230
+ }
231
+ }
232
+ workgroupBarrier();
233
+ }
234
+ }
build/webgpu/linear-attention.scalar.wgsl.jinja ADDED
@@ -0,0 +1,332 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro read_scalar(name, index, dtype) %}
2
+ {% if dtype == "float16" %}
3
+ f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
4
+ {% endmacro %}
5
+ {% macro write_scalar(expr, dtype) %}
6
+ {% if dtype == "float16" %}
7
+ f16({{ expr }}){% else %}{{ expr }}{% endif %}
8
+ {% endmacro -%}
9
+ {% macro emit_tiled_setup(dvGroups=1) %}
10
+ let head_dim_k = params.qPackedDim / params.qNumHeads;
11
+ let head_dim_v = params.vPackedDim / params.kvNumHeads;
12
+ let n_key_heads = params.kPackedDim / head_dim_k;
13
+ let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
14
+ let kv_per_key_head = params.kvNumHeads / n_key_heads;
15
+ {% if dvGroups == 1 %}
16
+ let dv_tiles = (head_dim_v + TILE_V - 1u) / TILE_V;
17
+ {% else %}
18
+ // Tile slots per workgroup-index step: each step covers DV_GROUPS value tiles.
19
+ let dv_tiles = ((head_dim_v + TILE_V - 1u) / TILE_V + DV_GROUPS - 1u) / DV_GROUPS;
20
+ {% endif %}
21
+ let scale = select(inverseSqrt(f32(head_dim_k)), params.scale, params.scale != 0.0);
22
+
23
+ // 2D-folded flat (batch*head*dv_tile) index: wg.y carries the high bits past
24
+ // the maxComputeWorkgroupsPerDimension dispatch limit. Reduces to wg.x when nwg.y == 1; the batch_idx >=
25
+ // params.batchSize guard drops the over-dispatched tail.
26
+ let workgroup_idx = wg.x + wg.y * nwg.x;
27
+ let dv_tile_idx = workgroup_idx % dv_tiles;
28
+ let bh = workgroup_idx / dv_tiles;
29
+ let head_idx = bh % params.kvNumHeads;
30
+ let batch_idx = bh / params.kvNumHeads;
31
+ if (batch_idx >= params.batchSize) {
32
+ return;
33
+ }
34
+
35
+ {% if dvGroups == 1 %}
36
+ let dv_start = dv_tile_idx * TILE_V;
37
+ {% else %}
38
+ let dv_start = (dv_tile_idx * DV_GROUPS + dv_group) * TILE_V;
39
+ {% endif %}
40
+ let packed_out = max(params.qNumHeads, params.kvNumHeads) * head_dim_v;
41
+ let key_head_idx = head_idx / kv_per_key_head;
42
+ {%- endmacro -%}
43
+ {% macro emit_query_groups(first_group) %}
44
+ for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
45
+ let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
46
+ let out_head = head_idx * heads_per_group + qg;
47
+ var q_val = 0.0;
48
+ {% if not useSubgroups %}
49
+ var local_pre: array<f32, TILE_V>;
50
+ {% endif %}
51
+ if (tid < head_dim_k) {
52
+ let q_idx = bt * params.qPackedDim + q_head * head_dim_k + tid;
53
+ q_val = {{ read_scalar("query", "q_idx", queryDtype) }};
54
+ }
55
+ {% if useSubgroups %}
56
+ let qg_subgroup_index = tid / sg_size;
57
+ let qg_subgroup_count = (WG + sg_size - 1u) / sg_size;
58
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
59
+ let sg_pre = subgroupAdd(state[j] * q_val);
60
+ if (sg_lid == 0u) {
61
+ red_preout[j * WG + qg_subgroup_index] = sg_pre;
62
+ }
63
+ }
64
+ workgroupBarrier();
65
+ if (WG > sg_size) {
66
+ if (tid < TILE_V) {
67
+ var pre_total = 0.0;
68
+ for (var i = 1u; i < qg_subgroup_count; i = i + 1u) {
69
+ pre_total = pre_total + red_preout[tid * WG + i];
70
+ }
71
+ red_preout[tid * WG] = red_preout[tid * WG] + pre_total;
72
+ }
73
+ workgroupBarrier();
74
+ }
75
+ {% else %}
76
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
77
+ red_preout[j * WG + tid] = state[j] * q_val;
78
+ }
79
+ workgroupBarrier();
80
+ for (var j = tid; j < TILE_V; j = j + WG) {
81
+ var pre_total = red_preout[j * WG];
82
+ for (var i = 1u; i < WG; i = i + 1u) {
83
+ pre_total = pre_total + red_preout[j * WG + i];
84
+ }
85
+ local_pre[j] = pre_total;
86
+ }
87
+ workgroupBarrier();
88
+ {% endif %}
89
+ {% if useSubgroups %}
90
+ if (tid == 0u) {
91
+ let out_base = bt * packed_out + out_head * head_dim_v + dv_start;
92
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
93
+ if (dv_start + j < head_dim_v) {
94
+ output[out_base + j] = {{ write_scalar("red_preout[j * WG] * scale", outputDtype) }};
95
+ }
96
+ }
97
+ }
98
+ workgroupBarrier();
99
+ {% else %}
100
+ let out_base = bt * packed_out + out_head * head_dim_v + dv_start;
101
+ for (var j = tid; j < TILE_V; j = j + WG) {
102
+ if (dv_start + j < head_dim_v) {
103
+ output[out_base + j] = {{ write_scalar("local_pre[j] * scale", outputDtype) }};
104
+ }
105
+ }
106
+ {% endif %}
107
+ }
108
+ {%- endmacro %}
109
+ {% set usesDecay = updateRule == "gated" or updateRule == "gated_delta" %}
110
+ {% set usesBeta = updateRule == "delta" or updateRule == "gated_delta" %}
111
+ {% if queryDtype == "float16" or stateDtype == "float16" %}
112
+ enable f16;
113
+ {% endif %}
114
+ {% if useSubgroups %}
115
+ enable subgroups;
116
+ {% endif %}
117
+ {{ env.wgsl.resourceDeclarations }}
118
+
119
+ const WG: u32 = {{ workgroupSize }}u;
120
+ const TILE_V: u32 = {{ tileV }}u;
121
+
122
+ {% if usesBeta %}var<workgroup> red_retrieved: array<f32, WG * TILE_V>;
123
+ {% endif %}var<workgroup> red_preout: array<f32, WG * TILE_V>;
124
+ {% if usesBeta %}var<workgroup> red_kq: array<f32, WG>;
125
+ var<workgroup> broadcast_delta: array<f32, TILE_V>;
126
+
127
+ {% endif %}
128
+ @compute @workgroup_size(WG, 1, 1)
129
+ fn main(
130
+ @builtin(workgroup_id) wg: vec3<u32>,
131
+ @builtin(num_workgroups) nwg: vec3<u32>,
132
+ @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
133
+ @builtin(subgroup_invocation_id) sg_lid: u32,
134
+ @builtin(subgroup_size) sg_size: u32{% endif %}
135
+ ) {
136
+ let tid = lid.x;
137
+ {{ emit_tiled_setup() }}
138
+
139
+ var state: array<f32, TILE_V>;
140
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
141
+ state[j] = 0.0;
142
+ }
143
+ {% if hasPastState %}
144
+
145
+ if (tid < head_dim_k) {
146
+ {% if hasStateWindow %}
147
+ // A windowed past_state is read only from slot stateWindow-1, the state after
148
+ // the last token of the previous call.
149
+ {% endif %}
150
+ let state_base = {% if hasStateWindow %}(params.stateWindow - 1u) * params.stateSlotStride + {% endif %}((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + tid) * head_dim_v + dv_start;
151
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
152
+ if (dv_start + j < head_dim_v) {
153
+ state[j] = {{ read_scalar("past_state", "state_base + j", stateDtype) }};
154
+ }
155
+ }
156
+ }
157
+
158
+ {% endif %}
159
+ {% if hasStateWindow %}
160
+ // Slots below max(0, stateWindow - seqLength) hold no token from this call.
161
+ if (tid < head_dim_k) {
162
+ for (var z = 0u; z + params.seqLength < params.stateWindow; z = z + 1u) {
163
+ let z_base = z * params.stateSlotStride
164
+ + ((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + tid) * head_dim_v + dv_start;
165
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
166
+ if (dv_start + j < head_dim_v) {
167
+ present_state[z_base + j] = {{ write_scalar("0.0", stateDtype) }};
168
+ }
169
+ }
170
+ }
171
+ }
172
+ {% endif %}
173
+ for (var t = 0u; t < params.seqLength; t = t + 1u) {
174
+ let bt = batch_idx * params.seqLength + t;
175
+
176
+ var k_val = 0.0;
177
+ if (tid < head_dim_k) {
178
+ let k_idx = bt * params.kPackedDim + key_head_idx * head_dim_k + tid;
179
+ k_val = {{ read_scalar("key", "k_idx", keyDtype) }};
180
+ }
181
+ {% if usesDecay %}
182
+
183
+ var decay_factor = 1.0;
184
+ if (params.decayPackedDim == params.kvNumHeads) {
185
+ decay_factor = exp({{ read_scalar("decay", "bt * params.kvNumHeads + head_idx", decayDtype) }});
186
+ } else if (tid < head_dim_k) {
187
+ let decay_idx = bt * params.decayPackedDim + head_idx * head_dim_k + tid;
188
+ decay_factor = exp({{ read_scalar("decay", "decay_idx", decayDtype) }});
189
+ }
190
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
191
+ state[j] = state[j] * decay_factor;
192
+ }
193
+
194
+ {% endif %}
195
+ {% if usesBeta %}
196
+ let q_head_0 = (head_idx * params.qNumHeads) / params.kvNumHeads;
197
+ let out_head_0 = head_idx * heads_per_group;
198
+ var q0_val = 0.0;
199
+ if (tid < head_dim_k) {
200
+ let q0_idx = bt * params.qPackedDim + q_head_0 * head_dim_k + tid;
201
+ q0_val = {{ read_scalar("query", "q0_idx", queryDtype) }};
202
+ }
203
+ {% if useSubgroups %}
204
+ let subgroup_index_b = tid / sg_size;
205
+ let subgroup_count_b = (WG + sg_size - 1u) / sg_size;
206
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
207
+ let sg_rp = subgroupAdd(vec2<f32>(state[j] * k_val, state[j] * q0_val));
208
+ if (sg_lid == 0u) {
209
+ red_retrieved[j * WG + subgroup_index_b] = sg_rp.x;
210
+ red_preout[j * WG + subgroup_index_b] = sg_rp.y;
211
+ }
212
+ }
213
+ let sg_kq = subgroupAdd(k_val * q0_val);
214
+ if (sg_lid == 0u) {
215
+ red_kq[subgroup_index_b] = sg_kq;
216
+ }
217
+ workgroupBarrier();
218
+ if (WG > sg_size) {
219
+ if (tid < TILE_V) {
220
+ var ret_total = 0.0;
221
+ var pre_total = 0.0;
222
+ for (var i = 1u; i < subgroup_count_b; i = i + 1u) {
223
+ ret_total = ret_total + red_retrieved[tid * WG + i];
224
+ pre_total = pre_total + red_preout[tid * WG + i];
225
+ }
226
+ red_retrieved[tid * WG] = red_retrieved[tid * WG] + ret_total;
227
+ red_preout[tid * WG] = red_preout[tid * WG] + pre_total;
228
+ }
229
+ if (tid == TILE_V) {
230
+ var kq_total = 0.0;
231
+ for (var i = 1u; i < subgroup_count_b; i = i + 1u) {
232
+ kq_total = kq_total + red_kq[i];
233
+ }
234
+ red_kq[0] = red_kq[0] + kq_total;
235
+ }
236
+ workgroupBarrier();
237
+ }
238
+ {% else %}
239
+ // Only the TILE_V output lanes consume these dot products. After one rendezvous,
240
+ // one lane folds each column and publishes it before the second barrier.
241
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
242
+ red_retrieved[j * WG + tid] = state[j] * k_val;
243
+ red_preout[j * WG + tid] = state[j] * q0_val;
244
+ }
245
+ red_kq[tid] = k_val * q0_val;
246
+ workgroupBarrier();
247
+
248
+ for (var j = tid; j < TILE_V; j = j + WG) {
249
+ var ret_total = red_retrieved[j * WG];
250
+ var pre_total = red_preout[j * WG];
251
+ for (var i = 1u; i < WG; i = i + 1u) {
252
+ ret_total = ret_total + red_retrieved[j * WG + i];
253
+ pre_total = pre_total + red_preout[j * WG + i];
254
+ }
255
+ red_retrieved[j * WG] = ret_total;
256
+ red_preout[j * WG] = pre_total;
257
+ }
258
+ if (tid == 0u) {
259
+ var kq_total = red_kq[0];
260
+ for (var i = 1u; i < WG; i = i + 1u) {
261
+ kq_total = kq_total + red_kq[i];
262
+ }
263
+ red_kq[0] = kq_total;
264
+ }
265
+ workgroupBarrier();
266
+ {% endif %}
267
+
268
+ if (tid == 0u) {
269
+ var beta_idx = bt * params.kvNumHeads + head_idx;
270
+ if (params.betaPackedDim == 1u) {
271
+ beta_idx = bt;
272
+ }
273
+ let beta_val = {{ read_scalar("beta", "beta_idx", betaDtype) }};
274
+ let v_base = bt * params.vPackedDim + head_idx * head_dim_v + dv_start;
275
+ let out_base = bt * packed_out + out_head_0 * head_dim_v + dv_start;
276
+ let kq_dot = red_kq[0];
277
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
278
+ if (dv_start + j < head_dim_v) {
279
+ let v_val = {{ read_scalar("value", "v_base + j", valueDtype) }};
280
+ let delta_j = beta_val * (v_val - red_retrieved[j * WG]);
281
+ broadcast_delta[j] = delta_j;
282
+ output[out_base + j] = {{ write_scalar("(red_preout[j * WG] + delta_j * kq_dot) * scale", outputDtype) }};
283
+ } else {
284
+ broadcast_delta[j] = 0.0;
285
+ }
286
+ }
287
+ }
288
+ workgroupBarrier();
289
+
290
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
291
+ state[j] = state[j] + k_val * broadcast_delta[j];
292
+ }
293
+
294
+ {{ emit_query_groups(1) }}
295
+ {% else %}
296
+ let v_base = bt * params.vPackedDim + head_idx * head_dim_v + dv_start;
297
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
298
+ if (dv_start + j < head_dim_v) {
299
+ state[j] = state[j] + k_val * {{ read_scalar("value", "v_base + j", valueDtype) }};
300
+ }
301
+ }
302
+
303
+ {{ emit_query_groups(0) }}
304
+ {% endif %}
305
+ {% if hasStateWindow %}
306
+ // Snapshot the state after this token into its window slot. Slot j holds the
307
+ // state after token (seqLength - stateWindow + j), so this token owns slot
308
+ // (t + stateWindow - seqLength) whenever that lands inside the window.
309
+ if (t + params.stateWindow >= params.seqLength && tid < head_dim_k) {
310
+ let win_slot = t + params.stateWindow - params.seqLength;
311
+ let win_base = win_slot * params.stateSlotStride
312
+ + ((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + tid) * head_dim_v + dv_start;
313
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
314
+ if (dv_start + j < head_dim_v) {
315
+ present_state[win_base + j] = {{ write_scalar("state[j]", stateDtype) }};
316
+ }
317
+ }
318
+ }
319
+ {% endif %}
320
+ }
321
+
322
+ {% if not hasStateWindow %}
323
+ if (tid < head_dim_k) {
324
+ let state_base = ((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + tid) * head_dim_v + dv_start;
325
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
326
+ if (dv_start + j < head_dim_v) {
327
+ present_state[state_base + j] = {{ write_scalar("state[j]", stateDtype) }};
328
+ }
329
+ }
330
+ }
331
+ {% endif %}
332
+ }
build/webgpu/linear-attention.serial.wgsl.jinja ADDED
@@ -0,0 +1,155 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro read_scalar(name, index, dtype) %}
2
+ {% if dtype == "float16" %}
3
+ f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
4
+ {% endmacro %}
5
+ {% macro write_scalar(expr, dtype) %}
6
+ {% if dtype == "float16" %}
7
+ f16({{ expr }}){% else %}{{ expr }}{% endif %}
8
+ {% endmacro -%}
9
+ {% macro emit_serial_query_groups(first_group) %}
10
+ for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
11
+ let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
12
+ let out_head = head_idx * heads_per_group + qg;
13
+ let q_base = bt * params.qPackedDim + q_head * HEAD_DIM_K;
14
+ var q_out = 0.0;
15
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
16
+ q_out = q_out + state[d] * {{ read_scalar("query", "q_base + d", queryDtype) }};
17
+ }
18
+ let out_idx = bt * packed_out + out_head * head_dim_v + dv_idx;
19
+ output[out_idx] = {{ write_scalar("q_out * scale", outputDtype) }};
20
+ }
21
+ {%- endmacro -%}
22
+ {% set gatedDeltaRule = updateRule == "gated_delta" %}
23
+ {% if queryDtype == "float16" or stateDtype == "float16" %}
24
+ enable f16;
25
+ {% endif %}
26
+ {{ env.wgsl.resourceDeclarations }}
27
+
28
+ const HEAD_DIM_K: u32 = {{ headDimK }}u;
29
+
30
+ // Barrier-free small-dk recurrence. One invocation owns one
31
+ // (batch, kv-head, value-dimension) state column and keeps its complete dk
32
+ // slice private across the sequence. This trades dk-lane parallelism for zero
33
+ // workgroup synchronization. Selection limits this route to small dk.
34
+ @compute @workgroup_size(1, 1, 1)
35
+ fn main(
36
+ @builtin(workgroup_id) wg: vec3<u32>,
37
+ @builtin(num_workgroups) nwg: vec3<u32>,
38
+ ) {
39
+ let flat_idx = wg.x + wg.y * nwg.x;
40
+ let head_dim_v = params.vPackedDim / params.kvNumHeads;
41
+ let dv_idx = flat_idx % head_dim_v;
42
+ let bh = flat_idx / head_dim_v;
43
+ let head_idx = bh % params.kvNumHeads;
44
+ let batch_idx = bh / params.kvNumHeads;
45
+ if (batch_idx >= params.batchSize) {
46
+ return;
47
+ }
48
+
49
+ let n_key_heads = params.kPackedDim / HEAD_DIM_K;
50
+ let kv_per_key_head = params.kvNumHeads / n_key_heads;
51
+ let key_head_idx = head_idx / kv_per_key_head;
52
+ // Standard GQA has qNumHeads >= kvNumHeads and emits one output head per query head. Inverse
53
+ // GQA has kvNumHeads > qNumHeads and emits one per KV head, with several KV heads sharing a
54
+ // query head. max() makes the group count 1
55
+ // in that case, and then these two formulas cover both layouts with no branch:
56
+ // q_head = (head_idx * qNumHeads) / kvNumHeads + qg
57
+ // out_head = head_idx * heads_per_group + qg
58
+ // In the standard layout the division is exact and the two agree; in the inverse layout the
59
+ // group count is 1, so q_head floors several KV heads onto one query head and out_head is the
60
+ // KV head itself.
61
+ let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
62
+ let packed_out = max(params.qNumHeads, params.kvNumHeads) * head_dim_v;
63
+ let scale = select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
64
+
65
+ var state: array<f32, HEAD_DIM_K>;
66
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
67
+ {% if hasPastState %}
68
+ let state_idx = {% if hasStateWindow %}(params.stateWindow - 1u) * params.stateSlotStride + {% endif %}((batch_idx * params.kvNumHeads + head_idx) * HEAD_DIM_K + d) * head_dim_v + dv_idx;
69
+ state[d] = {{ read_scalar("past_state", "state_idx", stateDtype) }};
70
+ {% else %}
71
+ state[d] = 0.0;
72
+ {% endif %}
73
+ }
74
+
75
+ var key_local: array<f32, HEAD_DIM_K>;
76
+ {% if hasStateWindow %}
77
+ // Slots below max(0, stateWindow - seqLength) hold no token from this call.
78
+ for (var z = 0u; z + params.seqLength < params.stateWindow; z = z + 1u) {
79
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
80
+ present_state[z * params.stateSlotStride + ((batch_idx * params.kvNumHeads + head_idx) * HEAD_DIM_K + d) * head_dim_v + dv_idx] = {{ write_scalar("0.0", stateDtype) }};
81
+ }
82
+ }
83
+ {% endif %}
84
+ for (var t = 0u; t < params.seqLength; t = t + 1u) {
85
+ let bt = batch_idx * params.seqLength + t;
86
+ let key_base = bt * params.kPackedDim + key_head_idx * HEAD_DIM_K;
87
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
88
+ key_local[d] = {{ read_scalar("key", "key_base + d", keyDtype) }};
89
+ }
90
+
91
+ {% if gatedDeltaRule %}
92
+ if (params.decayPackedDim == params.kvNumHeads) {
93
+ let decay_factor = exp({{ read_scalar("decay", "bt * params.kvNumHeads + head_idx", decayDtype) }});
94
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
95
+ state[d] = state[d] * decay_factor;
96
+ }
97
+ } else {
98
+ let decay_base = bt * params.decayPackedDim + head_idx * HEAD_DIM_K;
99
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
100
+ state[d] = state[d] * exp({{ read_scalar("decay", "decay_base + d", decayDtype) }});
101
+ }
102
+ }
103
+
104
+ let q_head_0 = (head_idx * params.qNumHeads) / params.kvNumHeads;
105
+ let out_head_0 = head_idx * heads_per_group;
106
+ let q0_base = bt * params.qPackedDim + q_head_0 * HEAD_DIM_K;
107
+ var retrieved = 0.0;
108
+ var preout = 0.0;
109
+ var kq_dot = 0.0;
110
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
111
+ let q_val = {{ read_scalar("query", "q0_base + d", queryDtype) }};
112
+ retrieved = retrieved + state[d] * key_local[d];
113
+ preout = preout + state[d] * q_val;
114
+ kq_dot = kq_dot + key_local[d] * q_val;
115
+ }
116
+ var beta_idx = bt * params.kvNumHeads + head_idx;
117
+ if (params.betaPackedDim == 1u) {
118
+ beta_idx = bt;
119
+ }
120
+ let beta_val = {{ read_scalar("beta", "beta_idx", betaDtype) }};
121
+ let v_idx = bt * params.vPackedDim + head_idx * head_dim_v + dv_idx;
122
+ let delta = beta_val * ({{ read_scalar("value", "v_idx", valueDtype) }} - retrieved);
123
+ let out_idx_0 = bt * packed_out + out_head_0 * head_dim_v + dv_idx;
124
+ output[out_idx_0] = {{ write_scalar("(preout + delta * kq_dot) * scale", outputDtype) }};
125
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
126
+ state[d] = state[d] + key_local[d] * delta;
127
+ }
128
+ {{ emit_serial_query_groups(1) }}
129
+ {% else %}
130
+ let v_idx = bt * params.vPackedDim + head_idx * head_dim_v + dv_idx;
131
+ let v_val = {{ read_scalar("value", "v_idx", valueDtype) }};
132
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
133
+ state[d] = state[d] + key_local[d] * v_val;
134
+ }
135
+ {{ emit_serial_query_groups(0) }}
136
+ {% endif %}
137
+ {% if hasStateWindow %}
138
+ // Slot j holds the state after token (seqLength - stateWindow + j), so this
139
+ // token owns slot (t + stateWindow - seqLength) when that lands in the window.
140
+ if (t + params.stateWindow >= params.seqLength) {
141
+ let win_slot = t + params.stateWindow - params.seqLength;
142
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
143
+ present_state[win_slot * params.stateSlotStride + ((batch_idx * params.kvNumHeads + head_idx) * HEAD_DIM_K + d) * head_dim_v + dv_idx] = {{ write_scalar("state[d]", stateDtype) }};
144
+ }
145
+ }
146
+ {% endif %}
147
+ }
148
+
149
+ {% if not hasStateWindow %}
150
+ for (var d = 0u; d < HEAD_DIM_K; d = d + 1u) {
151
+ let state_idx = ((batch_idx * params.kvNumHeads + head_idx) * HEAD_DIM_K + d) * head_dim_v + dv_idx;
152
+ present_state[state_idx] = {{ write_scalar("state[d]", stateDtype) }};
153
+ }
154
+ {% endif %}
155
+ }
build/webgpu/linear-attention.vec4.wgsl.jinja ADDED
@@ -0,0 +1,349 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {% macro read_scalar(name, index, dtype) %}
2
+ {% if dtype == "float16" %}
3
+ f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
4
+ {% endmacro %}
5
+ {% macro write_scalar(expr, dtype) %}
6
+ {% if dtype == "float16" %}
7
+ f16({{ expr }}){% else %}{{ expr }}{% endif %}
8
+ {% endmacro -%}
9
+ {% macro emit_tiled_setup(dvGroups=1) %}
10
+ let head_dim_k = params.qPackedDim / params.qNumHeads;
11
+ let head_dim_v = params.vPackedDim / params.kvNumHeads;
12
+ let n_key_heads = params.kPackedDim / head_dim_k;
13
+ let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
14
+ let kv_per_key_head = params.kvNumHeads / n_key_heads;
15
+ {% if dvGroups == 1 %}
16
+ let dv_tiles = (head_dim_v + TILE_V - 1u) / TILE_V;
17
+ {% else %}
18
+ // Tile slots per workgroup-index step: each step covers DV_GROUPS value tiles.
19
+ let dv_tiles = ((head_dim_v + TILE_V - 1u) / TILE_V + DV_GROUPS - 1u) / DV_GROUPS;
20
+ {% endif %}
21
+ let scale = select(inverseSqrt(f32(head_dim_k)), params.scale, params.scale != 0.0);
22
+
23
+ // 2D-folded flat (batch*head*dv_tile) index: wg.y carries the high bits past
24
+ // the maxComputeWorkgroupsPerDimension dispatch limit. Reduces to wg.x when nwg.y == 1; the batch_idx >=
25
+ // params.batchSize guard drops the over-dispatched tail.
26
+ let workgroup_idx = wg.x + wg.y * nwg.x;
27
+ let dv_tile_idx = workgroup_idx % dv_tiles;
28
+ let bh = workgroup_idx / dv_tiles;
29
+ let head_idx = bh % params.kvNumHeads;
30
+ let batch_idx = bh / params.kvNumHeads;
31
+ if (batch_idx >= params.batchSize) {
32
+ return;
33
+ }
34
+
35
+ {% if dvGroups == 1 %}
36
+ let dv_start = dv_tile_idx * TILE_V;
37
+ {% else %}
38
+ let dv_start = (dv_tile_idx * DV_GROUPS + dv_group) * TILE_V;
39
+ {% endif %}
40
+ let packed_out = max(params.qNumHeads, params.kvNumHeads) * head_dim_v;
41
+ let key_head_idx = head_idx / kv_per_key_head;
42
+ {%- endmacro -%}
43
+ {% macro emit_vec4_query_groups(first_group) %}
44
+ for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
45
+ let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
46
+ let out_head = head_idx * heads_per_group + qg;
47
+ var q_vec = vec4<f32>(0.0);
48
+ if (lane_active) {
49
+ let q_base = (bt * params.qPackedDim + q_head * head_dim_k) / VEC4_LANES + lane;
50
+ q_vec = vec4<f32>(query[q_base]);
51
+ }
52
+ var pre_qg: array<f32, TILE_V>;
53
+ {% if useSubgroups %}
54
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
55
+ pre_qg[j] = subgroupAdd(dot(state[j], q_vec));
56
+ }
57
+ {% else %}
58
+ {% for v in range(redVecs) %}
59
+ wg_fold[{{ v }}u * WG + tid] = {{ redT }}({% for c in range(redWidth) %}dot(state[{{ v * redWidth + c }}], q_vec){{ ", " if not loop.last else "" }}{% endfor %});
60
+ {% endfor %}
61
+ workgroupBarrier();
62
+ // One lane folds each staged slot. LANES is the reduction width, which small
63
+ // head_dim_k can drive below the slot count, so lanes stride over the slots.
64
+ for (var sl = lane; sl < {{ redVecs }}u; sl = sl + LANES) {
65
+ var t = {{ redT }}(0.0);
66
+ for (var i = 0u; i < LANES; i = i + 1u) {
67
+ t = t + wg_fold[sl * WG + dv_group * LANES + i];
68
+ }
69
+ wg_fold_out[dv_group * {{ redSlots }}u + sl] = t;
70
+ }
71
+ workgroupBarrier();
72
+ {% for j in range(tileV) %}
73
+ pre_qg[{{ j }}] = wg_fold_out[dv_group * {{ redSlots }}u + {{ (j / redWidth)|int }}u]{{ ("." ~ redComps[j % redWidth]) if redWidth > 1 else "" }};
74
+ {% endfor %}
75
+ {% endif %}
76
+ if (lane == 0u) {
77
+ let out_base_qg = bt * packed_out + out_head * head_dim_v + dv_start;
78
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
79
+ if (dv_start + j < head_dim_v) {
80
+ output[out_base_qg + j] = {{ write_scalar("pre_qg[j] * scale", queryDtype) }};
81
+ }
82
+ }
83
+ }
84
+ }
85
+ {%- endmacro %}
86
+ {% set usesDecay = updateRule == "gated" or updateRule == "gated_delta" %}
87
+ {% set redWidth = 4 if tileV % 4 == 0 else (2 if tileV % 2 == 0 else 1) %}
88
+ {% set redVecs = (tileV / redWidth)|int %}
89
+ {% set redT = ("vec" ~ redWidth ~ "<f32>") if redWidth > 1 else "f32" %}
90
+ {% set redComps = ["x", "y", "z", "w"] %}
91
+ {% set usesBeta = updateRule == "delta" or updateRule == "gated_delta" %}
92
+ {% set redSlots = (2 * redVecs + 1) if usesBeta else redVecs %}
93
+ {% if queryDtype == "float16" or stateDtype == "float16" %}
94
+ enable f16;
95
+ {% endif %}
96
+ {% if useSubgroups %}
97
+ enable subgroups;
98
+ {% endif %}
99
+ {{ env.wgsl.resourceDeclarations }}
100
+
101
+ // LANES threads cooperate on one value tile's reduction axis; DV_GROUPS such groups
102
+ // share a workgroup so the per-token key/query/decay stream is fetched once for all of
103
+ // them instead of once per value tile.
104
+ const LANES: u32 = {{ vec4Lanes }}u;
105
+ const DV_GROUPS: u32 = {{ dvGroups }}u;
106
+ const WG: u32 = LANES * DV_GROUPS;
107
+ const TILE_V: u32 = {{ tileV }}u;
108
+ const VEC4_LANES: u32 = 4u;
109
+ // query and key are bound as vec4, so a lane's four reduction components arrive in one
110
+ // 16-byte fetch instead of four scalar ones. Every row base this shader indexes is a
111
+ // multiple of four -- head_dim_k % 4 == 0 gates the family, and the packed dims are
112
+ // whole numbers of heads -- so the vec4 index is the row base in vec4 units plus the
113
+ // lane. The remaining scalar reads (decay's per-head arm, value, beta) are not aligned
114
+ // groups of four and keep their scalar view.
115
+ {% if not useSubgroups %}
116
+
117
+ // No-subgroup tier: shared-memory linear folds replace subgroupAdd. Every quantity a
118
+ // token reduces is staged before one barrier and read back after it, so the barrier
119
+ // count is a property of the token and not of how many quantities it reduces.
120
+ // A lane's TILE_V components are adjacent, so they travel as one {{ redT }} word and
121
+ // fold with {{ redWidth }}-wide adds -- the same per-component summation order as a
122
+ // scalar fold, at a {{ redWidth }}th of the shared-memory transactions.
123
+ // Folds stay within the caller's own lane group: groups own disjoint value tiles and
124
+ // must not see each other's partials.
125
+ var<workgroup> wg_fold: array<{{ redT }}, WG * {{ redSlots }}u>;
126
+ var<workgroup> wg_fold_out: array<{{ redT }}, DV_GROUPS * {{ redSlots }}u>;
127
+ {% endif %}
128
+ @compute @workgroup_size(WG, 1, 1)
129
+ fn main(
130
+ @builtin(workgroup_id) wg: vec3<u32>,
131
+ @builtin(num_workgroups) nwg: vec3<u32>,
132
+ @builtin(local_invocation_id) lid: vec3<u32>,
133
+ ) {
134
+ let tid = lid.x;
135
+ let lane = tid % LANES;
136
+ {% if dvGroups > 1 or not useSubgroups %}
137
+ let dv_group = tid / LANES;
138
+ {% endif %}
139
+ let dk_base = lane * VEC4_LANES;
140
+ {{ emit_tiled_setup(dvGroups=dvGroups) }}
141
+ let lane_active = dk_base < head_dim_k;
142
+
143
+ // state[j] holds 4 consecutive dk rows (the 4 components) for dv slot j.
144
+ var state: array<vec4<f32>, TILE_V>;
145
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
146
+ state[j] = vec4<f32>(0.0);
147
+ }
148
+ {% if hasPastState %}
149
+
150
+ if (lane_active) {
151
+ {% if hasStateWindow %}
152
+ // A windowed past_state is read only from slot stateWindow-1, the state after
153
+ // the last token of the previous call.
154
+ {% endif %}
155
+ let state_row_base = {% if hasStateWindow %}(params.stateWindow - 1u) * params.stateSlotStride + {% endif %}((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + dk_base) * head_dim_v + dv_start;
156
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
157
+ if (dv_start + j < head_dim_v) {
158
+ state[j] = vec4<f32>(
159
+ {{ read_scalar("past_state", "state_row_base + 0u * head_dim_v + j", stateDtype) }},
160
+ {{ read_scalar("past_state", "state_row_base + 1u * head_dim_v + j", stateDtype) }},
161
+ {{ read_scalar("past_state", "state_row_base + 2u * head_dim_v + j", stateDtype) }},
162
+ {{ read_scalar("past_state", "state_row_base + 3u * head_dim_v + j", stateDtype) }},
163
+ );
164
+ }
165
+ }
166
+ }
167
+
168
+ {% endif %}
169
+ {% if hasStateWindow %}
170
+ // Slots below max(0, stateWindow - seqLength) hold no token from this call.
171
+ if (lane_active) {
172
+ let zero_row_base = ((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + dk_base) * head_dim_v + dv_start;
173
+ for (var z = 0u; z + params.seqLength < params.stateWindow; z = z + 1u) {
174
+ let zb = z * params.stateSlotStride + zero_row_base;
175
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
176
+ if (dv_start + j < head_dim_v) {
177
+ present_state[zb + 0u * head_dim_v + j] = {{ write_scalar("0.0", stateDtype) }};
178
+ present_state[zb + 1u * head_dim_v + j] = {{ write_scalar("0.0", stateDtype) }};
179
+ present_state[zb + 2u * head_dim_v + j] = {{ write_scalar("0.0", stateDtype) }};
180
+ present_state[zb + 3u * head_dim_v + j] = {{ write_scalar("0.0", stateDtype) }};
181
+ }
182
+ }
183
+ }
184
+ }
185
+ {% endif %}
186
+ // A token's key and query do not depend on the recurrent state, so they are fetched one
187
+ // iteration ahead: the fetch for t+1 is issued before the reductions for t, and its memory
188
+ // latency overlaps the dependent chain instead of stalling in front of it. The state
189
+ // recurrence is what serializes this loop, and it leaves the load unit idle otherwise.
190
+ {% macro load_key(token) %}
191
+ if (lane_active) {
192
+ let k_base = ({{ token }} * params.kPackedDim + key_head_idx * head_dim_k) / VEC4_LANES + lane;
193
+ k_next = vec4<f32>(key[k_base]);
194
+ }
195
+ {%- endmacro %}
196
+ {% if usesBeta %}
197
+ {% macro load_query0(token) %}
198
+ if (lane_active) {
199
+ let q0_base = ({{ token }} * params.qPackedDim + q_head_0 * head_dim_k) / VEC4_LANES + lane;
200
+ q0_next = vec4<f32>(query[q0_base]);
201
+ }
202
+ {%- endmacro %}
203
+ let q_head_0 = (head_idx * params.qNumHeads) / params.kvNumHeads;
204
+ let out_head_0 = head_idx * heads_per_group;
205
+ var q0_next = vec4<f32>(0.0);
206
+ {% endif %}
207
+ var k_next = vec4<f32>(0.0);
208
+ let bt_first = batch_idx * params.seqLength;
209
+ {{ load_key("bt_first") }}
210
+ {% if usesBeta %}
211
+ {{ load_query0("bt_first") }}
212
+ {% endif %}
213
+ for (var t = 0u; t < params.seqLength; t = t + 1u) {
214
+ let bt = batch_idx * params.seqLength + t;
215
+
216
+ let k_vec = k_next;
217
+ {% if usesBeta %}
218
+ let q0_vec = q0_next;
219
+ {% endif %}
220
+ // Issued here, consumed on the next trip.
221
+ let bt_next = bt + 1u;
222
+ if (t + 1u < params.seqLength) {
223
+ {{ load_key("bt_next") }}
224
+ {% if usesBeta %}
225
+ {{ load_query0("bt_next") }}
226
+ {% endif %}
227
+ }
228
+ {% if usesDecay %}
229
+
230
+ var decay_vec = vec4<f32>(1.0);
231
+ if (params.decayPackedDim == params.kvNumHeads) {
232
+ let factor = exp({{ read_scalar("decay", "bt * params.kvNumHeads + head_idx", queryDtype) }});
233
+ decay_vec = vec4<f32>(factor);
234
+ } else if (lane_active) {
235
+ let decay_base = bt * params.decayPackedDim + head_idx * head_dim_k + dk_base;
236
+ decay_vec = vec4<f32>(
237
+ exp({{ read_scalar("decay", "decay_base + 0u", queryDtype) }}),
238
+ exp({{ read_scalar("decay", "decay_base + 1u", queryDtype) }}),
239
+ exp({{ read_scalar("decay", "decay_base + 2u", queryDtype) }}),
240
+ exp({{ read_scalar("decay", "decay_base + 3u", queryDtype) }}),
241
+ );
242
+ }
243
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
244
+ state[j] = state[j] * decay_vec;
245
+ }
246
+
247
+ {% endif %}
248
+ {% if usesBeta %}
249
+ var ret_vals: array<f32, TILE_V>;
250
+ var pre_vals: array<f32, TILE_V>;
251
+ {% if useSubgroups %}
252
+ let kq_dot = subgroupAdd(dot(k_vec, q0_vec));
253
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
254
+ ret_vals[j] = subgroupAdd(dot(state[j], k_vec));
255
+ pre_vals[j] = subgroupAdd(dot(state[j], q0_vec));
256
+ }
257
+ {% else %}
258
+ // kq and both state projections are known before any of them is consumed, so all
259
+ // three stage together and cross one barrier pair rather than two.
260
+ {% for v in range(redVecs) %}
261
+ wg_fold[{{ v }}u * WG + tid] = {{ redT }}({% for c in range(redWidth) %}dot(state[{{ v * redWidth + c }}], k_vec){{ ", " if not loop.last else "" }}{% endfor %});
262
+ wg_fold[{{ redVecs + v }}u * WG + tid] = {{ redT }}({% for c in range(redWidth) %}dot(state[{{ v * redWidth + c }}], q0_vec){{ ", " if not loop.last else "" }}{% endfor %});
263
+ {% endfor %}
264
+ wg_fold[{{ 2 * redVecs }}u * WG + tid] = {{ redT }}(dot(k_vec, q0_vec){% for c in range(redWidth - 1) %}, 0.0{% endfor %});
265
+ workgroupBarrier();
266
+ for (var sl = lane; sl < {{ redSlots }}u; sl = sl + LANES) {
267
+ var t = {{ redT }}(0.0);
268
+ for (var i = 0u; i < LANES; i = i + 1u) {
269
+ t = t + wg_fold[sl * WG + dv_group * LANES + i];
270
+ }
271
+ wg_fold_out[dv_group * {{ redSlots }}u + sl] = t;
272
+ }
273
+ workgroupBarrier();
274
+ let kq_dot = wg_fold_out[dv_group * {{ redSlots }}u + {{ 2 * redVecs }}u]{{ ("." ~ redComps[0]) if redWidth > 1 else "" }};
275
+ {% for j in range(tileV) %}
276
+ ret_vals[{{ j }}] = wg_fold_out[dv_group * {{ redSlots }}u + {{ (j / redWidth)|int }}u]{{ ("." ~ redComps[j % redWidth]) if redWidth > 1 else "" }};
277
+ pre_vals[{{ j }}] = wg_fold_out[dv_group * {{ redSlots }}u + {{ redVecs + (j / redWidth)|int }}u]{{ ("." ~ redComps[j % redWidth]) if redWidth > 1 else "" }};
278
+ {% endfor %}
279
+ {% endif %}
280
+
281
+ var beta_idx = bt * params.kvNumHeads + head_idx;
282
+ if (params.betaPackedDim == 1u) {
283
+ beta_idx = bt;
284
+ }
285
+ let beta_val = {{ read_scalar("beta", "beta_idx", queryDtype) }};
286
+ let v_base = bt * params.vPackedDim + head_idx * head_dim_v + dv_start;
287
+ let out_base = bt * packed_out + out_head_0 * head_dim_v + dv_start;
288
+ var deltas: array<f32, TILE_V>;
289
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
290
+ if (dv_start + j < head_dim_v) {
291
+ let v_val = {{ read_scalar("value", "v_base + j", queryDtype) }};
292
+ let delta_j = beta_val * (v_val - ret_vals[j]);
293
+ deltas[j] = delta_j;
294
+ if (lane == 0u) {
295
+ output[out_base + j] = {{ write_scalar("(pre_vals[j] + delta_j * kq_dot) * scale", queryDtype) }};
296
+ }
297
+ } else {
298
+ deltas[j] = 0.0;
299
+ }
300
+ }
301
+
302
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
303
+ state[j] = state[j] + k_vec * deltas[j];
304
+ }
305
+
306
+ {{ emit_vec4_query_groups(1) }}
307
+ {% else %}
308
+ let v_base = bt * params.vPackedDim + head_idx * head_dim_v + dv_start;
309
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
310
+ if (dv_start + j < head_dim_v) {
311
+ let v_val = {{ read_scalar("value", "v_base + j", queryDtype) }};
312
+ state[j] = state[j] + k_vec * v_val;
313
+ }
314
+ }
315
+
316
+ {{ emit_vec4_query_groups(0) }}
317
+ {% endif %}
318
+ {% if hasStateWindow %}
319
+ // Slot j holds the state after token (seqLength - stateWindow + j), so this
320
+ // token owns slot (t + stateWindow - seqLength) when that lands in the window.
321
+ if (t + params.stateWindow >= params.seqLength && lane_active) {
322
+ let win_base = (t + params.stateWindow - params.seqLength) * params.stateSlotStride
323
+ + ((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + dk_base) * head_dim_v + dv_start;
324
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
325
+ if (dv_start + j < head_dim_v) {
326
+ present_state[win_base + 0u * head_dim_v + j] = {{ write_scalar("state[j].x", stateDtype) }};
327
+ present_state[win_base + 1u * head_dim_v + j] = {{ write_scalar("state[j].y", stateDtype) }};
328
+ present_state[win_base + 2u * head_dim_v + j] = {{ write_scalar("state[j].z", stateDtype) }};
329
+ present_state[win_base + 3u * head_dim_v + j] = {{ write_scalar("state[j].w", stateDtype) }};
330
+ }
331
+ }
332
+ }
333
+ {% endif %}
334
+ }
335
+
336
+ {% if not hasStateWindow %}
337
+ if (lane_active) {
338
+ let state_row_base = ((batch_idx * params.kvNumHeads + head_idx) * head_dim_k + dk_base) * head_dim_v + dv_start;
339
+ for (var j = 0u; j < TILE_V; j = j + 1u) {
340
+ if (dv_start + j < head_dim_v) {
341
+ present_state[state_row_base + 0u * head_dim_v + j] = {{ write_scalar("state[j].x", stateDtype) }};
342
+ present_state[state_row_base + 1u * head_dim_v + j] = {{ write_scalar("state[j].y", stateDtype) }};
343
+ present_state[state_row_base + 2u * head_dim_v + j] = {{ write_scalar("state[j].z", stateDtype) }};
344
+ present_state[state_row_base + 3u * head_dim_v + j] = {{ write_scalar("state[j].w", stateDtype) }};
345
+ }
346
+ }
347
+ }
348
+ {% endif %}
349
+ }
build/webgpu/manifest.json ADDED
The diff for this file is too large to render. See raw diff
 
build/webgpu/metadata.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "com.microsoft.LinearAttention",
3
+ "id": "_com_microsoft_linearattention_webgpu_dc21501",
4
+ "version": 1,
5
+ "license": "Apache-2.0",
6
+ "backend": { "type": "webgpu" },
7
+ "digest": {
8
+ "algorithm": "sha256",
9
+ "files": {
10
+ "bench.json": "UnbCbJjzzX17ihI6vVaaXbbgxyNRzghqQ9jFtny3sZs=",
11
+ "chunk-out.wgsl.jinja": "n2BVpE+dJE8bWsb3Vv9EZalVcE6Bpr0+RYnUxCliSSQ=",
12
+ "chunk-prep.wgsl.jinja": "U9V7Cld2Zou0sX77ohI+0hmq+UZUsMQohiWrCLFrFm0=",
13
+ "chunk-scan.wgsl.jinja": "4yrz7CYz19GJAxJ5cBmc4XJt6aXa1ICUi+fg4ma2+mM=",
14
+ "chunk-ut.wgsl.jinja": "9KqJJCt5PwygvCrBjjy8oglG1ly4HiumH2Im/4Zlz5M=",
15
+ "linear-attention.scalar.wgsl.jinja": "mc7lY/Xgl1LkvRxiCeaSZ2DwldVlIeoTti2QFqhvCiI=",
16
+ "linear-attention.serial.wgsl.jinja": "p2AF3fRzazYsMPvGijvV3UyJhxvVxdZJyOhCYLRCJgI=",
17
+ "linear-attention.vec4.wgsl.jinja": "heG1tNzKu8g1M6PqKtDFED/CzGMDA/qWvlTuZ4Z557M=",
18
+ "manifest.json": "75pbsCWWrC10tm4NEaVyyoM7nDVgCqxwvj8Bh568M2g=",
19
+ "test.json": "NtZMMzB3k6UYTuxbvHod7FJlq+HqMV0oYuZi/qhW8sM="
20
+ }
21
+ },
22
+ "provenance": { "kernel": { "sha": "2e7068faf55e7f43df740015f6d1ee49391a41c5", "dirty": false } },
23
+ "webgpu": { "manifestSpec": "1.0", "specialized": true, "opPath": "ops/com.microsoft.LinearAttention" }
24
+ }
build/webgpu/test.json ADDED
The diff for this file is too large to render. See raw diff