sync 6fdf6301e2bb
Browse files- README.md +19 -2
- build/webgpu/bench.json +1277 -1
- build/webgpu/manifest.json +277 -59
- build/webgpu/metadata.json +15 -11
- build/webgpu/sparse-attention-sgmat.wgsl.jinja +87 -61
- build/webgpu/sparse-attention.wgsl.jinja +53 -67
- build/webgpu/sparse-kv-append.wgsl.jinja +8 -8
- build/webgpu/sparse-q-rotary.wgsl.jinja +8 -8
- build/webgpu/test.json +0 -0
README.md
CHANGED
|
@@ -60,6 +60,23 @@ Attributes and default values (overridable per request):
|
|
| 60 |
| `T` | `float32`, `float16` |
|
| 61 |
| `M` | `int32` |
|
| 62 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
## Device requirements
|
| 64 |
|
| 65 |
Some implementation variants require `subgroup-matrix` and `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.
|
|
@@ -69,7 +86,7 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
|
|
| 69 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 70 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 71 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 72 |
-
- [`bench.json`](build/webgpu/bench.json) — benchmark
|
| 73 |
- [`sparse-attention-sgmat.wgsl.jinja`](build/webgpu/sparse-attention-sgmat.wgsl.jinja)
|
| 74 |
- [`sparse-attention.wgsl.jinja`](build/webgpu/sparse-attention.wgsl.jinja)
|
| 75 |
- [`sparse-kv-append.wgsl.jinja`](build/webgpu/sparse-kv-append.wgsl.jinja)
|
|
@@ -78,7 +95,7 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
|
|
| 78 |
## Use with `@huggingface/kernels`
|
| 79 |
|
| 80 |
```sh
|
| 81 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 82 |
```
|
| 83 |
|
| 84 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
|
|
| 60 |
| `T` | `float32`, `float16` |
|
| 61 |
| `M` | `int32` |
|
| 62 |
|
| 63 |
+
## Implementation variants
|
| 64 |
+
|
| 65 |
+
One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
|
| 66 |
+
|
| 67 |
+
- `separate` — Portable online attention shares each key/value load across query rows. Spare lanes partition value accumulation, with the partition count bounded by the head width, workgroup size, and available workgroup storage. Single-partition staging requires one output column per lane and includes the staging array in its storage budget.
|
| 68 |
+
- `separate_sgmat` — Float32 subgroup-matrix online attention computes each score tile once, rescales output fragments through existing score scratch, and directly loads complete query tiles. Only a partial query tile is staged; key-tile width and fragment-bank capacity follow the reported workgroup-storage limit.
|
| 69 |
+
- `separate_sgmat_tail` — For a prompt, a thin query tail uses the shared portable template after complete matrix tiles. History calls retain the full matrix computation. Tail size and value partitions follow workgroup capacity. Heads narrower than its query tile use the full matrix path, since their shorter value contraction does not repay a separate portable dispatch.
|
| 70 |
+
- `separate_rotary` — Portable online attention shares each key/value load across query rows. Spare lanes partition value accumulation, with the partition count bounded by the head width, workgroup size, and available workgroup storage. Single-partition staging requires one output column per lane and includes the staging array in its storage budget.
|
| 71 |
+
- `separate_rotary_sgmat` — Float32 subgroup-matrix online attention computes each score tile once, rescales output fragments through existing score scratch, and directly loads complete query tiles. Only a partial query tile is staged; key-tile width and fragment-bank capacity follow the reported workgroup-storage limit.
|
| 72 |
+
- `separate_rotary_sgmat_tail` — For a prompt, a thin query tail uses the shared portable template after complete matrix tiles. History calls retain the full matrix computation. Tail size and value partitions follow workgroup capacity. Heads narrower than its query tile use the full matrix path, since their shorter value contraction does not repay a separate portable dispatch.
|
| 73 |
+
- `packed` — Portable online attention shares each key/value load across query rows. Spare lanes partition value accumulation, with the partition count bounded by the head width, workgroup size, and available workgroup storage. Single-partition staging requires one output column per lane and includes the staging array in its storage budget.
|
| 74 |
+
- `packed_sgmat` — Float32 subgroup-matrix online attention computes each score tile once, rescales output fragments through existing score scratch, and directly loads complete query tiles. Only a partial query tile is staged; key-tile width and fragment-bank capacity follow the reported workgroup-storage limit.
|
| 75 |
+
- `packed_sgmat_tail` — For a prompt, a thin query tail uses the shared portable template after complete matrix tiles. History calls retain the full matrix computation. Tail size and value partitions follow workgroup capacity. Heads narrower than its query tile use the full matrix path, since their shorter value contraction does not repay a separate portable dispatch.
|
| 76 |
+
- `packed_rotary` — Portable online attention shares each key/value load across query rows. Spare lanes partition value accumulation, with the partition count bounded by the head width, workgroup size, and available workgroup storage. Single-partition staging requires one output column per lane and includes the staging array in its storage budget.
|
| 77 |
+
- `packed_rotary_sgmat` — Float32 subgroup-matrix online attention computes each score tile once, rescales output fragments through existing score scratch, and directly loads complete query tiles. Only a partial query tile is staged; key-tile width and fragment-bank capacity follow the reported workgroup-storage limit.
|
| 78 |
+
- `packed_rotary_sgmat_tail` — For a prompt, a thin query tail uses the shared portable template after complete matrix tiles. History calls retain the full matrix computation. Tail size and value partitions follow workgroup capacity. Heads narrower than its query tile use the full matrix path, since their shorter value contraction does not repay a separate portable dispatch.
|
| 79 |
+
|
| 80 |
## Device requirements
|
| 81 |
|
| 82 |
Some implementation variants require `subgroup-matrix` and `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.
|
|
|
|
| 86 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 87 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 88 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 89 |
+
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
|
| 90 |
- [`sparse-attention-sgmat.wgsl.jinja`](build/webgpu/sparse-attention-sgmat.wgsl.jinja)
|
| 91 |
- [`sparse-attention.wgsl.jinja`](build/webgpu/sparse-attention.wgsl.jinja)
|
| 92 |
- [`sparse-kv-append.wgsl.jinja`](build/webgpu/sparse-kv-append.wgsl.jinja)
|
|
|
|
| 95 |
## Use with `@huggingface/kernels`
|
| 96 |
|
| 97 |
```sh
|
| 98 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 99 |
```
|
| 100 |
|
| 101 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
build/webgpu/bench.json
CHANGED
|
@@ -1,7 +1,8 @@
|
|
| 1 |
{
|
| 2 |
"fixtureArrays": {
|
| 3 |
"block_column_indices_t_pattern": [0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 8, 0, 2, 3, 4, 5, 6, 7, 8, 9, 0, 3, 4, 5, 6, 7, 8, 9, 10, 0, 4, 5, 6, 7, 8, 9, 10, 11, 0, 4, 5, 6, 7, 8, 9, 10, 11, 12, 0, 4, 6, 7, 8, 9, 10, 11, 12, 13, 0, 4, 7, 8, 9, 10, 11, 12, 13, 14, 0, 4, 8, 9, 10, 11, 12, 13, 14, 15, 0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7, 8, 2, 3, 4, 5, 6, 7, 8, 9, 3, 4, 5, 6, 7, 8, 9, 10, 3, 4, 5, 6, 7, 8, 9, 10, 11, 3, 5, 6, 7, 8, 9, 10, 11, 12, 3, 6, 7, 8, 9, 10, 11, 12, 13, 3, 7, 8, 9, 10, 11, 12, 13, 14, 3, 7, 8, 9, 10, 11, 12, 13, 14, 15, -1, -1, -1, -1, -1, -1, 0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7, 8, 2, 3, 4, 5, 6, 7, 8, 9, 2, 3, 4, 5, 6, 7, 8, 9, 10, 2, 4, 5, 6, 7, 8, 9, 10, 11, 2, 5, 6, 7, 8, 9, 10, 11, 12, 2, 6, 7, 8, 9, 10, 11, 12, 13, 2, 6, 7, 8, 9, 10, 11, 12, 13, 14, 2, 6, 8, 9, 10, 11, 12, 13, 14, 15, -1, -1, -1, -1, 0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7, 8, 1, 2, 3, 4, 5, 6, 7, 8, 9, 1, 3, 4, 5, 6, 7, 8, 9, 10, 1, 4, 5, 6, 7, 8, 9, 10, 11, 1, 5, 6, 7, 8, 9, 10, 11, 12, 1, 5, 6, 7, 8, 9, 10, 11, 12, 13, 1, 5, 7, 8, 9, 10, 11, 12, 13, 14, 1, 5, 8, 9, 10, 11, 12, 13, 14, 15, -1, -1],
|
| 4 |
-
"sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT": [0, 1, 3, 6, 10, 15, 21, 28, 36, 45, 54, 63, 72, 82, 92, 102, 112, 0, 1, 3, 6, 10, 15, 21, 28, 36, 44, 52, 60, 69, 78, 87, 96, 106, 0, 1, 3, 6, 10, 15, 21, 28, 36, 44, 52, 61, 70, 79, 88, 98, 108, 0, 1, 3, 6, 10, 15, 21, 28, 36, 44, 53, 62, 71, 80, 90, 100, 110]
|
|
|
|
| 5 |
},
|
| 6 |
"tunableSpace": { "WORKGROUP_SIZE": [32, 64, 128, 256], "APPEND_WORKGROUP_SIZE": [64, 128, 256] },
|
| 7 |
"cases": [
|
|
@@ -922,6 +923,1281 @@
|
|
| 922 |
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 256 }
|
| 923 |
},
|
| 924 |
"outputs": { "outputT": { "shape": [1, 128, 768], "dtype": "float32" } }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 925 |
}
|
| 926 |
]
|
| 927 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"fixtureArrays": {
|
| 3 |
"block_column_indices_t_pattern": [0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7, 8, 0, 2, 3, 4, 5, 6, 7, 8, 9, 0, 3, 4, 5, 6, 7, 8, 9, 10, 0, 4, 5, 6, 7, 8, 9, 10, 11, 0, 4, 5, 6, 7, 8, 9, 10, 11, 12, 0, 4, 6, 7, 8, 9, 10, 11, 12, 13, 0, 4, 7, 8, 9, 10, 11, 12, 13, 14, 0, 4, 8, 9, 10, 11, 12, 13, 14, 15, 0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7, 8, 2, 3, 4, 5, 6, 7, 8, 9, 3, 4, 5, 6, 7, 8, 9, 10, 3, 4, 5, 6, 7, 8, 9, 10, 11, 3, 5, 6, 7, 8, 9, 10, 11, 12, 3, 6, 7, 8, 9, 10, 11, 12, 13, 3, 7, 8, 9, 10, 11, 12, 13, 14, 3, 7, 8, 9, 10, 11, 12, 13, 14, 15, -1, -1, -1, -1, -1, -1, 0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7, 8, 2, 3, 4, 5, 6, 7, 8, 9, 2, 3, 4, 5, 6, 7, 8, 9, 10, 2, 4, 5, 6, 7, 8, 9, 10, 11, 2, 5, 6, 7, 8, 9, 10, 11, 12, 2, 6, 7, 8, 9, 10, 11, 12, 13, 2, 6, 7, 8, 9, 10, 11, 12, 13, 14, 2, 6, 8, 9, 10, 11, 12, 13, 14, 15, -1, -1, -1, -1, 0, 0, 1, 0, 1, 2, 0, 1, 2, 3, 0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 0, 1, 2, 3, 4, 5, 6, 7, 1, 2, 3, 4, 5, 6, 7, 8, 1, 2, 3, 4, 5, 6, 7, 8, 9, 1, 3, 4, 5, 6, 7, 8, 9, 10, 1, 4, 5, 6, 7, 8, 9, 10, 11, 1, 5, 6, 7, 8, 9, 10, 11, 12, 1, 5, 6, 7, 8, 9, 10, 11, 12, 13, 1, 5, 7, 8, 9, 10, 11, 12, 13, 14, 1, 5, 8, 9, 10, 11, 12, 13, 14, 15, -1, -1],
|
| 4 |
+
"sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT": [0, 1, 3, 6, 10, 15, 21, 28, 36, 45, 54, 63, 72, 82, 92, 102, 112, 0, 1, 3, 6, 10, 15, 21, 28, 36, 44, 52, 60, 69, 78, 87, 96, 106, 0, 1, 3, 6, 10, 15, 21, 28, 36, 44, 52, 61, 70, 79, 88, 98, 108, 0, 1, 3, 6, 10, 15, 21, 28, 36, 44, 53, 62, 71, 80, 90, 100, 110],
|
| 5 |
+
"boundary-value-partition-float32-d520-s1_input_blockColIndicesT": [0, 0, 1, 1, 2, 1, 2, 3, -1, 0, 0, 1, 0, 1, 2, 0, 2, 3]
|
| 6 |
},
|
| 7 |
"tunableSpace": { "WORKGROUP_SIZE": [32, 64, 128, 256], "APPEND_WORKGROUP_SIZE": [64, 128, 256] },
|
| 8 |
"cases": [
|
|
|
|
| 923 |
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 256 }
|
| 924 |
},
|
| 925 |
"outputs": { "outputT": { "shape": [1, 128, 768], "dtype": "float32" } }
|
| 926 |
+
},
|
| 927 |
+
{
|
| 928 |
+
"name": "boundary-value-partition-float32-d520-s1",
|
| 929 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sparse_block_size": 16 },
|
| 930 |
+
"inputs": {
|
| 931 |
+
"queryT": {
|
| 932 |
+
"dtype": "float32",
|
| 933 |
+
"shape": [2, 1, 2080],
|
| 934 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 0.5 }
|
| 935 |
+
},
|
| 936 |
+
"keyT": {
|
| 937 |
+
"dtype": "float32",
|
| 938 |
+
"shape": [2, 1, 1040],
|
| 939 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.17, "cosStep": 0.31, "scale": 0.5 }
|
| 940 |
+
},
|
| 941 |
+
"valueT": {
|
| 942 |
+
"dtype": "float32",
|
| 943 |
+
"shape": [2, 1, 1040],
|
| 944 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.5 }
|
| 945 |
+
},
|
| 946 |
+
"pastKeyT": {
|
| 947 |
+
"dtype": "float32",
|
| 948 |
+
"shape": [2, 2, 64, 520],
|
| 949 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.017, "scale": 0.25 }
|
| 950 |
+
},
|
| 951 |
+
"pastValueT": {
|
| 952 |
+
"dtype": "float32",
|
| 953 |
+
"shape": [2, 2, 64, 520],
|
| 954 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.013, "scale": 0.25 }
|
| 955 |
+
},
|
| 956 |
+
"blockRowIndicesT": {
|
| 957 |
+
"dtype": "int32",
|
| 958 |
+
"shape": [2, 5],
|
| 959 |
+
"data": { "kind": "values", "values": [0, 1, 3, 5, 8, 0, 1, 3, 6, 9] }
|
| 960 |
+
},
|
| 961 |
+
"blockColIndicesT": {
|
| 962 |
+
"dtype": "int32",
|
| 963 |
+
"shape": [2, 9],
|
| 964 |
+
"data": {
|
| 965 |
+
"kind": "values",
|
| 966 |
+
"values": { "$ref": "#/fixtureArrays/boundary-value-partition-float32-d520-s1_input_blockColIndicesT" }
|
| 967 |
+
}
|
| 968 |
+
},
|
| 969 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [41] } },
|
| 970 |
+
"keyTotalSequenceLengthsT": { "dtype": "int32", "shape": [2], "data": { "kind": "values", "values": [41, 34] } }
|
| 971 |
+
},
|
| 972 |
+
"outputs": { "outputT": { "dtype": "float32", "shape": [2, 1, 2080] } },
|
| 973 |
+
"preset": "model"
|
| 974 |
+
},
|
| 975 |
+
{
|
| 976 |
+
"name": "boundary-value-partition-float32-d520-s5",
|
| 977 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sparse_block_size": 16 },
|
| 978 |
+
"inputs": {
|
| 979 |
+
"queryT": {
|
| 980 |
+
"dtype": "float32",
|
| 981 |
+
"shape": [2, 5, 2080],
|
| 982 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 0.5 }
|
| 983 |
+
},
|
| 984 |
+
"keyT": {
|
| 985 |
+
"dtype": "float32",
|
| 986 |
+
"shape": [2, 5, 1040],
|
| 987 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.17, "cosStep": 0.31, "scale": 0.5 }
|
| 988 |
+
},
|
| 989 |
+
"valueT": {
|
| 990 |
+
"dtype": "float32",
|
| 991 |
+
"shape": [2, 5, 1040],
|
| 992 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.5 }
|
| 993 |
+
},
|
| 994 |
+
"pastKeyT": {
|
| 995 |
+
"dtype": "float32",
|
| 996 |
+
"shape": [2, 2, 64, 520],
|
| 997 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.017, "scale": 0.25 }
|
| 998 |
+
},
|
| 999 |
+
"pastValueT": {
|
| 1000 |
+
"dtype": "float32",
|
| 1001 |
+
"shape": [2, 2, 64, 520],
|
| 1002 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.013, "scale": 0.25 }
|
| 1003 |
+
},
|
| 1004 |
+
"blockRowIndicesT": {
|
| 1005 |
+
"dtype": "int32",
|
| 1006 |
+
"shape": [2, 5],
|
| 1007 |
+
"data": { "kind": "values", "values": [0, 1, 3, 5, 8, 0, 1, 3, 6, 9] }
|
| 1008 |
+
},
|
| 1009 |
+
"blockColIndicesT": {
|
| 1010 |
+
"dtype": "int32",
|
| 1011 |
+
"shape": [2, 9],
|
| 1012 |
+
"data": {
|
| 1013 |
+
"kind": "values",
|
| 1014 |
+
"values": { "$ref": "#/fixtureArrays/boundary-value-partition-float32-d520-s1_input_blockColIndicesT" }
|
| 1015 |
+
}
|
| 1016 |
+
},
|
| 1017 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [41] } },
|
| 1018 |
+
"keyTotalSequenceLengthsT": { "dtype": "int32", "shape": [2], "data": { "kind": "values", "values": [41, 34] } }
|
| 1019 |
+
},
|
| 1020 |
+
"outputs": { "outputT": { "dtype": "float32", "shape": [2, 5, 2080] } },
|
| 1021 |
+
"preset": "model"
|
| 1022 |
+
},
|
| 1023 |
+
{
|
| 1024 |
+
"name": "boundary-value-owner-float32-d144-wg32-s1",
|
| 1025 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sparse_block_size": 16 },
|
| 1026 |
+
"inputs": {
|
| 1027 |
+
"queryT": {
|
| 1028 |
+
"dtype": "float32",
|
| 1029 |
+
"shape": [2, 1, 576],
|
| 1030 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 0.5 }
|
| 1031 |
+
},
|
| 1032 |
+
"keyT": {
|
| 1033 |
+
"dtype": "float32",
|
| 1034 |
+
"shape": [2, 1, 288],
|
| 1035 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.17, "cosStep": 0.31, "scale": 0.5 }
|
| 1036 |
+
},
|
| 1037 |
+
"valueT": {
|
| 1038 |
+
"dtype": "float32",
|
| 1039 |
+
"shape": [2, 1, 288],
|
| 1040 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.5 }
|
| 1041 |
+
},
|
| 1042 |
+
"pastKeyT": {
|
| 1043 |
+
"dtype": "float32",
|
| 1044 |
+
"shape": [2, 2, 64, 144],
|
| 1045 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.017, "scale": 0.25 }
|
| 1046 |
+
},
|
| 1047 |
+
"pastValueT": {
|
| 1048 |
+
"dtype": "float32",
|
| 1049 |
+
"shape": [2, 2, 64, 144],
|
| 1050 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.013, "scale": 0.25 }
|
| 1051 |
+
},
|
| 1052 |
+
"blockRowIndicesT": {
|
| 1053 |
+
"dtype": "int32",
|
| 1054 |
+
"shape": [2, 5],
|
| 1055 |
+
"data": { "kind": "values", "values": [0, 1, 3, 5, 8, 0, 1, 3, 6, 9] }
|
| 1056 |
+
},
|
| 1057 |
+
"blockColIndicesT": {
|
| 1058 |
+
"dtype": "int32",
|
| 1059 |
+
"shape": [2, 9],
|
| 1060 |
+
"data": {
|
| 1061 |
+
"kind": "values",
|
| 1062 |
+
"values": { "$ref": "#/fixtureArrays/boundary-value-partition-float32-d520-s1_input_blockColIndicesT" }
|
| 1063 |
+
}
|
| 1064 |
+
},
|
| 1065 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [41] } },
|
| 1066 |
+
"keyTotalSequenceLengthsT": { "dtype": "int32", "shape": [2], "data": { "kind": "values", "values": [41, 34] } }
|
| 1067 |
+
},
|
| 1068 |
+
"outputs": { "outputT": { "dtype": "float32", "shape": [2, 1, 576] } },
|
| 1069 |
+
"tunables": { "WORKGROUP_SIZE": 32 },
|
| 1070 |
+
"preset": "model"
|
| 1071 |
+
},
|
| 1072 |
+
{
|
| 1073 |
+
"name": "boundary-value-owner-float32-d144-wg32-s5",
|
| 1074 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sparse_block_size": 16 },
|
| 1075 |
+
"inputs": {
|
| 1076 |
+
"queryT": {
|
| 1077 |
+
"dtype": "float32",
|
| 1078 |
+
"shape": [2, 5, 576],
|
| 1079 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 0.5 }
|
| 1080 |
+
},
|
| 1081 |
+
"keyT": {
|
| 1082 |
+
"dtype": "float32",
|
| 1083 |
+
"shape": [2, 5, 288],
|
| 1084 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.17, "cosStep": 0.31, "scale": 0.5 }
|
| 1085 |
+
},
|
| 1086 |
+
"valueT": {
|
| 1087 |
+
"dtype": "float32",
|
| 1088 |
+
"shape": [2, 5, 288],
|
| 1089 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.5 }
|
| 1090 |
+
},
|
| 1091 |
+
"pastKeyT": {
|
| 1092 |
+
"dtype": "float32",
|
| 1093 |
+
"shape": [2, 2, 64, 144],
|
| 1094 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.017, "scale": 0.25 }
|
| 1095 |
+
},
|
| 1096 |
+
"pastValueT": {
|
| 1097 |
+
"dtype": "float32",
|
| 1098 |
+
"shape": [2, 2, 64, 144],
|
| 1099 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.013, "scale": 0.25 }
|
| 1100 |
+
},
|
| 1101 |
+
"blockRowIndicesT": {
|
| 1102 |
+
"dtype": "int32",
|
| 1103 |
+
"shape": [2, 5],
|
| 1104 |
+
"data": { "kind": "values", "values": [0, 1, 3, 5, 8, 0, 1, 3, 6, 9] }
|
| 1105 |
+
},
|
| 1106 |
+
"blockColIndicesT": {
|
| 1107 |
+
"dtype": "int32",
|
| 1108 |
+
"shape": [2, 9],
|
| 1109 |
+
"data": {
|
| 1110 |
+
"kind": "values",
|
| 1111 |
+
"values": { "$ref": "#/fixtureArrays/boundary-value-partition-float32-d520-s1_input_blockColIndicesT" }
|
| 1112 |
+
}
|
| 1113 |
+
},
|
| 1114 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [41] } },
|
| 1115 |
+
"keyTotalSequenceLengthsT": { "dtype": "int32", "shape": [2], "data": { "kind": "values", "values": [41, 34] } }
|
| 1116 |
+
},
|
| 1117 |
+
"outputs": { "outputT": { "dtype": "float32", "shape": [2, 5, 576] } },
|
| 1118 |
+
"tunables": { "WORKGROUP_SIZE": 32 },
|
| 1119 |
+
"preset": "model"
|
| 1120 |
+
},
|
| 1121 |
+
{
|
| 1122 |
+
"name": "sparse-prompt-tail-b1-s127-h32kv8-d128-blk64",
|
| 1123 |
+
"preset": "model",
|
| 1124 |
+
"vars": {
|
| 1125 |
+
"dtype": "float32",
|
| 1126 |
+
"batch": 1,
|
| 1127 |
+
"seq": 127,
|
| 1128 |
+
"heads": 32,
|
| 1129 |
+
"kvHeads": 8,
|
| 1130 |
+
"headDim": 128,
|
| 1131 |
+
"qkPairs": 260096,
|
| 1132 |
+
"attendedKeys": 1024,
|
| 1133 |
+
"dtypeBytes": 4
|
| 1134 |
+
},
|
| 1135 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1136 |
+
"inputs": {
|
| 1137 |
+
"queryT": { "shape": [1, 127, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1138 |
+
"keyT": { "shape": [1, 127, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1139 |
+
"valueT": { "shape": [1, 127, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1140 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1141 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1142 |
+
"blockRowIndicesT": {
|
| 1143 |
+
"shape": [4, 17],
|
| 1144 |
+
"dtype": "int32",
|
| 1145 |
+
"data": {
|
| 1146 |
+
"kind": "values",
|
| 1147 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1148 |
+
}
|
| 1149 |
+
},
|
| 1150 |
+
"blockColIndicesT": {
|
| 1151 |
+
"shape": [4, 112],
|
| 1152 |
+
"dtype": "int32",
|
| 1153 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1154 |
+
},
|
| 1155 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [127] } },
|
| 1156 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 127 }
|
| 1157 |
+
},
|
| 1158 |
+
"outputs": { "outputT": { "shape": [1, 127, 4096], "dtype": "float32" } },
|
| 1159 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1160 |
+
},
|
| 1161 |
+
{
|
| 1162 |
+
"name": "sparse-prompt-tail-b1-s129-h32kv8-d128-blk64",
|
| 1163 |
+
"preset": "model",
|
| 1164 |
+
"vars": {
|
| 1165 |
+
"dtype": "float32",
|
| 1166 |
+
"batch": 1,
|
| 1167 |
+
"seq": 129,
|
| 1168 |
+
"heads": 32,
|
| 1169 |
+
"kvHeads": 8,
|
| 1170 |
+
"headDim": 128,
|
| 1171 |
+
"qkPairs": 268320,
|
| 1172 |
+
"attendedKeys": 1024,
|
| 1173 |
+
"dtypeBytes": 4
|
| 1174 |
+
},
|
| 1175 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1176 |
+
"inputs": {
|
| 1177 |
+
"queryT": { "shape": [1, 129, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1178 |
+
"keyT": { "shape": [1, 129, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1179 |
+
"valueT": { "shape": [1, 129, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1180 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1181 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1182 |
+
"blockRowIndicesT": {
|
| 1183 |
+
"shape": [4, 17],
|
| 1184 |
+
"dtype": "int32",
|
| 1185 |
+
"data": {
|
| 1186 |
+
"kind": "values",
|
| 1187 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1188 |
+
}
|
| 1189 |
+
},
|
| 1190 |
+
"blockColIndicesT": {
|
| 1191 |
+
"shape": [4, 112],
|
| 1192 |
+
"dtype": "int32",
|
| 1193 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1194 |
+
},
|
| 1195 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [129] } },
|
| 1196 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 129 }
|
| 1197 |
+
},
|
| 1198 |
+
"outputs": { "outputT": { "shape": [1, 129, 4096], "dtype": "float32" } },
|
| 1199 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1200 |
+
},
|
| 1201 |
+
{
|
| 1202 |
+
"name": "sparse-prompt-tail-b1-s511-h32kv8-d128-blk64",
|
| 1203 |
+
"preset": "model",
|
| 1204 |
+
"vars": {
|
| 1205 |
+
"dtype": "float32",
|
| 1206 |
+
"batch": 1,
|
| 1207 |
+
"seq": 511,
|
| 1208 |
+
"heads": 32,
|
| 1209 |
+
"kvHeads": 8,
|
| 1210 |
+
"headDim": 128,
|
| 1211 |
+
"qkPairs": 4186112,
|
| 1212 |
+
"attendedKeys": 1024,
|
| 1213 |
+
"dtypeBytes": 4
|
| 1214 |
+
},
|
| 1215 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1216 |
+
"inputs": {
|
| 1217 |
+
"queryT": { "shape": [1, 511, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1218 |
+
"keyT": { "shape": [1, 511, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1219 |
+
"valueT": { "shape": [1, 511, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1220 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1221 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1222 |
+
"blockRowIndicesT": {
|
| 1223 |
+
"shape": [4, 17],
|
| 1224 |
+
"dtype": "int32",
|
| 1225 |
+
"data": {
|
| 1226 |
+
"kind": "values",
|
| 1227 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1228 |
+
}
|
| 1229 |
+
},
|
| 1230 |
+
"blockColIndicesT": {
|
| 1231 |
+
"shape": [4, 112],
|
| 1232 |
+
"dtype": "int32",
|
| 1233 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1234 |
+
},
|
| 1235 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [511] } },
|
| 1236 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 511 }
|
| 1237 |
+
},
|
| 1238 |
+
"outputs": { "outputT": { "shape": [1, 511, 4096], "dtype": "float32" } },
|
| 1239 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1240 |
+
},
|
| 1241 |
+
{
|
| 1242 |
+
"name": "sparse-prompt-tail-b1-s513-h32kv8-d128-blk64",
|
| 1243 |
+
"preset": "model",
|
| 1244 |
+
"vars": {
|
| 1245 |
+
"dtype": "float32",
|
| 1246 |
+
"batch": 1,
|
| 1247 |
+
"seq": 513,
|
| 1248 |
+
"heads": 32,
|
| 1249 |
+
"kvHeads": 8,
|
| 1250 |
+
"headDim": 128,
|
| 1251 |
+
"qkPairs": 4217376,
|
| 1252 |
+
"attendedKeys": 1024,
|
| 1253 |
+
"dtypeBytes": 4
|
| 1254 |
+
},
|
| 1255 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1256 |
+
"inputs": {
|
| 1257 |
+
"queryT": { "shape": [1, 513, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1258 |
+
"keyT": { "shape": [1, 513, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1259 |
+
"valueT": { "shape": [1, 513, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1260 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1261 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1262 |
+
"blockRowIndicesT": {
|
| 1263 |
+
"shape": [4, 17],
|
| 1264 |
+
"dtype": "int32",
|
| 1265 |
+
"data": {
|
| 1266 |
+
"kind": "values",
|
| 1267 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1268 |
+
}
|
| 1269 |
+
},
|
| 1270 |
+
"blockColIndicesT": {
|
| 1271 |
+
"shape": [4, 112],
|
| 1272 |
+
"dtype": "int32",
|
| 1273 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1274 |
+
},
|
| 1275 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [513] } },
|
| 1276 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 513 }
|
| 1277 |
+
},
|
| 1278 |
+
"outputs": { "outputT": { "shape": [1, 513, 4096], "dtype": "float32" } },
|
| 1279 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1280 |
+
},
|
| 1281 |
+
{
|
| 1282 |
+
"name": "sparse-prompt-tail-b1-s1023-h32kv8-d128-blk64",
|
| 1283 |
+
"preset": "model",
|
| 1284 |
+
"vars": {
|
| 1285 |
+
"dtype": "float32",
|
| 1286 |
+
"batch": 1,
|
| 1287 |
+
"seq": 1023,
|
| 1288 |
+
"heads": 32,
|
| 1289 |
+
"kvHeads": 8,
|
| 1290 |
+
"headDim": 128,
|
| 1291 |
+
"qkPairs": 13234176,
|
| 1292 |
+
"attendedKeys": 1024,
|
| 1293 |
+
"dtypeBytes": 4
|
| 1294 |
+
},
|
| 1295 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1296 |
+
"inputs": {
|
| 1297 |
+
"queryT": { "shape": [1, 1023, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1298 |
+
"keyT": { "shape": [1, 1023, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1299 |
+
"valueT": { "shape": [1, 1023, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1300 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1301 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1302 |
+
"blockRowIndicesT": {
|
| 1303 |
+
"shape": [4, 17],
|
| 1304 |
+
"dtype": "int32",
|
| 1305 |
+
"data": {
|
| 1306 |
+
"kind": "values",
|
| 1307 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1308 |
+
}
|
| 1309 |
+
},
|
| 1310 |
+
"blockColIndicesT": {
|
| 1311 |
+
"shape": [4, 112],
|
| 1312 |
+
"dtype": "int32",
|
| 1313 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1314 |
+
},
|
| 1315 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [1023] } },
|
| 1316 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 1023 }
|
| 1317 |
+
},
|
| 1318 |
+
"outputs": { "outputT": { "shape": [1, 1023, 4096], "dtype": "float32" } },
|
| 1319 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1320 |
+
},
|
| 1321 |
+
{
|
| 1322 |
+
"name": "hybrid-neighbor-s65-h32kv8-d128",
|
| 1323 |
+
"preset": "model",
|
| 1324 |
+
"vars": {
|
| 1325 |
+
"dtype": "float32",
|
| 1326 |
+
"batch": 1,
|
| 1327 |
+
"seq": 65,
|
| 1328 |
+
"heads": 32,
|
| 1329 |
+
"kvHeads": 8,
|
| 1330 |
+
"headDim": 128,
|
| 1331 |
+
"qkPairs": 68640,
|
| 1332 |
+
"attendedKeys": 65,
|
| 1333 |
+
"dtypeBytes": 4
|
| 1334 |
+
},
|
| 1335 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1336 |
+
"inputs": {
|
| 1337 |
+
"queryT": { "shape": [1, 65, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1338 |
+
"keyT": { "shape": [1, 65, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1339 |
+
"valueT": { "shape": [1, 65, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1340 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1341 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1342 |
+
"blockRowIndicesT": {
|
| 1343 |
+
"shape": [4, 17],
|
| 1344 |
+
"dtype": "int32",
|
| 1345 |
+
"data": {
|
| 1346 |
+
"kind": "values",
|
| 1347 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1348 |
+
}
|
| 1349 |
+
},
|
| 1350 |
+
"blockColIndicesT": {
|
| 1351 |
+
"shape": [4, 112],
|
| 1352 |
+
"dtype": "int32",
|
| 1353 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1354 |
+
},
|
| 1355 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [65] } },
|
| 1356 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 65 }
|
| 1357 |
+
},
|
| 1358 |
+
"outputs": { "outputT": { "shape": [1, 65, 4096], "dtype": "float32" } },
|
| 1359 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1360 |
+
},
|
| 1361 |
+
{
|
| 1362 |
+
"name": "hybrid-neighbor-s66-h32kv8-d128",
|
| 1363 |
+
"preset": "model",
|
| 1364 |
+
"vars": {
|
| 1365 |
+
"dtype": "float32",
|
| 1366 |
+
"batch": 1,
|
| 1367 |
+
"seq": 66,
|
| 1368 |
+
"heads": 32,
|
| 1369 |
+
"kvHeads": 8,
|
| 1370 |
+
"headDim": 128,
|
| 1371 |
+
"qkPairs": 70752,
|
| 1372 |
+
"attendedKeys": 66,
|
| 1373 |
+
"dtypeBytes": 4
|
| 1374 |
+
},
|
| 1375 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1376 |
+
"inputs": {
|
| 1377 |
+
"queryT": { "shape": [1, 66, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1378 |
+
"keyT": { "shape": [1, 66, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1379 |
+
"valueT": { "shape": [1, 66, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1380 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1381 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1382 |
+
"blockRowIndicesT": {
|
| 1383 |
+
"shape": [4, 17],
|
| 1384 |
+
"dtype": "int32",
|
| 1385 |
+
"data": {
|
| 1386 |
+
"kind": "values",
|
| 1387 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1388 |
+
}
|
| 1389 |
+
},
|
| 1390 |
+
"blockColIndicesT": {
|
| 1391 |
+
"shape": [4, 112],
|
| 1392 |
+
"dtype": "int32",
|
| 1393 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1394 |
+
},
|
| 1395 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [66] } },
|
| 1396 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 66 }
|
| 1397 |
+
},
|
| 1398 |
+
"outputs": { "outputT": { "shape": [1, 66, 4096], "dtype": "float32" } },
|
| 1399 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1400 |
+
},
|
| 1401 |
+
{
|
| 1402 |
+
"name": "hybrid-neighbor-s67-h32kv8-d128",
|
| 1403 |
+
"preset": "model",
|
| 1404 |
+
"vars": {
|
| 1405 |
+
"dtype": "float32",
|
| 1406 |
+
"batch": 1,
|
| 1407 |
+
"seq": 67,
|
| 1408 |
+
"heads": 32,
|
| 1409 |
+
"kvHeads": 8,
|
| 1410 |
+
"headDim": 128,
|
| 1411 |
+
"qkPairs": 72896,
|
| 1412 |
+
"attendedKeys": 67,
|
| 1413 |
+
"dtypeBytes": 4
|
| 1414 |
+
},
|
| 1415 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1416 |
+
"inputs": {
|
| 1417 |
+
"queryT": { "shape": [1, 67, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1418 |
+
"keyT": { "shape": [1, 67, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1419 |
+
"valueT": { "shape": [1, 67, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1420 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1421 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1422 |
+
"blockRowIndicesT": {
|
| 1423 |
+
"shape": [4, 17],
|
| 1424 |
+
"dtype": "int32",
|
| 1425 |
+
"data": {
|
| 1426 |
+
"kind": "values",
|
| 1427 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1428 |
+
}
|
| 1429 |
+
},
|
| 1430 |
+
"blockColIndicesT": {
|
| 1431 |
+
"shape": [4, 112],
|
| 1432 |
+
"dtype": "int32",
|
| 1433 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1434 |
+
},
|
| 1435 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [67] } },
|
| 1436 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 67 }
|
| 1437 |
+
},
|
| 1438 |
+
"outputs": { "outputT": { "shape": [1, 67, 4096], "dtype": "float32" } },
|
| 1439 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1440 |
+
},
|
| 1441 |
+
{
|
| 1442 |
+
"name": "hybrid-neighbor-s68-h32kv8-d128",
|
| 1443 |
+
"preset": "model",
|
| 1444 |
+
"vars": {
|
| 1445 |
+
"dtype": "float32",
|
| 1446 |
+
"batch": 1,
|
| 1447 |
+
"seq": 68,
|
| 1448 |
+
"heads": 32,
|
| 1449 |
+
"kvHeads": 8,
|
| 1450 |
+
"headDim": 128,
|
| 1451 |
+
"qkPairs": 75072,
|
| 1452 |
+
"attendedKeys": 68,
|
| 1453 |
+
"dtypeBytes": 4
|
| 1454 |
+
},
|
| 1455 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1456 |
+
"inputs": {
|
| 1457 |
+
"queryT": { "shape": [1, 68, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1458 |
+
"keyT": { "shape": [1, 68, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1459 |
+
"valueT": { "shape": [1, 68, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1460 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1461 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1462 |
+
"blockRowIndicesT": {
|
| 1463 |
+
"shape": [4, 17],
|
| 1464 |
+
"dtype": "int32",
|
| 1465 |
+
"data": {
|
| 1466 |
+
"kind": "values",
|
| 1467 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1468 |
+
}
|
| 1469 |
+
},
|
| 1470 |
+
"blockColIndicesT": {
|
| 1471 |
+
"shape": [4, 112],
|
| 1472 |
+
"dtype": "int32",
|
| 1473 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1474 |
+
},
|
| 1475 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [68] } },
|
| 1476 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 68 }
|
| 1477 |
+
},
|
| 1478 |
+
"outputs": { "outputT": { "shape": [1, 68, 4096], "dtype": "float32" } },
|
| 1479 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1480 |
+
},
|
| 1481 |
+
{
|
| 1482 |
+
"name": "hybrid-neighbor-s130-h32kv8-d128",
|
| 1483 |
+
"preset": "model",
|
| 1484 |
+
"vars": {
|
| 1485 |
+
"dtype": "float32",
|
| 1486 |
+
"batch": 1,
|
| 1487 |
+
"seq": 130,
|
| 1488 |
+
"heads": 32,
|
| 1489 |
+
"kvHeads": 8,
|
| 1490 |
+
"headDim": 128,
|
| 1491 |
+
"qkPairs": 272480,
|
| 1492 |
+
"attendedKeys": 130,
|
| 1493 |
+
"dtypeBytes": 4
|
| 1494 |
+
},
|
| 1495 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1496 |
+
"inputs": {
|
| 1497 |
+
"queryT": { "shape": [1, 130, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1498 |
+
"keyT": { "shape": [1, 130, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1499 |
+
"valueT": { "shape": [1, 130, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1500 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1501 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1502 |
+
"blockRowIndicesT": {
|
| 1503 |
+
"shape": [4, 17],
|
| 1504 |
+
"dtype": "int32",
|
| 1505 |
+
"data": {
|
| 1506 |
+
"kind": "values",
|
| 1507 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1508 |
+
}
|
| 1509 |
+
},
|
| 1510 |
+
"blockColIndicesT": {
|
| 1511 |
+
"shape": [4, 112],
|
| 1512 |
+
"dtype": "int32",
|
| 1513 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1514 |
+
},
|
| 1515 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [130] } },
|
| 1516 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 130 }
|
| 1517 |
+
},
|
| 1518 |
+
"outputs": { "outputT": { "shape": [1, 130, 4096], "dtype": "float32" } },
|
| 1519 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1520 |
+
},
|
| 1521 |
+
{
|
| 1522 |
+
"name": "hybrid-neighbor-s131-h32kv8-d128",
|
| 1523 |
+
"preset": "model",
|
| 1524 |
+
"vars": {
|
| 1525 |
+
"dtype": "float32",
|
| 1526 |
+
"batch": 1,
|
| 1527 |
+
"seq": 131,
|
| 1528 |
+
"heads": 32,
|
| 1529 |
+
"kvHeads": 8,
|
| 1530 |
+
"headDim": 128,
|
| 1531 |
+
"qkPairs": 276672,
|
| 1532 |
+
"attendedKeys": 131,
|
| 1533 |
+
"dtypeBytes": 4
|
| 1534 |
+
},
|
| 1535 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1536 |
+
"inputs": {
|
| 1537 |
+
"queryT": { "shape": [1, 131, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1538 |
+
"keyT": { "shape": [1, 131, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1539 |
+
"valueT": { "shape": [1, 131, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1540 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1541 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1542 |
+
"blockRowIndicesT": {
|
| 1543 |
+
"shape": [4, 17],
|
| 1544 |
+
"dtype": "int32",
|
| 1545 |
+
"data": {
|
| 1546 |
+
"kind": "values",
|
| 1547 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1548 |
+
}
|
| 1549 |
+
},
|
| 1550 |
+
"blockColIndicesT": {
|
| 1551 |
+
"shape": [4, 112],
|
| 1552 |
+
"dtype": "int32",
|
| 1553 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1554 |
+
},
|
| 1555 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [131] } },
|
| 1556 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 131 }
|
| 1557 |
+
},
|
| 1558 |
+
"outputs": { "outputT": { "shape": [1, 131, 4096], "dtype": "float32" } },
|
| 1559 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1560 |
+
},
|
| 1561 |
+
{
|
| 1562 |
+
"name": "hybrid-neighbor-s132-h32kv8-d128",
|
| 1563 |
+
"preset": "model",
|
| 1564 |
+
"vars": {
|
| 1565 |
+
"dtype": "float32",
|
| 1566 |
+
"batch": 1,
|
| 1567 |
+
"seq": 132,
|
| 1568 |
+
"heads": 32,
|
| 1569 |
+
"kvHeads": 8,
|
| 1570 |
+
"headDim": 128,
|
| 1571 |
+
"qkPairs": 280896,
|
| 1572 |
+
"attendedKeys": 132,
|
| 1573 |
+
"dtypeBytes": 4
|
| 1574 |
+
},
|
| 1575 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1576 |
+
"inputs": {
|
| 1577 |
+
"queryT": { "shape": [1, 132, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1578 |
+
"keyT": { "shape": [1, 132, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1579 |
+
"valueT": { "shape": [1, 132, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1580 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1581 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1582 |
+
"blockRowIndicesT": {
|
| 1583 |
+
"shape": [4, 17],
|
| 1584 |
+
"dtype": "int32",
|
| 1585 |
+
"data": {
|
| 1586 |
+
"kind": "values",
|
| 1587 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1588 |
+
}
|
| 1589 |
+
},
|
| 1590 |
+
"blockColIndicesT": {
|
| 1591 |
+
"shape": [4, 112],
|
| 1592 |
+
"dtype": "int32",
|
| 1593 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1594 |
+
},
|
| 1595 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [132] } },
|
| 1596 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 132 }
|
| 1597 |
+
},
|
| 1598 |
+
"outputs": { "outputT": { "shape": [1, 132, 4096], "dtype": "float32" } },
|
| 1599 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1600 |
+
},
|
| 1601 |
+
{
|
| 1602 |
+
"name": "hybrid-neighbor-s512-h32kv8-d128",
|
| 1603 |
+
"preset": "model",
|
| 1604 |
+
"vars": {
|
| 1605 |
+
"dtype": "float32",
|
| 1606 |
+
"batch": 1,
|
| 1607 |
+
"seq": 512,
|
| 1608 |
+
"heads": 32,
|
| 1609 |
+
"kvHeads": 8,
|
| 1610 |
+
"headDim": 128,
|
| 1611 |
+
"qkPairs": 4202496,
|
| 1612 |
+
"attendedKeys": 512,
|
| 1613 |
+
"dtypeBytes": 4
|
| 1614 |
+
},
|
| 1615 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1616 |
+
"inputs": {
|
| 1617 |
+
"queryT": { "shape": [1, 512, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1618 |
+
"keyT": { "shape": [1, 512, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1619 |
+
"valueT": { "shape": [1, 512, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1620 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1621 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1622 |
+
"blockRowIndicesT": {
|
| 1623 |
+
"shape": [4, 17],
|
| 1624 |
+
"dtype": "int32",
|
| 1625 |
+
"data": {
|
| 1626 |
+
"kind": "values",
|
| 1627 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1628 |
+
}
|
| 1629 |
+
},
|
| 1630 |
+
"blockColIndicesT": {
|
| 1631 |
+
"shape": [4, 112],
|
| 1632 |
+
"dtype": "int32",
|
| 1633 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1634 |
+
},
|
| 1635 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [512] } },
|
| 1636 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 512 }
|
| 1637 |
+
},
|
| 1638 |
+
"outputs": { "outputT": { "shape": [1, 512, 4096], "dtype": "float32" } },
|
| 1639 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1640 |
+
},
|
| 1641 |
+
{
|
| 1642 |
+
"name": "hybrid-neighbor-s514-h32kv8-d128",
|
| 1643 |
+
"preset": "model",
|
| 1644 |
+
"vars": {
|
| 1645 |
+
"dtype": "float32",
|
| 1646 |
+
"batch": 1,
|
| 1647 |
+
"seq": 514,
|
| 1648 |
+
"heads": 32,
|
| 1649 |
+
"kvHeads": 8,
|
| 1650 |
+
"headDim": 128,
|
| 1651 |
+
"qkPairs": 4232288,
|
| 1652 |
+
"attendedKeys": 514,
|
| 1653 |
+
"dtypeBytes": 4
|
| 1654 |
+
},
|
| 1655 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1656 |
+
"inputs": {
|
| 1657 |
+
"queryT": { "shape": [1, 514, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1658 |
+
"keyT": { "shape": [1, 514, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1659 |
+
"valueT": { "shape": [1, 514, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1660 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1661 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1662 |
+
"blockRowIndicesT": {
|
| 1663 |
+
"shape": [4, 17],
|
| 1664 |
+
"dtype": "int32",
|
| 1665 |
+
"data": {
|
| 1666 |
+
"kind": "values",
|
| 1667 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1668 |
+
}
|
| 1669 |
+
},
|
| 1670 |
+
"blockColIndicesT": {
|
| 1671 |
+
"shape": [4, 112],
|
| 1672 |
+
"dtype": "int32",
|
| 1673 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1674 |
+
},
|
| 1675 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [514] } },
|
| 1676 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 514 }
|
| 1677 |
+
},
|
| 1678 |
+
"outputs": { "outputT": { "shape": [1, 514, 4096], "dtype": "float32" } },
|
| 1679 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1680 |
+
},
|
| 1681 |
+
{
|
| 1682 |
+
"name": "hybrid-neighbor-s515-h32kv8-d128",
|
| 1683 |
+
"preset": "model",
|
| 1684 |
+
"vars": {
|
| 1685 |
+
"dtype": "float32",
|
| 1686 |
+
"batch": 1,
|
| 1687 |
+
"seq": 515,
|
| 1688 |
+
"heads": 32,
|
| 1689 |
+
"kvHeads": 8,
|
| 1690 |
+
"headDim": 128,
|
| 1691 |
+
"qkPairs": 4247232,
|
| 1692 |
+
"attendedKeys": 515,
|
| 1693 |
+
"dtypeBytes": 4
|
| 1694 |
+
},
|
| 1695 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1696 |
+
"inputs": {
|
| 1697 |
+
"queryT": { "shape": [1, 515, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1698 |
+
"keyT": { "shape": [1, 515, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1699 |
+
"valueT": { "shape": [1, 515, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1700 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1701 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1702 |
+
"blockRowIndicesT": {
|
| 1703 |
+
"shape": [4, 17],
|
| 1704 |
+
"dtype": "int32",
|
| 1705 |
+
"data": {
|
| 1706 |
+
"kind": "values",
|
| 1707 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1708 |
+
}
|
| 1709 |
+
},
|
| 1710 |
+
"blockColIndicesT": {
|
| 1711 |
+
"shape": [4, 112],
|
| 1712 |
+
"dtype": "int32",
|
| 1713 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1714 |
+
},
|
| 1715 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [515] } },
|
| 1716 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 515 }
|
| 1717 |
+
},
|
| 1718 |
+
"outputs": { "outputT": { "shape": [1, 515, 4096], "dtype": "float32" } },
|
| 1719 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1720 |
+
},
|
| 1721 |
+
{
|
| 1722 |
+
"name": "hybrid-neighbor-s516-h32kv8-d128",
|
| 1723 |
+
"preset": "model",
|
| 1724 |
+
"vars": {
|
| 1725 |
+
"dtype": "float32",
|
| 1726 |
+
"batch": 1,
|
| 1727 |
+
"seq": 516,
|
| 1728 |
+
"heads": 32,
|
| 1729 |
+
"kvHeads": 8,
|
| 1730 |
+
"headDim": 128,
|
| 1731 |
+
"qkPairs": 4262208,
|
| 1732 |
+
"attendedKeys": 516,
|
| 1733 |
+
"dtypeBytes": 4
|
| 1734 |
+
},
|
| 1735 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1736 |
+
"inputs": {
|
| 1737 |
+
"queryT": { "shape": [1, 516, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1738 |
+
"keyT": { "shape": [1, 516, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1739 |
+
"valueT": { "shape": [1, 516, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1740 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1741 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1742 |
+
"blockRowIndicesT": {
|
| 1743 |
+
"shape": [4, 17],
|
| 1744 |
+
"dtype": "int32",
|
| 1745 |
+
"data": {
|
| 1746 |
+
"kind": "values",
|
| 1747 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1748 |
+
}
|
| 1749 |
+
},
|
| 1750 |
+
"blockColIndicesT": {
|
| 1751 |
+
"shape": [4, 112],
|
| 1752 |
+
"dtype": "int32",
|
| 1753 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1754 |
+
},
|
| 1755 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [516] } },
|
| 1756 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 516 }
|
| 1757 |
+
},
|
| 1758 |
+
"outputs": { "outputT": { "shape": [1, 516, 4096], "dtype": "float32" } },
|
| 1759 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1760 |
+
},
|
| 1761 |
+
{
|
| 1762 |
+
"name": "hybrid-neighbor-s517-h32kv8-d128",
|
| 1763 |
+
"preset": "model",
|
| 1764 |
+
"vars": {
|
| 1765 |
+
"dtype": "float32",
|
| 1766 |
+
"batch": 1,
|
| 1767 |
+
"seq": 517,
|
| 1768 |
+
"heads": 32,
|
| 1769 |
+
"kvHeads": 8,
|
| 1770 |
+
"headDim": 128,
|
| 1771 |
+
"qkPairs": 4277216,
|
| 1772 |
+
"attendedKeys": 517,
|
| 1773 |
+
"dtypeBytes": 4
|
| 1774 |
+
},
|
| 1775 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1776 |
+
"inputs": {
|
| 1777 |
+
"queryT": { "shape": [1, 517, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1778 |
+
"keyT": { "shape": [1, 517, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1779 |
+
"valueT": { "shape": [1, 517, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1780 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1781 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1782 |
+
"blockRowIndicesT": {
|
| 1783 |
+
"shape": [4, 17],
|
| 1784 |
+
"dtype": "int32",
|
| 1785 |
+
"data": {
|
| 1786 |
+
"kind": "values",
|
| 1787 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1788 |
+
}
|
| 1789 |
+
},
|
| 1790 |
+
"blockColIndicesT": {
|
| 1791 |
+
"shape": [4, 112],
|
| 1792 |
+
"dtype": "int32",
|
| 1793 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1794 |
+
},
|
| 1795 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [517] } },
|
| 1796 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 517 }
|
| 1797 |
+
},
|
| 1798 |
+
"outputs": { "outputT": { "shape": [1, 517, 4096], "dtype": "float32" } },
|
| 1799 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1800 |
+
},
|
| 1801 |
+
{
|
| 1802 |
+
"name": "hybrid-neighbor-s577-h32kv8-d128",
|
| 1803 |
+
"preset": "model",
|
| 1804 |
+
"vars": {
|
| 1805 |
+
"dtype": "float32",
|
| 1806 |
+
"batch": 1,
|
| 1807 |
+
"seq": 577,
|
| 1808 |
+
"heads": 32,
|
| 1809 |
+
"kvHeads": 8,
|
| 1810 |
+
"headDim": 128,
|
| 1811 |
+
"qkPairs": 5234720,
|
| 1812 |
+
"attendedKeys": 577,
|
| 1813 |
+
"dtypeBytes": 4
|
| 1814 |
+
},
|
| 1815 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1816 |
+
"inputs": {
|
| 1817 |
+
"queryT": { "shape": [1, 577, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1818 |
+
"keyT": { "shape": [1, 577, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1819 |
+
"valueT": { "shape": [1, 577, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1820 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1821 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1822 |
+
"blockRowIndicesT": {
|
| 1823 |
+
"shape": [4, 17],
|
| 1824 |
+
"dtype": "int32",
|
| 1825 |
+
"data": {
|
| 1826 |
+
"kind": "values",
|
| 1827 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1828 |
+
}
|
| 1829 |
+
},
|
| 1830 |
+
"blockColIndicesT": {
|
| 1831 |
+
"shape": [4, 112],
|
| 1832 |
+
"dtype": "int32",
|
| 1833 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1834 |
+
},
|
| 1835 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [577] } },
|
| 1836 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 577 }
|
| 1837 |
+
},
|
| 1838 |
+
"outputs": { "outputT": { "shape": [1, 577, 4096], "dtype": "float32" } },
|
| 1839 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1840 |
+
},
|
| 1841 |
+
{
|
| 1842 |
+
"name": "isolate-packed-rotary-prompt-s65",
|
| 1843 |
+
"preset": "model",
|
| 1844 |
+
"vars": {
|
| 1845 |
+
"dtype": "float32",
|
| 1846 |
+
"batch": 1,
|
| 1847 |
+
"seq": 65,
|
| 1848 |
+
"heads": 32,
|
| 1849 |
+
"kvHeads": 8,
|
| 1850 |
+
"headDim": 128,
|
| 1851 |
+
"qkPairs": 68640,
|
| 1852 |
+
"attendedKeys": 65,
|
| 1853 |
+
"dtypeBytes": 4
|
| 1854 |
+
},
|
| 1855 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64, "do_rotary": 1 },
|
| 1856 |
+
"inputs": {
|
| 1857 |
+
"queryT": { "shape": [1, 65, 6144], "dtype": "float32", "dist": "normal", "seed": 9330, "scale": 1 },
|
| 1858 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9333, "scale": 1 },
|
| 1859 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9334, "scale": 1 },
|
| 1860 |
+
"blockRowIndicesT": {
|
| 1861 |
+
"shape": [4, 17],
|
| 1862 |
+
"dtype": "int32",
|
| 1863 |
+
"data": {
|
| 1864 |
+
"kind": "values",
|
| 1865 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1866 |
+
}
|
| 1867 |
+
},
|
| 1868 |
+
"blockColIndicesT": {
|
| 1869 |
+
"shape": [4, 112],
|
| 1870 |
+
"dtype": "int32",
|
| 1871 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1872 |
+
},
|
| 1873 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [65] } },
|
| 1874 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 65 },
|
| 1875 |
+
"cosCacheT": { "shape": [1024, 64], "dtype": "float32", "dist": "normal", "seed": 9335, "scale": 1 },
|
| 1876 |
+
"sinCacheT": { "shape": [1024, 64], "dtype": "float32", "dist": "normal", "seed": 9336, "scale": 1 }
|
| 1877 |
+
},
|
| 1878 |
+
"outputs": { "outputT": { "shape": [1, 65, 4096], "dtype": "float32" } },
|
| 1879 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1880 |
+
},
|
| 1881 |
+
{
|
| 1882 |
+
"name": "isolate-separate-history-s65",
|
| 1883 |
+
"preset": "model",
|
| 1884 |
+
"vars": {
|
| 1885 |
+
"dtype": "float32",
|
| 1886 |
+
"batch": 1,
|
| 1887 |
+
"seq": 65,
|
| 1888 |
+
"heads": 32,
|
| 1889 |
+
"kvHeads": 8,
|
| 1890 |
+
"headDim": 128,
|
| 1891 |
+
"qkPairs": 1264000,
|
| 1892 |
+
"attendedKeys": 1020,
|
| 1893 |
+
"dtypeBytes": 4
|
| 1894 |
+
},
|
| 1895 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1896 |
+
"inputs": {
|
| 1897 |
+
"queryT": { "shape": [1, 65, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1898 |
+
"keyT": { "shape": [1, 65, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1899 |
+
"valueT": { "shape": [1, 65, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1900 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1901 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1902 |
+
"blockRowIndicesT": {
|
| 1903 |
+
"shape": [4, 17],
|
| 1904 |
+
"dtype": "int32",
|
| 1905 |
+
"data": {
|
| 1906 |
+
"kind": "values",
|
| 1907 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1908 |
+
}
|
| 1909 |
+
},
|
| 1910 |
+
"blockColIndicesT": {
|
| 1911 |
+
"shape": [4, 112],
|
| 1912 |
+
"dtype": "int32",
|
| 1913 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1914 |
+
},
|
| 1915 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [1020] } },
|
| 1916 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 1020 }
|
| 1917 |
+
},
|
| 1918 |
+
"outputs": { "outputT": { "shape": [1, 65, 4096], "dtype": "float32" } },
|
| 1919 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1920 |
+
},
|
| 1921 |
+
{
|
| 1922 |
+
"name": "isolate-separate-history-s129",
|
| 1923 |
+
"preset": "model",
|
| 1924 |
+
"vars": {
|
| 1925 |
+
"dtype": "float32",
|
| 1926 |
+
"batch": 1,
|
| 1927 |
+
"seq": 129,
|
| 1928 |
+
"heads": 32,
|
| 1929 |
+
"kvHeads": 8,
|
| 1930 |
+
"headDim": 128,
|
| 1931 |
+
"qkPairs": 2474880,
|
| 1932 |
+
"attendedKeys": 1020,
|
| 1933 |
+
"dtypeBytes": 4
|
| 1934 |
+
},
|
| 1935 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1936 |
+
"inputs": {
|
| 1937 |
+
"queryT": { "shape": [1, 129, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1938 |
+
"keyT": { "shape": [1, 129, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1939 |
+
"valueT": { "shape": [1, 129, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1940 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1941 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1942 |
+
"blockRowIndicesT": {
|
| 1943 |
+
"shape": [4, 17],
|
| 1944 |
+
"dtype": "int32",
|
| 1945 |
+
"data": {
|
| 1946 |
+
"kind": "values",
|
| 1947 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1948 |
+
}
|
| 1949 |
+
},
|
| 1950 |
+
"blockColIndicesT": {
|
| 1951 |
+
"shape": [4, 112],
|
| 1952 |
+
"dtype": "int32",
|
| 1953 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1954 |
+
},
|
| 1955 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [1020] } },
|
| 1956 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 1020 }
|
| 1957 |
+
},
|
| 1958 |
+
"outputs": { "outputT": { "shape": [1, 129, 4096], "dtype": "float32" } },
|
| 1959 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 1960 |
+
},
|
| 1961 |
+
{
|
| 1962 |
+
"name": "isolate-separate-history-s513",
|
| 1963 |
+
"preset": "model",
|
| 1964 |
+
"vars": {
|
| 1965 |
+
"dtype": "float32",
|
| 1966 |
+
"batch": 1,
|
| 1967 |
+
"seq": 513,
|
| 1968 |
+
"heads": 32,
|
| 1969 |
+
"kvHeads": 8,
|
| 1970 |
+
"headDim": 128,
|
| 1971 |
+
"qkPairs": 9052032,
|
| 1972 |
+
"attendedKeys": 1020,
|
| 1973 |
+
"dtypeBytes": 4
|
| 1974 |
+
},
|
| 1975 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 1976 |
+
"inputs": {
|
| 1977 |
+
"queryT": { "shape": [1, 513, 4096], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 1978 |
+
"keyT": { "shape": [1, 513, 1024], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 1979 |
+
"valueT": { "shape": [1, 513, 1024], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 1980 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 1981 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 1982 |
+
"blockRowIndicesT": {
|
| 1983 |
+
"shape": [4, 17],
|
| 1984 |
+
"dtype": "int32",
|
| 1985 |
+
"data": {
|
| 1986 |
+
"kind": "values",
|
| 1987 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 1988 |
+
}
|
| 1989 |
+
},
|
| 1990 |
+
"blockColIndicesT": {
|
| 1991 |
+
"shape": [4, 112],
|
| 1992 |
+
"dtype": "int32",
|
| 1993 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 1994 |
+
},
|
| 1995 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [1020] } },
|
| 1996 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 1020 }
|
| 1997 |
+
},
|
| 1998 |
+
"outputs": { "outputT": { "shape": [1, 513, 4096], "dtype": "float32" } },
|
| 1999 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 2000 |
+
},
|
| 2001 |
+
{
|
| 2002 |
+
"name": "hybrid-packed-rotary-s65-past955",
|
| 2003 |
+
"preset": "model",
|
| 2004 |
+
"vars": {
|
| 2005 |
+
"dtype": "float32",
|
| 2006 |
+
"batch": 1,
|
| 2007 |
+
"seq": 65,
|
| 2008 |
+
"heads": 32,
|
| 2009 |
+
"kvHeads": 8,
|
| 2010 |
+
"headDim": 128,
|
| 2011 |
+
"qkPairs": 1264000,
|
| 2012 |
+
"attendedKeys": 1020,
|
| 2013 |
+
"dtypeBytes": 4
|
| 2014 |
+
},
|
| 2015 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64, "do_rotary": 1 },
|
| 2016 |
+
"inputs": {
|
| 2017 |
+
"queryT": { "shape": [1, 65, 6144], "dtype": "float32", "dist": "normal", "seed": 9330, "scale": 1 },
|
| 2018 |
+
"pastKeyT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9333, "scale": 1 },
|
| 2019 |
+
"pastValueT": { "shape": [1, 8, 1024, 128], "dtype": "float32", "dist": "normal", "seed": 9334, "scale": 1 },
|
| 2020 |
+
"blockRowIndicesT": {
|
| 2021 |
+
"shape": [4, 17],
|
| 2022 |
+
"dtype": "int32",
|
| 2023 |
+
"data": {
|
| 2024 |
+
"kind": "values",
|
| 2025 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 2026 |
+
}
|
| 2027 |
+
},
|
| 2028 |
+
"blockColIndicesT": {
|
| 2029 |
+
"shape": [4, 112],
|
| 2030 |
+
"dtype": "int32",
|
| 2031 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 2032 |
+
},
|
| 2033 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [1020] } },
|
| 2034 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 1020 },
|
| 2035 |
+
"cosCacheT": { "shape": [1024, 64], "dtype": "float32", "dist": "normal", "seed": 9335, "scale": 1 },
|
| 2036 |
+
"sinCacheT": { "shape": [1024, 64], "dtype": "float32", "dist": "normal", "seed": 9336, "scale": 1 }
|
| 2037 |
+
},
|
| 2038 |
+
"outputs": { "outputT": { "shape": [1, 65, 4096], "dtype": "float32" } },
|
| 2039 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 2040 |
+
},
|
| 2041 |
+
{
|
| 2042 |
+
"name": "hybrid-d96-s131",
|
| 2043 |
+
"preset": "model",
|
| 2044 |
+
"vars": {
|
| 2045 |
+
"dtype": "float32",
|
| 2046 |
+
"batch": 1,
|
| 2047 |
+
"seq": 131,
|
| 2048 |
+
"heads": 32,
|
| 2049 |
+
"kvHeads": 8,
|
| 2050 |
+
"headDim": 96,
|
| 2051 |
+
"qkPairs": 276672,
|
| 2052 |
+
"attendedKeys": 131,
|
| 2053 |
+
"dtypeBytes": 4
|
| 2054 |
+
},
|
| 2055 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 2056 |
+
"inputs": {
|
| 2057 |
+
"queryT": { "shape": [1, 131, 3072], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 2058 |
+
"keyT": { "shape": [1, 131, 768], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 2059 |
+
"valueT": { "shape": [1, 131, 768], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 2060 |
+
"pastKeyT": { "shape": [1, 8, 1024, 96], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 2061 |
+
"pastValueT": { "shape": [1, 8, 1024, 96], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 2062 |
+
"blockRowIndicesT": {
|
| 2063 |
+
"shape": [4, 17],
|
| 2064 |
+
"dtype": "int32",
|
| 2065 |
+
"data": {
|
| 2066 |
+
"kind": "values",
|
| 2067 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 2068 |
+
}
|
| 2069 |
+
},
|
| 2070 |
+
"blockColIndicesT": {
|
| 2071 |
+
"shape": [4, 112],
|
| 2072 |
+
"dtype": "int32",
|
| 2073 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 2074 |
+
},
|
| 2075 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [131] } },
|
| 2076 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 131 }
|
| 2077 |
+
},
|
| 2078 |
+
"outputs": { "outputT": { "shape": [1, 131, 3072], "dtype": "float32" } },
|
| 2079 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 2080 |
+
},
|
| 2081 |
+
{
|
| 2082 |
+
"name": "hybrid-d32-s131",
|
| 2083 |
+
"preset": "model",
|
| 2084 |
+
"vars": {
|
| 2085 |
+
"dtype": "float32",
|
| 2086 |
+
"batch": 1,
|
| 2087 |
+
"seq": 131,
|
| 2088 |
+
"heads": 32,
|
| 2089 |
+
"kvHeads": 8,
|
| 2090 |
+
"headDim": 32,
|
| 2091 |
+
"qkPairs": 276672,
|
| 2092 |
+
"attendedKeys": 131,
|
| 2093 |
+
"dtypeBytes": 4
|
| 2094 |
+
},
|
| 2095 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 2096 |
+
"inputs": {
|
| 2097 |
+
"queryT": { "shape": [1, 131, 1024], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 2098 |
+
"keyT": { "shape": [1, 131, 256], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 2099 |
+
"valueT": { "shape": [1, 131, 256], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 2100 |
+
"pastKeyT": { "shape": [1, 8, 1024, 32], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 2101 |
+
"pastValueT": { "shape": [1, 8, 1024, 32], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 2102 |
+
"blockRowIndicesT": {
|
| 2103 |
+
"shape": [4, 17],
|
| 2104 |
+
"dtype": "int32",
|
| 2105 |
+
"data": {
|
| 2106 |
+
"kind": "values",
|
| 2107 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 2108 |
+
}
|
| 2109 |
+
},
|
| 2110 |
+
"blockColIndicesT": {
|
| 2111 |
+
"shape": [4, 112],
|
| 2112 |
+
"dtype": "int32",
|
| 2113 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 2114 |
+
},
|
| 2115 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [131] } },
|
| 2116 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 131 }
|
| 2117 |
+
},
|
| 2118 |
+
"outputs": { "outputT": { "shape": [1, 131, 1024], "dtype": "float32" } },
|
| 2119 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 2120 |
+
},
|
| 2121 |
+
{
|
| 2122 |
+
"name": "hybrid-d64-s131",
|
| 2123 |
+
"preset": "model",
|
| 2124 |
+
"vars": {
|
| 2125 |
+
"dtype": "float32",
|
| 2126 |
+
"batch": 1,
|
| 2127 |
+
"seq": 131,
|
| 2128 |
+
"heads": 32,
|
| 2129 |
+
"kvHeads": 8,
|
| 2130 |
+
"headDim": 64,
|
| 2131 |
+
"qkPairs": 276672,
|
| 2132 |
+
"attendedKeys": 131,
|
| 2133 |
+
"dtypeBytes": 4
|
| 2134 |
+
},
|
| 2135 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 2136 |
+
"inputs": {
|
| 2137 |
+
"queryT": { "shape": [1, 131, 2048], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 2138 |
+
"keyT": { "shape": [1, 131, 512], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 2139 |
+
"valueT": { "shape": [1, 131, 512], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 2140 |
+
"pastKeyT": { "shape": [1, 8, 1024, 64], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 2141 |
+
"pastValueT": { "shape": [1, 8, 1024, 64], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 2142 |
+
"blockRowIndicesT": {
|
| 2143 |
+
"shape": [4, 17],
|
| 2144 |
+
"dtype": "int32",
|
| 2145 |
+
"data": {
|
| 2146 |
+
"kind": "values",
|
| 2147 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 2148 |
+
}
|
| 2149 |
+
},
|
| 2150 |
+
"blockColIndicesT": {
|
| 2151 |
+
"shape": [4, 112],
|
| 2152 |
+
"dtype": "int32",
|
| 2153 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 2154 |
+
},
|
| 2155 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [131] } },
|
| 2156 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 131 }
|
| 2157 |
+
},
|
| 2158 |
+
"outputs": { "outputT": { "shape": [1, 131, 2048], "dtype": "float32" } },
|
| 2159 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] }
|
| 2160 |
+
},
|
| 2161 |
+
{
|
| 2162 |
+
"name": "hybrid-d32-s131-narrow",
|
| 2163 |
+
"preset": "model",
|
| 2164 |
+
"vars": {
|
| 2165 |
+
"dtype": "float32",
|
| 2166 |
+
"batch": 1,
|
| 2167 |
+
"seq": 131,
|
| 2168 |
+
"heads": 32,
|
| 2169 |
+
"kvHeads": 8,
|
| 2170 |
+
"headDim": 32,
|
| 2171 |
+
"qkPairs": 276672,
|
| 2172 |
+
"attendedKeys": 131,
|
| 2173 |
+
"dtypeBytes": 4
|
| 2174 |
+
},
|
| 2175 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "sparse_block_size": 64 },
|
| 2176 |
+
"inputs": {
|
| 2177 |
+
"queryT": { "shape": [1, 131, 1024], "dtype": "float32", "dist": "normal", "seed": 9300, "scale": 1 },
|
| 2178 |
+
"keyT": { "shape": [1, 131, 256], "dtype": "float32", "dist": "normal", "seed": 9301, "scale": 1 },
|
| 2179 |
+
"valueT": { "shape": [1, 131, 256], "dtype": "float32", "dist": "normal", "seed": 9302, "scale": 1 },
|
| 2180 |
+
"pastKeyT": { "shape": [1, 8, 1024, 32], "dtype": "float32", "dist": "normal", "seed": 9303, "scale": 1 },
|
| 2181 |
+
"pastValueT": { "shape": [1, 8, 1024, 32], "dtype": "float32", "dist": "normal", "seed": 9304, "scale": 1 },
|
| 2182 |
+
"blockRowIndicesT": {
|
| 2183 |
+
"shape": [4, 17],
|
| 2184 |
+
"dtype": "int32",
|
| 2185 |
+
"data": {
|
| 2186 |
+
"kind": "values",
|
| 2187 |
+
"values": { "$ref": "#/fixtureArrays/sparse-prompt-b1-s1024-h32kv8-d128-blk64_input_blockRowIndicesT" }
|
| 2188 |
+
}
|
| 2189 |
+
},
|
| 2190 |
+
"blockColIndicesT": {
|
| 2191 |
+
"shape": [4, 112],
|
| 2192 |
+
"dtype": "int32",
|
| 2193 |
+
"data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/block_column_indices_t_pattern" } }
|
| 2194 |
+
},
|
| 2195 |
+
"totalSequenceLengthT": { "shape": [1], "dtype": "int32", "data": { "kind": "values", "values": [131] } },
|
| 2196 |
+
"keyTotalSequenceLengthsT": { "shape": [1], "dtype": "int32", "dist": "constant", "value": 131 }
|
| 2197 |
+
},
|
| 2198 |
+
"outputs": { "outputT": { "shape": [1, 131, 1024], "dtype": "float32" } },
|
| 2199 |
+
"bench": { "metrics": [{ "type": "gflops", "value": "4 * args.qkPairs * args.headDim" }] },
|
| 2200 |
+
"tunables": { "NARROW_MIN_WORKGROUPS": 1 }
|
| 2201 |
}
|
| 2202 |
]
|
| 2203 |
}
|
build/webgpu/manifest.json
CHANGED
|
@@ -54,9 +54,10 @@
|
|
| 54 |
"sparseBlockSize": "attrs.sparse_block_size",
|
| 55 |
"headSize": "dim(shapes.pastKeyT, 3)",
|
| 56 |
"headVec": "headSize / 4",
|
| 57 |
-
"
|
|
|
|
| 58 |
"sparseQueryTileCap": "max(1, floor((device.limits.maxComputeWorkgroupStorageSize / 4 - sparseWidthBound) / (2 * headSize + 3 * sparseWidthBound)))",
|
| 59 |
-
"sparseQueryTileWant": "min(
|
| 60 |
"sparseQueryTile": "1 if seqLen <= 1 else (16 if sparseQueryTileWant >= 16 and seqLen >= 16 else (8 if sparseQueryTileWant >= 8 and seqLen >= 8 else (4 if sparseQueryTileWant >= 4 and seqLen >= 4 else (2 if sparseQueryTileWant >= 2 and seqLen >= 2 else 1))))",
|
| 61 |
"sparseQueryTiles": "ceilDiv(seqLen, sparseQueryTile)",
|
| 62 |
"sparseAttnWorkgroups": "sparseQueryTiles * batchSize * numHeads",
|
|
@@ -82,93 +83,98 @@
|
|
| 82 |
"rotaryPairOk": "not doRotary or (present.cosCacheT and present.sinCacheT and ranks.cosCacheT == 2 and ranks.sinCacheT == 2 and rotaryHalf % 8 == 0 and rotaryDim <= headSize and sameShape(shapes.sinCacheT, shapes.cosCacheT) and tensorDtypes.cosCacheT == tensorDtypes.queryT and tensorDtypes.sinCacheT == tensorDtypes.queryT)",
|
| 83 |
"blockIndexShapeOk": "ranks.blockRowIndicesT == 2 and ranks.blockColIndicesT == 2 and dim(shapes.blockColIndicesT, 0) == numLayout and maxBlocks >= 1 and maxNnz >= 0 and maxNnz <= maxBlocks * maxBlocks and tensorDtypes.blockRowIndicesT == \"int32\" and tensorDtypes.blockColIndicesT == \"int32\"",
|
| 84 |
"scheduleShapeOk": "(ranks.totalSequenceLengthT == 0 or ranks.totalSequenceLengthT == 1) and numel(shapes.totalSequenceLengthT) == 1 and ranks.keyTotalSequenceLengthsT == 1 and dim(shapes.keyTotalSequenceLengthsT, 0) == batchSize and tensorDtypes.totalSequenceLengthT == \"int32\" and tensorDtypes.keyTotalSequenceLengthsT == \"int32\"",
|
| 85 |
-
"
|
|
|
|
| 86 |
"contract": "ranks.queryT == 3 and ranks.outputT == 3 and (tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(dtypes.T) and tensorDtypes.pastKeyT == tensorDtypes.queryT and tensorDtypes.pastValueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT and numHeads >= 1 and kvNumHeads >= 1 and numHeads % kvNumHeads == 0 and headSize >= 8 and headSize % 8 == 0 and (not doRotary or headSize % 16 == 0) and numLayout >= 1 and numHeads % numLayout == 0 and (sparseBlockSize == 16 or sparseBlockSize == 32 or sparseBlockSize == 64 or sparseBlockSize == 128) and cacheShapeOk and queryShapeOk and kvShapeOk and kvPairOk and rotaryPairOk and blockIndexShapeOk and scheduleShapeOk and dim(shapes.outputT, 0) == batchSize and dim(shapes.outputT, 1) == seqLen and dim(shapes.outputT, 2) == qHidden",
|
| 87 |
"packedContract": "contract and packedQkv and not useRotary",
|
| 88 |
"packedRotaryContract": "contract and packedQkv and useRotary",
|
| 89 |
"separateContract": "contract and not packedQkv and not useRotary",
|
| 90 |
"separateRotaryContract": "contract and not packedQkv and useRotary",
|
| 91 |
-
"
|
| 92 |
-
"
|
| 93 |
-
"
|
| 94 |
-
"
|
|
|
|
|
|
|
|
|
|
| 95 |
"sparseSgmatTileK": "sparseSgmatTileN / 2",
|
| 96 |
-
"sparseSgmatLdsBytes": "(64 * sparseSgmatTileK + 64 * sparseSgmatTileN + 64 *
|
| 97 |
-
"sparseSgmatGeometryOk": "
|
| 98 |
"sparseSgmatOk": "tensorDtypes.queryT == \"float32\" and seqLen >= 64 and sparseBlockSize % 64 == 0 and headSize % 32 == 0 and headSize <= 128 and maxCacheSeq % 64 == 0 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and sparseSgmatGeometryOk",
|
| 99 |
"scalar": "dtypes.T",
|
| 100 |
"cacheVec": "\"vec4<f16>\" if dtypes.T == \"f16\" else \"vec4<f32>\"",
|
| 101 |
"attnWorkgroup": "sparseAttnWorkgroup",
|
| 102 |
"usesRotary": "useRotary",
|
| 103 |
-
"appendWorkgroupSize": "tunables.APPEND_WORKGROUP_SIZE"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 104 |
},
|
| 105 |
"when": ["geometryOk"],
|
| 106 |
"bindings": {
|
| 107 |
-
"new_key": { "arg": "keyT", "
|
| 108 |
-
"new_value": { "arg": "valueT", "
|
| 109 |
-
"present_key": { "arg": "pastKeyT", "
|
| 110 |
-
"present_value": { "arg": "pastValueT", "
|
| 111 |
-
"key_total_sequence_lengths": {
|
| 112 |
-
|
| 113 |
-
"buffer": "read-only-storage",
|
| 114 |
-
"elementType": "i32"
|
| 115 |
-
},
|
| 116 |
-
"total_sequence_length": { "arg": "totalSequenceLengthT", "buffer": "read-only-storage", "elementType": "i32" },
|
| 117 |
"params": {
|
| 118 |
-
"buffer": "uniform",
|
| 119 |
"struct": [
|
| 120 |
{ "name": "batchSize", "type": "u32", "value": "batchSize" },
|
| 121 |
{ "name": "seqLen", "type": "u32", "value": "seqLen" }
|
| 122 |
]
|
| 123 |
},
|
| 124 |
-
"cos_cache": { "arg": "cosCacheT", "
|
| 125 |
-
"sin_cache": { "arg": "sinCacheT", "
|
| 126 |
-
"packed_qkv": { "arg": "queryT", "
|
| 127 |
-
"query": { "arg": "queryT", "
|
| 128 |
-
"
|
| 129 |
"arg": "pastKeyT",
|
| 130 |
"name": "present_key",
|
| 131 |
"buffer": "read-only-storage",
|
| 132 |
"elementType": "$cacheVec"
|
| 133 |
},
|
| 134 |
-
"
|
| 135 |
"arg": "pastValueT",
|
| 136 |
"name": "present_value",
|
| 137 |
"buffer": "read-only-storage",
|
| 138 |
"elementType": "$cacheVec"
|
| 139 |
},
|
| 140 |
-
"block_row_indices": { "arg": "blockRowIndicesT", "
|
| 141 |
-
"block_col_indices": { "arg": "blockColIndicesT", "
|
| 142 |
-
"output": { "arg": "outputT", "
|
| 143 |
-
"
|
| 144 |
"name": "params",
|
| 145 |
-
"buffer": "uniform",
|
| 146 |
"struct": [
|
| 147 |
{ "name": "seqLen", "type": "u32", "value": "seqLen" },
|
| 148 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 149 |
]
|
| 150 |
},
|
| 151 |
"q_rotary": { "scratch": "QRotary", "buffer": "read-only-storage", "elementType": "f32" },
|
| 152 |
-
"
|
| 153 |
"arg": "pastKeyT",
|
| 154 |
"name": "present_key",
|
| 155 |
"buffer": "read-only-storage",
|
| 156 |
"elementType": "$scalar"
|
| 157 |
},
|
| 158 |
-
"
|
| 159 |
"arg": "pastValueT",
|
| 160 |
"name": "present_value",
|
| 161 |
"buffer": "read-only-storage",
|
| 162 |
"elementType": "$scalar"
|
| 163 |
},
|
| 164 |
-
"
|
| 165 |
},
|
| 166 |
"variants": [
|
| 167 |
{
|
| 168 |
"id": "separate",
|
| 169 |
"priority": 0,
|
| 170 |
"when": ["separateContract"],
|
| 171 |
-
"derive": { "vStageWorthIt": "sparseVStageWorthIt" },
|
| 172 |
"passes": [
|
| 173 |
{
|
| 174 |
"id": "append",
|
|
@@ -186,8 +192,8 @@
|
|
| 186 |
"id": "attention",
|
| 187 |
"name": "SparseAttention.Attention",
|
| 188 |
"shader": "sparse-attention.wgsl.jinja",
|
| 189 |
-
"derive": { "qTile": "sparseQueryTile" },
|
| 190 |
-
"bindings": ["query", "
|
| 191 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 192 |
}
|
| 193 |
]
|
|
@@ -217,18 +223,65 @@
|
|
| 217 |
"id": "attention",
|
| 218 |
"name": "SparseAttention.AttentionSgmat",
|
| 219 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 220 |
-
"bindings": ["query", "
|
| 221 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 222 |
-
"
|
| 223 |
}
|
| 224 |
],
|
| 225 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
| 226 |
},
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 227 |
{
|
| 228 |
"id": "separate_rotary",
|
| 229 |
"priority": 10,
|
| 230 |
"when": ["separateRotaryContract"],
|
| 231 |
-
"derive": { "vStageWorthIt": "sparseVStageWorthIt" },
|
| 232 |
"intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }],
|
| 233 |
"passes": [
|
| 234 |
{
|
|
@@ -247,7 +300,7 @@
|
|
| 247 |
"id": "qrotary",
|
| 248 |
"name": "SparseAttention.QueryRotary",
|
| 249 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 250 |
-
"bindings": ["query", "cos_cache", "sin_cache", "
|
| 251 |
"dispatch": {
|
| 252 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 253 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
@@ -258,8 +311,8 @@
|
|
| 258 |
"id": "attention",
|
| 259 |
"name": "SparseAttention.Attention",
|
| 260 |
"shader": "sparse-attention.wgsl.jinja",
|
| 261 |
-
"derive": { "qTile": "sparseQueryTile" },
|
| 262 |
-
"bindings": ["q_rotary", "
|
| 263 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 264 |
}
|
| 265 |
]
|
|
@@ -290,7 +343,7 @@
|
|
| 290 |
"id": "qrotary",
|
| 291 |
"name": "SparseAttention.QueryRotary",
|
| 292 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 293 |
-
"bindings": ["query", "cos_cache", "sin_cache", "
|
| 294 |
"dispatch": {
|
| 295 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 296 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
@@ -301,18 +354,77 @@
|
|
| 301 |
"id": "attention",
|
| 302 |
"name": "SparseAttention.AttentionSgmat",
|
| 303 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 304 |
-
"bindings": ["q_rotary", "
|
| 305 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 306 |
-
"
|
| 307 |
}
|
| 308 |
],
|
| 309 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
| 310 |
},
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 311 |
{
|
| 312 |
"id": "packed",
|
| 313 |
"priority": 0,
|
| 314 |
"when": ["packedContract"],
|
| 315 |
-
"derive": { "vStageWorthIt": "sparseVStageWorthIt" },
|
| 316 |
"passes": [
|
| 317 |
{
|
| 318 |
"id": "append",
|
|
@@ -330,8 +442,8 @@
|
|
| 330 |
"id": "attention",
|
| 331 |
"name": "SparseAttention.Attention",
|
| 332 |
"shader": "sparse-attention.wgsl.jinja",
|
| 333 |
-
"derive": { "qTile": "sparseQueryTile" },
|
| 334 |
-
"bindings": ["query", "
|
| 335 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 336 |
}
|
| 337 |
]
|
|
@@ -361,18 +473,65 @@
|
|
| 361 |
"id": "attention",
|
| 362 |
"name": "SparseAttention.AttentionSgmat",
|
| 363 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 364 |
-
"bindings": ["query", "
|
| 365 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 366 |
-
"
|
| 367 |
}
|
| 368 |
],
|
| 369 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
| 370 |
},
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 371 |
{
|
| 372 |
"id": "packed_rotary",
|
| 373 |
"priority": 10,
|
| 374 |
"when": ["packedRotaryContract"],
|
| 375 |
-
"derive": { "vStageWorthIt": "sparseVStageWorthIt" },
|
| 376 |
"intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }],
|
| 377 |
"passes": [
|
| 378 |
{
|
|
@@ -391,7 +550,7 @@
|
|
| 391 |
"id": "qrotary",
|
| 392 |
"name": "SparseAttention.QueryRotary",
|
| 393 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 394 |
-
"bindings": ["query", "cos_cache", "sin_cache", "
|
| 395 |
"dispatch": {
|
| 396 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 397 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
@@ -402,8 +561,8 @@
|
|
| 402 |
"id": "attention",
|
| 403 |
"name": "SparseAttention.Attention",
|
| 404 |
"shader": "sparse-attention.wgsl.jinja",
|
| 405 |
-
"derive": { "qTile": "sparseQueryTile" },
|
| 406 |
-
"bindings": ["q_rotary", "
|
| 407 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 408 |
}
|
| 409 |
]
|
|
@@ -434,7 +593,7 @@
|
|
| 434 |
"id": "qrotary",
|
| 435 |
"name": "SparseAttention.QueryRotary",
|
| 436 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 437 |
-
"bindings": ["query", "cos_cache", "sin_cache", "
|
| 438 |
"dispatch": {
|
| 439 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 440 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
@@ -445,12 +604,71 @@
|
|
| 445 |
"id": "attention",
|
| 446 |
"name": "SparseAttention.AttentionSgmat",
|
| 447 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 448 |
-
"bindings": ["q_rotary", "
|
| 449 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 450 |
-
"
|
| 451 |
}
|
| 452 |
],
|
| 453 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 454 |
}
|
| 455 |
]
|
| 456 |
}
|
|
|
|
| 54 |
"sparseBlockSize": "attrs.sparse_block_size",
|
| 55 |
"headSize": "dim(shapes.pastKeyT, 3)",
|
| 56 |
"headVec": "headSize / 4",
|
| 57 |
+
"sparseRequestedQueryTile": "max(tunables.QUERY_TILE, 8) if device.features.has(\"subgroups\") and wave32Effective and not device.features.has(\"shader-f16\") and ceilDiv(seqLen, 8) * batchSize * numHeads >= tunables.NARROW_MIN_WORKGROUPS else tunables.QUERY_TILE",
|
| 58 |
+
"sparseWidthBound": "min(256, max(32, pow2ceil(headVec))) if ceilDiv(seqLen, max(1, min(sparseBlockSize, sparseRequestedQueryTile))) * batchSize * numHeads >= tunables.NARROW_MIN_WORKGROUPS else max(256, tunables.WORKGROUP_SIZE)",
|
| 59 |
"sparseQueryTileCap": "max(1, floor((device.limits.maxComputeWorkgroupStorageSize / 4 - sparseWidthBound) / (2 * headSize + 3 * sparseWidthBound)))",
|
| 60 |
+
"sparseQueryTileWant": "min(sparseRequestedQueryTile, min(sparseBlockSize, sparseQueryTileCap))",
|
| 61 |
"sparseQueryTile": "1 if seqLen <= 1 else (16 if sparseQueryTileWant >= 16 and seqLen >= 16 else (8 if sparseQueryTileWant >= 8 and seqLen >= 8 else (4 if sparseQueryTileWant >= 4 and seqLen >= 4 else (2 if sparseQueryTileWant >= 2 and seqLen >= 2 else 1))))",
|
| 62 |
"sparseQueryTiles": "ceilDiv(seqLen, sparseQueryTile)",
|
| 63 |
"sparseAttnWorkgroups": "sparseQueryTiles * batchSize * numHeads",
|
|
|
|
| 83 |
"rotaryPairOk": "not doRotary or (present.cosCacheT and present.sinCacheT and ranks.cosCacheT == 2 and ranks.sinCacheT == 2 and rotaryHalf % 8 == 0 and rotaryDim <= headSize and sameShape(shapes.sinCacheT, shapes.cosCacheT) and tensorDtypes.cosCacheT == tensorDtypes.queryT and tensorDtypes.sinCacheT == tensorDtypes.queryT)",
|
| 84 |
"blockIndexShapeOk": "ranks.blockRowIndicesT == 2 and ranks.blockColIndicesT == 2 and dim(shapes.blockColIndicesT, 0) == numLayout and maxBlocks >= 1 and maxNnz >= 0 and maxNnz <= maxBlocks * maxBlocks and tensorDtypes.blockRowIndicesT == \"int32\" and tensorDtypes.blockColIndicesT == \"int32\"",
|
| 85 |
"scheduleShapeOk": "(ranks.totalSequenceLengthT == 0 or ranks.totalSequenceLengthT == 1) and numel(shapes.totalSequenceLengthT) == 1 and ranks.keyTotalSequenceLengthsT == 1 and dim(shapes.keyTotalSequenceLengthsT, 0) == batchSize and tensorDtypes.totalSequenceLengthT == \"int32\" and tensorDtypes.keyTotalSequenceLengthsT == \"int32\"",
|
| 86 |
+
"sparseAttnBaseLdsBytes": "(2 * sparseQueryTile * headSize + (3 * sparseQueryTile + 1) * sparseAttnWorkgroup) * 4",
|
| 87 |
+
"geometryOk": "tunables.WORKGROUP_SIZE >= 1 and floor(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and pow2ceil(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and tunables.WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX and tunables.APPEND_WORKGROUP_SIZE >= 1 and floor(tunables.APPEND_WORKGROUP_SIZE) == tunables.APPEND_WORKGROUP_SIZE and tunables.APPEND_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.APPEND_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX and sparseQueryTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and batchSize * numHeads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(ceilDiv(qRotaryElements, tunables.APPEND_WORKGROUP_SIZE), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and sparseAttnBaseLdsBytes <= device.limits.maxComputeWorkgroupStorageSize and sparseAttnWorkgroup <= device.limits.maxComputeInvocationsPerWorkgroup and sparseAttnWorkgroup <= device.limits.maxComputeWorkgroupSizeX",
|
| 88 |
"contract": "ranks.queryT == 3 and ranks.outputT == 3 and (tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(dtypes.T) and tensorDtypes.pastKeyT == tensorDtypes.queryT and tensorDtypes.pastValueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT and numHeads >= 1 and kvNumHeads >= 1 and numHeads % kvNumHeads == 0 and headSize >= 8 and headSize % 8 == 0 and (not doRotary or headSize % 16 == 0) and numLayout >= 1 and numHeads % numLayout == 0 and (sparseBlockSize == 16 or sparseBlockSize == 32 or sparseBlockSize == 64 or sparseBlockSize == 128) and cacheShapeOk and queryShapeOk and kvShapeOk and kvPairOk and rotaryPairOk and blockIndexShapeOk and scheduleShapeOk and dim(shapes.outputT, 0) == batchSize and dim(shapes.outputT, 1) == seqLen and dim(shapes.outputT, 2) == qHidden",
|
| 89 |
"packedContract": "contract and packedQkv and not useRotary",
|
| 90 |
"packedRotaryContract": "contract and packedQkv and useRotary",
|
| 91 |
"separateContract": "contract and not packedQkv and not useRotary",
|
| 92 |
"separateRotaryContract": "contract and not packedQkv and useRotary",
|
| 93 |
+
"sparseValueParts": "max(1, min(floor(sparseAttnWorkgroup / headVec), floor((device.limits.maxComputeWorkgroupStorageSize - sparseAttnBaseLdsBytes) / (sparseQueryTile * headSize * 4))))",
|
| 94 |
+
"sparseVStageWorthIt": "sparseAttnWorkgroups <= tunables.V_STAGE_MAX_WORKGROUPS and headVec <= sparseAttnWorkgroup and sparseAttnBaseLdsBytes + 16 * headSize * 4 <= device.limits.maxComputeWorkgroupStorageSize",
|
| 95 |
+
"sparseSgmatWorkgroup": "256",
|
| 96 |
+
"sparseSgmatTileM": "64",
|
| 97 |
+
"sgmatQueryTiles": "ceilDiv(seqLen, sparseSgmatTileM)",
|
| 98 |
+
"sgmatDirectQuery": "seqLen % sparseSgmatTileM == 0",
|
| 99 |
+
"sparseSgmatTileN": "64 if (64 * 32 + 64 * 64 + 64 * 3 + 128 * 2) * 4 <= device.limits.maxComputeWorkgroupStorageSize else 32",
|
| 100 |
"sparseSgmatTileK": "sparseSgmatTileN / 2",
|
| 101 |
+
"sparseSgmatLdsBytes": "(64 * sparseSgmatTileK + 64 * sparseSgmatTileN + 64 * 3 + 128 * 2) * 4",
|
| 102 |
+
"sparseSgmatGeometryOk": "sparseSgmatWorkgroup <= device.limits.maxComputeInvocationsPerWorkgroup and sparseSgmatWorkgroup <= device.limits.maxComputeWorkgroupSizeX and sgmatQueryTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and batchSize * numHeads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and sparseSgmatLdsBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 103 |
"sparseSgmatOk": "tensorDtypes.queryT == \"float32\" and seqLen >= 64 and sparseBlockSize % 64 == 0 and headSize % 32 == 0 and headSize <= 128 and maxCacheSeq % 64 == 0 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and sparseSgmatGeometryOk",
|
| 104 |
"scalar": "dtypes.T",
|
| 105 |
"cacheVec": "\"vec4<f16>\" if dtypes.T == \"f16\" else \"vec4<f32>\"",
|
| 106 |
"attnWorkgroup": "sparseAttnWorkgroup",
|
| 107 |
"usesRotary": "useRotary",
|
| 108 |
+
"appendWorkgroupSize": "tunables.APPEND_WORKGROUP_SIZE",
|
| 109 |
+
"sparseTailRows": "seqLen % sparseSgmatTileM",
|
| 110 |
+
"sparsePrefixRows": "seqLen - sparseTailRows",
|
| 111 |
+
"sparseTailWorkgroup": "min(256, max(32, pow2ceil(headVec))) if batchSize * numHeads >= tunables.NARROW_MIN_WORKGROUPS else tunables.WORKGROUP_SIZE",
|
| 112 |
+
"sparseTailBaseLdsBytes": "(2 * sparseTailRows * headSize + (3 * sparseTailRows + 1) * sparseTailWorkgroup) * 4",
|
| 113 |
+
"sparseTailValueParts": "max(1, min(floor(sparseTailWorkgroup / headVec), floor((device.limits.maxComputeWorkgroupStorageSize - sparseTailBaseLdsBytes) / (max(1, sparseTailRows) * headSize * 4))))",
|
| 114 |
+
"sparseTailVStageWorthIt": "batchSize * numHeads <= tunables.V_STAGE_MAX_WORKGROUPS and headVec <= sparseTailWorkgroup and sparseTailBaseLdsBytes + 16 * headSize * 4 <= device.limits.maxComputeWorkgroupStorageSize",
|
| 115 |
+
"sparseTailOk": "sparseTailRows > 0 and sparseTailRows <= sparseQueryTile and sparsePrefixRows >= sparseSgmatTileM and sparseTailBaseLdsBytes <= device.limits.maxComputeWorkgroupStorageSize and sparseTailWorkgroup <= device.limits.maxComputeInvocationsPerWorkgroup and sparseTailWorkgroup <= device.limits.maxComputeWorkgroupSizeX"
|
| 116 |
},
|
| 117 |
"when": ["geometryOk"],
|
| 118 |
"bindings": {
|
| 119 |
+
"new_key": { "arg": "keyT", "elementType": "$scalar" },
|
| 120 |
+
"new_value": { "arg": "valueT", "elementType": "$scalar" },
|
| 121 |
+
"present_key": { "arg": "pastKeyT", "elementType": "$scalar" },
|
| 122 |
+
"present_value": { "arg": "pastValueT", "elementType": "$scalar" },
|
| 123 |
+
"key_total_sequence_lengths": { "arg": "keyTotalSequenceLengthsT", "elementType": "i32" },
|
| 124 |
+
"total_sequence_length": { "arg": "totalSequenceLengthT", "elementType": "i32" },
|
|
|
|
|
|
|
|
|
|
|
|
|
| 125 |
"params": {
|
|
|
|
| 126 |
"struct": [
|
| 127 |
{ "name": "batchSize", "type": "u32", "value": "batchSize" },
|
| 128 |
{ "name": "seqLen", "type": "u32", "value": "seqLen" }
|
| 129 |
]
|
| 130 |
},
|
| 131 |
+
"cos_cache": { "arg": "cosCacheT", "elementType": "$scalar" },
|
| 132 |
+
"sin_cache": { "arg": "sinCacheT", "elementType": "$scalar" },
|
| 133 |
+
"packed_qkv": { "arg": "queryT", "elementType": "$scalar" },
|
| 134 |
+
"query": { "arg": "queryT", "elementType": "$scalar" },
|
| 135 |
+
"present_key_packed": {
|
| 136 |
"arg": "pastKeyT",
|
| 137 |
"name": "present_key",
|
| 138 |
"buffer": "read-only-storage",
|
| 139 |
"elementType": "$cacheVec"
|
| 140 |
},
|
| 141 |
+
"present_value_packed": {
|
| 142 |
"arg": "pastValueT",
|
| 143 |
"name": "present_value",
|
| 144 |
"buffer": "read-only-storage",
|
| 145 |
"elementType": "$cacheVec"
|
| 146 |
},
|
| 147 |
+
"block_row_indices": { "arg": "blockRowIndicesT", "elementType": "i32" },
|
| 148 |
+
"block_col_indices": { "arg": "blockColIndicesT", "elementType": "i32" },
|
| 149 |
+
"output": { "arg": "outputT", "elementType": "$scalar" },
|
| 150 |
+
"params_packed": {
|
| 151 |
"name": "params",
|
|
|
|
| 152 |
"struct": [
|
| 153 |
{ "name": "seqLen", "type": "u32", "value": "seqLen" },
|
| 154 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 155 |
]
|
| 156 |
},
|
| 157 |
"q_rotary": { "scratch": "QRotary", "buffer": "read-only-storage", "elementType": "f32" },
|
| 158 |
+
"present_key_scalar": {
|
| 159 |
"arg": "pastKeyT",
|
| 160 |
"name": "present_key",
|
| 161 |
"buffer": "read-only-storage",
|
| 162 |
"elementType": "$scalar"
|
| 163 |
},
|
| 164 |
+
"present_value_scalar": {
|
| 165 |
"arg": "pastValueT",
|
| 166 |
"name": "present_value",
|
| 167 |
"buffer": "read-only-storage",
|
| 168 |
"elementType": "$scalar"
|
| 169 |
},
|
| 170 |
+
"q_rotary_f32": { "scratch": "QRotary", "name": "q_rotary", "elementType": "f32" }
|
| 171 |
},
|
| 172 |
"variants": [
|
| 173 |
{
|
| 174 |
"id": "separate",
|
| 175 |
"priority": 0,
|
| 176 |
"when": ["separateContract"],
|
| 177 |
+
"derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" },
|
| 178 |
"passes": [
|
| 179 |
{
|
| 180 |
"id": "append",
|
|
|
|
| 192 |
"id": "attention",
|
| 193 |
"name": "SparseAttention.Attention",
|
| 194 |
"shader": "sparse-attention.wgsl.jinja",
|
| 195 |
+
"derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" },
|
| 196 |
+
"bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 197 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 198 |
}
|
| 199 |
]
|
|
|
|
| 223 |
"id": "attention",
|
| 224 |
"name": "SparseAttention.AttentionSgmat",
|
| 225 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 226 |
+
"bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 227 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 228 |
+
"derive": { "splitQueryTail": "false" }
|
| 229 |
}
|
| 230 |
],
|
| 231 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
| 232 |
},
|
| 233 |
+
{
|
| 234 |
+
"id": "separate_sgmat_tail",
|
| 235 |
+
"priority": 21,
|
| 236 |
+
"when": ["separateContract", "sparseSgmatOk", "sparseTailOk"],
|
| 237 |
+
"requires": {
|
| 238 |
+
"features": ["subgroups", "chromium-experimental-subgroup-matrix"],
|
| 239 |
+
"subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
|
| 240 |
+
},
|
| 241 |
+
"passes": [
|
| 242 |
+
{
|
| 243 |
+
"id": "append",
|
| 244 |
+
"name": "SparseAttention.Append",
|
| 245 |
+
"shader": "sparse-kv-append.wgsl.jinja",
|
| 246 |
+
"derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" },
|
| 247 |
+
"bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 248 |
+
"dispatch": {
|
| 249 |
+
"x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 250 |
+
"y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 251 |
+
"z": 1
|
| 252 |
+
}
|
| 253 |
+
},
|
| 254 |
+
{
|
| 255 |
+
"id": "attention",
|
| 256 |
+
"name": "SparseAttention.AttentionSgmat",
|
| 257 |
+
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 258 |
+
"bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 259 |
+
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 260 |
+
"derive": { "splitQueryTail": "true" }
|
| 261 |
+
},
|
| 262 |
+
{
|
| 263 |
+
"id": "tail",
|
| 264 |
+
"name": "SparseAttention.Tail",
|
| 265 |
+
"shader": "sparse-attention.wgsl.jinja",
|
| 266 |
+
"derive": {
|
| 267 |
+
"qTile": "sparseTailRows",
|
| 268 |
+
"queryOffset": "sparsePrefixRows",
|
| 269 |
+
"attnWorkgroup": "sparseTailWorkgroup",
|
| 270 |
+
"valueParts": "sparseTailValueParts",
|
| 271 |
+
"vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1",
|
| 272 |
+
"promptTail": "true"
|
| 273 |
+
},
|
| 274 |
+
"bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 275 |
+
"dispatch": { "x": "1", "y": "batchSize * numHeads" }
|
| 276 |
+
}
|
| 277 |
+
],
|
| 278 |
+
"demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"]
|
| 279 |
+
},
|
| 280 |
{
|
| 281 |
"id": "separate_rotary",
|
| 282 |
"priority": 10,
|
| 283 |
"when": ["separateRotaryContract"],
|
| 284 |
+
"derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" },
|
| 285 |
"intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }],
|
| 286 |
"passes": [
|
| 287 |
{
|
|
|
|
| 300 |
"id": "qrotary",
|
| 301 |
"name": "SparseAttention.QueryRotary",
|
| 302 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 303 |
+
"bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 304 |
"dispatch": {
|
| 305 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 306 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
|
|
| 311 |
"id": "attention",
|
| 312 |
"name": "SparseAttention.Attention",
|
| 313 |
"shader": "sparse-attention.wgsl.jinja",
|
| 314 |
+
"derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" },
|
| 315 |
+
"bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 316 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 317 |
}
|
| 318 |
]
|
|
|
|
| 343 |
"id": "qrotary",
|
| 344 |
"name": "SparseAttention.QueryRotary",
|
| 345 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 346 |
+
"bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 347 |
"dispatch": {
|
| 348 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 349 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
|
|
| 354 |
"id": "attention",
|
| 355 |
"name": "SparseAttention.AttentionSgmat",
|
| 356 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 357 |
+
"bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 358 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 359 |
+
"derive": { "splitQueryTail": "false" }
|
| 360 |
}
|
| 361 |
],
|
| 362 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
| 363 |
},
|
| 364 |
+
{
|
| 365 |
+
"id": "separate_rotary_sgmat_tail",
|
| 366 |
+
"priority": 31,
|
| 367 |
+
"when": ["separateRotaryContract", "sparseSgmatOk", "sparseTailOk"],
|
| 368 |
+
"requires": {
|
| 369 |
+
"features": ["subgroups", "chromium-experimental-subgroup-matrix"],
|
| 370 |
+
"subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
|
| 371 |
+
},
|
| 372 |
+
"intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }],
|
| 373 |
+
"passes": [
|
| 374 |
+
{
|
| 375 |
+
"id": "append",
|
| 376 |
+
"name": "SparseAttention.Append",
|
| 377 |
+
"shader": "sparse-kv-append.wgsl.jinja",
|
| 378 |
+
"derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" },
|
| 379 |
+
"bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"],
|
| 380 |
+
"dispatch": {
|
| 381 |
+
"x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 382 |
+
"y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 383 |
+
"z": 1
|
| 384 |
+
}
|
| 385 |
+
},
|
| 386 |
+
{
|
| 387 |
+
"id": "qrotary",
|
| 388 |
+
"name": "SparseAttention.QueryRotary",
|
| 389 |
+
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 390 |
+
"bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 391 |
+
"dispatch": {
|
| 392 |
+
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 393 |
+
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 394 |
+
"z": 1
|
| 395 |
+
}
|
| 396 |
+
},
|
| 397 |
+
{
|
| 398 |
+
"id": "attention",
|
| 399 |
+
"name": "SparseAttention.AttentionSgmat",
|
| 400 |
+
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 401 |
+
"bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 402 |
+
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 403 |
+
"derive": { "splitQueryTail": "true" }
|
| 404 |
+
},
|
| 405 |
+
{
|
| 406 |
+
"id": "tail",
|
| 407 |
+
"name": "SparseAttention.Tail",
|
| 408 |
+
"shader": "sparse-attention.wgsl.jinja",
|
| 409 |
+
"derive": {
|
| 410 |
+
"qTile": "sparseTailRows",
|
| 411 |
+
"queryOffset": "sparsePrefixRows",
|
| 412 |
+
"attnWorkgroup": "sparseTailWorkgroup",
|
| 413 |
+
"valueParts": "sparseTailValueParts",
|
| 414 |
+
"vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1",
|
| 415 |
+
"promptTail": "true"
|
| 416 |
+
},
|
| 417 |
+
"bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 418 |
+
"dispatch": { "x": "1", "y": "batchSize * numHeads" }
|
| 419 |
+
}
|
| 420 |
+
],
|
| 421 |
+
"demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"]
|
| 422 |
+
},
|
| 423 |
{
|
| 424 |
"id": "packed",
|
| 425 |
"priority": 0,
|
| 426 |
"when": ["packedContract"],
|
| 427 |
+
"derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" },
|
| 428 |
"passes": [
|
| 429 |
{
|
| 430 |
"id": "append",
|
|
|
|
| 442 |
"id": "attention",
|
| 443 |
"name": "SparseAttention.Attention",
|
| 444 |
"shader": "sparse-attention.wgsl.jinja",
|
| 445 |
+
"derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" },
|
| 446 |
+
"bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 447 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 448 |
}
|
| 449 |
]
|
|
|
|
| 473 |
"id": "attention",
|
| 474 |
"name": "SparseAttention.AttentionSgmat",
|
| 475 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 476 |
+
"bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 477 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 478 |
+
"derive": { "splitQueryTail": "false" }
|
| 479 |
}
|
| 480 |
],
|
| 481 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
| 482 |
},
|
| 483 |
+
{
|
| 484 |
+
"id": "packed_sgmat_tail",
|
| 485 |
+
"priority": 21,
|
| 486 |
+
"when": ["packedContract", "sparseSgmatOk", "sparseTailOk"],
|
| 487 |
+
"requires": {
|
| 488 |
+
"features": ["subgroups", "chromium-experimental-subgroup-matrix"],
|
| 489 |
+
"subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
|
| 490 |
+
},
|
| 491 |
+
"passes": [
|
| 492 |
+
{
|
| 493 |
+
"id": "append",
|
| 494 |
+
"name": "SparseAttention.Append",
|
| 495 |
+
"shader": "sparse-kv-append.wgsl.jinja",
|
| 496 |
+
"derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" },
|
| 497 |
+
"bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 498 |
+
"dispatch": {
|
| 499 |
+
"x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 500 |
+
"y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 501 |
+
"z": 1
|
| 502 |
+
}
|
| 503 |
+
},
|
| 504 |
+
{
|
| 505 |
+
"id": "attention",
|
| 506 |
+
"name": "SparseAttention.AttentionSgmat",
|
| 507 |
+
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 508 |
+
"bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 509 |
+
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 510 |
+
"derive": { "splitQueryTail": "true" }
|
| 511 |
+
},
|
| 512 |
+
{
|
| 513 |
+
"id": "tail",
|
| 514 |
+
"name": "SparseAttention.Tail",
|
| 515 |
+
"shader": "sparse-attention.wgsl.jinja",
|
| 516 |
+
"derive": {
|
| 517 |
+
"qTile": "sparseTailRows",
|
| 518 |
+
"queryOffset": "sparsePrefixRows",
|
| 519 |
+
"attnWorkgroup": "sparseTailWorkgroup",
|
| 520 |
+
"valueParts": "sparseTailValueParts",
|
| 521 |
+
"vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1",
|
| 522 |
+
"promptTail": "true"
|
| 523 |
+
},
|
| 524 |
+
"bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 525 |
+
"dispatch": { "x": "1", "y": "batchSize * numHeads" }
|
| 526 |
+
}
|
| 527 |
+
],
|
| 528 |
+
"demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"]
|
| 529 |
+
},
|
| 530 |
{
|
| 531 |
"id": "packed_rotary",
|
| 532 |
"priority": 10,
|
| 533 |
"when": ["packedRotaryContract"],
|
| 534 |
+
"derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" },
|
| 535 |
"intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }],
|
| 536 |
"passes": [
|
| 537 |
{
|
|
|
|
| 550 |
"id": "qrotary",
|
| 551 |
"name": "SparseAttention.QueryRotary",
|
| 552 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 553 |
+
"bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 554 |
"dispatch": {
|
| 555 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 556 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
|
|
| 561 |
"id": "attention",
|
| 562 |
"name": "SparseAttention.Attention",
|
| 563 |
"shader": "sparse-attention.wgsl.jinja",
|
| 564 |
+
"derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" },
|
| 565 |
+
"bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 566 |
"dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" }
|
| 567 |
}
|
| 568 |
]
|
|
|
|
| 593 |
"id": "qrotary",
|
| 594 |
"name": "SparseAttention.QueryRotary",
|
| 595 |
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 596 |
+
"bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 597 |
"dispatch": {
|
| 598 |
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 599 |
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
|
|
|
| 604 |
"id": "attention",
|
| 605 |
"name": "SparseAttention.AttentionSgmat",
|
| 606 |
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 607 |
+
"bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 608 |
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 609 |
+
"derive": { "splitQueryTail": "false" }
|
| 610 |
}
|
| 611 |
],
|
| 612 |
"demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"]
|
| 613 |
+
},
|
| 614 |
+
{
|
| 615 |
+
"id": "packed_rotary_sgmat_tail",
|
| 616 |
+
"priority": 31,
|
| 617 |
+
"when": ["packedRotaryContract", "sparseSgmatOk", "sparseTailOk"],
|
| 618 |
+
"requires": {
|
| 619 |
+
"features": ["subgroups", "chromium-experimental-subgroup-matrix"],
|
| 620 |
+
"subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
|
| 621 |
+
},
|
| 622 |
+
"intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }],
|
| 623 |
+
"passes": [
|
| 624 |
+
{
|
| 625 |
+
"id": "append",
|
| 626 |
+
"name": "SparseAttention.Append",
|
| 627 |
+
"shader": "sparse-kv-append.wgsl.jinja",
|
| 628 |
+
"derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" },
|
| 629 |
+
"bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"],
|
| 630 |
+
"dispatch": {
|
| 631 |
+
"x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 632 |
+
"y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 633 |
+
"z": 1
|
| 634 |
+
}
|
| 635 |
+
},
|
| 636 |
+
{
|
| 637 |
+
"id": "qrotary",
|
| 638 |
+
"name": "SparseAttention.QueryRotary",
|
| 639 |
+
"shader": "sparse-q-rotary.wgsl.jinja",
|
| 640 |
+
"bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"],
|
| 641 |
+
"dispatch": {
|
| 642 |
+
"x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 643 |
+
"y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)",
|
| 644 |
+
"z": 1
|
| 645 |
+
}
|
| 646 |
+
},
|
| 647 |
+
{
|
| 648 |
+
"id": "attention",
|
| 649 |
+
"name": "SparseAttention.AttentionSgmat",
|
| 650 |
+
"shader": "sparse-attention-sgmat.wgsl.jinja",
|
| 651 |
+
"bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 652 |
+
"dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
|
| 653 |
+
"derive": { "splitQueryTail": "true" }
|
| 654 |
+
},
|
| 655 |
+
{
|
| 656 |
+
"id": "tail",
|
| 657 |
+
"name": "SparseAttention.Tail",
|
| 658 |
+
"shader": "sparse-attention.wgsl.jinja",
|
| 659 |
+
"derive": {
|
| 660 |
+
"qTile": "sparseTailRows",
|
| 661 |
+
"queryOffset": "sparsePrefixRows",
|
| 662 |
+
"attnWorkgroup": "sparseTailWorkgroup",
|
| 663 |
+
"valueParts": "sparseTailValueParts",
|
| 664 |
+
"vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1",
|
| 665 |
+
"promptTail": "true"
|
| 666 |
+
},
|
| 667 |
+
"bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"],
|
| 668 |
+
"dispatch": { "x": "1", "y": "batchSize * numHeads" }
|
| 669 |
+
}
|
| 670 |
+
],
|
| 671 |
+
"demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"]
|
| 672 |
}
|
| 673 |
]
|
| 674 |
}
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,33 +1,37 @@
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.SparseAttention",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
-
"bench.json": "
|
| 11 |
-
"manifest.json": "
|
| 12 |
-
"sparse-attention-sgmat.wgsl.jinja": "
|
| 13 |
-
"sparse-attention.wgsl.jinja": "
|
| 14 |
-
"sparse-kv-append.wgsl.jinja": "
|
| 15 |
-
"sparse-q-rotary.wgsl.jinja": "
|
| 16 |
-
"test.json": "
|
| 17 |
}
|
| 18 |
},
|
| 19 |
-
"provenance": { "kernel": { "sha": "
|
| 20 |
"webgpu": {
|
| 21 |
-
"manifestSpec": "2.
|
| 22 |
"variants": {
|
| 23 |
"separate": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 24 |
"separate_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
|
|
|
| 25 |
"separate_rotary": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
| 26 |
"separate_rotary_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
|
|
|
| 27 |
"packed": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 28 |
"packed_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
|
|
|
| 29 |
"packed_rotary": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
| 30 |
-
"packed_rotary_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"]
|
|
|
|
| 31 |
}
|
| 32 |
}
|
| 33 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.SparseAttention",
|
| 3 |
+
"id": "_com_microsoft_sparseattention_webgpu_c2657b2",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
+
"bench.json": "4NNfSy0Pz7J1r5bacbFNwgbyHBbeJ3OYn3iajMdCQ64=",
|
| 11 |
+
"manifest.json": "8xLb9NHTZ4tCa1RARx5gvMP/WuEU4tsurxnNVELHnPc=",
|
| 12 |
+
"sparse-attention-sgmat.wgsl.jinja": "bUA2Sjw4qg7s/2KvOkUddx7N6v718rkkQNG7yQko+gI=",
|
| 13 |
+
"sparse-attention.wgsl.jinja": "rrRJNjo17muFxfSrY5YgMA9zkYih1zSDLYwZqDbViKQ=",
|
| 14 |
+
"sparse-kv-append.wgsl.jinja": "P1nJPEk4RcBb6GcajLnqUqMRV15GsBkIwCfLWBls6YE=",
|
| 15 |
+
"sparse-q-rotary.wgsl.jinja": "fTkf6Ht1LF4ygO2DviYl8kYqoGd7LKL3bSD1XBG+1sg=",
|
| 16 |
+
"test.json": "wyUajoZ/2iaiyzrdWd5+dceFdPQAmEAy/ol6vKTe0os="
|
| 17 |
}
|
| 18 |
},
|
| 19 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 20 |
"webgpu": {
|
| 21 |
+
"manifestSpec": "2.1",
|
| 22 |
"variants": {
|
| 23 |
"separate": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 24 |
"separate_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 25 |
+
"separate_sgmat_tail": ["sparse-attention-sgmat.wgsl.jinja", "sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 26 |
"separate_rotary": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
| 27 |
"separate_rotary_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
| 28 |
+
"separate_rotary_sgmat_tail": ["sparse-attention-sgmat.wgsl.jinja", "sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
| 29 |
"packed": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 30 |
"packed_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 31 |
+
"packed_sgmat_tail": ["sparse-attention-sgmat.wgsl.jinja", "sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja"],
|
| 32 |
"packed_rotary": ["sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
| 33 |
+
"packed_rotary_sgmat": ["sparse-attention-sgmat.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"],
|
| 34 |
+
"packed_rotary_sgmat_tail": ["sparse-attention-sgmat.wgsl.jinja", "sparse-attention.wgsl.jinja", "sparse-kv-append.wgsl.jinja", "sparse-q-rotary.wgsl.jinja"]
|
| 35 |
}
|
| 36 |
}
|
| 37 |
}
|
build/webgpu/sparse-attention-sgmat.wgsl.jinja
CHANGED
|
@@ -10,8 +10,7 @@ fn past_sequence_length(batch: u32) -> u32 {
|
|
| 10 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 11 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 12 |
}
|
| 13 |
-
{%
|
| 14 |
-
|
| 15 |
enable subgroups;
|
| 16 |
{% if pinSubgroupSize32 %}
|
| 17 |
enable subgroup_size_control;
|
|
@@ -19,19 +18,18 @@ enable subgroup_size_control;
|
|
| 19 |
enable chromium_experimental_subgroup_matrix;
|
| 20 |
diagnostic(off, chromium.subgroup_matrix_uniformity);
|
| 21 |
|
| 22 |
-
|
| 23 |
{{ env.wgsl.resourceDeclarations }}
|
| 24 |
|
| 25 |
// Subgroup-matrix attention over 64-query tiles. Each workgroup processes one
|
| 26 |
-
//
|
| 27 |
-
// as
|
| 28 |
-
//
|
| 29 |
-
// with a partial final tile stage Q with zero padding. The manifest bounds tile widths by workgroup storage.
|
| 30 |
//
|
| 31 |
-
//
|
| 32 |
-
//
|
| 33 |
-
//
|
| 34 |
-
//
|
|
|
|
| 35 |
const Q_HEADS: u32 = {{ numHeads }}u;
|
| 36 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 37 |
const HEAD_DIM: u32 = {{ headSize }}u;
|
|
@@ -47,7 +45,7 @@ const Q_STRIDE: u32 = {{ packedStride if packedQkv else numHeads * headSize }}u;
|
|
| 47 |
// Eight 32-lane subgroups form a 4x2 grid. The device storage budget chooses
|
| 48 |
// 64 or 32 key columns and a matching head-dimension staging width. Both divide
|
| 49 |
// the admitted sparse-block and head dimensions without a key or head tail.
|
| 50 |
-
const TILE_M: u32 =
|
| 51 |
const TILE_N: u32 = {{ sparseSgmatTileN }}u;
|
| 52 |
const TILE_K: u32 = {{ sparseSgmatTileK }}u;
|
| 53 |
const SUB_TILES: u32 = {{ (sparseBlockSize / sparseSgmatTileN) | int }}u;
|
|
@@ -74,49 +72,41 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 74 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 75 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 76 |
}
|
|
|
|
| 77 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 78 |
return exp(shifted_value(value, maxValue));
|
| 79 |
}
|
| 80 |
|
| 81 |
-
// Q staging for the score GEMM;
|
| 82 |
-
//
|
| 83 |
var<workgroup> tile_q: array<f32, {{ 64 * sparseSgmatTileK }}>;
|
| 84 |
// Tile probabilities for the P.V GEMM; the output epilogue aliases it as the
|
| 85 |
// result-fragment scratch once the last key tile's readers are done.
|
| 86 |
var<workgroup> prob_tile: array<f32, {{ 64 * sparseSgmatTileN }}>;
|
| 87 |
var<workgroup> row_m: array<f32, 64>;
|
| 88 |
var<workgroup> row_d: array<f32, 64>;
|
|
|
|
| 89 |
// Per-key-tile row partials, one slot per (row, subgroup column group).
|
| 90 |
var<workgroup> part_m: array<f32, 128>;
|
| 91 |
var<workgroup> part_d: array<f32, 128>;
|
| 92 |
|
| 93 |
-
{%
|
| 94 |
-
fn scale_value() -> f32 {
|
| 95 |
if (params.scale != 0.0) { return params.scale; }
|
| 96 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 97 |
}
|
| 98 |
|
| 99 |
-
|
| 100 |
{{ sparse_schedule() }}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
|
| 102 |
-
{% set queryMatrixSource = "tile_q" if not sgmatDirectQuery else ("q_rotary" if usesRotary else "query") %}
|
| 103 |
-
{% set queryMatrixStride = "TILE_K" if not sgmatDirectQuery else ("HEAD_DIM" if usesRotary else "Q_STRIDE") %}
|
| 104 |
-
{% macro query_matrix_offset(rb) -%}
|
| 105 |
-
{% if not sgmatDirectQuery -%}
|
| 106 |
-
(base_a + {{ rb * 8 }}u) * TILE_K + step
|
| 107 |
-
{%- elif usesRotary -%}
|
| 108 |
-
((batch * Q_HEADS + head) * params.seqLen + tile0 + base_a + {{ rb * 8 }}u) * HEAD_DIM + k_base + step
|
| 109 |
-
{%- else -%}
|
| 110 |
-
(batch * params.seqLen + tile0 + base_a + {{ rb * 8 }}u) * Q_STRIDE + head * HEAD_DIM + k_base + step
|
| 111 |
-
{%- endif %}
|
| 112 |
-
{%- endmacro %}
|
| 113 |
-
|
| 114 |
-
{% macro score_tile() %}
|
| 115 |
// S = Q.K^T, retaining the same sequence of 8-wide matrix operations.
|
| 116 |
// Fully populated query tiles need neither staging nor K-loop barriers.
|
| 117 |
-
//
|
| 118 |
for (var k_base = 0u; k_base < HEAD_DIM; k_base += TILE_K) {
|
| 119 |
-
{% if not
|
| 120 |
{
|
| 121 |
let a_row = li / 4u;
|
| 122 |
let a_col = (li % 4u) * {{ (sparseSgmatTileK / 4) | int }}u;
|
|
@@ -140,7 +130,7 @@ fn scale_value() -> f32 {
|
|
| 140 |
for (var step = 0u; step < TILE_K; step += 8u) {
|
| 141 |
{% for rb in range(2) %}
|
| 142 |
let mat_a{{ rb }} = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
|
| 143 |
-
&{{ queryMatrixSource }}, {{
|
| 144 |
);
|
| 145 |
{% endfor %}
|
| 146 |
{% for cb in range(scoreColBlocks) %}
|
|
@@ -156,13 +146,13 @@ fn scale_value() -> f32 {
|
|
| 156 |
{% endfor %}
|
| 157 |
{% endfor %}
|
| 158 |
}
|
| 159 |
-
{% if not
|
| 160 |
workgroupBarrier();
|
| 161 |
{% endif %}
|
| 162 |
}
|
| 163 |
{% endmacro %}
|
| 164 |
|
| 165 |
-
{% macro sweep(
|
| 166 |
// Consecutive queries span at most two mask rows, and every query of a row
|
| 167 |
// selects the same blocks, so one sweep per row covers the tile; a query
|
| 168 |
// contributes only to its own row's tiles.
|
|
@@ -203,20 +193,28 @@ fn scale_value() -> f32 {
|
|
| 203 |
var mat_s{{ rb }}{{ cb }}: subgroup_matrix_result<f32, 8, 8>;
|
| 204 |
{% endfor %}
|
| 205 |
{% endfor %}
|
| 206 |
-
{
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 207 |
{% for rb in range(2) %}
|
| 208 |
{% if rb > 0 %}
|
| 209 |
// The banks alias the Q staging tile; the previous row block's readers
|
| 210 |
// must finish before this one overwrites them.
|
| 211 |
workgroupBarrier();
|
| 212 |
{% endif %}
|
| 213 |
-
|
| 214 |
// All four lanes of a quad carry the same score row (row_in_block is
|
| 215 |
// lane / 4), so the accumulator below is a partial over one row and
|
| 216 |
// the butterfly merging it is quad-uniform.
|
| 217 |
var tile_stat_m{{ rb }} = -FLT_MAX;
|
| 218 |
var tile_stat_d{{ rb }} = 0.0;
|
| 219 |
-
|
| 220 |
{% for cb in range(scoreColBlocks) %}
|
| 221 |
subgroupMatrixStore<row_major>(
|
| 222 |
&tile_q, (subgroup * {{ scoreColBlocks }}u + {{ cb }}u) * 64u, mat_s{{ rb }}{{ cb }}, 8u
|
|
@@ -229,28 +227,21 @@ fn scale_value() -> f32 {
|
|
| 229 |
let key = key_base + base_b + {{ cb * 8 }}u + col_in_block + pair;
|
| 230 |
let q_abs = q_abs0 + r;
|
| 231 |
let allowed = r < rows_live && q_abs / SPARSE_BLOCK == mask_row && key <= q_abs;
|
| 232 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 233 |
if (allowed) {
|
| 234 |
-
let scored = tile_q[
|
| 235 |
-
(subgroup * {{ scoreColBlocks }}u + {{ cb }}u) * 64u + row_in_block * 8u + col_in_block + pair
|
| 236 |
-
] * scale;
|
| 237 |
let prev_m = tile_stat_m{{ rb }};
|
| 238 |
tile_stat_m{{ rb }} = max(tile_stat_m{{ rb }}, scored);
|
| 239 |
tile_stat_d{{ rb }} = tile_stat_d{{ rb }} * exp_shift(prev_m, tile_stat_m{{ rb }})
|
| 240 |
+ exp_shift(scored, tile_stat_m{{ rb }});
|
| 241 |
}
|
| 242 |
-
|
| 243 |
-
var prob = 0.0;
|
| 244 |
-
if (allowed) {
|
| 245 |
-
prob = exp_shift(tile_q[
|
| 246 |
-
(subgroup * {{ scoreColBlocks }}u + {{ cb }}u) * 64u + row_in_block * 8u + col_in_block + pair
|
| 247 |
-
] * scale, row_m[r]);
|
| 248 |
-
}
|
| 249 |
-
prob_tile[r * TILE_N + base_b + {{ cb * 8 }}u + col_in_block + pair] = prob;
|
| 250 |
-
{% endif %}
|
| 251 |
}
|
| 252 |
{% endfor %}
|
| 253 |
-
|
| 254 |
// Butterfly the quad unconditionally: a lane whose row ran past the
|
| 255 |
// query tail carries the exact identity (-FLT_MAX, 0), which merges to
|
| 256 |
// a no-op, and a subgroup shuffle under a partial guard would not be
|
|
@@ -270,9 +261,9 @@ fn scale_value() -> f32 {
|
|
| 270 |
part_m[stat_row * 2u + subtile_idx] = tile_stat_m{{ rb }};
|
| 271 |
part_d[stat_row * 2u + subtile_idx] = tile_stat_d{{ rb }};
|
| 272 |
}
|
| 273 |
-
|
| 274 |
{% endfor %}
|
| 275 |
-
|
| 276 |
workgroupBarrier();
|
| 277 |
// Fold both column groups' partials into the running row statistics,
|
| 278 |
// in the same (max, rescale, add) form as the per-lane walk: an empty
|
|
@@ -288,11 +279,44 @@ fn scale_value() -> f32 {
|
|
| 288 |
merged_d = merged_d * exp_shift(merged_m, new_m) + d2 * exp_shift(m2, new_m);
|
| 289 |
merged_m = new_m;
|
| 290 |
}
|
|
|
|
| 291 |
row_m[li] = merged_m;
|
| 292 |
row_d[li] = merged_d;
|
| 293 |
}
|
| 294 |
workgroupBarrier();
|
| 295 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 296 |
workgroupBarrier();
|
| 297 |
// O += P.V. P streams from shared memory; V rows are contiguous cache
|
| 298 |
// rows, loaded directly as right-hand fragments. A masked or padded
|
|
@@ -319,19 +343,22 @@ fn scale_value() -> f32 {
|
|
| 319 |
}
|
| 320 |
// Orders this tile's prob_tile reads before the next tile rewrites it.
|
| 321 |
workgroupBarrier();
|
| 322 |
-
{% endif %}
|
| 323 |
}
|
| 324 |
}
|
| 325 |
}
|
| 326 |
{% endmacro %}
|
| 327 |
|
| 328 |
-
@compute @workgroup_size(
|
| 329 |
fn main(
|
| 330 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 331 |
@builtin(local_invocation_index) li: u32,
|
| 332 |
@builtin(subgroup_invocation_id) lane: u32
|
| 333 |
) {
|
| 334 |
let tile0 = wg.x * TILE_M;
|
|
|
|
|
|
|
|
|
|
|
|
|
| 335 |
let head = wg.y % Q_HEADS;
|
| 336 |
let batch = wg.y / Q_HEADS;
|
| 337 |
|
|
@@ -367,14 +394,13 @@ fn main(
|
|
| 367 |
{% endfor %}
|
| 368 |
{% endfor %}
|
| 369 |
|
| 370 |
-
for (var r = li; r < TILE_M; r +=
|
| 371 |
row_m[r] = -FLT_MAX;
|
| 372 |
row_d[r] = 0.0;
|
| 373 |
}
|
| 374 |
workgroupBarrier();
|
| 375 |
|
| 376 |
-
{{ sweep(
|
| 377 |
-
{{ sweep("apply") }}
|
| 378 |
|
| 379 |
// Normalize by the final denominators and store. Publish only as many
|
| 380 |
// fragment columns per batch as fit prob_tile, keeping the smaller key tile's
|
|
@@ -413,7 +439,7 @@ fn main(
|
|
| 413 |
if (!(row_d[r] > 0.0)) {
|
| 414 |
let key_bound = q_abs0 + r + 1u;
|
| 415 |
let cache_base = (batch * KV_HEADS + kv_head) * MAX_CACHE_SEQ * HEAD_DIM;
|
| 416 |
-
for (var d = li; d < HEAD_DIM; d +=
|
| 417 |
var total = 0.0;
|
| 418 |
for (var key = 0u; key < key_bound; key++) {
|
| 419 |
total += f32(present_value[cache_base + key * HEAD_DIM + d]);
|
|
|
|
| 10 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 11 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 12 |
}
|
| 13 |
+
{% endmacro %}
|
|
|
|
| 14 |
enable subgroups;
|
| 15 |
{% if pinSubgroupSize32 %}
|
| 16 |
enable subgroup_size_control;
|
|
|
|
| 18 |
enable chromium_experimental_subgroup_matrix;
|
| 19 |
diagnostic(off, chromium.subgroup_matrix_uniformity);
|
| 20 |
|
|
|
|
| 21 |
{{ env.wgsl.resourceDeclarations }}
|
| 22 |
|
| 23 |
// Subgroup-matrix attention over 64-query tiles. Each workgroup processes one
|
| 24 |
+
// (batch, query tile, query head). K/V rows and complete query tiles load
|
| 25 |
+
// directly as matrix fragments; only the final partial query tile stages Q.
|
| 26 |
+
// The device workgroup-storage budget bounds the key and staging tile widths.
|
|
|
|
| 27 |
//
|
| 28 |
+
// One score sweep updates each row's online softmax statistics and rescales its
|
| 29 |
+
// accumulated P.V fragments before adding the current tile. The score scratch
|
| 30 |
+
// doubles as a bounded bank for fragment rescaling, avoiding a second QK sweep
|
| 31 |
+
// or a full output staging array. Causal bounds, duplicate CSR columns, dense
|
| 32 |
+
// rows, and all-masked rows retain the shared sparse-attention semantics.
|
| 33 |
const Q_HEADS: u32 = {{ numHeads }}u;
|
| 34 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 35 |
const HEAD_DIM: u32 = {{ headSize }}u;
|
|
|
|
| 45 |
// Eight 32-lane subgroups form a 4x2 grid. The device storage budget chooses
|
| 46 |
// 64 or 32 key columns and a matching head-dimension staging width. Both divide
|
| 47 |
// the admitted sparse-block and head dimensions without a key or head tail.
|
| 48 |
+
const TILE_M: u32 = {{ sparseSgmatTileM }}u;
|
| 49 |
const TILE_N: u32 = {{ sparseSgmatTileN }}u;
|
| 50 |
const TILE_K: u32 = {{ sparseSgmatTileK }}u;
|
| 51 |
const SUB_TILES: u32 = {{ (sparseBlockSize / sparseSgmatTileN) | int }}u;
|
|
|
|
| 72 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 73 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 74 |
}
|
| 75 |
+
|
| 76 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 77 |
return exp(shifted_value(value, maxValue));
|
| 78 |
}
|
| 79 |
|
| 80 |
+
// Q staging for the score GEMM; score epilogues and output rescaling alias
|
| 81 |
+
// these banks (8 subgroups x scoreColBlocks banks x 64 elements).
|
| 82 |
var<workgroup> tile_q: array<f32, {{ 64 * sparseSgmatTileK }}>;
|
| 83 |
// Tile probabilities for the P.V GEMM; the output epilogue aliases it as the
|
| 84 |
// result-fragment scratch once the last key tile's readers are done.
|
| 85 |
var<workgroup> prob_tile: array<f32, {{ 64 * sparseSgmatTileN }}>;
|
| 86 |
var<workgroup> row_m: array<f32, 64>;
|
| 87 |
var<workgroup> row_d: array<f32, 64>;
|
| 88 |
+
var<workgroup> row_c: array<f32, 64>;
|
| 89 |
// Per-key-tile row partials, one slot per (row, subgroup column group).
|
| 90 |
var<workgroup> part_m: array<f32, 128>;
|
| 91 |
var<workgroup> part_d: array<f32, 128>;
|
| 92 |
|
| 93 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 94 |
if (params.scale != 0.0) { return params.scale; }
|
| 95 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 96 |
}
|
| 97 |
|
|
|
|
| 98 |
{{ sparse_schedule() }}
|
| 99 |
+
{% macro score_tile(direct) %}
|
| 100 |
+
{% set queryMatrixSource = "tile_q" if not direct else ("q_rotary" if usesRotary else "query") %}
|
| 101 |
+
{% set queryMatrixStride = "TILE_K" if not direct else ("HEAD_DIM" if usesRotary else "Q_STRIDE") %}
|
| 102 |
+
{% set queryRowOrigin = "" if not direct else ("(batch * Q_HEADS + head) * params.seqLen + tile0 + " if usesRotary else "batch * params.seqLen + tile0 + ") %}
|
| 103 |
+
{% set queryColOrigin = "" if not direct else ("k_base + " if usesRotary else "head * HEAD_DIM + k_base + ") %}
|
| 104 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
// S = Q.K^T, retaining the same sequence of 8-wide matrix operations.
|
| 106 |
// Fully populated query tiles need neither staging nor K-loop barriers.
|
| 107 |
+
// Only a partial query tile needs zero-padded staging.
|
| 108 |
for (var k_base = 0u; k_base < HEAD_DIM; k_base += TILE_K) {
|
| 109 |
+
{% if not direct %}
|
| 110 |
{
|
| 111 |
let a_row = li / 4u;
|
| 112 |
let a_col = (li % 4u) * {{ (sparseSgmatTileK / 4) | int }}u;
|
|
|
|
| 130 |
for (var step = 0u; step < TILE_K; step += 8u) {
|
| 131 |
{% for rb in range(2) %}
|
| 132 |
let mat_a{{ rb }} = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
|
| 133 |
+
&{{ queryMatrixSource }}, ({{ queryRowOrigin }}base_a + {{ rb * 8 }}u) * {{ queryMatrixStride }} + {{ queryColOrigin }}step, {{ queryMatrixStride }}
|
| 134 |
);
|
| 135 |
{% endfor %}
|
| 136 |
{% for cb in range(scoreColBlocks) %}
|
|
|
|
| 146 |
{% endfor %}
|
| 147 |
{% endfor %}
|
| 148 |
}
|
| 149 |
+
{% if not direct %}
|
| 150 |
workgroupBarrier();
|
| 151 |
{% endif %}
|
| 152 |
}
|
| 153 |
{% endmacro %}
|
| 154 |
|
| 155 |
+
{% macro sweep() %}
|
| 156 |
// Consecutive queries span at most two mask rows, and every query of a row
|
| 157 |
// selects the same blocks, so one sweep per row covers the tile; a query
|
| 158 |
// contributes only to its own row's tiles.
|
|
|
|
| 193 |
var mat_s{{ rb }}{{ cb }}: subgroup_matrix_result<f32, 8, 8>;
|
| 194 |
{% endfor %}
|
| 195 |
{% endfor %}
|
| 196 |
+
{% if sgmatDirectQuery %}
|
| 197 |
+
{{ score_tile(true) }}
|
| 198 |
+
{% else %}
|
| 199 |
+
if (rows_live == TILE_M) {
|
| 200 |
+
{{ score_tile(true) }}
|
| 201 |
+
} else {
|
| 202 |
+
{{ score_tile(false) }}
|
| 203 |
+
}
|
| 204 |
+
{% endif %}
|
| 205 |
{% for rb in range(2) %}
|
| 206 |
{% if rb > 0 %}
|
| 207 |
// The banks alias the Q staging tile; the previous row block's readers
|
| 208 |
// must finish before this one overwrites them.
|
| 209 |
workgroupBarrier();
|
| 210 |
{% endif %}
|
| 211 |
+
|
| 212 |
// All four lanes of a quad carry the same score row (row_in_block is
|
| 213 |
// lane / 4), so the accumulator below is a partial over one row and
|
| 214 |
// the butterfly merging it is quad-uniform.
|
| 215 |
var tile_stat_m{{ rb }} = -FLT_MAX;
|
| 216 |
var tile_stat_d{{ rb }} = 0.0;
|
| 217 |
+
|
| 218 |
{% for cb in range(scoreColBlocks) %}
|
| 219 |
subgroupMatrixStore<row_major>(
|
| 220 |
&tile_q, (subgroup * {{ scoreColBlocks }}u + {{ cb }}u) * 64u, mat_s{{ rb }}{{ cb }}, 8u
|
|
|
|
| 227 |
let key = key_base + base_b + {{ cb * 8 }}u + col_in_block + pair;
|
| 228 |
let q_abs = q_abs0 + r;
|
| 229 |
let allowed = r < rows_live && q_abs / SPARSE_BLOCK == mask_row && key <= q_abs;
|
| 230 |
+
|
| 231 |
+
let scored = tile_q[
|
| 232 |
+
(subgroup * {{ scoreColBlocks }}u + {{ cb }}u) * 64u + row_in_block * 8u + col_in_block + pair
|
| 233 |
+
] * scale;
|
| 234 |
+
prob_tile[r * TILE_N + base_b + {{ cb * 8 }}u + col_in_block + pair] = scored;
|
| 235 |
if (allowed) {
|
|
|
|
|
|
|
|
|
|
| 236 |
let prev_m = tile_stat_m{{ rb }};
|
| 237 |
tile_stat_m{{ rb }} = max(tile_stat_m{{ rb }}, scored);
|
| 238 |
tile_stat_d{{ rb }} = tile_stat_d{{ rb }} * exp_shift(prev_m, tile_stat_m{{ rb }})
|
| 239 |
+ exp_shift(scored, tile_stat_m{{ rb }});
|
| 240 |
}
|
| 241 |
+
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 242 |
}
|
| 243 |
{% endfor %}
|
| 244 |
+
|
| 245 |
// Butterfly the quad unconditionally: a lane whose row ran past the
|
| 246 |
// query tail carries the exact identity (-FLT_MAX, 0), which merges to
|
| 247 |
// a no-op, and a subgroup shuffle under a partial guard would not be
|
|
|
|
| 261 |
part_m[stat_row * 2u + subtile_idx] = tile_stat_m{{ rb }};
|
| 262 |
part_d[stat_row * 2u + subtile_idx] = tile_stat_d{{ rb }};
|
| 263 |
}
|
| 264 |
+
|
| 265 |
{% endfor %}
|
| 266 |
+
|
| 267 |
workgroupBarrier();
|
| 268 |
// Fold both column groups' partials into the running row statistics,
|
| 269 |
// in the same (max, rescale, add) form as the per-lane walk: an empty
|
|
|
|
| 279 |
merged_d = merged_d * exp_shift(merged_m, new_m) + d2 * exp_shift(m2, new_m);
|
| 280 |
merged_m = new_m;
|
| 281 |
}
|
| 282 |
+
row_c[li] = exp_shift(row_m[li], merged_m);
|
| 283 |
row_m[li] = merged_m;
|
| 284 |
row_d[li] = merged_d;
|
| 285 |
}
|
| 286 |
workgroupBarrier();
|
| 287 |
+
|
| 288 |
+
// Rescale the persistent output fragments through the existing score
|
| 289 |
+
// scratch, using only as many banks as this device's key tile provides.
|
| 290 |
+
{% for rb in range(2) %}
|
| 291 |
+
{% for cbBase in range(0, pvColBlocks, scoreColBlocks) %}
|
| 292 |
+
{% set cbEnd = pvColBlocks if pvColBlocks < cbBase + scoreColBlocks else cbBase + scoreColBlocks %}
|
| 293 |
+
{% for cb in range(cbBase, cbEnd) %}
|
| 294 |
+
subgroupMatrixStore<row_major>(&tile_q,
|
| 295 |
+
(subgroup * {{ scoreColBlocks }}u + {{ cb - cbBase }}u) * 64u, mat_o{{ rb }}{{ cb }}, 8u);
|
| 296 |
+
{% endfor %}
|
| 297 |
+
workgroupBarrier();
|
| 298 |
+
{% for cb in range(cbBase, cbEnd) %}
|
| 299 |
+
for (var pair = 0u; pair < 2u; pair++) {
|
| 300 |
+
let index = (subgroup * {{ scoreColBlocks }}u + {{ cb - cbBase }}u) * 64u + row_in_block * 8u + col_in_block + pair;
|
| 301 |
+
tile_q[index] *= row_c[base_a + {{ rb * 8 }}u + row_in_block];
|
| 302 |
+
}
|
| 303 |
+
{% endfor %}
|
| 304 |
+
workgroupBarrier();
|
| 305 |
+
{% for cb in range(cbBase, cbEnd) %}
|
| 306 |
+
mat_o{{ rb }}{{ cb }} = subgroupMatrixLoad<subgroup_matrix_result<f32, 8, 8>, row_major>(
|
| 307 |
+
&tile_q, (subgroup * {{ scoreColBlocks }}u + {{ cb - cbBase }}u) * 64u, 8u);
|
| 308 |
+
{% endfor %}
|
| 309 |
+
workgroupBarrier();
|
| 310 |
+
{% endfor %}
|
| 311 |
+
{% endfor %}
|
| 312 |
+
// The single score sweep supplied each row's new normalizer above.
|
| 313 |
+
for (var i = li; i < TILE_M * TILE_N; i += {{ sparseSgmatWorkgroup }}u) {
|
| 314 |
+
let r = i / TILE_N;
|
| 315 |
+
let key = key_base + i % TILE_N;
|
| 316 |
+
let q_abs = q_abs0 + r;
|
| 317 |
+
let allowed = r < rows_live && q_abs / SPARSE_BLOCK == mask_row && key <= q_abs;
|
| 318 |
+
prob_tile[i] = select(0.0, exp_shift(prob_tile[i], row_m[r]), allowed);
|
| 319 |
+
}
|
| 320 |
workgroupBarrier();
|
| 321 |
// O += P.V. P streams from shared memory; V rows are contiguous cache
|
| 322 |
// rows, loaded directly as right-hand fragments. A masked or padded
|
|
|
|
| 343 |
}
|
| 344 |
// Orders this tile's prob_tile reads before the next tile rewrites it.
|
| 345 |
workgroupBarrier();
|
|
|
|
| 346 |
}
|
| 347 |
}
|
| 348 |
}
|
| 349 |
{% endmacro %}
|
| 350 |
|
| 351 |
+
@compute @workgroup_size({{ sparseSgmatWorkgroup }}, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
|
| 352 |
fn main(
|
| 353 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 354 |
@builtin(local_invocation_index) li: u32,
|
| 355 |
@builtin(subgroup_invocation_id) lane: u32
|
| 356 |
) {
|
| 357 |
let tile0 = wg.x * TILE_M;
|
| 358 |
+
{% if splitQueryTail %}
|
| 359 |
+
// Only a prompt delegates its partial query tile to the portable tail pass.
|
| 360 |
+
if (tile0 >= {{ sparsePrefixRows }}u && u32(total_sequence_length[0]) == params.seqLen) { return; }
|
| 361 |
+
{% endif %}
|
| 362 |
let head = wg.y % Q_HEADS;
|
| 363 |
let batch = wg.y / Q_HEADS;
|
| 364 |
|
|
|
|
| 394 |
{% endfor %}
|
| 395 |
{% endfor %}
|
| 396 |
|
| 397 |
+
for (var r = li; r < TILE_M; r += {{ sparseSgmatWorkgroup }}u) {
|
| 398 |
row_m[r] = -FLT_MAX;
|
| 399 |
row_d[r] = 0.0;
|
| 400 |
}
|
| 401 |
workgroupBarrier();
|
| 402 |
|
| 403 |
+
{{ sweep() }}
|
|
|
|
| 404 |
|
| 405 |
// Normalize by the final denominators and store. Publish only as many
|
| 406 |
// fragment columns per batch as fit prob_tile, keeping the smaller key tile's
|
|
|
|
| 439 |
if (!(row_d[r] > 0.0)) {
|
| 440 |
let key_bound = q_abs0 + r + 1u;
|
| 441 |
let cache_base = (batch * KV_HEADS + kv_head) * MAX_CACHE_SEQ * HEAD_DIM;
|
| 442 |
+
for (var d = li; d < HEAD_DIM; d += {{ sparseSgmatWorkgroup }}u) {
|
| 443 |
var total = 0.0;
|
| 444 |
for (var key = 0u; key < key_bound; key++) {
|
| 445 |
total += f32(present_value[cache_base + key * HEAD_DIM + d]);
|
build/webgpu/sparse-attention.wgsl.jinja
CHANGED
|
@@ -9,8 +9,7 @@ fn past_sequence_length(batch: u32) -> u32 {
|
|
| 9 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 10 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 11 |
}
|
| 12 |
-
{%
|
| 13 |
-
|
| 14 |
{{ env.wgsl.resourceDeclarations }}
|
| 15 |
|
| 16 |
// com.microsoft.SparseAttention, attention pass.
|
|
@@ -66,6 +65,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 66 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 67 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 68 |
}
|
|
|
|
| 69 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 70 |
return exp(shifted_value(value, maxValue));
|
| 71 |
}
|
|
@@ -79,6 +79,10 @@ var<workgroup> probs: array<f32, Q_TILE * WG>;
|
|
| 79 |
const V_STAGE_KEYS: u32 = 16u;
|
| 80 |
var<workgroup> v_stage: array<vec4<f32>, V_STAGE_KEYS * HEAD_VEC>;
|
| 81 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 82 |
// One resolved cache row base per key of the current tile, so the value accumulation
|
| 83 |
// re-reads a base instead of re-walking the column list per head dimension.
|
| 84 |
var<workgroup> key_rows: array<u32, WG>;
|
|
@@ -90,59 +94,6 @@ var<workgroup> key_rows: array<u32, WG>;
|
|
| 90 |
// Both the subgroup and portable barrier-tree engines return the same merged
|
| 91 |
// pair to every invocation. Repeated merges require a workgroup barrier between
|
| 92 |
// calls before their shared partial storage is reused.
|
| 93 |
-
{% set combineSubgroups = combineSubgroups is defined and combineSubgroups %}
|
| 94 |
-
{% if combineSubgroups %}
|
| 95 |
-
// Cross-subgroup merge that assumes nothing about which invocations share a
|
| 96 |
-
// subgroup or how many subgroups there are: each subgroup's elected lane
|
| 97 |
-
// publishes the subgroup pair in the slot at its OWN invocation index and sets
|
| 98 |
-
// that index's bit in a workgroup bitmask; thread 0 then folds exactly the
|
| 99 |
-
// published slots, in ascending index order (the online (m, d) merge is not
|
| 100 |
-
// float-associative, so the order is fixed), and clears the mask for the next
|
| 101 |
-
// call as it reads it. Workgroup memory starts zeroed, so the mask needs no
|
| 102 |
-
// setup. Same three collectives as a single-subgroup reduce, two barriers.
|
| 103 |
-
var<workgroup> partialM: array<f32, WG>;
|
| 104 |
-
var<workgroup> partialD: array<f32, WG>;
|
| 105 |
-
var<workgroup> leaderMask: array<atomic<u32>, (WG + 31u) / 32u>;
|
| 106 |
-
var<workgroup> combinedMD: vec2<f32>;
|
| 107 |
-
|
| 108 |
-
// When the whole workgroup is one subgroup the subgroup reduce already covers
|
| 109 |
-
// it (no barriers, no shared state). `subgroup_size` is the size of the current
|
| 110 |
-
// subgroup and uniform, so the test is exact and may guard the barriers below.
|
| 111 |
-
fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
|
| 112 |
-
let sgM = subgroupMax(m);
|
| 113 |
-
// A lane with no elements contributes d == 0 (exact identity). A +inf
|
| 114 |
-
// element made exp(inf - inf) = NaN stick in that lane's d; a NaN element
|
| 115 |
-
// landed in d via exp(NaN); both survive the merge and are detected by the
|
| 116 |
-
// code after the reduction.
|
| 117 |
-
let sgD = subgroupAdd(d * exp_shift(m, sgM));
|
| 118 |
-
if (sgSize == WG) {
|
| 119 |
-
return vec2<f32>(sgM, sgD);
|
| 120 |
-
}
|
| 121 |
-
if (subgroupElect()) {
|
| 122 |
-
partialM[lidx] = sgM;
|
| 123 |
-
partialD[lidx] = sgD;
|
| 124 |
-
atomicOr(&leaderMask[lidx / 32u], 1u << (lidx % 32u));
|
| 125 |
-
}
|
| 126 |
-
workgroupBarrier();
|
| 127 |
-
if (lidx == 0u) {
|
| 128 |
-
var accM = -FLT_MAX;
|
| 129 |
-
var accD = 0.0;
|
| 130 |
-
for (var w = 0u; w < (WG + 31u) / 32u; w = w + 1u) {
|
| 131 |
-
var bits = atomicExchange(&leaderMask[w], 0u);
|
| 132 |
-
while (bits != 0u) {
|
| 133 |
-
let slot = w * 32u + firstTrailingBit(bits);
|
| 134 |
-
bits = bits & (bits - 1u);
|
| 135 |
-
let mNew = max(accM, partialM[slot]);
|
| 136 |
-
accD = accD * exp_shift(accM, mNew) + partialD[slot] * exp_shift(partialM[slot], mNew);
|
| 137 |
-
accM = mNew;
|
| 138 |
-
}
|
| 139 |
-
}
|
| 140 |
-
combinedMD = vec2<f32>(accM, accD);
|
| 141 |
-
}
|
| 142 |
-
workgroupBarrier();
|
| 143 |
-
return combinedMD;
|
| 144 |
-
}
|
| 145 |
-
{% else %}
|
| 146 |
{% set mdStreamed = mdStreams is defined %}
|
| 147 |
{% set mdStreams = mdStreams if mdStreams is defined else 1 %}
|
| 148 |
{% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
|
|
@@ -207,21 +158,20 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
|
| 207 |
return merged;
|
| 208 |
}
|
| 209 |
{% endif %}
|
| 210 |
-
{% endif %}
|
| 211 |
|
| 212 |
-
|
| 213 |
-
{% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
|
| 214 |
-
fn scale_value() -> f32 {
|
| 215 |
if (params.scale != 0.0) { return params.scale; }
|
| 216 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 217 |
}
|
| 218 |
|
| 219 |
-
|
| 220 |
{{ sparse_schedule() }}
|
| 221 |
-
|
| 222 |
@compute @workgroup_size(WG, 1, 1)
|
| 223 |
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
|
| 224 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 225 |
let head = wg.y % Q_HEADS;
|
| 226 |
let batch = wg.y / Q_HEADS;
|
| 227 |
let tid = lid.x;
|
|
@@ -276,7 +226,12 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
|
|
| 276 |
// sweep of its own row, which is why its online state is never merged across rows.
|
| 277 |
{% if qTile > 1 %}
|
| 278 |
let row_first = q_abs_0 / SPARSE_BLOCK;
|
|
|
|
|
|
|
|
|
|
|
|
|
| 279 |
let row_last = mask_row_{{ qTile - 1 }};
|
|
|
|
| 280 |
for (var mask_row = row_first; mask_row <= row_last; mask_row = mask_row + 1u) {
|
| 281 |
{% else %}
|
| 282 |
{
|
|
@@ -364,11 +319,42 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
|
|
| 364 |
{% endfor %}
|
| 365 |
workgroupBarrier();
|
| 366 |
|
| 367 |
-
//
|
| 368 |
-
//
|
| 369 |
-
//
|
| 370 |
-
//
|
| 371 |
-
{% if
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 372 |
let tileCount = min(WG, slot_count - tileBase);
|
| 373 |
{% for j in range(qTile) %}
|
| 374 |
var vSum_{{ j }} = vec4<f32>(0.0);
|
|
|
|
| 9 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 10 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 11 |
}
|
| 12 |
+
{% endmacro %}
|
|
|
|
| 13 |
{{ env.wgsl.resourceDeclarations }}
|
| 14 |
|
| 15 |
// com.microsoft.SparseAttention, attention pass.
|
|
|
|
| 65 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 66 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 67 |
}
|
| 68 |
+
|
| 69 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 70 |
return exp(shifted_value(value, maxValue));
|
| 71 |
}
|
|
|
|
| 79 |
const V_STAGE_KEYS: u32 = 16u;
|
| 80 |
var<workgroup> v_stage: array<vec4<f32>, V_STAGE_KEYS * HEAD_VEC>;
|
| 81 |
{% endif %}
|
| 82 |
+
{% if valueParts > 1 %}
|
| 83 |
+
// Each lane owns one (key partition, output vec4); query streams share V loads.
|
| 84 |
+
var<workgroup> value_partials: array<vec4<f32>, {{ valueParts * headVec * qTile | int }}>;
|
| 85 |
+
{% endif %}
|
| 86 |
// One resolved cache row base per key of the current tile, so the value accumulation
|
| 87 |
// re-reads a base instead of re-walking the column list per head dimension.
|
| 88 |
var<workgroup> key_rows: array<u32, WG>;
|
|
|
|
| 94 |
// Both the subgroup and portable barrier-tree engines return the same merged
|
| 95 |
// pair to every invocation. Repeated merges require a workgroup barrier between
|
| 96 |
// calls before their shared partial storage is reused.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
{% set mdStreamed = mdStreams is defined %}
|
| 98 |
{% set mdStreams = mdStreams if mdStreams is defined else 1 %}
|
| 99 |
{% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
|
|
|
|
| 158 |
return merged;
|
| 159 |
}
|
| 160 |
{% endif %}
|
|
|
|
| 161 |
|
| 162 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
|
|
|
| 163 |
if (params.scale != 0.0) { return params.scale; }
|
| 164 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 165 |
}
|
| 166 |
|
|
|
|
| 167 |
{{ sparse_schedule() }}
|
|
|
|
| 168 |
@compute @workgroup_size(WG, 1, 1)
|
| 169 |
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
|
| 170 |
+
{% if promptTail %}
|
| 171 |
+
// The matrix pass retains every query when this call includes past history.
|
| 172 |
+
if (u32(total_sequence_length[0]) != params.seqLen) { return; }
|
| 173 |
+
{% endif %}
|
| 174 |
+
let tile0 = wg.x * Q_TILE{% if queryOffset > 0 %} + {{ queryOffset }}u{% endif %};
|
| 175 |
let head = wg.y % Q_HEADS;
|
| 176 |
let batch = wg.y / Q_HEADS;
|
| 177 |
let tid = lid.x;
|
|
|
|
| 226 |
// sweep of its own row, which is why its online state is never merged across rows.
|
| 227 |
{% if qTile > 1 %}
|
| 228 |
let row_first = q_abs_0 / SPARSE_BLOCK;
|
| 229 |
+
{% if (seqLen - queryOffset) % qTile != 0 %}
|
| 230 |
+
// Inactive query lanes must not extend traversal past the final CSR row.
|
| 231 |
+
let row_last = (past + min(tile0 + Q_TILE, params.seqLen) - 1u) / SPARSE_BLOCK;
|
| 232 |
+
{% else %}
|
| 233 |
let row_last = mask_row_{{ qTile - 1 }};
|
| 234 |
+
{% endif %}
|
| 235 |
for (var mask_row = row_first; mask_row <= row_last; mask_row = mask_row + 1u) {
|
| 236 |
{% else %}
|
| 237 |
{
|
|
|
|
| 319 |
{% endfor %}
|
| 320 |
workgroupBarrier();
|
| 321 |
|
| 322 |
+
// One value vector serves every query. When head columns underfill the
|
| 323 |
+
// workgroup, spare lanes sum disjoint contiguous key ranges and publish
|
| 324 |
+
// partials; each output owner then folds them in key-range order. The
|
| 325 |
+
// manifest limits this storage and keeps a single-owner fallback.
|
| 326 |
+
{% if valueParts > 1 %}
|
| 327 |
+
let tileCount = min(WG, slot_count - tileBase);
|
| 328 |
+
let valueDim = tid % HEAD_VEC;
|
| 329 |
+
let valuePart = tid / HEAD_VEC;
|
| 330 |
+
if (valuePart < {{ valueParts }}u) {
|
| 331 |
+
let first = tileCount * valuePart / {{ valueParts }}u;
|
| 332 |
+
let last = tileCount * (valuePart + 1u) / {{ valueParts }}u;
|
| 333 |
+
{% for j in range(qTile) %}
|
| 334 |
+
var valueSum_{{ j }} = vec4<f32>(0.0);
|
| 335 |
+
{% endfor %}
|
| 336 |
+
for (var i = first; i < last; i++) {
|
| 337 |
+
let vv = vec4<f32>(present_value[key_rows[i] / 4u + valueDim]);
|
| 338 |
+
{% for j in range(qTile) %}
|
| 339 |
+
valueSum_{{ j }} += probs[{{ j }}u * WG + i] * vv;
|
| 340 |
+
{% endfor %}
|
| 341 |
+
}
|
| 342 |
+
{% for j in range(qTile) %}
|
| 343 |
+
value_partials[{{ j * valueParts * headVec | int }}u + tid] = valueSum_{{ j }};
|
| 344 |
+
{% endfor %}
|
| 345 |
+
}
|
| 346 |
+
workgroupBarrier();
|
| 347 |
+
if (tid < HEAD_VEC) {
|
| 348 |
+
{% for j in range(qTile) %}
|
| 349 |
+
var valueSum_{{ j }} = value_partials[{{ j * valueParts * headVec | int }}u + tid];
|
| 350 |
+
{% for p in range(1, valueParts) %}
|
| 351 |
+
valueSum_{{ j }} += value_partials[{{ (j * valueParts + p) * headVec | int }}u + tid];
|
| 352 |
+
{% endfor %}
|
| 353 |
+
running_out[{{ j }}u * HEAD_VEC + tid] = running_out[{{ j }}u * HEAD_VEC + tid] * correction_{{ j }} + valueSum_{{ j }};
|
| 354 |
+
{% endfor %}
|
| 355 |
+
}
|
| 356 |
+
workgroupBarrier();
|
| 357 |
+
{% elif vStageWorthIt %}
|
| 358 |
let tileCount = min(WG, slot_count - tileBase);
|
| 359 |
{% for j in range(qTile) %}
|
| 360 |
var vSum_{{ j }} = vec4<f32>(0.0);
|
build/webgpu/sparse-kv-append.wgsl.jinja
CHANGED
|
@@ -9,7 +9,7 @@ fn past_sequence_length(batch: u32) -> u32 {
|
|
| 9 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 10 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 11 |
}
|
| 12 |
-
{%
|
| 13 |
{% macro sparse_rotary(interleaved) %}
|
| 14 |
// Which cos/sin entry a component uses, and which member of its rotation pair it is.
|
| 15 |
// The two layouts differ only here: the NeoX split pairs d with d + ROTARY_HALF, and the
|
|
@@ -44,8 +44,12 @@ fn rotary_is_first(d: u32) -> bool {
|
|
| 44 |
fn rotary_value(own: f32, partner: f32, cs: f32, sn: f32, first: bool) -> f32 {
|
| 45 |
return select(own * cs + partner * sn, own * cs - partner * sn, first);
|
| 46 |
}
|
| 47 |
-
{%
|
| 48 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
{{ env.wgsl.resourceDeclarations }}
|
| 50 |
|
| 51 |
// com.microsoft.SparseAttention, KV append pass.
|
|
@@ -72,15 +76,11 @@ const ROTARY_DIM: u32 = {{ rotaryDim }}u;
|
|
| 72 |
|
| 73 |
{{ sparse_schedule() }}
|
| 74 |
{% if usesRotary %}
|
| 75 |
-
|
| 76 |
{{ sparse_rotary(rotaryInterleaved) }}
|
| 77 |
{% endif %}
|
| 78 |
-
|
| 79 |
@compute @workgroup_size(WG, 1, 1)
|
| 80 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 81 |
-
|
| 82 |
-
// per-axis dispatch fold width. Reduces to gid.x when the dispatch does not fold.
|
| 83 |
-
let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
|
| 84 |
let count = params.batchSize * KV_HEADS * params.seqLen * HEAD_DIM;
|
| 85 |
if (index >= count) {
|
| 86 |
return;
|
|
|
|
| 9 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 10 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 11 |
}
|
| 12 |
+
{% endmacro %}
|
| 13 |
{% macro sparse_rotary(interleaved) %}
|
| 14 |
// Which cos/sin entry a component uses, and which member of its rotation pair it is.
|
| 15 |
// The two layouts differ only here: the NeoX split pairs d with d + ROTARY_HALF, and the
|
|
|
|
| 44 |
fn rotary_value(own: f32, partner: f32, cs: f32, sn: f32, first: bool) -> f32 {
|
| 45 |
return select(own * cs + partner * sn, own * cs - partner * sn, first);
|
| 46 |
}
|
| 47 |
+
{% endmacro %}
|
| 48 |
+
{% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
|
| 49 |
+
{% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
|
| 50 |
+
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 51 |
+
// per-axis workgroup fold width.
|
| 52 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
|
| 53 |
{{ env.wgsl.resourceDeclarations }}
|
| 54 |
|
| 55 |
// com.microsoft.SparseAttention, KV append pass.
|
|
|
|
| 76 |
|
| 77 |
{{ sparse_schedule() }}
|
| 78 |
{% if usesRotary %}
|
|
|
|
| 79 |
{{ sparse_rotary(rotaryInterleaved) }}
|
| 80 |
{% endif %}
|
|
|
|
| 81 |
@compute @workgroup_size(WG, 1, 1)
|
| 82 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 83 |
+
{{ flat_index_2d("WG", "index", "") }}
|
|
|
|
|
|
|
| 84 |
let count = params.batchSize * KV_HEADS * params.seqLen * HEAD_DIM;
|
| 85 |
if (index >= count) {
|
| 86 |
return;
|
build/webgpu/sparse-q-rotary.wgsl.jinja
CHANGED
|
@@ -9,7 +9,7 @@ fn past_sequence_length(batch: u32) -> u32 {
|
|
| 9 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 10 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 11 |
}
|
| 12 |
-
{%
|
| 13 |
{% macro sparse_rotary(interleaved) %}
|
| 14 |
// Which cos/sin entry a component uses, and which member of its rotation pair it is.
|
| 15 |
// The two layouts differ only here: the NeoX split pairs d with d + ROTARY_HALF, and the
|
|
@@ -44,8 +44,12 @@ fn rotary_is_first(d: u32) -> bool {
|
|
| 44 |
fn rotary_value(own: f32, partner: f32, cs: f32, sn: f32, first: bool) -> f32 {
|
| 45 |
return select(own * cs + partner * sn, own * cs - partner * sn, first);
|
| 46 |
}
|
| 47 |
-
{%
|
| 48 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 49 |
{{ env.wgsl.resourceDeclarations }}
|
| 50 |
|
| 51 |
// com.microsoft.SparseAttention, query rotary pass.
|
|
@@ -61,14 +65,10 @@ const Q_STRIDE: u32 = {{ packedStride if packedQkv else numHeads * headSize }}u;
|
|
| 61 |
const WG: u32 = {{ appendWorkgroupSize }}u;
|
| 62 |
|
| 63 |
{{ sparse_schedule() }}
|
| 64 |
-
|
| 65 |
{{ sparse_rotary(rotaryInterleaved) }}
|
| 66 |
-
|
| 67 |
@compute @workgroup_size(WG, 1, 1)
|
| 68 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 69 |
-
|
| 70 |
-
// per-axis dispatch fold width. Reduces to gid.x when the dispatch does not fold.
|
| 71 |
-
let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
|
| 72 |
let count = params.batchSize * Q_HEADS * params.seqLen * HEAD_DIM;
|
| 73 |
if (index >= count) {
|
| 74 |
return;
|
|
|
|
| 9 |
let total = u32(key_total_sequence_lengths[batch]);
|
| 10 |
return select(0u, total - params.seqLen, total >= params.seqLen);
|
| 11 |
}
|
| 12 |
+
{% endmacro %}
|
| 13 |
{% macro sparse_rotary(interleaved) %}
|
| 14 |
// Which cos/sin entry a component uses, and which member of its rotation pair it is.
|
| 15 |
// The two layouts differ only here: the NeoX split pairs d with d + ROTARY_HALF, and the
|
|
|
|
| 44 |
fn rotary_value(own: f32, partner: f32, cs: f32, sn: f32, first: bool) -> f32 {
|
| 45 |
return select(own * cs + partner * sn, own * cs - partner * sn, first);
|
| 46 |
}
|
| 47 |
+
{% endmacro %}
|
| 48 |
+
{% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
|
| 49 |
+
{% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
|
| 50 |
+
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 51 |
+
// per-axis workgroup fold width.
|
| 52 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
|
| 53 |
{{ env.wgsl.resourceDeclarations }}
|
| 54 |
|
| 55 |
// com.microsoft.SparseAttention, query rotary pass.
|
|
|
|
| 65 |
const WG: u32 = {{ appendWorkgroupSize }}u;
|
| 66 |
|
| 67 |
{{ sparse_schedule() }}
|
|
|
|
| 68 |
{{ sparse_rotary(rotaryInterleaved) }}
|
|
|
|
| 69 |
@compute @workgroup_size(WG, 1, 1)
|
| 70 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 71 |
+
{{ flat_index_2d("WG", "index", "") }}
|
|
|
|
|
|
|
| 72 |
let count = params.batchSize * Q_HEADS * params.seqLen * HEAD_DIM;
|
| 73 |
if (index >= count) {
|
| 74 |
return;
|
build/webgpu/test.json
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|