sync 2e7068faf55e
Browse files- README.md +91 -0
- build/webgpu/bench.json +323 -0
- build/webgpu/chunk-out.wgsl.jinja +200 -0
- build/webgpu/chunk-prep.wgsl.jinja +47 -0
- build/webgpu/chunk-scan.wgsl.jinja +216 -0
- build/webgpu/chunk-ut.wgsl.jinja +234 -0
- build/webgpu/linear-attention.scalar.wgsl.jinja +332 -0
- build/webgpu/linear-attention.serial.wgsl.jinja +155 -0
- build/webgpu/linear-attention.vec4.wgsl.jinja +349 -0
- build/webgpu/manifest.json +0 -0
- build/webgpu/metadata.json +24 -0
- build/webgpu/test.json +0 -0
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
|
|
|