Xenova HF Staff commited on
Commit
6d43358
·
verified ·
1 Parent(s): f03331e

sync 6fdf6301e2bb

Browse files
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 + tuning cases
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.2
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
- "sparseWidthBound": "max(256, tunables.WORKGROUP_SIZE)",
 
58
  "sparseQueryTileCap": "max(1, floor((device.limits.maxComputeWorkgroupStorageSize / 4 - sparseWidthBound) / (2 * headSize + 3 * sparseWidthBound)))",
59
- "sparseQueryTileWant": "min(tunables.QUERY_TILE, min(sparseBlockSize, sparseQueryTileCap))",
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
- "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 (2 * sparseQueryTile * headSize + (3 * sparseQueryTile + 1) * sparseAttnWorkgroup) * 4 <= device.limits.maxComputeWorkgroupStorageSize and sparseAttnWorkgroup <= device.limits.maxComputeInvocationsPerWorkgroup and sparseAttnWorkgroup <= device.limits.maxComputeWorkgroupSizeX",
 
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
- "sparseVStageWorthIt": "sparseQueryTiles * batchSize * numHeads <= tunables.V_STAGE_MAX_WORKGROUPS",
92
- "sgmatQueryTiles": "ceilDiv(seqLen, 64)",
93
- "sgmatDirectQuery": "seqLen % 64 == 0",
94
- "sparseSgmatTileN": "64 if (64 * 32 + 64 * 64 + 64 * 2 + 128 * 2) * 4 <= device.limits.maxComputeWorkgroupStorageSize else 32",
 
 
 
95
  "sparseSgmatTileK": "sparseSgmatTileN / 2",
96
- "sparseSgmatLdsBytes": "(64 * sparseSgmatTileK + 64 * sparseSgmatTileN + 64 * 2 + 128 * 2) * 4",
97
- "sparseSgmatGeometryOk": "256 <= device.limits.maxComputeInvocationsPerWorkgroup and 256 <= device.limits.maxComputeWorkgroupSizeX and sgmatQueryTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and batchSize * numHeads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and sparseSgmatLdsBytes <= device.limits.maxComputeWorkgroupStorageSize",
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", "buffer": "read-only-storage", "elementType": "$scalar" },
108
- "new_value": { "arg": "valueT", "buffer": "read-only-storage", "elementType": "$scalar" },
109
- "present_key": { "arg": "pastKeyT", "buffer": "storage", "elementType": "$scalar" },
110
- "present_value": { "arg": "pastValueT", "buffer": "storage", "elementType": "$scalar" },
111
- "key_total_sequence_lengths": {
112
- "arg": "keyTotalSequenceLengthsT",
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", "buffer": "read-only-storage", "elementType": "$scalar" },
125
- "sin_cache": { "arg": "sinCacheT", "buffer": "read-only-storage", "elementType": "$scalar" },
126
- "packed_qkv": { "arg": "queryT", "buffer": "read-only-storage", "elementType": "$scalar" },
127
- "query": { "arg": "queryT", "buffer": "read-only-storage", "elementType": "$scalar" },
128
- "present_key_2": {
129
  "arg": "pastKeyT",
130
  "name": "present_key",
131
  "buffer": "read-only-storage",
132
  "elementType": "$cacheVec"
133
  },
134
- "present_value_2": {
135
  "arg": "pastValueT",
136
  "name": "present_value",
137
  "buffer": "read-only-storage",
138
  "elementType": "$cacheVec"
139
  },
140
- "block_row_indices": { "arg": "blockRowIndicesT", "buffer": "read-only-storage", "elementType": "i32" },
141
- "block_col_indices": { "arg": "blockColIndicesT", "buffer": "read-only-storage", "elementType": "i32" },
142
- "output": { "arg": "outputT", "buffer": "storage", "elementType": "$scalar" },
143
- "params_2": {
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
- "present_key_3": {
153
  "arg": "pastKeyT",
154
  "name": "present_key",
155
  "buffer": "read-only-storage",
156
  "elementType": "$scalar"
157
  },
158
- "present_value_3": {
159
  "arg": "pastValueT",
160
  "name": "present_value",
161
  "buffer": "read-only-storage",
162
  "elementType": "$scalar"
163
  },
164
- "q_rotary_2": { "scratch": "QRotary", "name": "q_rotary", "buffer": "storage", "elementType": "f32" }
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", "present_key_2", "present_value_2", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
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", "present_key_3", "present_value_3", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
221
  "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
222
- "subgroupCollectivesWidth": 32
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", "q_rotary_2", "key_total_sequence_lengths", "total_sequence_length", "params"],
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", "present_key_2", "present_value_2", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
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", "q_rotary_2", "key_total_sequence_lengths", "total_sequence_length", "params"],
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", "present_key_3", "present_value_3", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
305
  "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
306
- "subgroupCollectivesWidth": 32
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", "present_key_2", "present_value_2", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
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", "present_key_3", "present_value_3", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
365
  "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
366
- "subgroupCollectivesWidth": 32
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", "q_rotary_2", "key_total_sequence_lengths", "total_sequence_length", "params"],
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", "present_key_2", "present_value_2", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
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", "q_rotary_2", "key_total_sequence_lengths", "total_sequence_length", "params"],
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", "present_key_3", "present_value_3", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_2"],
449
  "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" },
450
- "subgroupCollectivesWidth": 32
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": "_com_microsoft_sparseattention_webgpu_9e87250",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "nMASpXh23ARMsnca9ti0EsTpERZGnOaFX7OI9QEKmcc=",
11
- "manifest.json": "pGu41bfBMFQF8i3CC/V4Elx5QtkSWNUtxk4Z6+gSFsM=",
12
- "sparse-attention-sgmat.wgsl.jinja": "At5cbiFuunZkmu32JwzR5KDqNs6/D2VfJb89Op7HljM=",
13
- "sparse-attention.wgsl.jinja": "4VYhFW1x7bbvfGHKcFwP/sptJwhPShWMSe383RasMEY=",
14
- "sparse-kv-append.wgsl.jinja": "BHKpoS1Ekd524fT8Ft3A7OBCaCXcCe8VpJzXy5q9lA4=",
15
- "sparse-q-rotary.wgsl.jinja": "1eKB4VVovC4N6CeIJW+S27RUk8VMkDUSYpHp/BXTfRY=",
16
- "test.json": "lURO8tA6ovXlCiXrg3GjYlxvgcn4OdJ/A1EJ93voYq4="
17
  }
18
  },
19
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
20
  "webgpu": {
21
- "manifestSpec": "2.0",
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
- {%- endmacro %}
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
- // `(batch, query tile, query head)` tuple. Selected sparse blocks are traversed
27
- // as 64-query score tiles; their contiguous K/V cache rows load directly as matrix
28
- // fragments. Sequences of complete query tiles also load Q directly; sequences
29
- // with a partial final tile stage Q with zero padding. The manifest bounds tile widths by workgroup storage.
30
  //
31
- // The first sweep folds score tiles into per-row `(max, denominator)` softmax
32
- // statistics. The second recomputes the same scores, applies the completed
33
- // normalization, and accumulates P.V into result fragments. Causal bounds,
34
- // duplicate CSR columns, dense rows, and all-masked rows are handled explicitly.
 
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 = 64u;
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; the score epilogues alias it as the
82
- // fragment-store scratch (8 subgroups x scoreColBlocks banks x 64 elements).
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
- {% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
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
- // A sequence with a partial query tile retains zero-padded staging.
118
  for (var k_base = 0u; k_base < HEAD_DIM; k_base += TILE_K) {
119
- {% if not sgmatDirectQuery %}
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 }}, {{ query_matrix_offset(rb) }}, {{ queryMatrixStride }}
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 sgmatDirectQuery %}
160
  workgroupBarrier();
161
  {% endif %}
162
  }
163
  {% endmacro %}
164
 
165
- {% macro sweep(phase) %}
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
- {{ score_tile() }}
 
 
 
 
 
 
 
 
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
- {% if phase == "stats" %}
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
- {% endif %}
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
- {% if phase == "stats" %}
 
 
 
 
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
- {% else %}
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
- {% if phase == "stats" %}
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
- {% endif %}
274
  {% endfor %}
275
- {% if phase == "stats" %}
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
- {% else %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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(256, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
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 += 256u) {
371
  row_m[r] = -FLT_MAX;
372
  row_d[r] = 0.0;
373
  }
374
  workgroupBarrier();
375
 
376
- {{ sweep("stats") }}
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 += 256u) {
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
- {%- endmacro %}
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
- let tile0 = wg.x * Q_TILE;
 
 
 
 
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
- // running_out[j][d] is owned by the same thread across every tile (tid = d mod WG),
368
- // so this rescale-and-accumulate needs no further synchronization. One value vector
369
- // serves every query, which is the other half of the tile's reuse; a key outside a
370
- // query's causal bound carries prob 0 and is multiplied away.
371
- {% if vStageWorthIt %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- {%- 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,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
- {%- endmacro %}
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
- // 2D-folded flat index: gid.y carries the high bits past the
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
- {%- 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,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
- {%- endmacro %}
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
- // 2D-folded flat index: gid.y carries the high bits past the
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