sync 6fdf6301e2bb
Browse files- README.md +16 -28
- build/webgpu/attention-rank4-tiled.wgsl.jinja +1 -25
- build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja +5 -34
- build/webgpu/attn-flash-decode-splitk.wgsl.jinja +21 -109
- build/webgpu/attn-flash-online.wgsl.jinja +15 -95
- build/webgpu/attn-flash-prefill-cluster.wgsl.jinja +112 -100
- build/webgpu/attn-flash-q32-broadcast.wgsl.jinja +47 -8
- build/webgpu/attn-materialized-rowstats-combine-f32.wgsl.jinja +8 -2
- build/webgpu/attn-materialized-sgmat-f32.wgsl.jinja +37 -64
- build/webgpu/attn-online-scalar.wgsl.jinja +53 -28
- build/webgpu/bench.json +603 -6
- build/webgpu/gqa-attention.wgsl.jinja +55 -38
- build/webgpu/gqa-present.wgsl.jinja +70 -42
- build/webgpu/gqa-qprep.wgsl.jinja +7 -5
- build/webgpu/manifest.json +0 -0
- build/webgpu/metadata.json +27 -37
- build/webgpu/test.json +0 -0
README.md
CHANGED
|
@@ -87,19 +87,16 @@ One implementation is selected per call from the device capabilities, the reques
|
|
| 87 |
- `window_shift_materialized_sgmat_f32` — Materializes causal chunked-prefill attention with a windowed cache. The shift pass compacts surviving rows and appends the chunk; score and apply passes enforce both the sliding-window floor and causal bound.
|
| 88 |
- `share_append_materialized_sgmat_f32` — Materializes causal chunked-prefill attention while updating a shared-capacity cache in place. Score and apply passes bound every causal tile by the live length from `seqlens_k`, not the buffer capacity, so right-padded batches remain left-aligned.
|
| 89 |
- `new_kv_share_append_split` — Retains the existing cache and appends new key/value rows in separate passes before portable attention. It avoids rebuilding unchanged cache positions when the past and present allocations share capacity.
|
|
|
|
| 90 |
- `qkv_present_tiled_nosg` — Portable tiled prefill route that computes attention online and writes the present cache separately. It is used when the flash shape is valid but no suitable subgroup route is admissible.
|
|
|
|
|
|
|
| 91 |
- `quant_int8_decode_splitk` — Appends int8-quantized key/value rows, partitions cached decode across the key axis, and merges partial online-softmax results. Used when one workgroup per head exposes too little independent work.
|
| 92 |
- `qkv_present_flash_splitk` — Partitions direct Q/K/V attention across the key axis, merges partial online-softmax results, and writes the present cache separately. It serves short-query shapes needing more key-axis parallelism.
|
| 93 |
- `qkv_present_flash_cluster` — Computes clustered online-softmax prefill directly from Q/K/V and writes the present cache separately. The family provides subgroup and portable reductions for the same tiled algorithm.
|
| 94 |
-
- `quant_int8_decode_splitk_nosg` — Appends int8-quantized key/value rows, partitions cached decode across the key axis, and merges partial online-softmax results. Used when one workgroup per head exposes too little independent work.
|
| 95 |
-
- `qkv_present_flash_splitk_nosg` — Partitions direct Q/K/V attention across the key axis, merges partial online-softmax results, and writes the present cache separately. It serves short-query shapes needing more key-axis parallelism.
|
| 96 |
-
- `qkv_present_flash_cluster_nosg` — Computes clustered online-softmax prefill directly from Q/K/V and writes the present cache separately. The family provides subgroup and portable reductions for the same tiled algorithm.
|
| 97 |
- `past_kv_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 98 |
- `new_kv_past_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 99 |
- `window_shift_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 100 |
-
- `past_kv_decode_splitk_nosg` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 101 |
-
- `new_kv_past_decode_splitk_nosg` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 102 |
-
- `window_shift_decode_splitk_nosg` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 103 |
- `past_kv_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 104 |
- `new_kv_past_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 105 |
- `past_kv_rotary_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
|
@@ -107,21 +104,14 @@ One implementation is selected per call from the device capabilities, the reques
|
|
| 107 |
- `past_kv_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 108 |
- `past_kv_bias_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 109 |
- `window_shift_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 110 |
-
- `
|
| 111 |
-
- `
|
| 112 |
-
- `
|
| 113 |
-
- `past_kv_softcap_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 114 |
-
- `past_kv_headsink_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 115 |
-
- `past_kv_bias_headsink_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 116 |
-
- `window_shift_flash_prefill_nosg` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 117 |
- `quant_int8_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
|
| 118 |
- `quant_int4_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
|
| 119 |
-
- `quant_int8_flash_prefill_nosg` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
|
| 120 |
-
- `quant_int4_flash_prefill_nosg` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
|
| 121 |
- `share_append_split_decode_splitk` — Retains shared-capacity cache rows, appends new rows separately, then partitions decode across the key axis. It avoids rebuilding unchanged cache positions while exposing split-key parallelism.
|
| 122 |
-
- `share_append_split_decode_splitk_nosg` — Retains shared-capacity cache rows, appends new rows separately, then partitions decode across the key axis. It avoids rebuilding unchanged cache positions while exposing split-key parallelism.
|
| 123 |
- `share_append_split_flash_prefill` — Retains shared-capacity cache rows, appends new rows separately, then applies clustered causal prefill. It avoids rebuilding unchanged cache positions and includes subgroup and portable forms.
|
| 124 |
-
- `
|
| 125 |
|
| 126 |
## Device requirements
|
| 127 |
|
|
@@ -132,7 +122,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
|
|
| 132 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 133 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 134 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 135 |
-
- [`bench.json`](build/webgpu/bench.json) — benchmark
|
| 136 |
- [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
|
| 137 |
- [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
|
| 138 |
- [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
|
|
@@ -149,7 +139,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
|
|
| 149 |
## Use with `@huggingface/kernels`
|
| 150 |
|
| 151 |
```sh
|
| 152 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 153 |
```
|
| 154 |
|
| 155 |
Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
|
|
@@ -170,18 +160,16 @@ import { getKernel } from "@huggingface/kernels";
|
|
| 170 |
const kernel = await getKernel("webgpu-kernels/com.microsoft.GroupQueryAttention", { version: 1 });
|
| 171 |
// Explicit destinations request optional results or supply metadata that cannot be inferred.
|
| 172 |
const { outputT, presentKeyT, presentValueT } = await kernel({
|
| 173 |
-
queryT: { data: queryTData, shape: [
|
| 174 |
-
keyT: { data: keyTData, shape: [
|
| 175 |
-
valueT: { data: valueTData, shape: [
|
| 176 |
-
|
| 177 |
-
pastValueT: { data: pastValueTData, shape: [2, 1, 8, 8] },
|
| 178 |
-
seqlensKT: { data: seqlensKTData, shape: [2] },
|
| 179 |
totalSequenceLengthT: { data: totalSequenceLengthTData, shape: [1] },
|
| 180 |
}, {
|
| 181 |
-
attrs: { num_heads:
|
| 182 |
outputs: {
|
| 183 |
-
presentKeyT: { shape: [
|
| 184 |
-
presentValueT: { shape: [
|
| 185 |
},
|
| 186 |
});
|
| 187 |
```
|
|
|
|
| 87 |
- `window_shift_materialized_sgmat_f32` — Materializes causal chunked-prefill attention with a windowed cache. The shift pass compacts surviving rows and appends the chunk; score and apply passes enforce both the sliding-window floor and causal bound.
|
| 88 |
- `share_append_materialized_sgmat_f32` — Materializes causal chunked-prefill attention while updating a shared-capacity cache in place. Score and apply passes bound every causal tile by the live length from `seqlens_k`, not the buffer capacity, so right-padded batches remain left-aligned.
|
| 89 |
- `new_kv_share_append_split` — Retains the existing cache and appends new key/value rows in separate passes before portable attention. It avoids rebuilding unchanged cache positions when the past and present allocations share capacity.
|
| 90 |
+
- `new_kv_share_append_bidirectional_split` — Bidirectional (`causal = 0`) attention over a shared-capacity cache: the new key/value rows are appended exactly as in the causal sibling, then every query attends to the whole live range `[0, seqlens_k[b] + 1)` instead of stopping at its own position. ONNX Runtime forbids `causal = 0` with a local window, so the window floor never applies here.
|
| 91 |
- `qkv_present_tiled_nosg` — Portable tiled prefill route that computes attention online and writes the present cache separately. It is used when the flash shape is valid but no suitable subgroup route is admissible.
|
| 92 |
+
- `qkv_present_flash_causal` — Causal prefill without a past cache and without a sliding window (`causal=1`, `local_window_size=-1`): the spelling onnxruntime's GroupQueryAttention tests use for a first prompt. The twin of `qkv_present_flash` with the causal ceiling enabled; the key sequence equals the query sequence, so the upper-left ceiling is the right one.
|
| 93 |
+
- `qkv_present_causal` — Causal prefill without a past cache and without a sliding window (`causal=1`, `local_window_size=-1`): the spelling onnxruntime's GroupQueryAttention tests use for a first prompt. The twin of `qkv_present` with the causal ceiling enabled; the key sequence equals the query sequence, so the upper-left ceiling is the right one.
|
| 94 |
- `quant_int8_decode_splitk` — Appends int8-quantized key/value rows, partitions cached decode across the key axis, and merges partial online-softmax results. Used when one workgroup per head exposes too little independent work.
|
| 95 |
- `qkv_present_flash_splitk` — Partitions direct Q/K/V attention across the key axis, merges partial online-softmax results, and writes the present cache separately. It serves short-query shapes needing more key-axis parallelism.
|
| 96 |
- `qkv_present_flash_cluster` — Computes clustered online-softmax prefill directly from Q/K/V and writes the present cache separately. The family provides subgroup and portable reductions for the same tiled algorithm.
|
|
|
|
|
|
|
|
|
|
| 97 |
- `past_kv_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 98 |
- `new_kv_past_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
| 99 |
- `window_shift_decode_splitk` — Updates the float cache in copy, merge, or window-shift mode, then partitions cached decode across the key axis and combines partial online-softmax results. The family includes subgroup and portable forms.
|
|
|
|
|
|
|
|
|
|
| 100 |
- `past_kv_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 101 |
- `new_kv_past_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 102 |
- `past_kv_rotary_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
|
|
|
| 104 |
- `past_kv_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 105 |
- `past_kv_bias_headsink_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 106 |
- `window_shift_flash_prefill` — Updates the float cache in the selected mode and applies clustered online-softmax causal prefill with the configured cache, mask, rotary, softcap, bias, or head-sink features.
|
| 107 |
+
- `window_shift_rotary_append` — Windowed cache with rotary embeddings. The shift pass slides surviving rows down by the eviction count and clears what left the window; the append pass writes the new key/value rows and rotates K at its ABSOLUTE position (window origin plus row, not the cache row); attention rotates Q at the same absolute position and masks cache-relative.
|
| 108 |
+
- `window_shift_rotary_headsink_append` — Windowed rotary cache with a per-head softmax sink (the gpt-oss attention layer). The sink contributes `exp(sink - m)` to the denominator once per query row and nothing to the numerator.
|
| 109 |
+
- `window_shift_rotary_bias_append` — Windowed rotary cache with an additive attention bias. The bias row is indexed by ABSOLUTE key position with the tensor's own last dimension as the stride (ONNX spells it `total_sequence_length`, which keeps counting past the cache capacity), so resident cache row j reads bias column `origin + j`.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 110 |
- `quant_int8_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
|
| 111 |
- `quant_int4_flash_prefill` — Builds an int8 or int4 cache and applies clustered online-softmax prefill directly from the packed values. The family supplies subgroup and portable reductions.
|
|
|
|
|
|
|
| 112 |
- `share_append_split_decode_splitk` — Retains shared-capacity cache rows, appends new rows separately, then partitions decode across the key axis. It avoids rebuilding unchanged cache positions while exposing split-key parallelism.
|
|
|
|
| 113 |
- `share_append_split_flash_prefill` — Retains shared-capacity cache rows, appends new rows separately, then applies clustered causal prefill. It avoids rebuilding unchanged cache positions and includes subgroup and portable forms.
|
| 114 |
+
- `share_append_bidirectional_flash_prefill` — Retains and appends a shared-capacity cache, then applies clustered bidirectional prefill to the full live key range.
|
| 115 |
|
| 116 |
## Device requirements
|
| 117 |
|
|
|
|
| 122 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 123 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 124 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 125 |
+
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
|
| 126 |
- [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
|
| 127 |
- [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
|
| 128 |
- [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
|
|
|
|
| 139 |
## Use with `@huggingface/kernels`
|
| 140 |
|
| 141 |
```sh
|
| 142 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 143 |
```
|
| 144 |
|
| 145 |
Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
|
|
|
|
| 160 |
const kernel = await getKernel("webgpu-kernels/com.microsoft.GroupQueryAttention", { version: 1 });
|
| 161 |
// Explicit destinations request optional results or supply metadata that cannot be inferred.
|
| 162 |
const { outputT, presentKeyT, presentValueT } = await kernel({
|
| 163 |
+
queryT: { data: queryTData, shape: [1, 2, 8] },
|
| 164 |
+
keyT: { data: keyTData, shape: [1, 2, 8] },
|
| 165 |
+
valueT: { data: valueTData, shape: [1, 2, 8] },
|
| 166 |
+
seqlensKT: { data: seqlensKTData, shape: [1] },
|
|
|
|
|
|
|
| 167 |
totalSequenceLengthT: { data: totalSequenceLengthTData, shape: [1] },
|
| 168 |
}, {
|
| 169 |
+
attrs: { num_heads: 1, kv_num_heads: 1 },
|
| 170 |
outputs: {
|
| 171 |
+
presentKeyT: { shape: [1, 1, 2, 8], dtype: "float32" },
|
| 172 |
+
presentValueT: { shape: [1, 1, 2, 8], dtype: "float32" },
|
| 173 |
},
|
| 174 |
});
|
| 175 |
```
|
build/webgpu/attention-rank4-tiled.wgsl.jinja
CHANGED
|
@@ -1,4 +1,4 @@
|
|
| 1 |
-
{% set
|
| 2 |
{{ env.wgsl.resourceDeclarations }}
|
| 3 |
|
| 4 |
// Thread-per-query online-softmax flash fallback for rank-4 (BNSH) attention with
|
|
@@ -16,13 +16,11 @@ const BLOCK_M: u32 = {{ blockM }}u;
|
|
| 16 |
var<workgroup> acc: array<f32, {{ vHeadCap * blockM }}>;
|
| 17 |
|
| 18 |
{% set ATTN_SCALE_DIM = "params.headSize" %}
|
| 19 |
-
{% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
|
| 20 |
fn scale_value() -> f32 {
|
| 21 |
if (params.scale != 0.0) { return params.scale; }
|
| 22 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 23 |
}
|
| 24 |
|
| 25 |
-
|
| 26 |
fn kv_head(q_head: u32) -> u32 {
|
| 27 |
return q_head / (params.qHeads / params.kvHeads);
|
| 28 |
}
|
|
@@ -51,14 +49,8 @@ fn main(
|
|
| 51 |
if (qs >= params.qSeq) { return; }
|
| 52 |
let kh = kv_head(qh);
|
| 53 |
|
| 54 |
-
{% if layout == "bsh" %}
|
| 55 |
// Token-major packed QKV: row = (batch*seq + token)*HIDDEN + head*head_dim.
|
| 56 |
let qBase = (batch * params.qSeq + qs) * params.qHidden + qh * params.headSize;
|
| 57 |
-
{% else %}
|
| 58 |
-
let qBase = ((batch * params.qHeads + qh) * params.qSeq + qs) * params.headSize;
|
| 59 |
-
let kBase = (batch * params.kvHeads + kh) * params.kvSeq;
|
| 60 |
-
let vBase = (batch * params.kvHeads + kh) * params.kvSeq;
|
| 61 |
-
{% endif %}
|
| 62 |
let scale = scale_value();
|
| 63 |
|
| 64 |
for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
|
|
@@ -76,21 +68,13 @@ fn main(
|
|
| 76 |
|
| 77 |
for (var ks: u32 = 0u; ks < maxKj; ks = ks + 1u) {
|
| 78 |
var masked = false;
|
| 79 |
-
{% if hasMask and maskIsBool %}
|
| 80 |
-
let mIdxB = batch * params.maskBatchStride + qh * params.maskHeadStride + qs * params.maskSeqStride + ks;
|
| 81 |
-
masked = attn_mask[mIdxB] == 0u;
|
| 82 |
-
{% endif %}
|
| 83 |
// A rejected bool-mask key has zero softmax mass. Skipping it is safe here:
|
| 84 |
// this thread-per-query kernel has no barriers inside the key loop.
|
| 85 |
if (masked) { continue; }
|
| 86 |
// Adjacent query threads load the same K/V row for this key.
|
| 87 |
var score: f32 = -3.4028234663852886e38;
|
| 88 |
if (!masked) {
|
| 89 |
-
{% if layout == "bsh" %}
|
| 90 |
let kRow = (batch * params.kvSeq + ks) * params.kvHidden + kh * params.headSize;
|
| 91 |
-
{% else %}
|
| 92 |
-
let kRow = (kBase + ks) * params.headSize;
|
| 93 |
-
{% endif %}
|
| 94 |
var dot: f32 = 0.0;
|
| 95 |
for (var d: u32 = 0u; d < params.headSize; d = d + 1u) {
|
| 96 |
dot = dot + f32(q[qBase + d]) * f32(k[kRow + d]);
|
|
@@ -110,22 +94,14 @@ fn main(
|
|
| 110 |
let weight = exp(score - next_max);
|
| 111 |
running_max = next_max;
|
| 112 |
running_denom = running_denom * prev_scale + weight;
|
| 113 |
-
{% if layout == "bsh" %}
|
| 114 |
let vRow = (batch * params.kvSeq + ks) * params.vHidden + kh * params.vHeadSize;
|
| 115 |
-
{% else %}
|
| 116 |
-
let vRow = (vBase + ks) * params.vHeadSize;
|
| 117 |
-
{% endif %}
|
| 118 |
for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
|
| 119 |
acc[d * BLOCK_M + tid] = acc[d * BLOCK_M + tid] * prev_scale + weight * f32(v[vRow + d]);
|
| 120 |
}
|
| 121 |
}
|
| 122 |
|
| 123 |
let inv_denom = select(0.0, 1.0 / running_denom, running_denom > 0.0);
|
| 124 |
-
{% if layout == "bsh" %}
|
| 125 |
let yBase = (batch * params.qSeq + qs) * (params.qHeads * params.vHeadSize) + qh * params.vHeadSize;
|
| 126 |
-
{% else %}
|
| 127 |
-
let yBase = ((batch * params.qHeads + qh) * params.qSeq + qs) * params.vHeadSize;
|
| 128 |
-
{% endif %}
|
| 129 |
for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
|
| 130 |
y[yBase + d] = {{ scalar }}(acc[d * BLOCK_M + tid] * inv_denom);
|
| 131 |
}
|
|
|
|
| 1 |
+
{% set maskIsBool = false %}
|
| 2 |
{{ env.wgsl.resourceDeclarations }}
|
| 3 |
|
| 4 |
// Thread-per-query online-softmax flash fallback for rank-4 (BNSH) attention with
|
|
|
|
| 16 |
var<workgroup> acc: array<f32, {{ vHeadCap * blockM }}>;
|
| 17 |
|
| 18 |
{% set ATTN_SCALE_DIM = "params.headSize" %}
|
|
|
|
| 19 |
fn scale_value() -> f32 {
|
| 20 |
if (params.scale != 0.0) { return params.scale; }
|
| 21 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 22 |
}
|
| 23 |
|
|
|
|
| 24 |
fn kv_head(q_head: u32) -> u32 {
|
| 25 |
return q_head / (params.qHeads / params.kvHeads);
|
| 26 |
}
|
|
|
|
| 49 |
if (qs >= params.qSeq) { return; }
|
| 50 |
let kh = kv_head(qh);
|
| 51 |
|
|
|
|
| 52 |
// Token-major packed QKV: row = (batch*seq + token)*HIDDEN + head*head_dim.
|
| 53 |
let qBase = (batch * params.qSeq + qs) * params.qHidden + qh * params.headSize;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 54 |
let scale = scale_value();
|
| 55 |
|
| 56 |
for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
|
|
|
|
| 68 |
|
| 69 |
for (var ks: u32 = 0u; ks < maxKj; ks = ks + 1u) {
|
| 70 |
var masked = false;
|
|
|
|
|
|
|
|
|
|
|
|
|
| 71 |
// A rejected bool-mask key has zero softmax mass. Skipping it is safe here:
|
| 72 |
// this thread-per-query kernel has no barriers inside the key loop.
|
| 73 |
if (masked) { continue; }
|
| 74 |
// Adjacent query threads load the same K/V row for this key.
|
| 75 |
var score: f32 = -3.4028234663852886e38;
|
| 76 |
if (!masked) {
|
|
|
|
| 77 |
let kRow = (batch * params.kvSeq + ks) * params.kvHidden + kh * params.headSize;
|
|
|
|
|
|
|
|
|
|
| 78 |
var dot: f32 = 0.0;
|
| 79 |
for (var d: u32 = 0u; d < params.headSize; d = d + 1u) {
|
| 80 |
dot = dot + f32(q[qBase + d]) * f32(k[kRow + d]);
|
|
|
|
| 94 |
let weight = exp(score - next_max);
|
| 95 |
running_max = next_max;
|
| 96 |
running_denom = running_denom * prev_scale + weight;
|
|
|
|
| 97 |
let vRow = (batch * params.kvSeq + ks) * params.vHidden + kh * params.vHeadSize;
|
|
|
|
|
|
|
|
|
|
| 98 |
for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
|
| 99 |
acc[d * BLOCK_M + tid] = acc[d * BLOCK_M + tid] * prev_scale + weight * f32(v[vRow + d]);
|
| 100 |
}
|
| 101 |
}
|
| 102 |
|
| 103 |
let inv_denom = select(0.0, 1.0 / running_denom, running_denom > 0.0);
|
|
|
|
| 104 |
let yBase = (batch * params.qSeq + qs) * (params.qHeads * params.vHeadSize) + qh * params.vHeadSize;
|
|
|
|
|
|
|
|
|
|
| 105 |
for (var d: u32 = 0u; d < params.vHeadSize; d = d + 1u) {
|
| 106 |
y[yBase + d] = {{ scalar }}(acc[d * BLOCK_M + tid] * inv_denom);
|
| 107 |
}
|
build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja
CHANGED
|
@@ -1,9 +1,6 @@
|
|
| 1 |
-
{%
|
| 2 |
-
|
| 3 |
-
{%
|
| 4 |
-
{% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
|
| 5 |
-
{% if hasGate is not defined %}{% set hasGate = false %}{% endif %}
|
| 6 |
-
{% set qSeq = qSeq | default(0) %}
|
| 7 |
{{ env.wgsl.resourceDeclarations }}
|
| 8 |
|
| 9 |
// Split-K flash decode, pass 2 of 2. Combines the per-split un-normalized online
|
|
@@ -25,9 +22,6 @@ enable f16;
|
|
| 25 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
| 26 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 27 |
const NUM_SPLITS: u32 = {{ numSplits }}u;
|
| 28 |
-
{% if splitQueries %}
|
| 29 |
-
const Q_SEQ: u32 = {{ qSeq }}u;
|
| 30 |
-
{% endif %}
|
| 31 |
{% if layout == "bsh" %}
|
| 32 |
const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u;
|
| 33 |
{% endif %}
|
|
@@ -57,6 +51,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 57 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 58 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 59 |
}
|
|
|
|
| 60 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 61 |
return exp(shifted_value(value, maxValue));
|
| 62 |
}
|
|
@@ -65,21 +60,14 @@ fn main(
|
|
| 65 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 66 |
@builtin(local_invocation_id) lid: vec3<u32>
|
| 67 |
) {
|
| 68 |
-
{% if splitQueries %}
|
| 69 |
-
let queryToken = wg.x;
|
| 70 |
-
{% endif %}
|
| 71 |
let h = wg.y;
|
| 72 |
let b = wg.z;
|
| 73 |
let d4 = lid.x;
|
| 74 |
-
if (h >= Q_HEADS || d4 >= HEAD_DIM_V4
|
| 75 |
return;
|
| 76 |
}
|
| 77 |
|
| 78 |
-
{% if splitQueries %}
|
| 79 |
-
let mdBase = ((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS;
|
| 80 |
-
{% else %}
|
| 81 |
let mdBase = (b * Q_HEADS + h) * NUM_SPLITS;
|
| 82 |
-
{% endif %}
|
| 83 |
var globalMax = -FLT_MAX;
|
| 84 |
for (var s: u32 = 0u; s < NUM_SPLITS; s = s + 1u) {
|
| 85 |
globalMax = max(globalMax, partial_stats[mdBase + s].x);
|
|
@@ -90,28 +78,16 @@ fn main(
|
|
| 90 |
let stats = partial_stats[mdBase + s];
|
| 91 |
let w = exp_shift(stats.x, globalMax);
|
| 92 |
globalDenom = globalDenom + stats.y * w;
|
| 93 |
-
{% if splitQueries %}
|
| 94 |
-
let pBase = (((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + s) * HEAD_DIM_V4;
|
| 95 |
-
{% else %}
|
| 96 |
let pBase = ((b * Q_HEADS + h) * NUM_SPLITS + s) * HEAD_DIM_V4;
|
| 97 |
-
{% endif %}
|
| 98 |
outV = outV + partial_out[pBase + d4] * w;
|
| 99 |
}
|
| 100 |
|
| 101 |
{% if layout == "bsh" %}
|
| 102 |
-
{% if splitQueries %}
|
| 103 |
-
let qBaseV4 = (b * Q_SEQ + queryToken) * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
|
| 104 |
-
{% else %}
|
| 105 |
let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
|
| 106 |
-
{% endif %}
|
| 107 |
{% elif layout == "layer_cache" %}
|
| 108 |
let qBaseV4 = h * HEAD_DIM_V4;
|
| 109 |
-
{% else %}
|
| 110 |
-
{% if splitQueries %}
|
| 111 |
-
let qBaseV4 = ((b * Q_HEADS + h) * Q_SEQ + queryToken) * HEAD_DIM_V4;
|
| 112 |
{% else %}
|
| 113 |
let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
|
| 114 |
-
{% endif %}
|
| 115 |
{% endif %}
|
| 116 |
// A split that covers no keys writes vec2(-FLT_MAX, 0), so when every split is empty —
|
| 117 |
// a query whose window admits nothing — globalMax stays -FLT_MAX, exp_shift(x, x) is 1,
|
|
@@ -122,11 +98,6 @@ fn main(
|
|
| 122 |
// V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
|
| 123 |
let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
|
| 124 |
outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
|
| 125 |
-
{% endif %}
|
| 126 |
-
{% if hasGate %}
|
| 127 |
-
// The gated route multiplies the normalized attention output elementwise by its gate.
|
| 128 |
-
let gateV = vec4<f32>(gate[qBaseV4 + d4]);
|
| 129 |
-
outValue = outValue * (vec4<f32>(1.0) / (vec4<f32>(1.0) + exp(-gateV)));
|
| 130 |
{% endif %}
|
| 131 |
output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
|
| 132 |
}
|
|
|
|
| 1 |
+
{% set qHiddenV4 = qHiddenV4 | default(0) %}
|
| 2 |
+
{% set qHidden = qHidden | default(0) %}
|
| 3 |
+
{% set hasBias = hasBias is defined and hasBias %}
|
|
|
|
|
|
|
|
|
|
| 4 |
{{ env.wgsl.resourceDeclarations }}
|
| 5 |
|
| 6 |
// Split-K flash decode, pass 2 of 2. Combines the per-split un-normalized online
|
|
|
|
| 22 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
| 23 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 24 |
const NUM_SPLITS: u32 = {{ numSplits }}u;
|
|
|
|
|
|
|
|
|
|
| 25 |
{% if layout == "bsh" %}
|
| 26 |
const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u;
|
| 27 |
{% endif %}
|
|
|
|
| 51 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 52 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 53 |
}
|
| 54 |
+
|
| 55 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 56 |
return exp(shifted_value(value, maxValue));
|
| 57 |
}
|
|
|
|
| 60 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 61 |
@builtin(local_invocation_id) lid: vec3<u32>
|
| 62 |
) {
|
|
|
|
|
|
|
|
|
|
| 63 |
let h = wg.y;
|
| 64 |
let b = wg.z;
|
| 65 |
let d4 = lid.x;
|
| 66 |
+
if (h >= Q_HEADS || d4 >= HEAD_DIM_V4) {
|
| 67 |
return;
|
| 68 |
}
|
| 69 |
|
|
|
|
|
|
|
|
|
|
| 70 |
let mdBase = (b * Q_HEADS + h) * NUM_SPLITS;
|
|
|
|
| 71 |
var globalMax = -FLT_MAX;
|
| 72 |
for (var s: u32 = 0u; s < NUM_SPLITS; s = s + 1u) {
|
| 73 |
globalMax = max(globalMax, partial_stats[mdBase + s].x);
|
|
|
|
| 78 |
let stats = partial_stats[mdBase + s];
|
| 79 |
let w = exp_shift(stats.x, globalMax);
|
| 80 |
globalDenom = globalDenom + stats.y * w;
|
|
|
|
|
|
|
|
|
|
| 81 |
let pBase = ((b * Q_HEADS + h) * NUM_SPLITS + s) * HEAD_DIM_V4;
|
|
|
|
| 82 |
outV = outV + partial_out[pBase + d4] * w;
|
| 83 |
}
|
| 84 |
|
| 85 |
{% if layout == "bsh" %}
|
|
|
|
|
|
|
|
|
|
| 86 |
let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
|
|
|
|
| 87 |
{% elif layout == "layer_cache" %}
|
| 88 |
let qBaseV4 = h * HEAD_DIM_V4;
|
|
|
|
|
|
|
|
|
|
| 89 |
{% else %}
|
| 90 |
let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
|
|
|
|
| 91 |
{% endif %}
|
| 92 |
// A split that covers no keys writes vec2(-FLT_MAX, 0), so when every split is empty —
|
| 93 |
// a query whose window admits nothing — globalMax stays -FLT_MAX, exp_shift(x, x) is 1,
|
|
|
|
| 98 |
// V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
|
| 99 |
let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
|
| 100 |
outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
{% endif %}
|
| 102 |
output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
|
| 103 |
}
|
build/webgpu/attn-flash-decode-splitk.wgsl.jinja
CHANGED
|
@@ -1,21 +1,22 @@
|
|
| 1 |
{% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
|
| 2 |
-
{% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
|
| 3 |
{% if quantizedCache is not defined %}{% set quantizedCache = false %}{% endif %}
|
| 4 |
{% if cacheSeqlens is not defined %}{% set cacheSeqlens = false %}{% endif %}
|
| 5 |
{% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
|
| 6 |
-
{%
|
| 7 |
-
{% set
|
| 8 |
-
{% set
|
| 9 |
-
{% set scale = scale | default("0.0") %}
|
| 10 |
{% if hasRotary is not defined %}{% set hasRotary = false %}{% endif %}
|
| 11 |
-
{% set
|
|
|
|
|
|
|
| 12 |
{% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
|
| 13 |
{% if useSubgroups %}
|
| 14 |
enable subgroups;
|
| 15 |
{% endif %}
|
| 16 |
{{ env.wgsl.resourceDeclarations }}
|
| 17 |
|
| 18 |
-
// Split-K flash attention
|
|
|
|
| 19 |
// handles decode and short-query, long-context prefill inputs.
|
| 20 |
//
|
| 21 |
// The non-split flash decode launches only `batch * numHeads` workgroups, each
|
|
@@ -50,9 +51,6 @@ const ATTN_SCALE: f32 = {{ scale }};
|
|
| 50 |
{% endif %}
|
| 51 |
const WG: u32 = {{ splitKWorkgroupSize }}u;
|
| 52 |
const NUM_SPLITS: u32 = {{ numSplits }}u;
|
| 53 |
-
{% if splitQueries %}
|
| 54 |
-
const Q_SEQ: u32 = {{ qSeq }}u;
|
| 55 |
-
{% endif %}
|
| 56 |
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 57 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
| 58 |
// accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
|
|
@@ -72,6 +70,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 72 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 73 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 74 |
}
|
|
|
|
| 75 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 76 |
return exp(shifted_value(value, maxValue));
|
| 77 |
}
|
|
@@ -79,7 +78,7 @@ fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
|
| 79 |
var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
|
| 80 |
var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
|
| 81 |
var<workgroup> probs: array<f32, WG>;
|
| 82 |
-
{% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) %}
|
| 83 |
{% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
|
| 84 |
{% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
|
| 85 |
{% if coopQk %}
|
|
@@ -148,41 +147,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
|
|
| 148 |
return combinedMD;
|
| 149 |
}
|
| 150 |
{% else %}
|
| 151 |
-
{% set
|
| 152 |
-
{% set mdStreams = mdStreams if mdStreams is defined else 1 %}
|
| 153 |
-
{% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
|
| 154 |
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 155 |
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
| 156 |
-
{% if mdStreamed %}
|
| 157 |
-
|
| 158 |
-
// In-place fold of {{ mdStreams }} streams. Input partials occupy
|
| 159 |
-
// partialM/partialD; stream s returns its merged pair in slot s * WG.
|
| 160 |
-
fn combine_partials_streams(lidx: u32) {
|
| 161 |
-
workgroupBarrier();
|
| 162 |
-
var stride = WG / 2u;
|
| 163 |
-
loop {
|
| 164 |
-
if (stride == 0u) {
|
| 165 |
-
break;
|
| 166 |
-
}
|
| 167 |
-
if (lidx < stride) {
|
| 168 |
-
{% for s in range(mdStreams) %}
|
| 169 |
-
{
|
| 170 |
-
let slot = {{ s }}u * WG + lidx;
|
| 171 |
-
let m1 = partialM[slot];
|
| 172 |
-
let d1 = partialD[slot];
|
| 173 |
-
let m2 = partialM[slot + stride];
|
| 174 |
-
let d2 = partialD[slot + stride];
|
| 175 |
-
let mNew = max(m1, m2);
|
| 176 |
-
partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
|
| 177 |
-
partialM[slot] = mNew;
|
| 178 |
-
}
|
| 179 |
-
{% endfor %}
|
| 180 |
-
}
|
| 181 |
-
workgroupBarrier();
|
| 182 |
-
stride = stride / 2u;
|
| 183 |
-
}
|
| 184 |
-
}
|
| 185 |
-
{% else %}
|
| 186 |
|
| 187 |
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 188 |
partialM[lidx] = m;
|
|
@@ -212,12 +179,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
|
| 212 |
return merged;
|
| 213 |
}
|
| 214 |
{% endif %}
|
| 215 |
-
{% endif %}
|
| 216 |
-
|
| 217 |
|
| 218 |
{% if layout == "layer_cache" %}{% set ATTN_SCALE_OVERRIDE = "ATTN_SCALE" %}{% endif %}
|
| 219 |
-
{%
|
| 220 |
-
fn scale_value() -> f32 {
|
| 221 |
{% if ATTN_SCALE_OVERRIDE is defined %}
|
| 222 |
return {{ ATTN_SCALE_OVERRIDE }};
|
| 223 |
{% else %}
|
|
@@ -226,7 +190,6 @@ fn scale_value() -> f32 {
|
|
| 226 |
{% endif %}
|
| 227 |
}
|
| 228 |
|
| 229 |
-
|
| 230 |
{% if quantizedCache %}
|
| 231 |
{% macro emit_quant_scale4(kind, scaleBuffer) %}
|
| 232 |
fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
|
|
@@ -240,28 +203,12 @@ fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
|
|
| 240 |
{{ scaleBuffer }}[base + 2u],
|
| 241 |
{{ scaleBuffer }}[base + 3u]
|
| 242 |
);
|
| 243 |
-
}
|
| 244 |
-
{%
|
| 245 |
-
{%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
|
| 246 |
{{ emit_quant_scale4(kind, scaleBuffer) }}
|
| 247 |
fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
|
| 248 |
-
{%- if format == "int8" %}
|
| 249 |
return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
|
| 250 |
-
{%
|
| 251 |
-
// Two elements cover this vec4: each carries two +8-biased nibbles, low first.
|
| 252 |
-
let rowBase = indexV4 - d4;
|
| 253 |
-
let lo = {{ buffer }}[rowBase + d4 * 2u];
|
| 254 |
-
let hi = {{ buffer }}[rowBase + d4 * 2u + 1u];
|
| 255 |
-
let nibbles = vec4<i32>(
|
| 256 |
-
i32(lo & 0xFu), i32((lo >> 4u) & 0xFu),
|
| 257 |
-
i32(hi & 0xFu), i32((hi >> 4u) & 0xFu)
|
| 258 |
-
);
|
| 259 |
-
let signed = nibbles - vec4<i32>(8);
|
| 260 |
-
return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
|
| 261 |
-
{%- endif %}
|
| 262 |
-
}
|
| 263 |
-
{%- endmacro %}
|
| 264 |
-
|
| 265 |
{{ emit_quant_load4("int8", "key", "key", "k_scale") }}
|
| 266 |
{{ emit_quant_load4("int8", "value", "value", "v_scale") }}
|
| 267 |
{% else %}
|
|
@@ -279,9 +226,11 @@ fn load_value4(indexV4: u32) -> vec4<f32> {
|
|
| 279 |
// query row before the Q.K dots; the K bias adds a constant to every key score
|
| 280 |
// that softmax cancels, so it is skipped; the V bias is token-independent and
|
| 281 |
// is applied once in the merge pass after the final normalize.
|
|
|
|
|
|
|
| 282 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 283 |
let offset = base + d4 * 4u;
|
| 284 |
-
return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
|
| 285 |
}
|
| 286 |
|
| 287 |
{% endif %}
|
|
@@ -297,15 +246,10 @@ fn main(
|
|
| 297 |
// choose an intermediate width at pipeline execution time.
|
| 298 |
if (sgSize == 0u || sgSize > WG || WG % sgSize != 0u) { return; }
|
| 299 |
{% endif %}
|
| 300 |
-
{% if splitQueries %}
|
| 301 |
-
let queryToken = wg.x / NUM_SPLITS;
|
| 302 |
-
let split = wg.x % NUM_SPLITS;
|
| 303 |
-
{% else %}
|
| 304 |
let split = wg.x;
|
| 305 |
-
{% endif %}
|
| 306 |
let h = wg.y;
|
| 307 |
let b = wg.z;
|
| 308 |
-
if (h >= Q_HEADS || split >= NUM_SPLITS{% if
|
| 309 |
return;
|
| 310 |
}
|
| 311 |
let tid = lid.x;
|
|
@@ -317,7 +261,7 @@ fn main(
|
|
| 317 |
{% if cacheSeqlens %}
|
| 318 |
// Buffer-sharing caches retain their capacity in the physical BNSH stride;
|
| 319 |
// seqlens_k supplies the active end independently for each batch.
|
| 320 |
-
let kvSeq = min(cacheSeq, u32(seqlens_k[
|
| 321 |
{% else %}
|
| 322 |
let kvSeq = cacheSeq;
|
| 323 |
{% endif %}
|
|
@@ -325,23 +269,15 @@ fn main(
|
|
| 325 |
|
| 326 |
// Query row (decode uses token zero; short-query prefill folds the token into wg.x).
|
| 327 |
{% if layout == "bsh" %}
|
| 328 |
-
{% if splitQueries %}
|
| 329 |
-
let qBaseV4 = (b * Q_SEQ + queryToken) * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
|
| 330 |
-
{% else %}
|
| 331 |
let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
|
| 332 |
-
{% endif %}
|
| 333 |
let kvBaseV4 = b * kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
|
| 334 |
let kvTokenStrideV4 = KV_HIDDEN_V4;
|
| 335 |
{% elif layout == "layer_cache" %}
|
| 336 |
let qBaseV4 = h * HEAD_DIM_V4;
|
| 337 |
let kvBaseV4 = (LAYER * CACHE_LEN * KV_HEADS + hKv) * HEAD_DIM_V4;
|
| 338 |
let kvTokenStrideV4 = KV_HEADS * HEAD_DIM_V4;
|
| 339 |
-
{% else %}
|
| 340 |
-
{% if splitQueries %}
|
| 341 |
-
let qBaseV4 = ((b * Q_HEADS + h) * Q_SEQ + queryToken) * HEAD_DIM_V4;
|
| 342 |
{% else %}
|
| 343 |
let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
|
| 344 |
-
{% endif %}
|
| 345 |
let kvBaseV4 = (b * KV_HEADS + hKv) * cacheSeq * HEAD_DIM_V4;
|
| 346 |
let kvTokenStrideV4 = HEAD_DIM_V4;
|
| 347 |
{% endif %}
|
|
@@ -437,27 +373,12 @@ fn main(
|
|
| 437 |
{% endif %}
|
| 438 |
{% if hasMask %}
|
| 439 |
if (keyAllowed) {
|
| 440 |
-
{% if splitQueries %}
|
| 441 |
-
let maskQuery = queryToken;
|
| 442 |
-
{% else %}
|
| 443 |
let maskQuery = 0u;
|
| 444 |
-
{% endif %}
|
| 445 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
|
| 446 |
-
{% if maskIsBool %}
|
| 447 |
-
// A rejected bool-mask key contributes no probability mass. The merge
|
| 448 |
-
// pass already maps a zero global denominator to an all-zero output row.
|
| 449 |
-
if (attn_mask[maskIndex] == 0u) {
|
| 450 |
-
keyAllowed = false;
|
| 451 |
-
score = -FLT_MAX;
|
| 452 |
-
dPart = 0.0;
|
| 453 |
-
}
|
| 454 |
-
{% else %}
|
| 455 |
score = score + f32(attn_mask[maskIndex]);
|
| 456 |
-
{% endif %}
|
| 457 |
m = score;
|
| 458 |
}
|
| 459 |
-
{% endif %}
|
| 460 |
-
let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
|
| 461 |
|
| 462 |
// Merge one key tile's online-softmax (maximum, denominator) partial into the
|
| 463 |
// running state, then store the per-key probabilities consumed by V accumulation.
|
|
@@ -473,7 +394,6 @@ fn main(
|
|
| 473 |
probs[tid] = prob;
|
| 474 |
workgroupBarrier();
|
| 475 |
|
| 476 |
-
|
| 477 |
{% if jSplitV %}
|
| 478 |
// j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
|
| 479 |
// J_GROUPS for dim block d4v into a register, then the groups combine
|
|
@@ -521,20 +441,12 @@ fn main(
|
|
| 521 |
}
|
| 522 |
|
| 523 |
// Emit un-normalized partials for (b, h, split): the merge pass divides.
|
| 524 |
-
{% if splitQueries %}
|
| 525 |
-
let partialBase = (((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4;
|
| 526 |
-
{% else %}
|
| 527 |
let partialBase = ((b * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4;
|
| 528 |
-
{% endif %}
|
| 529 |
for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
|
| 530 |
partial_out[partialBase + d4] = running_out[d4];
|
| 531 |
}
|
| 532 |
if (tid == 0u) {
|
| 533 |
-
{% if splitQueries %}
|
| 534 |
-
let mdBase = ((b * Q_SEQ + queryToken) * Q_HEADS + h) * NUM_SPLITS + split;
|
| 535 |
-
{% else %}
|
| 536 |
let mdBase = (b * Q_HEADS + h) * NUM_SPLITS + split;
|
| 537 |
-
{% endif %}
|
| 538 |
// (max, denom) travel together to the merge, so they share one buffer as an
|
| 539 |
// interleaved vec2 rather than costing two bindings. Interleaved, not two
|
| 540 |
// halves, so the index needs no region size — and the merge reads both
|
|
|
|
| 1 |
{% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
|
|
|
|
| 2 |
{% if quantizedCache is not defined %}{% set quantizedCache = false %}{% endif %}
|
| 3 |
{% if cacheSeqlens is not defined %}{% set cacheSeqlens = false %}{% endif %}
|
| 4 |
{% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
|
| 5 |
+
{% set cacheSeqlensIndex = "b" %}{% set layer = 0 %}
|
| 6 |
+
{% set cacheLen = 0 %}
|
| 7 |
+
{% set scale = "0.0" %}
|
|
|
|
| 8 |
{% if hasRotary is not defined %}{% set hasRotary = false %}{% endif %}
|
| 9 |
+
{% set qHiddenV4 = qHiddenV4 | default(0) %}
|
| 10 |
+
{% set kvHiddenV4 = kvHiddenV4 | default(0) %}
|
| 11 |
+
{% set hasBias = hasBias is defined and hasBias %}
|
| 12 |
{% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
|
| 13 |
{% if useSubgroups %}
|
| 14 |
enable subgroups;
|
| 15 |
{% endif %}
|
| 16 |
{{ env.wgsl.resourceDeclarations }}
|
| 17 |
|
| 18 |
+
// Split-K flash attention. Single-partition direct output skips the merge pass.
|
| 19 |
+
// Otherwise this is the first of two passes. This geometry
|
| 20 |
// handles decode and short-query, long-context prefill inputs.
|
| 21 |
//
|
| 22 |
// The non-split flash decode launches only `batch * numHeads` workgroups, each
|
|
|
|
| 51 |
{% endif %}
|
| 52 |
const WG: u32 = {{ splitKWorkgroupSize }}u;
|
| 53 |
const NUM_SPLITS: u32 = {{ numSplits }}u;
|
|
|
|
|
|
|
|
|
|
| 54 |
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 55 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
| 56 |
// accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
|
|
|
|
| 70 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 71 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 72 |
}
|
| 73 |
+
|
| 74 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 75 |
return exp(shifted_value(value, maxValue));
|
| 76 |
}
|
|
|
|
| 78 |
var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
|
| 79 |
var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
|
| 80 |
var<workgroup> probs: array<f32, WG>;
|
| 81 |
+
{% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) and (allowCooperativeQk if allowCooperativeQk is defined else true) %}
|
| 82 |
{% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
|
| 83 |
{% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
|
| 84 |
{% if coopQk %}
|
|
|
|
| 147 |
return combinedMD;
|
| 148 |
}
|
| 149 |
{% else %}
|
| 150 |
+
{% set mdExtent = "WG" %}
|
|
|
|
|
|
|
| 151 |
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 152 |
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 153 |
|
| 154 |
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 155 |
partialM[lidx] = m;
|
|
|
|
| 179 |
return merged;
|
| 180 |
}
|
| 181 |
{% endif %}
|
|
|
|
|
|
|
| 182 |
|
| 183 |
{% if layout == "layer_cache" %}{% set ATTN_SCALE_OVERRIDE = "ATTN_SCALE" %}{% endif %}
|
| 184 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 185 |
{% if ATTN_SCALE_OVERRIDE is defined %}
|
| 186 |
return {{ ATTN_SCALE_OVERRIDE }};
|
| 187 |
{% else %}
|
|
|
|
| 190 |
{% endif %}
|
| 191 |
}
|
| 192 |
|
|
|
|
| 193 |
{% if quantizedCache %}
|
| 194 |
{% macro emit_quant_scale4(kind, scaleBuffer) %}
|
| 195 |
fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
|
|
|
|
| 203 |
{{ scaleBuffer }}[base + 2u],
|
| 204 |
{{ scaleBuffer }}[base + 3u]
|
| 205 |
);
|
| 206 |
+
}{% endmacro %}
|
| 207 |
+
{% macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
|
|
|
|
| 208 |
{{ emit_quant_scale4(kind, scaleBuffer) }}
|
| 209 |
fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
|
|
|
|
| 210 |
return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
|
| 211 |
+
}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 212 |
{{ emit_quant_load4("int8", "key", "key", "k_scale") }}
|
| 213 |
{{ emit_quant_load4("int8", "value", "value", "v_scale") }}
|
| 214 |
{% else %}
|
|
|
|
| 226 |
// query row before the Q.K dots; the K bias adds a constant to every key score
|
| 227 |
// that softmax cancels, so it is skipped; the V bias is token-independent and
|
| 228 |
// is applied once in the merge pass after the final normalize.
|
| 229 |
+
{% set BW = "" %}
|
| 230 |
+
{% set BC = "" %}
|
| 231 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 232 |
let offset = base + d4 * 4u;
|
| 233 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 234 |
}
|
| 235 |
|
| 236 |
{% endif %}
|
|
|
|
| 246 |
// choose an intermediate width at pipeline execution time.
|
| 247 |
if (sgSize == 0u || sgSize > WG || WG % sgSize != 0u) { return; }
|
| 248 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 249 |
let split = wg.x;
|
|
|
|
| 250 |
let h = wg.y;
|
| 251 |
let b = wg.z;
|
| 252 |
+
if (h >= Q_HEADS || split >= NUM_SPLITS{% if layout == "layer_cache" %} || params.past_len >= CACHE_LEN{% endif %}) {
|
| 253 |
return;
|
| 254 |
}
|
| 255 |
let tid = lid.x;
|
|
|
|
| 261 |
{% if cacheSeqlens %}
|
| 262 |
// Buffer-sharing caches retain their capacity in the physical BNSH stride;
|
| 263 |
// seqlens_k supplies the active end independently for each batch.
|
| 264 |
+
let kvSeq = min(cacheSeq, u32(seqlens_k[{{ cacheSeqlensIndex }}]) + 1u);
|
| 265 |
{% else %}
|
| 266 |
let kvSeq = cacheSeq;
|
| 267 |
{% endif %}
|
|
|
|
| 269 |
|
| 270 |
// Query row (decode uses token zero; short-query prefill folds the token into wg.x).
|
| 271 |
{% if layout == "bsh" %}
|
|
|
|
|
|
|
|
|
|
| 272 |
let qBaseV4 = b * Q_HIDDEN_V4 + h * HEAD_DIM_V4;
|
|
|
|
| 273 |
let kvBaseV4 = b * kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
|
| 274 |
let kvTokenStrideV4 = KV_HIDDEN_V4;
|
| 275 |
{% elif layout == "layer_cache" %}
|
| 276 |
let qBaseV4 = h * HEAD_DIM_V4;
|
| 277 |
let kvBaseV4 = (LAYER * CACHE_LEN * KV_HEADS + hKv) * HEAD_DIM_V4;
|
| 278 |
let kvTokenStrideV4 = KV_HEADS * HEAD_DIM_V4;
|
|
|
|
|
|
|
|
|
|
| 279 |
{% else %}
|
| 280 |
let qBaseV4 = (b * Q_HEADS + h) * HEAD_DIM_V4;
|
|
|
|
| 281 |
let kvBaseV4 = (b * KV_HEADS + hKv) * cacheSeq * HEAD_DIM_V4;
|
| 282 |
let kvTokenStrideV4 = HEAD_DIM_V4;
|
| 283 |
{% endif %}
|
|
|
|
| 373 |
{% endif %}
|
| 374 |
{% if hasMask %}
|
| 375 |
if (keyAllowed) {
|
|
|
|
|
|
|
|
|
|
| 376 |
let maskQuery = 0u;
|
|
|
|
| 377 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 378 |
score = score + f32(attn_mask[maskIndex]);
|
|
|
|
| 379 |
m = score;
|
| 380 |
}
|
| 381 |
+
{% endif %} let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
|
|
|
|
| 382 |
|
| 383 |
// Merge one key tile's online-softmax (maximum, denominator) partial into the
|
| 384 |
// running state, then store the per-key probabilities consumed by V accumulation.
|
|
|
|
| 394 |
probs[tid] = prob;
|
| 395 |
workgroupBarrier();
|
| 396 |
|
|
|
|
| 397 |
{% if jSplitV %}
|
| 398 |
// j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
|
| 399 |
// J_GROUPS for dim block d4v into a register, then the groups combine
|
|
|
|
| 441 |
}
|
| 442 |
|
| 443 |
// Emit un-normalized partials for (b, h, split): the merge pass divides.
|
|
|
|
|
|
|
|
|
|
| 444 |
let partialBase = ((b * Q_HEADS + h) * NUM_SPLITS + split) * HEAD_DIM_V4;
|
|
|
|
| 445 |
for (var d4: u32 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
|
| 446 |
partial_out[partialBase + d4] = running_out[d4];
|
| 447 |
}
|
| 448 |
if (tid == 0u) {
|
|
|
|
|
|
|
|
|
|
| 449 |
let mdBase = (b * Q_HEADS + h) * NUM_SPLITS + split;
|
|
|
|
| 450 |
// (max, denom) travel together to the merge, so they share one buffer as an
|
| 451 |
// interleaved vec2 rather than costing two bindings. Interleaved, not two
|
| 452 |
// halves, so the index needs no region size — and the merge reads both
|
build/webgpu/attn-flash-online.wgsl.jinja
CHANGED
|
@@ -1,6 +1,4 @@
|
|
| 1 |
-
{% if combineSubgroups %}
|
| 2 |
enable subgroups;
|
| 3 |
-
{% endif %}
|
| 4 |
{{ env.wgsl.resourceDeclarations }}
|
| 5 |
|
| 6 |
// Flash-style tiled online-softmax attention for vec4-aligned head dimensions.
|
|
@@ -28,8 +26,8 @@ const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u;
|
|
| 28 |
const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
|
| 29 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 30 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 31 |
-
{% set qHeads = "
|
| 32 |
-
{% set kvHeads = "
|
| 33 |
const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
|
| 34 |
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 35 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
|
@@ -50,6 +48,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 50 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 51 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 52 |
}
|
|
|
|
| 53 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 54 |
return exp(shifted_value(value, maxValue));
|
| 55 |
}
|
|
@@ -63,8 +62,6 @@ var<workgroup> probs: array<f32, WG>;
|
|
| 63 |
// Both the subgroup and portable barrier-tree engines return the same merged
|
| 64 |
// pair to every invocation. Repeated merges require a workgroup barrier between
|
| 65 |
// calls before their shared partial storage is reused.
|
| 66 |
-
{% set combineSubgroups = combineSubgroups is defined and combineSubgroups %}
|
| 67 |
-
{% if combineSubgroups %}
|
| 68 |
// Cross-subgroup merge that assumes nothing about which invocations share a
|
| 69 |
// subgroup or how many subgroups there are: each subgroup's elected lane
|
| 70 |
// publishes the subgroup pair in the slot at its OWN invocation index and sets
|
|
@@ -115,96 +112,28 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
|
|
| 115 |
workgroupBarrier();
|
| 116 |
return combinedMD;
|
| 117 |
}
|
| 118 |
-
{% else %}
|
| 119 |
-
{% set mdStreamed = mdStreams is defined %}
|
| 120 |
-
{% set mdStreams = mdStreams if mdStreams is defined else 1 %}
|
| 121 |
-
{% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
|
| 122 |
-
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 123 |
-
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
| 124 |
-
{% if mdStreamed %}
|
| 125 |
-
|
| 126 |
-
// In-place fold of {{ mdStreams }} streams. Input partials occupy
|
| 127 |
-
// partialM/partialD; stream s returns its merged pair in slot s * WG.
|
| 128 |
-
fn combine_partials_streams(lidx: u32) {
|
| 129 |
-
workgroupBarrier();
|
| 130 |
-
var stride = WG / 2u;
|
| 131 |
-
loop {
|
| 132 |
-
if (stride == 0u) {
|
| 133 |
-
break;
|
| 134 |
-
}
|
| 135 |
-
if (lidx < stride) {
|
| 136 |
-
{% for s in range(mdStreams) %}
|
| 137 |
-
{
|
| 138 |
-
let slot = {{ s }}u * WG + lidx;
|
| 139 |
-
let m1 = partialM[slot];
|
| 140 |
-
let d1 = partialD[slot];
|
| 141 |
-
let m2 = partialM[slot + stride];
|
| 142 |
-
let d2 = partialD[slot + stride];
|
| 143 |
-
let mNew = max(m1, m2);
|
| 144 |
-
partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
|
| 145 |
-
partialM[slot] = mNew;
|
| 146 |
-
}
|
| 147 |
-
{% endfor %}
|
| 148 |
-
}
|
| 149 |
-
workgroupBarrier();
|
| 150 |
-
stride = stride / 2u;
|
| 151 |
-
}
|
| 152 |
-
}
|
| 153 |
-
{% else %}
|
| 154 |
-
|
| 155 |
-
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 156 |
-
partialM[lidx] = m;
|
| 157 |
-
partialD[lidx] = d;
|
| 158 |
-
workgroupBarrier();
|
| 159 |
-
var stride = WG / 2u;
|
| 160 |
-
loop {
|
| 161 |
-
if (stride == 0u) {
|
| 162 |
-
break;
|
| 163 |
-
}
|
| 164 |
-
if (lidx < stride) {
|
| 165 |
-
let m1 = partialM[lidx];
|
| 166 |
-
let d1 = partialD[lidx];
|
| 167 |
-
let m2 = partialM[lidx + stride];
|
| 168 |
-
let d2 = partialD[lidx + stride];
|
| 169 |
-
let mNew = max(m1, m2);
|
| 170 |
-
partialD[lidx] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
|
| 171 |
-
partialM[lidx] = mNew;
|
| 172 |
-
}
|
| 173 |
-
workgroupBarrier();
|
| 174 |
-
stride = stride / 2u;
|
| 175 |
-
}
|
| 176 |
-
let merged = vec2<f32>(partialM[0], partialD[0]);
|
| 177 |
-
// Trailing barrier so back-to-back calls cannot race a next call's partial
|
| 178 |
-
// stores against this call's reads of slot 0.
|
| 179 |
-
workgroupBarrier();
|
| 180 |
-
return merged;
|
| 181 |
-
}
|
| 182 |
-
{% endif %}
|
| 183 |
-
{% endif %}
|
| 184 |
-
|
| 185 |
|
| 186 |
// An explicit-zero specialization bakes the scale as 0. Otherwise,
|
| 187 |
// params.scale == 0 encodes an omitted scale and selects 1/sqrt(headDim).
|
| 188 |
-
{%
|
| 189 |
-
fn scale_value() -> f32 {
|
| 190 |
if (params.scale != 0.0) { return params.scale; }
|
| 191 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 192 |
}
|
| 193 |
-
|
| 194 |
{% if hasBias %}
|
| 195 |
|
|
|
|
|
|
|
| 196 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 197 |
let offset = base + d4 * 4u;
|
| 198 |
-
return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
|
| 199 |
}
|
| 200 |
|
| 201 |
{% endif %}
|
| 202 |
@compute @workgroup_size(WG, 1, 1)
|
| 203 |
fn main(
|
| 204 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 205 |
-
@builtin(local_invocation_id) lid: vec3<u32>
|
| 206 |
-
@builtin(subgroup_size) sgSize: u32
|
| 207 |
-
) {
|
| 208 |
let qi = wg.x;
|
| 209 |
let h = wg.y;
|
| 210 |
let b = wg.z;
|
|
@@ -243,7 +172,11 @@ fn main(
|
|
| 243 |
// Causal upper bound: query qi attends only keys 0..qi, so stop after the tile
|
| 244 |
// containing qi and skip the unattended tail. Non-causal keeps the full kvSeq
|
| 245 |
// sweep.
|
|
|
|
|
|
|
|
|
|
| 246 |
var keyBoundV = params.kvSeq;
|
|
|
|
| 247 |
var keyFloor: u32 = 0u;
|
| 248 |
{% if hasWindow %}
|
| 249 |
// Sliding window: query qi sits at absolute position p = kvSeq - qSeq
|
|
@@ -269,7 +202,7 @@ fn main(
|
|
| 269 |
var score = -FLT_MAX;
|
| 270 |
var m = -FLT_MAX;
|
| 271 |
var dPart = 0.0;
|
| 272 |
-
var keyAllowed = kj < keyBound{% if hasWindow %} && kj >= keyFloor{% endif %};
|
| 273 |
if (keyAllowed) {
|
| 274 |
let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
|
| 275 |
{% if hasMask %}
|
|
@@ -277,17 +210,6 @@ fn main(
|
|
| 277 |
// [q, k] masks set batch/head strides to 0).
|
| 278 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qi * params.maskSeqStride + kj;
|
| 279 |
{% endif %}
|
| 280 |
-
{% if hasMask and maskIsBool %}
|
| 281 |
-
if (attn_mask[maskIndex] == 0u) {
|
| 282 |
-
keyAllowed = false;
|
| 283 |
-
} else {
|
| 284 |
-
var acc: f32 = 0.0;
|
| 285 |
-
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 286 |
-
acc = acc + dot(q_shared[d4], vec4<f32>(key[kRowV4 + d4]));
|
| 287 |
-
}
|
| 288 |
-
score = acc * scale;
|
| 289 |
-
}
|
| 290 |
-
{% else %}
|
| 291 |
var acc: f32 = 0.0;
|
| 292 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 293 |
acc = acc + dot(q_shared[d4], vec4<f32>(key[kRowV4 + d4]));
|
|
@@ -295,12 +217,11 @@ fn main(
|
|
| 295 |
score = acc * scale;
|
| 296 |
{% if hasMask %}
|
| 297 |
score = score + f32(attn_mask[maskIndex]);
|
| 298 |
-
{% endif %}
|
| 299 |
{% endif %}
|
| 300 |
m = score;
|
| 301 |
dPart = select(0.0, 1.0, keyAllowed);
|
| 302 |
}
|
| 303 |
-
let tile = combine_partials(m, dPart, tid
|
| 304 |
|
| 305 |
// Online merge of the tile into the running state (softmax-online rule).
|
| 306 |
// Merge one key tile's online-softmax (maximum, denominator) partial into the
|
|
@@ -317,7 +238,6 @@ fn main(
|
|
| 317 |
probs[tid] = prob;
|
| 318 |
workgroupBarrier();
|
| 319 |
|
| 320 |
-
|
| 321 |
// running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
|
| 322 |
// tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
|
| 323 |
let tileCount = min(WG, keyBound - kjBase);
|
|
|
|
|
|
|
| 1 |
enable subgroups;
|
|
|
|
| 2 |
{{ env.wgsl.resourceDeclarations }}
|
| 3 |
|
| 4 |
// Flash-style tiled online-softmax attention for vec4-aligned head dimensions.
|
|
|
|
| 26 |
const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
|
| 27 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 28 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 29 |
+
{% set qHeads = "Q_HEADS" %}
|
| 30 |
+
{% set kvHeads = "KV_HEADS" %}
|
| 31 |
const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
|
| 32 |
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 33 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
|
|
|
| 48 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 49 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 50 |
}
|
| 51 |
+
|
| 52 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 53 |
return exp(shifted_value(value, maxValue));
|
| 54 |
}
|
|
|
|
| 62 |
// Both the subgroup and portable barrier-tree engines return the same merged
|
| 63 |
// pair to every invocation. Repeated merges require a workgroup barrier between
|
| 64 |
// calls before their shared partial storage is reused.
|
|
|
|
|
|
|
| 65 |
// Cross-subgroup merge that assumes nothing about which invocations share a
|
| 66 |
// subgroup or how many subgroups there are: each subgroup's elected lane
|
| 67 |
// publishes the subgroup pair in the slot at its OWN invocation index and sets
|
|
|
|
| 112 |
workgroupBarrier();
|
| 113 |
return combinedMD;
|
| 114 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
|
| 116 |
// An explicit-zero specialization bakes the scale as 0. Otherwise,
|
| 117 |
// params.scale == 0 encodes an omitted scale and selects 1/sqrt(headDim).
|
| 118 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 119 |
if (params.scale != 0.0) { return params.scale; }
|
| 120 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 121 |
}
|
|
|
|
| 122 |
{% if hasBias %}
|
| 123 |
|
| 124 |
+
{% set BW = "" %}
|
| 125 |
+
{% set BC = "" %}
|
| 126 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 127 |
let offset = base + d4 * 4u;
|
| 128 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 129 |
}
|
| 130 |
|
| 131 |
{% endif %}
|
| 132 |
@compute @workgroup_size(WG, 1, 1)
|
| 133 |
fn main(
|
| 134 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 135 |
+
@builtin(local_invocation_id) lid: vec3<u32>,
|
| 136 |
+
@builtin(subgroup_size) sgSize: u32) {
|
|
|
|
| 137 |
let qi = wg.x;
|
| 138 |
let h = wg.y;
|
| 139 |
let b = wg.z;
|
|
|
|
| 172 |
// Causal upper bound: query qi attends only keys 0..qi, so stop after the tile
|
| 173 |
// containing qi and skip the unattended tail. Non-causal keeps the full kvSeq
|
| 174 |
// sweep.
|
| 175 |
+
{% if hasCausal %}
|
| 176 |
+
var keyBoundV = select(params.kvSeq, min(params.kvSeq, qi + 1u), params.isCausal != 0u);
|
| 177 |
+
{% else %}
|
| 178 |
var keyBoundV = params.kvSeq;
|
| 179 |
+
{% endif %}
|
| 180 |
var keyFloor: u32 = 0u;
|
| 181 |
{% if hasWindow %}
|
| 182 |
// Sliding window: query qi sits at absolute position p = kvSeq - qSeq
|
|
|
|
| 202 |
var score = -FLT_MAX;
|
| 203 |
var m = -FLT_MAX;
|
| 204 |
var dPart = 0.0;
|
| 205 |
+
var keyAllowed = kj < keyBound{% if hasCausal %} && (params.isCausal == 0u || kj <= qi){% endif %}{% if hasWindow %} && kj >= keyFloor{% endif %};
|
| 206 |
if (keyAllowed) {
|
| 207 |
let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
|
| 208 |
{% if hasMask %}
|
|
|
|
| 210 |
// [q, k] masks set batch/head strides to 0).
|
| 211 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qi * params.maskSeqStride + kj;
|
| 212 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 213 |
var acc: f32 = 0.0;
|
| 214 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 215 |
acc = acc + dot(q_shared[d4], vec4<f32>(key[kRowV4 + d4]));
|
|
|
|
| 217 |
score = acc * scale;
|
| 218 |
{% if hasMask %}
|
| 219 |
score = score + f32(attn_mask[maskIndex]);
|
|
|
|
| 220 |
{% endif %}
|
| 221 |
m = score;
|
| 222 |
dPart = select(0.0, 1.0, keyAllowed);
|
| 223 |
}
|
| 224 |
+
let tile = combine_partials(m, dPart, tid, sgSize);
|
| 225 |
|
| 226 |
// Online merge of the tile into the running state (softmax-online rule).
|
| 227 |
// Merge one key tile's online-softmax (maximum, denominator) partial into the
|
|
|
|
| 238 |
probs[tid] = prob;
|
| 239 |
workgroupBarrier();
|
| 240 |
|
|
|
|
| 241 |
// running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
|
| 242 |
// tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
|
| 243 |
let tileCount = min(WG, keyBound - kjBase);
|
build/webgpu/attn-flash-prefill-cluster.wgsl.jinja
CHANGED
|
@@ -1,22 +1,18 @@
|
|
| 1 |
-
{% set
|
| 2 |
-
{% set
|
| 3 |
-
{% set
|
| 4 |
-
{% set
|
| 5 |
-
{% set
|
| 6 |
-
{% set
|
| 7 |
-
{% set
|
| 8 |
-
{% set
|
| 9 |
-
{% set KEY = "cache_keys" if sourceProfile == 1 else "key" %}
|
| 10 |
-
{% set VALUE = "cache_values" if sourceProfile == 1 else "value" %}
|
| 11 |
-
{% set OUTPUT = "attn_out" if sourceProfile == 1 else "output" %}
|
| 12 |
-
{% if sourceProfile == 1 %}{% set ATTN_SCALE_OVERRIDE = scaling %}{% endif %}
|
| 13 |
{% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
|
| 14 |
{% if batchNoSgReduction is not defined %}{% set batchNoSgReduction = false %}{% endif %}
|
|
|
|
| 15 |
{% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
|
| 16 |
-
{%
|
| 17 |
-
{% if maskIsBool is not defined %}{% set maskIsBool = false %}{% endif %}
|
| 18 |
-
{% if hasSoftcap is not defined %}{% set hasSoftcap = false %}{% endif %}
|
| 19 |
{% if hasHeadSink is not defined %}{% set hasHeadSink = false %}{% endif %}
|
|
|
|
| 20 |
{% if quantCacheFormat is not defined %}{% set quantCacheFormat = "" %}{% endif %}
|
| 21 |
{% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
|
| 22 |
// A windowed cache binds a fixed CAPACITY but keeps only the most recent
|
|
@@ -24,7 +20,14 @@
|
|
| 24 |
// count, which is still the right batch stride but the wrong attention bound, so
|
| 25 |
// the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
|
| 26 |
// in both roles.
|
| 27 |
-
{%
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 28 |
{% macro score_expr(part) %}{% if hasSoftcap %}params.softcap * tanh(clamp(({{ part }} * SCALE) / params.softcap, -30.0, 30.0)){% else %}{{ part }} * SCALE{% endif %}{% endmacro %}
|
| 29 |
{% if useSubgroups %}
|
| 30 |
enable subgroups;
|
|
@@ -40,11 +43,30 @@ enable subgroups;
|
|
| 40 |
{% set MASK_TILE_TYPE = "u32" if MASK_IS_INT else "f32" %}
|
| 41 |
{% set MASK_TILE_LOAD = "attn_mask[maskIndex]" if MASK_IS_INT else "f32(attn_mask[maskIndex])" %}
|
| 42 |
{% set MASK_TILE_ZERO = "0u" if MASK_IS_INT else "0.0" %}
|
| 43 |
-
{% set MASK_ELEMENT = "maskValue" if STAGE_MASK else "attn_mask[maskIndex]" %}
|
| 44 |
{% set MASK_ADDITIVE = "maskValue" if STAGE_MASK else "f32(attn_mask[maskIndex])" %}
|
| 45 |
{% set SLICE_COUNT = ((headDimV4 / LPQ) | int) %}
|
| 46 |
{% set QPL = (2 if (not usesF16 and SLICE_COUNT <= 4 and TILE_K <= 16 and TILE_Q % 2 == 0) else 1) if useSubgroups else 1 %}
|
| 47 |
{% macro qn(name, qi) %}{{ name }}{% if qi > 0 %}_q{{ qi }}{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 48 |
{% set ROPE_LANE_XOR = ((LPQ / 2) | int) %}
|
| 49 |
{% set QL = qLayout if qLayout is defined else layout %}
|
| 50 |
{% set KL = kvLayout if kvLayout is defined else layout %}
|
|
@@ -69,8 +91,7 @@ const HEAD_DIM: u32 = {{ headDim }}u;
|
|
| 69 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
| 70 |
{% if QL != "bhsd" %}
|
| 71 |
const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
|
| 72 |
-
{%
|
| 73 |
-
{% endif %}{% endif %}
|
| 74 |
{% if KL != "bhsd" %}
|
| 75 |
const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
|
| 76 |
{% endif %}
|
|
@@ -91,16 +112,7 @@ const WG: u32 = (TILE_Q / QPL) * LPQ;
|
|
| 91 |
{% else %}
|
| 92 |
const WG: u32 = TILE_Q * LPQ;
|
| 93 |
{% endif %}
|
| 94 |
-
{% if MASK_IS_INT %}
|
| 95 |
-
// Key-keep masks in contrib attention use a finite low logit for a rejected
|
| 96 |
-
// key. Logical ONNX bool masks use the exclusion sentinel instead, so a fully
|
| 97 |
-
// masked row has zero mass. Keeping one declaration shape for both mask modes
|
| 98 |
-
// lets both mask modes share the score loop below.
|
| 99 |
-
const NEG_INF: f32 = -3.4028234663852886e38;
|
| 100 |
-
const MASK_NEG: f32 = {{ "-1e38" if maskIsKeyKeep else "-3.4028234663852886e38" }};
|
| 101 |
-
{% else %}
|
| 102 |
const NEG_INF: f32 = -3.4028234663852886e38;
|
| 103 |
-
{% endif %}
|
| 104 |
|
| 105 |
var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
| 106 |
var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
|
@@ -121,16 +133,10 @@ var<workgroup> red: array<f32, WG>;
|
|
| 121 |
{% endif %}
|
| 122 |
{% endif %}
|
| 123 |
|
| 124 |
-
{%
|
| 125 |
-
fn scale_value() -> f32 {
|
| 126 |
-
{% if ATTN_SCALE_OVERRIDE is defined %}
|
| 127 |
-
return {{ ATTN_SCALE_OVERRIDE }};
|
| 128 |
-
{% else %}
|
| 129 |
if (params.scale != 0.0) { return params.scale; }
|
| 130 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 131 |
-
{% endif %}
|
| 132 |
}
|
| 133 |
-
|
| 134 |
{% if quantCacheFormat %}
|
| 135 |
// A quantized cache is dequantized once per key into the staged tile, then read
|
| 136 |
// by all TILE_Q queries in the workgroup. The unpack cost is amortized over the
|
|
@@ -147,14 +153,13 @@ fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
|
|
| 147 |
{{ scaleBuffer }}[base + 2u],
|
| 148 |
{{ scaleBuffer }}[base + 3u]
|
| 149 |
);
|
| 150 |
-
}
|
| 151 |
-
{%
|
| 152 |
-
{%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
|
| 153 |
{{ emit_quant_scale4(kind, scaleBuffer) }}
|
| 154 |
fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
|
| 155 |
-
{%
|
| 156 |
return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
|
| 157 |
-
{%
|
| 158 |
// Two elements cover this vec4: each carries two +8-biased nibbles, low first.
|
| 159 |
let rowBase = indexV4 - d4;
|
| 160 |
let lo = {{ buffer }}[rowBase + d4 * 2u];
|
|
@@ -165,14 +170,21 @@ fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
|
|
| 165 |
);
|
| 166 |
let signed = nibbles - vec4<i32>(8);
|
| 167 |
return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
|
| 168 |
-
{%
|
| 169 |
-
}
|
| 170 |
-
{%- endmacro %}
|
| 171 |
-
|
| 172 |
{{ emit_quant_load4(quantCacheFormat, "key", KEY, "k_scale") }}
|
| 173 |
{{ emit_quant_load4(quantCacheFormat, "value", VALUE, "v_scale") }}
|
| 174 |
{% endif %}
|
| 175 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 176 |
@compute @workgroup_size(WG, 1, 1)
|
| 177 |
fn main(
|
| 178 |
@builtin(workgroup_id) wg: vec3<u32>,
|
|
@@ -211,6 +223,9 @@ fn main(
|
|
| 211 |
{% endif %}
|
| 212 |
{% for c in range(SLICE_COUNT) %}
|
| 213 |
var {{ qn("qr" ~ c, qi) }} = vec4<f32>({{ QUERY }}[{{ qn("qBase4", qi) }} + {{ c }}u]);
|
|
|
|
|
|
|
|
|
|
| 214 |
var {{ qn("o" ~ c, qi) }} = vec4<f32>(0.0);
|
| 215 |
{% endfor %}
|
| 216 |
{% endfor %}
|
|
@@ -390,9 +405,13 @@ fn main(
|
|
| 390 |
// TILE_K is a small shader constant; the loop updates the named q/o slices in place.
|
| 391 |
{% if not useSubgroups and batchNoSgReduction %}
|
| 392 |
// First publish every key's partial dot without intervening barriers.
|
|
|
|
| 393 |
var s: array<f32, TILE_K>;
|
|
|
|
| 394 |
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
|
|
|
| 395 |
s[kk] = NEG_INF;
|
|
|
|
| 396 |
var part: f32 = 0.0;
|
| 397 |
let kb = kk * HEAD_DIM_V4 + lane8 * SLICE;
|
| 398 |
{% for c in range(SLICE_COUNT) %}
|
|
@@ -422,26 +441,41 @@ fn main(
|
|
| 422 |
workgroupBarrier();
|
| 423 |
// The barrier at the start of the next K tile protects this scratch before
|
| 424 |
// it is overwritten.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 425 |
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 426 |
let kj = kStart + kk;
|
| 427 |
let part = red[kk * WG + tid - lane8];
|
|
|
|
|
|
|
|
|
|
| 428 |
if (kj >= minKj && kj < maxKj) {
|
| 429 |
{% if hasMask %}
|
| 430 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
| 431 |
-
{% if maskIsBool %}
|
| 432 |
-
if (attn_mask[maskIndex] != 0u) {
|
| 433 |
-
s[kk] = {{ score_expr("part") }};
|
| 434 |
-
} else {
|
| 435 |
-
s[kk] = MASK_NEG;
|
| 436 |
-
}
|
| 437 |
-
{% else %}
|
| 438 |
s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
|
| 439 |
-
{% endif %}
|
| 440 |
{% else %}
|
| 441 |
s[kk] = {{ score_expr("part") }};
|
| 442 |
{% endif %}
|
| 443 |
}
|
| 444 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 445 |
{% elif QPL > 1 %}
|
| 446 |
{% for qi in range(QPL) %}
|
| 447 |
var {{ qn("s", qi) }}: array<f32, TILE_K>;
|
|
@@ -481,21 +515,7 @@ fn main(
|
|
| 481 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 482 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
|
| 483 |
{% endif %}
|
| 484 |
-
{% if maskIsKeyKeep %}
|
| 485 |
-
// A broadcast key mask uses 1 for a retained key and 0 for padding.
|
| 486 |
-
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
|
| 487 |
-
{% elif maskIsBool %}
|
| 488 |
-
// Logical bool: a rejected key contributes no softmax mass. Leaving
|
| 489 |
-
// the initialized NEG_INF sentinel in place makes a fully masked row
|
| 490 |
-
// land on the zero-denominator output guard below.
|
| 491 |
-
if ({{ MASK_ELEMENT }} != 0u) {
|
| 492 |
-
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
|
| 493 |
-
} else {
|
| 494 |
-
{{ qn("s", qi) }}[kk] = MASK_NEG;
|
| 495 |
-
}
|
| 496 |
-
{% else %}
|
| 497 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
|
| 498 |
-
{% endif %}
|
| 499 |
{% else %}
|
| 500 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
|
| 501 |
{% endif %}
|
|
@@ -541,21 +561,7 @@ fn main(
|
|
| 541 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 542 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
| 543 |
{% endif %}
|
| 544 |
-
{% if maskIsKeyKeep %}
|
| 545 |
-
// A broadcast key mask uses 1 for a retained key and 0 for padding.
|
| 546 |
-
s[kk] = {{ score_expr("part") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
|
| 547 |
-
{% elif maskIsBool %}
|
| 548 |
-
// Logical bool: a rejected key contributes no softmax mass. Leaving
|
| 549 |
-
// the initialized NEG_INF sentinel in place makes a fully masked row
|
| 550 |
-
// land on the zero-denominator output guard below.
|
| 551 |
-
if ({{ MASK_ELEMENT }} != 0u) {
|
| 552 |
-
s[kk] = {{ score_expr("part") }};
|
| 553 |
-
} else {
|
| 554 |
-
s[kk] = MASK_NEG;
|
| 555 |
-
}
|
| 556 |
-
{% else %}
|
| 557 |
s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
|
| 558 |
-
{% endif %}
|
| 559 |
{% else %}
|
| 560 |
s[kk] = {{ score_expr("part") }};
|
| 561 |
{% endif %}
|
|
@@ -563,30 +569,36 @@ fn main(
|
|
| 563 |
}
|
| 564 |
{% endif %}
|
| 565 |
|
| 566 |
-
|
| 567 |
-
|
| 568 |
-
{%
|
| 569 |
-
var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
|
| 570 |
-
for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
|
| 571 |
-
{{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
|
| 572 |
-
}
|
| 573 |
-
let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
|
| 574 |
-
let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
|
| 575 |
-
var {{ qn("pSum", qi) }}: f32 = 0.0;
|
| 576 |
-
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 577 |
-
let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
|
| 578 |
-
{{ qn("s", qi) }}[kk] = pk;
|
| 579 |
-
{{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
|
| 580 |
-
}
|
| 581 |
-
{{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
|
| 582 |
-
{{ qn("m", qi) }} = {{ qn("newMax", qi) }};
|
| 583 |
-
{% endfor %}
|
| 584 |
// A boundary tile can address V rows outside a query's attended range, and
|
| 585 |
// a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
|
| 586 |
// NaN, the guarded loop selects the V operand away for range-excluded keys.
|
| 587 |
// Interior tiles use the unguarded FMA chain. Mask exclusion applies to
|
| 588 |
// materialized V rows and does not require this range guard.
|
| 589 |
-
{% if
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 590 |
let tileInterior = {% for qi in range(QPL) %}(kStart >= {{ qn("minKj", qi) }} && kStart + TILE_K <= {{ qn("maxKj", qi) }}){% if not loop.last %} && {% endif %}{% endfor %};
|
| 591 |
{% for c in range(SLICE_COUNT) %}
|
| 592 |
{
|
|
@@ -635,7 +647,7 @@ fn main(
|
|
| 635 |
kStart = kStart + TILE_K;
|
| 636 |
}
|
| 637 |
|
| 638 |
-
{% macro attention_value(c, qi) %}{{ qn("o" ~ c, qi) }} * inv{% endmacro %}
|
| 639 |
{% for qi in range(QPL) %}
|
| 640 |
if ({{ qn("qValid", qi) }}) {
|
| 641 |
{% if QL == "bhsd" %}
|
|
|
|
| 1 |
+
{% set QSEQ = "params.qSeq" %}
|
| 2 |
+
{% set KVSEQ = "params.kvSeq" %}
|
| 3 |
+
{% set IS_CAUSAL = "params.isCausal" %}
|
| 4 |
+
{% set Q_STRIDE = "Q_HIDDEN_V4" %}
|
| 5 |
+
{% set QUERY = "query" %}
|
| 6 |
+
{% set KEY = "key" %}
|
| 7 |
+
{% set VALUE = "value" %}
|
| 8 |
+
{% set OUTPUT = "output" %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 9 |
{% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
|
| 10 |
{% if batchNoSgReduction is not defined %}{% set batchNoSgReduction = false %}{% endif %}
|
| 11 |
+
{% set SHARE_PROB = shareNoSgProb is defined and shareNoSgProb %}
|
| 12 |
{% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
|
| 13 |
+
{% set maskIsKeyKeep = false %}{% set maskIsBool = false %}{% if hasSoftcap is not defined %}{% set hasSoftcap = false %}{% endif %}
|
|
|
|
|
|
|
| 14 |
{% if hasHeadSink is not defined %}{% set hasHeadSink = false %}{% endif %}
|
| 15 |
+
{% set Q_HIDDEN = qHidden | default(0) %}
|
| 16 |
{% if quantCacheFormat is not defined %}{% set quantCacheFormat = "" %}{% endif %}
|
| 17 |
{% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
|
| 18 |
// A windowed cache binds a fixed CAPACITY but keeps only the most recent
|
|
|
|
| 20 |
// count, which is still the right batch stride but the wrong attention bound, so
|
| 21 |
// the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
|
| 22 |
// in both roles.
|
| 23 |
+
{% if hasKeyLimit is not defined %}{% set hasKeyLimit = false %}{% endif %}
|
| 24 |
+
{% if useSeqlens %}
|
| 25 |
+
{% set KVA = "kvActive" %}
|
| 26 |
+
{% elif hasKeyLimit %}
|
| 27 |
+
{% set KVA = "select(params.kvSeq, min(params.keyLimit, params.kvSeq), params.keyLimit > 0u)" %}
|
| 28 |
+
{% else %}
|
| 29 |
+
{% set KVA = KVSEQ %}
|
| 30 |
+
{% endif %}
|
| 31 |
{% macro score_expr(part) %}{% if hasSoftcap %}params.softcap * tanh(clamp(({{ part }} * SCALE) / params.softcap, -30.0, 30.0)){% else %}{{ part }} * SCALE{% endif %}{% endmacro %}
|
| 32 |
{% if useSubgroups %}
|
| 33 |
enable subgroups;
|
|
|
|
| 43 |
{% set MASK_TILE_TYPE = "u32" if MASK_IS_INT else "f32" %}
|
| 44 |
{% set MASK_TILE_LOAD = "attn_mask[maskIndex]" if MASK_IS_INT else "f32(attn_mask[maskIndex])" %}
|
| 45 |
{% set MASK_TILE_ZERO = "0u" if MASK_IS_INT else "0.0" %}
|
|
|
|
| 46 |
{% set MASK_ADDITIVE = "maskValue" if STAGE_MASK else "f32(attn_mask[maskIndex])" %}
|
| 47 |
{% set SLICE_COUNT = ((headDimV4 / LPQ) | int) %}
|
| 48 |
{% set QPL = (2 if (not usesF16 and SLICE_COUNT <= 4 and TILE_K <= 16 and TILE_Q % 2 == 0) else 1) if useSubgroups else 1 %}
|
| 49 |
{% macro qn(name, qi) %}{{ name }}{% if qi > 0 %}_q{{ qi }}{% endif %}{% endmacro %}
|
| 50 |
+
{% macro emit_tile_softmax() %}
|
| 51 |
+
// Per-thread online softmax over the tile. s[kk] is reused to hold the
|
| 52 |
+
// exponentiated probabilities for the PV accumulation below.
|
| 53 |
+
{% for qi in range(QPL) %}
|
| 54 |
+
var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
|
| 55 |
+
for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
|
| 56 |
+
{{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
|
| 57 |
+
}
|
| 58 |
+
let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
|
| 59 |
+
let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
|
| 60 |
+
var {{ qn("pSum", qi) }}: f32 = 0.0;
|
| 61 |
+
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 62 |
+
let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
|
| 63 |
+
{{ qn("s", qi) }}[kk] = pk;
|
| 64 |
+
{{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
|
| 65 |
+
}
|
| 66 |
+
{{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
|
| 67 |
+
{{ qn("m", qi) }} = {{ qn("newMax", qi) }};
|
| 68 |
+
{% endfor %}
|
| 69 |
+
{% endmacro %}
|
| 70 |
{% set ROPE_LANE_XOR = ((LPQ / 2) | int) %}
|
| 71 |
{% set QL = qLayout if qLayout is defined else layout %}
|
| 72 |
{% set KL = kvLayout if kvLayout is defined else layout %}
|
|
|
|
| 91 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
| 92 |
{% if QL != "bhsd" %}
|
| 93 |
const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
|
| 94 |
+
{% endif %}
|
|
|
|
| 95 |
{% if KL != "bhsd" %}
|
| 96 |
const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
|
| 97 |
{% endif %}
|
|
|
|
| 112 |
{% else %}
|
| 113 |
const WG: u32 = TILE_Q * LPQ;
|
| 114 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
const NEG_INF: f32 = -3.4028234663852886e38;
|
|
|
|
| 116 |
|
| 117 |
var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
| 118 |
var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
|
|
|
| 133 |
{% endif %}
|
| 134 |
{% endif %}
|
| 135 |
|
| 136 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
|
|
|
|
|
|
|
|
|
| 137 |
if (params.scale != 0.0) { return params.scale; }
|
| 138 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
|
|
|
| 139 |
}
|
|
|
|
| 140 |
{% if quantCacheFormat %}
|
| 141 |
// A quantized cache is dequantized once per key into the staged tile, then read
|
| 142 |
// by all TILE_Q queries in the workgroup. The unpack cost is amortized over the
|
|
|
|
| 153 |
{{ scaleBuffer }}[base + 2u],
|
| 154 |
{{ scaleBuffer }}[base + 3u]
|
| 155 |
);
|
| 156 |
+
}{% endmacro %}
|
| 157 |
+
{% macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
|
|
|
|
| 158 |
{{ emit_quant_scale4(kind, scaleBuffer) }}
|
| 159 |
fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
|
| 160 |
+
{% if format == "int8" %}
|
| 161 |
return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
|
| 162 |
+
{% else %}
|
| 163 |
// Two elements cover this vec4: each carries two +8-biased nibbles, low first.
|
| 164 |
let rowBase = indexV4 - d4;
|
| 165 |
let lo = {{ buffer }}[rowBase + d4 * 2u];
|
|
|
|
| 170 |
);
|
| 171 |
let signed = nibbles - vec4<i32>(8);
|
| 172 |
return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
|
| 173 |
+
{% endif %}
|
| 174 |
+
}{% endmacro %}
|
|
|
|
|
|
|
| 175 |
{{ emit_quant_load4(quantCacheFormat, "key", KEY, "k_scale") }}
|
| 176 |
{{ emit_quant_load4(quantCacheFormat, "value", VALUE, "v_scale") }}
|
| 177 |
{% endif %}
|
| 178 |
|
| 179 |
+
{% if hasBias %}
|
| 180 |
+
{% set BW = "" %}
|
| 181 |
+
{% set BC = "" %}
|
| 182 |
+
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 183 |
+
let offset = base + d4 * 4u;
|
| 184 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 185 |
+
}
|
| 186 |
+
|
| 187 |
+
{% endif %}
|
| 188 |
@compute @workgroup_size(WG, 1, 1)
|
| 189 |
fn main(
|
| 190 |
@builtin(workgroup_id) wg: vec3<u32>,
|
|
|
|
| 223 |
{% endif %}
|
| 224 |
{% for c in range(SLICE_COUNT) %}
|
| 225 |
var {{ qn("qr" ~ c, qi) }} = vec4<f32>({{ QUERY }}[{{ qn("qBase4", qi) }} + {{ c }}u]);
|
| 226 |
+
{% if hasBias %}
|
| 227 |
+
{{ qn("qr" ~ c, qi) }} = {{ qn("qr" ~ c, qi) }} + load_bias4(h * HEAD_DIM, lane8 * SLICE + {{ c }}u);
|
| 228 |
+
{% endif %}
|
| 229 |
var {{ qn("o" ~ c, qi) }} = vec4<f32>(0.0);
|
| 230 |
{% endfor %}
|
| 231 |
{% endfor %}
|
|
|
|
| 405 |
// TILE_K is a small shader constant; the loop updates the named q/o slices in place.
|
| 406 |
{% if not useSubgroups and batchNoSgReduction %}
|
| 407 |
// First publish every key's partial dot without intervening barriers.
|
| 408 |
+
{% if not SHARE_PROB %}
|
| 409 |
var s: array<f32, TILE_K>;
|
| 410 |
+
{% endif %}
|
| 411 |
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 412 |
+
{% if not SHARE_PROB %}
|
| 413 |
s[kk] = NEG_INF;
|
| 414 |
+
{% endif %}
|
| 415 |
var part: f32 = 0.0;
|
| 416 |
let kb = kk * HEAD_DIM_V4 + lane8 * SLICE;
|
| 417 |
{% for c in range(SLICE_COUNT) %}
|
|
|
|
| 441 |
workgroupBarrier();
|
| 442 |
// The barrier at the start of the next K tile protects this scratch before
|
| 443 |
// it is overwritten.
|
| 444 |
+
{% if SHARE_PROB %}
|
| 445 |
+
// One lane computes the query softmax, reusing the reduced-score slots for
|
| 446 |
+
// probabilities. The discarded partial lanes hold the three online states.
|
| 447 |
+
if (lane8 == 0u) {
|
| 448 |
+
var s: array<f32, TILE_K>;
|
| 449 |
+
{% endif %}
|
| 450 |
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 451 |
let kj = kStart + kk;
|
| 452 |
let part = red[kk * WG + tid - lane8];
|
| 453 |
+
{% if SHARE_PROB %}
|
| 454 |
+
s[kk] = NEG_INF;
|
| 455 |
+
{% endif %}
|
| 456 |
if (kj >= minKj && kj < maxKj) {
|
| 457 |
{% if hasMask %}
|
| 458 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 459 |
s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
|
|
|
|
| 460 |
{% else %}
|
| 461 |
s[kk] = {{ score_expr("part") }};
|
| 462 |
{% endif %}
|
| 463 |
}
|
| 464 |
}
|
| 465 |
+
{% if SHARE_PROB %}
|
| 466 |
+
{{ emit_tile_softmax() }}
|
| 467 |
+
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 468 |
+
red[kk * WG + cbase] = s[kk];
|
| 469 |
+
}
|
| 470 |
+
red[cbase + 1u] = corr;
|
| 471 |
+
red[cbase + 2u] = m;
|
| 472 |
+
red[cbase + 3u] = l;
|
| 473 |
+
}
|
| 474 |
+
workgroupBarrier();
|
| 475 |
+
let corr = red[cbase + 1u];
|
| 476 |
+
m = red[cbase + 2u];
|
| 477 |
+
l = red[cbase + 3u];
|
| 478 |
+
{% endif %}
|
| 479 |
{% elif QPL > 1 %}
|
| 480 |
{% for qi in range(QPL) %}
|
| 481 |
var {{ qn("s", qi) }}: array<f32, TILE_K>;
|
|
|
|
| 515 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 516 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
|
| 517 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 518 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
|
|
|
|
| 519 |
{% else %}
|
| 520 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
|
| 521 |
{% endif %}
|
|
|
|
| 561 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 562 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
| 563 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 564 |
s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
|
|
|
|
| 565 |
{% else %}
|
| 566 |
s[kk] = {{ score_expr("part") }};
|
| 567 |
{% endif %}
|
|
|
|
| 569 |
}
|
| 570 |
{% endif %}
|
| 571 |
|
| 572 |
+
{% if not SHARE_PROB %}
|
| 573 |
+
{{ emit_tile_softmax() }}
|
| 574 |
+
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 575 |
// A boundary tile can address V rows outside a query's attended range, and
|
| 576 |
// a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
|
| 577 |
// NaN, the guarded loop selects the V operand away for range-excluded keys.
|
| 578 |
// Interior tiles use the unguarded FMA chain. Mask exclusion applies to
|
| 579 |
// materialized V rows and does not require this range guard.
|
| 580 |
+
{% if SHARE_PROB %}
|
| 581 |
+
{% macro shared_pv(masked) %}
|
| 582 |
+
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 583 |
+
let weight = red[kk * WG + cbase];
|
| 584 |
+
let vb = kk * HEAD_DIM_V4 + lane8 * SLICE;
|
| 585 |
+
{% for c in range(SLICE_COUNT) %}
|
| 586 |
+
o{{ c }} = o{{ c }} + weight * {% if masked %}select(vec4<f32>(), {% endif %}vec4<f32>(v_tile[vb + {{ c }}u]){% if masked %}, weight != 0.0){% endif %};
|
| 587 |
+
{% endfor %}
|
| 588 |
+
}
|
| 589 |
+
{% endmacro %}
|
| 590 |
+
// Consume each shared probability once across all output slices. This
|
| 591 |
+
// retains the key order of every output accumulator without a private tile.
|
| 592 |
+
{% for c in range(SLICE_COUNT) %}
|
| 593 |
+
o{{ c }} = o{{ c }} * corr;
|
| 594 |
+
{% endfor %}
|
| 595 |
+
let tileInterior = kStart >= minKj && kStart + TILE_K <= maxKj;
|
| 596 |
+
if (tileInterior) {
|
| 597 |
+
{{ shared_pv(false) }}
|
| 598 |
+
} else {
|
| 599 |
+
{{ shared_pv(true) }}
|
| 600 |
+
}
|
| 601 |
+
{% elif QPL > 1 %}
|
| 602 |
let tileInterior = {% for qi in range(QPL) %}(kStart >= {{ qn("minKj", qi) }} && kStart + TILE_K <= {{ qn("maxKj", qi) }}){% if not loop.last %} && {% endif %}{% endfor %};
|
| 603 |
{% for c in range(SLICE_COUNT) %}
|
| 604 |
{
|
|
|
|
| 647 |
kStart = kStart + TILE_K;
|
| 648 |
}
|
| 649 |
|
| 650 |
+
{% macro attention_value(c, qi) %}{{ qn("o" ~ c, qi) }} * inv{% if hasBias %} + load_bias4(2u * {{ Q_HIDDEN }}u + h * HEAD_DIM, lane8 * SLICE + {{ c }}u){% endif %}{% endmacro %}
|
| 651 |
{% for qi in range(QPL) %}
|
| 652 |
if ({{ qn("qValid", qi) }}) {
|
| 653 |
{% if QL == "bhsd" %}
|
build/webgpu/attn-flash-q32-broadcast.wgsl.jinja
CHANGED
|
@@ -1,4 +1,5 @@
|
|
| 1 |
-
// Register-resident flash prefill uses one 32-lane subgroup per workgroup
|
|
|
|
| 2 |
// lane owns a query row and keeps its Q slice and f32 output in registers. Lanes
|
| 3 |
// cooperatively load K/V and broadcast them with subgroupShuffle, so each query
|
| 4 |
// computes q·k and p·v without cross-lane reductions or workgroup storage.
|
|
@@ -10,10 +11,15 @@
|
|
| 10 |
{% set ST = "f16" if usesF16 else "f32" %}
|
| 11 |
{% set components = ["x", "y", "z", "w"] %}
|
| 12 |
{% set USE_SUBGROUPS = useSubgroups if useSubgroups is defined else true %}
|
|
|
|
| 13 |
{% set Q_STEP = qStep if qStep is defined else 32 %}
|
|
|
|
| 14 |
{% if USE_SUBGROUPS %}
|
| 15 |
enable subgroups;
|
| 16 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
| 17 |
{{ env.wgsl.resourceDeclarations }}
|
| 18 |
|
| 19 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
|
@@ -43,8 +49,7 @@ fn scale_value() -> f32 {
|
|
| 43 |
}
|
| 44 |
|
| 45 |
|
| 46 |
-
|
| 47 |
-
@compute @workgroup_size({{ Q_STEP }}, 1, 1)
|
| 48 |
fn main(
|
| 49 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 50 |
@builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
|
|
@@ -85,8 +90,19 @@ fn main(
|
|
| 85 |
// Causal key ceiling. The key-loop bound (kvEnd) uses the workgroup's LAST query
|
| 86 |
// so every lane shares a uniform trip count (subgroup ops stay reconverged);
|
| 87 |
// each lane masks its own keys past myMaxKj to NEG_INF.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 88 |
let kvEnd = params.kvSeq;
|
| 89 |
let myMaxKj = params.kvSeq;
|
|
|
|
| 90 |
let kvBase = b * params.kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
|
| 91 |
let kvKeyStride = KV_HIDDEN_V4;
|
| 92 |
|
|
@@ -174,6 +190,13 @@ fn main(
|
|
| 174 |
{% endfor %}
|
| 175 |
previous_max = new_max;
|
| 176 |
previous_denom = denom;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 177 |
|
| 178 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 179 |
{% if USE_SUBGROUPS %}
|
|
@@ -190,21 +213,37 @@ fn main(
|
|
| 190 |
}
|
| 191 |
{% endif %}
|
| 192 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 193 |
var acc: vec4<f32> = vec4<f32>(0.0);
|
|
|
|
| 194 |
{% for g in range(qkGroups) %}
|
| 195 |
{% for lane in range(4) %}
|
| 196 |
-
{% if USE_SUBGROUPS %}
|
| 197 |
-
{%
|
| 198 |
-
|
|
|
|
| 199 |
{% else %}
|
| 200 |
-
|
| 201 |
{% endif %}
|
|
|
|
|
|
|
| 202 |
{% else %}
|
| 203 |
-
acc = acc + vec4<f32>(
|
| 204 |
{% endif %}
|
| 205 |
{% endfor %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 206 |
{% endfor %}
|
|
|
|
|
|
|
|
|
|
| 207 |
o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
|
|
|
|
| 208 |
}
|
| 209 |
{% if not USE_SUBGROUPS %}
|
| 210 |
workgroupBarrier();
|
|
|
|
| 1 |
+
// Register-resident flash prefill uses one 32-lane subgroup per workgroup (fixed
|
| 2 |
+
// by the adapter, or pinned where the adapter can compile exactly 32 lanes). Each
|
| 3 |
// lane owns a query row and keeps its Q slice and f32 output in registers. Lanes
|
| 4 |
// cooperatively load K/V and broadcast them with subgroupShuffle, so each query
|
| 5 |
// computes q·k and p·v without cross-lane reductions or workgroup storage.
|
|
|
|
| 11 |
{% set ST = "f16" if usesF16 else "f32" %}
|
| 12 |
{% set components = ["x", "y", "z", "w"] %}
|
| 13 |
{% set USE_SUBGROUPS = useSubgroups if useSubgroups is defined else true %}
|
| 14 |
+
{% set PIN_SUBGROUP_32 = USE_SUBGROUPS and pinSubgroupSize32 is defined and pinSubgroupSize32 %}
|
| 15 |
{% set Q_STEP = qStep if qStep is defined else 32 %}
|
| 16 |
+
{% set CAUSAL = false if (hasCausal is defined and not hasCausal) else true %}
|
| 17 |
{% if USE_SUBGROUPS %}
|
| 18 |
enable subgroups;
|
| 19 |
{% endif %}
|
| 20 |
+
{% if PIN_SUBGROUP_32 %}
|
| 21 |
+
enable subgroup_size_control;
|
| 22 |
+
{% endif %}
|
| 23 |
{{ env.wgsl.resourceDeclarations }}
|
| 24 |
|
| 25 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
|
|
|
| 49 |
}
|
| 50 |
|
| 51 |
|
| 52 |
+
@compute @workgroup_size({{ Q_STEP }}, 1, 1){{ " @subgroup_size(32)" if PIN_SUBGROUP_32 else "" }}
|
|
|
|
| 53 |
fn main(
|
| 54 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 55 |
@builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
|
|
|
|
| 90 |
// Causal key ceiling. The key-loop bound (kvEnd) uses the workgroup's LAST query
|
| 91 |
// so every lane shares a uniform trip count (subgroup ops stay reconverged);
|
| 92 |
// each lane masks its own keys past myMaxKj to NEG_INF.
|
| 93 |
+
{% if CAUSAL %}
|
| 94 |
+
{% set FIXED_CAUSAL = fixedCausal is defined and fixedCausal %}
|
| 95 |
+
{% macro key_ceiling(query) %}
|
| 96 |
+
{% if FIXED_CAUSAL %}
|
| 97 |
+
min({{ query }} + 1u, params.kvSeq){% else %}
|
| 98 |
+
select(params.kvSeq, min({{ query }} + 1u, params.kvSeq), params.isCausal != 0u){% endif %}{% endmacro %}
|
| 99 |
+
let lastQ = min(wg.x * Q_STEP + Q_STEP - 1u, params.qSeq - 1u);
|
| 100 |
+
let kvEnd = {{ key_ceiling("lastQ") }};
|
| 101 |
+
let myMaxKj = {{ key_ceiling("qi") }};
|
| 102 |
+
{% else %}
|
| 103 |
let kvEnd = params.kvSeq;
|
| 104 |
let myMaxKj = params.kvSeq;
|
| 105 |
+
{% endif %}
|
| 106 |
let kvBase = b * params.kvSeq * KV_HIDDEN_V4 + hKv * HEAD_DIM_V4;
|
| 107 |
let kvKeyStride = KV_HIDDEN_V4;
|
| 108 |
|
|
|
|
| 190 |
{% endfor %}
|
| 191 |
previous_max = new_max;
|
| 192 |
previous_denom = denom;
|
| 193 |
+
{% set PV_HALF = ST == "f16" and not (precisePv | default(false)) %}
|
| 194 |
+
{% set PV_FLUSH_GROUPS = 4 %}
|
| 195 |
+
{% if PV_HALF %}
|
| 196 |
+
{% for g in range(qkGroups) %}
|
| 197 |
+
let p{{ g }} = vec4<f16>(qk{{ g }});
|
| 198 |
+
{% endfor %}
|
| 199 |
+
{% endif %}
|
| 200 |
|
| 201 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 202 |
{% if USE_SUBGROUPS %}
|
|
|
|
| 213 |
}
|
| 214 |
{% endif %}
|
| 215 |
{% endif %}
|
| 216 |
+
{% if PV_HALF %}
|
| 217 |
+
var acc_f32: vec4<f32> = vec4<f32>(0.0);
|
| 218 |
+
var acc: vec4<f16> = vec4<f16>(0.0);
|
| 219 |
+
{% else %}
|
| 220 |
var acc: vec4<f32> = vec4<f32>(0.0);
|
| 221 |
+
{% endif %}
|
| 222 |
{% for g in range(qkGroups) %}
|
| 223 |
{% for lane in range(4) %}
|
| 224 |
+
{% if not USE_SUBGROUPS %}
|
| 225 |
+
{% set pvValue = "valueTile[d4 * K_STEP + " ~ (g * 4 + lane) ~ "u]" %}
|
| 226 |
+
{% elif g < 8 %}
|
| 227 |
+
{% set pvValue = "subgroupShuffle(v_local0, " ~ (g * 4 + lane) ~ "u)" %}
|
| 228 |
{% else %}
|
| 229 |
+
{% set pvValue = "subgroupShuffle(v_local1, " ~ ((g - 8) * 4 + lane) ~ "u)" %}
|
| 230 |
{% endif %}
|
| 231 |
+
{% if PV_HALF %}
|
| 232 |
+
acc = fma({{ pvValue }}, vec4<f16>(p{{ g }}.{{ components[lane] }}), acc);
|
| 233 |
{% else %}
|
| 234 |
+
acc = acc + vec4<f32>({{ pvValue }}) * qk{{ g }}.{{ components[lane] }};
|
| 235 |
{% endif %}
|
| 236 |
{% endfor %}
|
| 237 |
+
{% if PV_HALF and (g + 1) % PV_FLUSH_GROUPS == 0 %}
|
| 238 |
+
acc_f32 = acc_f32 + vec4<f32>(acc);
|
| 239 |
+
acc = vec4<f16>(0.0);
|
| 240 |
+
{% endif %}
|
| 241 |
{% endfor %}
|
| 242 |
+
{% if PV_HALF %}
|
| 243 |
+
o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc_f32;
|
| 244 |
+
{% else %}
|
| 245 |
o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
|
| 246 |
+
{% endif %}
|
| 247 |
}
|
| 248 |
{% if not USE_SUBGROUPS %}
|
| 249 |
workgroupBarrier();
|
build/webgpu/attn-materialized-rowstats-combine-f32.wgsl.jinja
CHANGED
|
@@ -12,6 +12,12 @@
|
|
| 12 |
// `maxOnly` means the producer published only a row max per slot because
|
| 13 |
// computing the denominator there would double the exp count. Fold maxima and
|
| 14 |
// leave the denominator to the apply pass, which sees every row element anyway.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 15 |
{{ env.wgsl.resourceDeclarations }}
|
| 16 |
|
| 17 |
const SLOTS: u32 = {{ statSlots }}u;
|
|
@@ -37,6 +43,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 37 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 38 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 39 |
}
|
|
|
|
| 40 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 41 |
return exp(shifted_value(value, maxValue));
|
| 42 |
}
|
|
@@ -45,8 +52,7 @@ fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
|
| 45 |
fn main(
|
| 46 |
@builtin(global_invocation_id) gid: vec3<u32>
|
| 47 |
) {
|
| 48 |
-
|
| 49 |
-
if (row >= params.rows) { return; }
|
| 50 |
|
| 51 |
// `row` already runs over (batch, head, query) together, and the partial
|
| 52 |
// layout puts that same product one axis out from the slot, so the stride
|
|
|
|
| 12 |
// `maxOnly` means the producer published only a row max per slot because
|
| 13 |
// computing the denominator there would double the exp count. Fold maxima and
|
| 14 |
// leave the denominator to the apply pass, which sees every row element anyway.
|
| 15 |
+
{% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
|
| 16 |
+
{% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
|
| 17 |
+
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 18 |
+
// per-axis workgroup fold width.
|
| 19 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
|
| 20 |
+
if ({{ name }} >= {{ bound }}) { return; }{% endmacro %}
|
| 21 |
{{ env.wgsl.resourceDeclarations }}
|
| 22 |
|
| 23 |
const SLOTS: u32 = {{ statSlots }}u;
|
|
|
|
| 43 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 44 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 45 |
}
|
| 46 |
+
|
| 47 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 48 |
return exp(shifted_value(value, maxValue));
|
| 49 |
}
|
|
|
|
| 52 |
fn main(
|
| 53 |
@builtin(global_invocation_id) gid: vec3<u32>
|
| 54 |
) {
|
| 55 |
+
{{ flat_index_2d("WG", "row", "params.rows", guardInline=true) }}
|
|
|
|
| 56 |
|
| 57 |
// `row` already runs over (batch, head, query) together, and the partial
|
| 58 |
// layout puts that same product one axis out from the slot, so the stride
|
build/webgpu/attn-materialized-sgmat-f32.wgsl.jinja
CHANGED
|
@@ -1,7 +1,5 @@
|
|
| 1 |
{% set MT = "f16" if (operandF16 is defined and operandF16) else "f32" %}
|
| 2 |
-
{%
|
| 3 |
-
enable f16;
|
| 4 |
-
{% endif %}
|
| 5 |
enable subgroups;
|
| 6 |
{% if pinSubgroupSize32 %}
|
| 7 |
enable subgroup_size_control;
|
|
@@ -9,18 +7,16 @@ enable subgroup_size_control;
|
|
| 9 |
enable chromium_experimental_subgroup_matrix;
|
| 10 |
diagnostic(off, chromium.subgroup_matrix_uniformity);
|
| 11 |
|
| 12 |
-
|
| 13 |
{{ env.wgsl.resourceDeclarations }}
|
| 14 |
|
| 15 |
-
{% set layout = layout | default("bsh") %}
|
| 16 |
-
{% set headMajor = layout == "bhsd" %}
|
| 17 |
{% set kvHeadMajor = (kvLayout | default(layout)) == "bhsd" %}
|
| 18 |
-
{% set CAUSAL_UPPER_LEFT =
|
| 19 |
{% set CAUSAL = (causalRightAlign is defined and causalRightAlign) or CAUSAL_UPPER_LEFT %}
|
| 20 |
-
{% macro q_index(row, d) %}
|
| 21 |
{% macro kv_index(seq, d) %}{% if kvHeadMajor %}((b * KV_HEADS + h_kv) * params.kvSeq + {{ seq }}) * HEAD_DIM + {{ d }}{% else %}(b * params.kvSeq + {{ seq }}) * KV_HIDDEN + h_kv * HEAD_DIM + {{ d }}{% endif %}{% endmacro %}
|
| 22 |
-
{% set OUT_ROW_STRIDE = "
|
| 23 |
{% set KV_ROW_STRIDE = "HEAD_DIM" if kvHeadMajor else "KV_HIDDEN" %}
|
|
|
|
| 24 |
{% set scorePhase = phase == "score" %}
|
| 25 |
{% set SCORE_BIAS = scorePhase and scoreBias is defined and scoreBias %}
|
| 26 |
{% set SCORE_WINDOW = CAUSAL and scoreWindow is defined and scoreWindow %}
|
|
@@ -28,45 +24,38 @@ diagnostic(off, chromium.subgroup_matrix_uniformity);
|
|
| 28 |
{% set FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
|
| 29 |
{% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
|
| 30 |
{% macro score_value(index, guard) %}
|
| 31 |
-
{% if
|
| 32 |
-
{% if PRIVATE_ROW_STATS %}
|
| 33 |
-
select(0.0, exp_shift(scores[
|
| 34 |
-
{%- else %}
|
| 35 |
-
select(0.0, exp_shift(scores[{{ index }}], softmax_m[{{ guard[0] }}]) / softmax_d[{{ guard[0] }}], {{ guard[1] }})
|
| 36 |
-
{%- endif %}
|
| 37 |
-
{% else %}
|
| 38 |
-
select(0.0, scores[{{ index }}], {{ guard[1] }})
|
| 39 |
-
{%- endif %}
|
| 40 |
-
{% endmacro %}
|
| 41 |
{% set TILE_M_VALUE = materializedSgmatQueryTile %}
|
| 42 |
{% set TILE_N_VALUE = materializedSgmatKeyTile %}
|
| 43 |
{% set TILE_K_VALUE = materializedSgmatInnerTile %}
|
| 44 |
-
{% set SUB_ROWS_VALUE =
|
| 45 |
-
{% set SUB_COLS_VALUE =
|
| 46 |
-
{% set ROW_BLOCKS =
|
| 47 |
-
{% set COL_BLOCKS =
|
| 48 |
{% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
|
| 49 |
{% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
|
| 50 |
{% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
|
| 51 |
{% set WORKGROUP_THREADS = SUBGROUP_COUNT * 32 %}
|
| 52 |
{% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
|
| 53 |
{% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
|
| 54 |
-
{% set DIRECT_SCORE_STORE =
|
| 55 |
-
{% set DIRECT_APPLY_STORE =
|
| 56 |
{% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
|
| 57 |
{% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
|
| 58 |
and materializedSgmatRuntimeDirectStore
|
| 59 |
-
and (scorePhase or (not hasBias and MT == "f32"))
|
| 60 |
and not EMIT_ROW_STATS %}
|
| 61 |
{% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
|
| 62 |
-
{%
|
|
|
|
|
|
|
| 63 |
|
| 64 |
const HEADS: u32 = {{ qNumHeads }}u;
|
| 65 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 66 |
const HEAD_DIM: u32 = {{ headDim }}u;
|
| 67 |
-
{% if layout != "bhsd" or hasBias %}
|
| 68 |
const HIDDEN: u32 = {{ qHidden }}u;
|
| 69 |
-
{% endif %}
|
| 70 |
{% if hasBias %}
|
| 71 |
/* Packed [Q; K; V] bias. Q folds into the query tile before the GEMM. K is
|
| 72 |
* omitted: expanding (q + bq).(k + bk) leaves a (q + bq).bk term that is
|
|
@@ -108,14 +97,12 @@ var<workgroup> softmax_m: array<f32, {{ TILE_M_VALUE }}>;
|
|
| 108 |
var<workgroup> softmax_d: array<f32, {{ TILE_M_VALUE }}>;
|
| 109 |
{% endif %}
|
| 110 |
{% if FUSED_SOFTMAX or EMIT_ROW_STATS %}
|
| 111 |
-
{% set stableUsage = stableHelperUsage if stableHelperUsage is defined else "all" %}
|
| 112 |
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 113 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
| 114 |
// accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
|
| 115 |
// a zero final denominator according to their public semantics. Using -inf
|
| 116 |
// here changes +inf-row behavior.
|
| 117 |
const FLT_MAX: f32 = 3.4028234663852886e38;
|
| 118 |
-
{% if stableUsage != "constant" %}
|
| 119 |
|
| 120 |
fn is_finite_f32(value: f32) -> bool {
|
| 121 |
return select(false, value <= FLT_MAX, value >= -FLT_MAX);
|
|
@@ -129,14 +116,10 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 129 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 130 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 131 |
}
|
| 132 |
-
{%- endif %}
|
| 133 |
-
{% if stableUsage == "all" %}
|
| 134 |
|
| 135 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 136 |
return exp(shifted_value(value, maxValue));
|
| 137 |
}
|
| 138 |
-
{%- endif %}
|
| 139 |
-
|
| 140 |
{% endif %}
|
| 141 |
|
| 142 |
@compute @workgroup_size({{ WORKGROUP_THREADS }}, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
|
|
@@ -177,7 +160,7 @@ fn main(
|
|
| 177 |
let stat_row = m_base + idx / {{ SUBGROUP_COLS }}u;
|
| 178 |
if (stat_row < params.qSeq) {
|
| 179 |
let slot = wg.x * {{ SUBGROUP_COLS }}u + idx % {{ SUBGROUP_COLS }}u;
|
| 180 |
-
let out_index = ((
|
| 181 |
scorePartials[out_index] = -FLT_MAX;
|
| 182 |
scorePartials[out_index + 1u] = 0.0;
|
| 183 |
}
|
|
@@ -211,30 +194,18 @@ fn main(
|
|
| 211 |
{% endif %}
|
| 212 |
{% endif %}
|
| 213 |
{% if FUSED_SOFTMAX %}
|
| 214 |
-
{% if PRIVATE_ROW_STATS %}
|
| 215 |
-
// In the admitted BM64/BN64/BK32/WG256 loader, four adjacent lanes own the
|
| 216 |
-
// same query row for every reduction tile. Keep that row's constants private:
|
| 217 |
-
// this removes both 512 bytes of workgroup storage and the initialization
|
| 218 |
-
// barrier while preserving the exact exp/divide sequence of the shared-memory
|
| 219 |
-
// row-stats arm.
|
| 220 |
-
let private_stat_row =
|
| 221 |
-
(b * HEADS + h) * params.qSeq + min(m_base + li / 4u, params.qSeq - 1u);
|
| 222 |
-
let private_softmax_m = rowStats[private_stat_row * 2u];
|
| 223 |
-
let private_softmax_d = rowStats[private_stat_row * 2u + 1u];
|
| 224 |
-
{% else %}
|
| 225 |
// One row-stats pair per query row of the tile. A query tail clamps to the last
|
| 226 |
// real row rather than reading past the buffer; those lanes are discarded by the
|
| 227 |
// staging guard anyway, and the clamp keeps the denominator non-zero.
|
| 228 |
for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
|
| 229 |
-
let stat_row =
|
| 230 |
softmax_m[r] = rowStats[stat_row * 2u];
|
| 231 |
softmax_d[r] = rowStats[stat_row * 2u + 1u];
|
| 232 |
}
|
| 233 |
workgroupBarrier();
|
| 234 |
-
{% endif %}
|
| 235 |
{% endif %}
|
| 236 |
for (var k_base = {% if SCORE_WINDOW and not scorePhase %}inner_start{% else %}0u{% endif %}; k_base < {% if CAUSAL and not scorePhase %}inner_bound{% else %}inner{% endif %}; k_base += TILE_K) {
|
| 237 |
-
{% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" %}
|
| 238 |
// Full interior PV tiles can be loaded directly from storage. Query,
|
| 239 |
// reduction, and output-dimension tails use the guarded shared path below.
|
| 240 |
if (
|
|
@@ -244,7 +215,7 @@ fn main(
|
|
| 244 |
) {
|
| 245 |
for (var step = 0u; step < TILE_K; step += 8u) {
|
| 246 |
{% for row_block in range(ROW_BLOCKS) %}
|
| 247 |
-
let score_offset{{ row_block }} =
|
| 248 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
|
| 249 |
var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
|
| 250 |
subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
|
|
@@ -294,7 +265,7 @@ fn main(
|
|
| 294 |
tile_A[a_row * TILE_K + a_col + i] = loaded;
|
| 295 |
{% endif %}
|
| 296 |
{% else %}
|
| 297 |
-
let score_base =
|
| 298 |
tile_A[a_row * TILE_K + a_col + i] =
|
| 299 |
{{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["a_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
|
| 300 |
{% endif %}
|
|
@@ -309,20 +280,20 @@ fn main(
|
|
| 309 |
{% if headDim % 32 == 0 %}
|
| 310 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 311 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 312 |
-
|
| 313 |
col < params.kvSeq
|
| 314 |
);
|
| 315 |
{% else %}
|
| 316 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 317 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 318 |
-
loaded =
|
| 319 |
}
|
| 320 |
tile_B[b_row * TILE_K + b_col + i] = loaded;
|
| 321 |
{% endif %}
|
| 322 |
{% else %}
|
| 323 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 324 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 325 |
-
|
| 326 |
k < params.kvSeq && col < HEAD_DIM
|
| 327 |
);
|
| 328 |
{% endif %}
|
|
@@ -341,13 +312,13 @@ fn main(
|
|
| 341 |
loaded = {{ q_tile_value(q_index("row", "k")) }};
|
| 342 |
}
|
| 343 |
{% elif FUSED_SOFTMAX %}
|
| 344 |
-
let score_base =
|
| 345 |
let loaded =
|
| 346 |
{{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["tile_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
|
| 347 |
{% else %}
|
| 348 |
var loaded = 0.0;
|
| 349 |
if (row < params.qSeq && k < params.kvSeq) {
|
| 350 |
-
let score_base =
|
| 351 |
loaded = scores[score_base + row * params.kvSeq + k];
|
| 352 |
}
|
| 353 |
{% endif %}
|
|
@@ -362,12 +333,12 @@ fn main(
|
|
| 362 |
{% if scorePhase %}
|
| 363 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 364 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 365 |
-
loaded =
|
| 366 |
}
|
| 367 |
{% else %}
|
| 368 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 369 |
if (k < params.kvSeq && col < HEAD_DIM) {
|
| 370 |
-
loaded =
|
| 371 |
}
|
| 372 |
{% endif %}
|
| 373 |
tile_B[idx] = loaded;
|
|
@@ -410,7 +381,7 @@ fn main(
|
|
| 410 |
{% for col_block in range(COL_BLOCKS) %}
|
| 411 |
{% if scorePhase %}
|
| 412 |
let output_offset{{ row_block }}{{ col_block }} =
|
| 413 |
-
|
| 414 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
|
| 415 |
+ n_base + base_B + {{ col_block * 8 }}u;
|
| 416 |
subgroupMatrixStore<row_major>(
|
|
@@ -467,7 +438,9 @@ fn main(
|
|
| 467 |
+ row_in_block * 8u + col_in_block + pair
|
| 468 |
];
|
| 469 |
{% if scorePhase %}
|
|
|
|
| 470 |
let scale = {{ attentionScaleExpression }};
|
|
|
|
| 471 |
{% if CAUSAL %}
|
| 472 |
var scored = result * scale;
|
| 473 |
{% if SCORE_BIAS %}
|
|
@@ -485,10 +458,10 @@ fn main(
|
|
| 485 |
if (i32(col) + i32(params.windowSize) <= kv_causal_off + i32(row)) { scored = -FLT_MAX; }
|
| 486 |
{% endif %}
|
| 487 |
{% else %}
|
| 488 |
-
let scored = result * scale;
|
| 489 |
{% endif %}
|
| 490 |
scores[
|
| 491 |
-
|
| 492 |
] = scored;
|
| 493 |
{% if EMIT_ROW_STATS %}
|
| 494 |
// Softmax sees the STORED value, so the statistics have to be taken on it
|
|
@@ -502,7 +475,7 @@ fn main(
|
|
| 502 |
{% if hasBias %}
|
| 503 |
// V bias row base: skip the packed Q and K blocks, then index this head.
|
| 504 |
{% endif %}
|
| 505 |
-
output[{{ q_index("row", "col") }}] = {{ "f16(" if
|
| 506 |
{% endif %}
|
| 507 |
}
|
| 508 |
}
|
|
@@ -530,7 +503,7 @@ fn main(
|
|
| 530 |
// eight consecutive pairs, and the combine pass reads a slot's whole column
|
| 531 |
// of rows contiguously.
|
| 532 |
let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
|
| 533 |
-
let out_index = ((
|
| 534 |
scorePartials[out_index] = stat_m{{ row_block }};
|
| 535 |
scorePartials[out_index + 1u] = stat_d{{ row_block }};
|
| 536 |
}
|
|
|
|
| 1 |
{% set MT = "f16" if (operandF16 is defined and operandF16) else "f32" %}
|
| 2 |
+
{% set ST = "f16" if (storageF16 is defined and storageF16) else MT %}
|
|
|
|
|
|
|
| 3 |
enable subgroups;
|
| 4 |
{% if pinSubgroupSize32 %}
|
| 5 |
enable subgroup_size_control;
|
|
|
|
| 7 |
enable chromium_experimental_subgroup_matrix;
|
| 8 |
diagnostic(off, chromium.subgroup_matrix_uniformity);
|
| 9 |
|
|
|
|
| 10 |
{{ env.wgsl.resourceDeclarations }}
|
| 11 |
|
|
|
|
|
|
|
| 12 |
{% set kvHeadMajor = (kvLayout | default(layout)) == "bhsd" %}
|
| 13 |
+
{% set CAUSAL_UPPER_LEFT = false %}
|
| 14 |
{% set CAUSAL = (causalRightAlign is defined and causalRightAlign) or CAUSAL_UPPER_LEFT %}
|
| 15 |
+
{% macro q_index(row, d) %}(b * params.qSeq + {{ row }}) * HIDDEN + h * HEAD_DIM + {{ d }}{% endmacro %}
|
| 16 |
{% macro kv_index(seq, d) %}{% if kvHeadMajor %}((b * KV_HEADS + h_kv) * params.kvSeq + {{ seq }}) * HEAD_DIM + {{ d }}{% else %}(b * params.kvSeq + {{ seq }}) * KV_HIDDEN + h_kv * HEAD_DIM + {{ d }}{% endif %}{% endmacro %}
|
| 17 |
+
{% set OUT_ROW_STRIDE = "HIDDEN" %}
|
| 18 |
{% set KV_ROW_STRIDE = "HEAD_DIM" if kvHeadMajor else "KV_HIDDEN" %}
|
| 19 |
+
{% set SCRATCH_HEAD = "(b * HEADS + h)" %}
|
| 20 |
{% set scorePhase = phase == "score" %}
|
| 21 |
{% set SCORE_BIAS = scorePhase and scoreBias is defined and scoreBias %}
|
| 22 |
{% set SCORE_WINDOW = CAUSAL and scoreWindow is defined and scoreWindow %}
|
|
|
|
| 24 |
{% set FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
|
| 25 |
{% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
|
| 26 |
{% macro score_value(index, guard) %}
|
| 27 |
+
{% set rowMax = "private_softmax_m" if PRIVATE_ROW_STATS else "softmax_m[" ~ guard[0] ~ "]" %}
|
| 28 |
+
{% set rowDenom = "private_softmax_d" if PRIVATE_ROW_STATS else "softmax_d[" ~ guard[0] ~ "]" %}
|
| 29 |
+
select(0.0, {{ "exp_shift(scores[" ~ index ~ "], " ~ rowMax ~ ") / " ~ rowDenom if FUSED_SOFTMAX else "scores[" ~ index ~ "]" }}, {{ guard[1] }}){% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
{% set TILE_M_VALUE = materializedSgmatQueryTile %}
|
| 31 |
{% set TILE_N_VALUE = materializedSgmatKeyTile %}
|
| 32 |
{% set TILE_K_VALUE = materializedSgmatInnerTile %}
|
| 33 |
+
{% set SUB_ROWS_VALUE = 16 %}
|
| 34 |
+
{% set SUB_COLS_VALUE = 32 %}
|
| 35 |
+
{% set ROW_BLOCKS = 2 %}
|
| 36 |
+
{% set COL_BLOCKS = 4 %}
|
| 37 |
{% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
|
| 38 |
{% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
|
| 39 |
{% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
|
| 40 |
{% set WORKGROUP_THREADS = SUBGROUP_COUNT * 32 %}
|
| 41 |
{% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
|
| 42 |
{% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
|
| 43 |
+
{% set DIRECT_SCORE_STORE = false %}
|
| 44 |
+
{% set DIRECT_APPLY_STORE = false %}
|
| 45 |
{% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
|
| 46 |
{% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
|
| 47 |
and materializedSgmatRuntimeDirectStore
|
| 48 |
+
and (scorePhase or (not hasBias and MT == "f32" and ST == "f32"))
|
| 49 |
and not EMIT_ROW_STATS %}
|
| 50 |
{% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
|
| 51 |
+
{% set SCALE_IN_Q = DIRECT_SCORE_STORE or (RUNTIME_DIRECT_STORE and scorePhase) %}
|
| 52 |
+
{% macro operand_load(name, index) %}{% if ST != MT %}{{ MT }}({% endif %}{{ name }}[{{ index }}]{% if ST != MT %}){% endif %}{% endmacro %}
|
| 53 |
+
{% macro q_tile_value(index) %}{% if hasBias %}({{ operand_load("query", index) }} + bias[h * HEAD_DIM + k]){% else %}{{ operand_load("query", index) }}{% endif %}{% if SCALE_IN_Q %} * score_scale{% endif %}{% endmacro %}
|
| 54 |
|
| 55 |
const HEADS: u32 = {{ qNumHeads }}u;
|
| 56 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 57 |
const HEAD_DIM: u32 = {{ headDim }}u;
|
|
|
|
| 58 |
const HIDDEN: u32 = {{ qHidden }}u;
|
|
|
|
| 59 |
{% if hasBias %}
|
| 60 |
/* Packed [Q; K; V] bias. Q folds into the query tile before the GEMM. K is
|
| 61 |
* omitted: expanding (q + bq).(k + bk) leaves a (q + bq).bk term that is
|
|
|
|
| 97 |
var<workgroup> softmax_d: array<f32, {{ TILE_M_VALUE }}>;
|
| 98 |
{% endif %}
|
| 99 |
{% if FUSED_SOFTMAX or EMIT_ROW_STATS %}
|
|
|
|
| 100 |
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 101 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
| 102 |
// accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
|
| 103 |
// a zero final denominator according to their public semantics. Using -inf
|
| 104 |
// here changes +inf-row behavior.
|
| 105 |
const FLT_MAX: f32 = 3.4028234663852886e38;
|
|
|
|
| 106 |
|
| 107 |
fn is_finite_f32(value: f32) -> bool {
|
| 108 |
return select(false, value <= FLT_MAX, value >= -FLT_MAX);
|
|
|
|
| 116 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 117 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 118 |
}
|
|
|
|
|
|
|
| 119 |
|
| 120 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 121 |
return exp(shifted_value(value, maxValue));
|
| 122 |
}
|
|
|
|
|
|
|
| 123 |
{% endif %}
|
| 124 |
|
| 125 |
@compute @workgroup_size({{ WORKGROUP_THREADS }}, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
|
|
|
|
| 160 |
let stat_row = m_base + idx / {{ SUBGROUP_COLS }}u;
|
| 161 |
if (stat_row < params.qSeq) {
|
| 162 |
let slot = wg.x * {{ SUBGROUP_COLS }}u + idx % {{ SUBGROUP_COLS }}u;
|
| 163 |
+
let out_index = (({{ SCRATCH_HEAD }} * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
|
| 164 |
scorePartials[out_index] = -FLT_MAX;
|
| 165 |
scorePartials[out_index + 1u] = 0.0;
|
| 166 |
}
|
|
|
|
| 194 |
{% endif %}
|
| 195 |
{% endif %}
|
| 196 |
{% if FUSED_SOFTMAX %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 197 |
// One row-stats pair per query row of the tile. A query tail clamps to the last
|
| 198 |
// real row rather than reading past the buffer; those lanes are discarded by the
|
| 199 |
// staging guard anyway, and the clamp keeps the denominator non-zero.
|
| 200 |
for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
|
| 201 |
+
let stat_row = {{ SCRATCH_HEAD }} * params.qSeq + min(m_base + r, params.qSeq - 1u);
|
| 202 |
softmax_m[r] = rowStats[stat_row * 2u];
|
| 203 |
softmax_d[r] = rowStats[stat_row * 2u + 1u];
|
| 204 |
}
|
| 205 |
workgroupBarrier();
|
|
|
|
| 206 |
{% endif %}
|
| 207 |
for (var k_base = {% if SCORE_WINDOW and not scorePhase %}inner_start{% else %}0u{% endif %}; k_base < {% if CAUSAL and not scorePhase %}inner_bound{% else %}inner{% endif %}; k_base += TILE_K) {
|
| 208 |
+
{% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" and ST == "f32" %}
|
| 209 |
// Full interior PV tiles can be loaded directly from storage. Query,
|
| 210 |
// reduction, and output-dimension tails use the guarded shared path below.
|
| 211 |
if (
|
|
|
|
| 215 |
) {
|
| 216 |
for (var step = 0u; step < TILE_K; step += 8u) {
|
| 217 |
{% for row_block in range(ROW_BLOCKS) %}
|
| 218 |
+
let score_offset{{ row_block }} = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
|
| 219 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
|
| 220 |
var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
|
| 221 |
subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
|
|
|
|
| 265 |
tile_A[a_row * TILE_K + a_col + i] = loaded;
|
| 266 |
{% endif %}
|
| 267 |
{% else %}
|
| 268 |
+
let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
|
| 269 |
tile_A[a_row * TILE_K + a_col + i] =
|
| 270 |
{{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["a_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
|
| 271 |
{% endif %}
|
|
|
|
| 280 |
{% if headDim % 32 == 0 %}
|
| 281 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 282 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 283 |
+
{{ operand_load("key", kv_index("col", "k")) }},
|
| 284 |
col < params.kvSeq
|
| 285 |
);
|
| 286 |
{% else %}
|
| 287 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 288 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 289 |
+
loaded = {{ operand_load("key", kv_index("col", "k")) }};
|
| 290 |
}
|
| 291 |
tile_B[b_row * TILE_K + b_col + i] = loaded;
|
| 292 |
{% endif %}
|
| 293 |
{% else %}
|
| 294 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 295 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 296 |
+
{{ operand_load("value", kv_index("k", "col")) }},
|
| 297 |
k < params.kvSeq && col < HEAD_DIM
|
| 298 |
);
|
| 299 |
{% endif %}
|
|
|
|
| 312 |
loaded = {{ q_tile_value(q_index("row", "k")) }};
|
| 313 |
}
|
| 314 |
{% elif FUSED_SOFTMAX %}
|
| 315 |
+
let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
|
| 316 |
let loaded =
|
| 317 |
{{ "f16(" if MT == "f16" else "" }}{{ score_value("score_base + row * params.kvSeq + k", ["tile_row", "row < params.qSeq && k < params.kvSeq"]) }}{{ ")" if MT == "f16" else "" }};
|
| 318 |
{% else %}
|
| 319 |
var loaded = 0.0;
|
| 320 |
if (row < params.qSeq && k < params.kvSeq) {
|
| 321 |
+
let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
|
| 322 |
loaded = scores[score_base + row * params.kvSeq + k];
|
| 323 |
}
|
| 324 |
{% endif %}
|
|
|
|
| 333 |
{% if scorePhase %}
|
| 334 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 335 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 336 |
+
loaded = {{ operand_load("key", kv_index("col", "k")) }};
|
| 337 |
}
|
| 338 |
{% else %}
|
| 339 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 340 |
if (k < params.kvSeq && col < HEAD_DIM) {
|
| 341 |
+
loaded = {{ operand_load("value", kv_index("k", "col")) }};
|
| 342 |
}
|
| 343 |
{% endif %}
|
| 344 |
tile_B[idx] = loaded;
|
|
|
|
| 381 |
{% for col_block in range(COL_BLOCKS) %}
|
| 382 |
{% if scorePhase %}
|
| 383 |
let output_offset{{ row_block }}{{ col_block }} =
|
| 384 |
+
{{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
|
| 385 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
|
| 386 |
+ n_base + base_B + {{ col_block * 8 }}u;
|
| 387 |
subgroupMatrixStore<row_major>(
|
|
|
|
| 438 |
+ row_in_block * 8u + col_in_block + pair
|
| 439 |
];
|
| 440 |
{% if scorePhase %}
|
| 441 |
+
{% if not SCALE_IN_Q %}
|
| 442 |
let scale = {{ attentionScaleExpression }};
|
| 443 |
+
{% endif %}
|
| 444 |
{% if CAUSAL %}
|
| 445 |
var scored = result * scale;
|
| 446 |
{% if SCORE_BIAS %}
|
|
|
|
| 458 |
if (i32(col) + i32(params.windowSize) <= kv_causal_off + i32(row)) { scored = -FLT_MAX; }
|
| 459 |
{% endif %}
|
| 460 |
{% else %}
|
| 461 |
+
let scored = result{% if not SCALE_IN_Q %} * scale{% endif %};
|
| 462 |
{% endif %}
|
| 463 |
scores[
|
| 464 |
+
{{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq + row * params.kvSeq + col
|
| 465 |
] = scored;
|
| 466 |
{% if EMIT_ROW_STATS %}
|
| 467 |
// Softmax sees the STORED value, so the statistics have to be taken on it
|
|
|
|
| 475 |
{% if hasBias %}
|
| 476 |
// V bias row base: skip the packed Q and K blocks, then index this head.
|
| 477 |
{% endif %}
|
| 478 |
+
output[{{ q_index("row", "col") }}] = {{ "f16(" if ST == "f16" else "" }}result{{ ")" if ST == "f16" else "" }}{% if hasBias %} + bias[2u * HIDDEN + h * HEAD_DIM + col]{% endif %};
|
| 479 |
{% endif %}
|
| 480 |
}
|
| 481 |
}
|
|
|
|
| 503 |
// eight consecutive pairs, and the combine pass reads a slot's whole column
|
| 504 |
// of rows contiguously.
|
| 505 |
let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
|
| 506 |
+
let out_index = (({{ SCRATCH_HEAD }} * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
|
| 507 |
scorePartials[out_index] = stat_m{{ row_block }};
|
| 508 |
scorePartials[out_index + 1u] = stat_d{{ row_block }};
|
| 509 |
}
|
build/webgpu/attn-online-scalar.wgsl.jinja
CHANGED
|
@@ -1,14 +1,17 @@
|
|
| 1 |
{{ env.wgsl.resourceDeclarations }}
|
| 2 |
-
{% set MASK_BATCH = "batch * params.maskBatchStride + "
|
| 3 |
|
| 4 |
// Online-softmax attention fallback with no feature requirements: one
|
| 5 |
// workgroup per (batch, head, query token) walks the keys serially; the
|
| 6 |
-
// workgroup cooperates on each q·k dot (tree reduction) and on the
|
| 7 |
-
// V accumulator, with the online rescale applied per key. This path
|
| 8 |
-
// no subgroup or subgroup-matrix features.
|
| 9 |
// Layout: rank-3 token-major [batch, seq, heads * headDim].
|
| 10 |
// An omitted scale uses 1/sqrt(headDim); explicit zero remains zero through
|
| 11 |
// the scaleIsExplicitZero specialization.
|
|
|
|
|
|
|
|
|
|
| 12 |
{% if hasMask %}
|
| 13 |
// Additive per-score attention bias with broadcast strides: a stride of 0
|
| 14 |
// collapses that axis (batch and/or head broadcast).
|
|
@@ -18,18 +21,39 @@ const Q_HIDDEN: u32 = {{ qHidden }}u;
|
|
| 18 |
const KV_HIDDEN: u32 = {{ kvHidden }}u;
|
| 19 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 20 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 21 |
-
{% set qHeads = "
|
| 22 |
-
{% set kvHeads = "
|
| 23 |
-
{% set scale = scale | default("0.0") %}
|
| 24 |
const WG: u32 = {{ workgroupSize }}u;
|
| 25 |
|
| 26 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
var<workgroup> running_max: f32;
|
| 28 |
var<workgroup> running_denom: f32;
|
| 29 |
var<workgroup> running_out: array<f32, HEAD_DIM>;
|
| 30 |
var<workgroup> previous_scale: f32;
|
| 31 |
-
{% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
|
| 32 |
-
fn {{ name }}(value:
|
| 33 |
{{ buffer }}[tid] = value;
|
| 34 |
workgroupBarrier();
|
| 35 |
// Ceil-halving keeps every lane when the workgroup size is not a power of
|
|
@@ -39,11 +63,7 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
|
|
| 39 |
loop {
|
| 40 |
let half = (n + 1u) / 2u;
|
| 41 |
if (tid < n - half) {
|
| 42 |
-
{
|
| 43 |
-
{{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
|
| 44 |
-
{% else %}
|
| 45 |
-
{{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
|
| 46 |
-
{% endif %}
|
| 47 |
}
|
| 48 |
workgroupBarrier();
|
| 49 |
n = half;
|
|
@@ -55,22 +75,16 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
|
|
| 55 |
// slot 0 here, so the next call's first store must not run until all lanes have read it.
|
| 56 |
// `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
|
| 57 |
let reduced = {{ buffer }}[0];
|
| 58 |
-
{% if trailingBarrier %}
|
| 59 |
workgroupBarrier();
|
| 60 |
-
{% endif %}
|
| 61 |
return reduced;
|
| 62 |
}
|
| 63 |
{% endmacro %}
|
| 64 |
-
|
| 65 |
-
{
|
| 66 |
-
|
| 67 |
-
{% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
|
| 68 |
-
fn scale_value() -> f32 {
|
| 69 |
if (params.scale != 0.0) { return params.scale; }
|
| 70 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 71 |
}
|
| 72 |
|
| 73 |
-
|
| 74 |
@compute @workgroup_size(WG, 1, 1)
|
| 75 |
fn main(
|
| 76 |
@builtin(workgroup_id) wg: vec3<u32>,
|
|
@@ -110,6 +124,9 @@ fn main(
|
|
| 110 |
|
| 111 |
var maxKj = params.kvSeq;
|
| 112 |
var minKj: u32 = 0u;
|
|
|
|
|
|
|
|
|
|
| 113 |
{% if hasWindow %}
|
| 114 |
// local_window_size: the query at relative index
|
| 115 |
// query_token sits at absolute position p = kvSeq - qSeq + query_token, so it
|
|
@@ -124,21 +141,26 @@ fn main(
|
|
| 124 |
|
| 125 |
for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
|
| 126 |
let kRow = kvBase + key_token * kvTokenStride;
|
| 127 |
-
var partial_dot = 0.0;
|
| 128 |
for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
|
| 129 |
var q_value = f32(query[qBase + d]);
|
| 130 |
var k_value = f32(key[kRow + d]);
|
| 131 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 132 |
}
|
| 133 |
|
| 134 |
-
//
|
| 135 |
// workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
|
| 136 |
// through shared memory by the barrier below).
|
| 137 |
{% if hasMask %}
|
| 138 |
let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
|
| 139 |
-
let score =
|
| 140 |
{% else %}
|
| 141 |
-
let score =
|
| 142 |
{% endif %}
|
| 143 |
if (tid == 0u) {
|
| 144 |
let next_max = max(running_max, score);
|
|
@@ -153,6 +175,9 @@ fn main(
|
|
| 153 |
let probability_numerator = exp(score - running_max);
|
| 154 |
for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
|
| 155 |
var v_value = f32(value[kRow + d]);
|
|
|
|
|
|
|
|
|
|
| 156 |
running_out[d] = running_out[d] * previous_scale + probability_numerator * v_value;
|
| 157 |
}
|
| 158 |
workgroupBarrier();
|
|
|
|
| 1 |
{{ env.wgsl.resourceDeclarations }}
|
| 2 |
+
{% set MASK_BATCH = "batch * params.maskBatchStride + " %}
|
| 3 |
|
| 4 |
// Online-softmax attention fallback with no feature requirements: one
|
| 5 |
// workgroup per (batch, head, query token) walks the keys serially; the
|
| 6 |
+
// workgroup cooperates on each q·k dot (compensated tree reduction) and on the
|
| 7 |
+
// running V accumulator, with the online rescale applied per key. This path
|
| 8 |
+
// requires no subgroup or subgroup-matrix features.
|
| 9 |
// Layout: rank-3 token-major [batch, seq, heads * headDim].
|
| 10 |
// An omitted scale uses 1/sqrt(headDim); explicit zero remains zero through
|
| 11 |
// the scaleIsExplicitZero specialization.
|
| 12 |
+
{% if hasBias %}
|
| 13 |
+
// Packed [Q; K; V] bias rows are applied during the serial key walk.
|
| 14 |
+
{% endif %}
|
| 15 |
{% if hasMask %}
|
| 16 |
// Additive per-score attention bias with broadcast strides: a stride of 0
|
| 17 |
// collapses that axis (batch and/or head broadcast).
|
|
|
|
| 21 |
const KV_HIDDEN: u32 = {{ kvHidden }}u;
|
| 22 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 23 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 24 |
+
{% set qHeads = "Q_HEADS" %}
|
| 25 |
+
{% set kvHeads = "KV_HEADS" %}
|
|
|
|
| 26 |
const WG: u32 = {{ workgroupSize }}u;
|
| 27 |
|
| 28 |
+
{% set dotType = "f32" %}
|
| 29 |
+
// Retain product and addition residuals across a dot product. Explicit fma
|
| 30 |
+
// boundaries preserve the addition error transform under reassociation.
|
| 31 |
+
struct DotAccumulator {
|
| 32 |
+
hi: {{ dotType }},
|
| 33 |
+
lo: {{ dotType }},
|
| 34 |
+
}
|
| 35 |
+
fn dot_accumulate(acc: DotAccumulator, a: {{ dotType }}, b: {{ dotType }}) -> DotAccumulator {
|
| 36 |
+
let product = fma(a, b, {{ dotType }}(0.0));
|
| 37 |
+
let productError = fma(a, b, -product);
|
| 38 |
+
let sum = fma(acc.hi, {{ dotType }}(1.0), product);
|
| 39 |
+
let bv = fma({{ dotType }}(-1.0), acc.hi, sum);
|
| 40 |
+
let av = fma({{ dotType }}(-1.0), bv, sum);
|
| 41 |
+
let ae = fma({{ dotType }}(-1.0), av, acc.hi);
|
| 42 |
+
let be = fma({{ dotType }}(-1.0), bv, product);
|
| 43 |
+
let error = fma(fma(fma(ae, {{ dotType }}(1.0), be), {{ dotType }}(1.0), acc.lo), {{ dotType }}(1.0), productError);
|
| 44 |
+
let hi = fma(sum, {{ dotType }}(1.0), error);
|
| 45 |
+
return DotAccumulator(hi, fma({{ dotType }}(-1.0), fma({{ dotType }}(-1.0), sum, hi), error));
|
| 46 |
+
}
|
| 47 |
+
fn dot_value(acc: DotAccumulator) -> {{ dotType }} {
|
| 48 |
+
return fma(acc.hi, {{ dotType }}(1.0), acc.lo);
|
| 49 |
+
}
|
| 50 |
+
var<workgroup> partial: array<DotAccumulator, WG>;
|
| 51 |
var<workgroup> running_max: f32;
|
| 52 |
var<workgroup> running_denom: f32;
|
| 53 |
var<workgroup> running_out: array<f32, HEAD_DIM>;
|
| 54 |
var<workgroup> previous_scale: f32;
|
| 55 |
+
{% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
|
| 56 |
+
fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
|
| 57 |
{{ buffer }}[tid] = value;
|
| 58 |
workgroupBarrier();
|
| 59 |
// Ceil-halving keeps every lane when the workgroup size is not a power of
|
|
|
|
| 63 |
loop {
|
| 64 |
let half = (n + 1u) / 2u;
|
| 65 |
if (tid < n - half) {
|
| 66 |
+
{{ buffer }}[tid] = dot_accumulate(dot_accumulate({{ buffer }}[tid], {{ buffer }}[tid + half].hi, 1.0), {{ buffer }}[tid + half].lo, 1.0);
|
|
|
|
|
|
|
|
|
|
|
|
|
| 67 |
}
|
| 68 |
workgroupBarrier();
|
| 69 |
n = half;
|
|
|
|
| 75 |
// slot 0 here, so the next call's first store must not run until all lanes have read it.
|
| 76 |
// `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
|
| 77 |
let reduced = {{ buffer }}[0];
|
|
|
|
| 78 |
workgroupBarrier();
|
|
|
|
| 79 |
return reduced;
|
| 80 |
}
|
| 81 |
{% endmacro %}
|
| 82 |
+
{{ wgsl_tree_reduce_f32("reduce_dot", "compensated", "partial", "WG", valueType="DotAccumulator") }}
|
| 83 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
|
|
|
|
|
|
| 84 |
if (params.scale != 0.0) { return params.scale; }
|
| 85 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 86 |
}
|
| 87 |
|
|
|
|
| 88 |
@compute @workgroup_size(WG, 1, 1)
|
| 89 |
fn main(
|
| 90 |
@builtin(workgroup_id) wg: vec3<u32>,
|
|
|
|
| 124 |
|
| 125 |
var maxKj = params.kvSeq;
|
| 126 |
var minKj: u32 = 0u;
|
| 127 |
+
{% if hasCausal %}
|
| 128 |
+
maxKj = min(maxKj, select(params.kvSeq, query_token + 1u, params.isCausal != 0u));
|
| 129 |
+
{% endif %}
|
| 130 |
{% if hasWindow %}
|
| 131 |
// local_window_size: the query at relative index
|
| 132 |
// query_token sits at absolute position p = kvSeq - qSeq + query_token, so it
|
|
|
|
| 141 |
|
| 142 |
for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
|
| 143 |
let kRow = kvBase + key_token * kvTokenStride;
|
| 144 |
+
var partial_dot = DotAccumulator(0.0, 0.0);
|
| 145 |
for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
|
| 146 |
var q_value = f32(query[qBase + d]);
|
| 147 |
var k_value = f32(key[kRow + d]);
|
| 148 |
+
{% if hasBias %}
|
| 149 |
+
let channel = h * HEAD_DIM + d;
|
| 150 |
+
q_value = q_value + f32(bias[channel]);
|
| 151 |
+
k_value = k_value + f32(bias[Q_HIDDEN + channel]);
|
| 152 |
+
{% endif %}
|
| 153 |
+
partial_dot = dot_accumulate(partial_dot, q_value, k_value);
|
| 154 |
}
|
| 155 |
|
| 156 |
+
// reduce_dot returns the same partial[0] to every lane, so `score` is already
|
| 157 |
// workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
|
| 158 |
// through shared memory by the barrier below).
|
| 159 |
{% if hasMask %}
|
| 160 |
let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
|
| 161 |
+
let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value() + f32(attn_mask[maskIndex]);
|
| 162 |
{% else %}
|
| 163 |
+
let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value();
|
| 164 |
{% endif %}
|
| 165 |
if (tid == 0u) {
|
| 166 |
let next_max = max(running_max, score);
|
|
|
|
| 175 |
let probability_numerator = exp(score - running_max);
|
| 176 |
for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
|
| 177 |
var v_value = f32(value[kRow + d]);
|
| 178 |
+
{% if hasBias %}
|
| 179 |
+
v_value = v_value + f32(bias[2u * Q_HIDDEN + h * HEAD_DIM + d]);
|
| 180 |
+
{% endif %}
|
| 181 |
running_out[d] = running_out[d] * previous_scale + probability_numerator * v_value;
|
| 182 |
}
|
| 183 |
workgroupBarrier();
|
build/webgpu/bench.json
CHANGED
|
@@ -710,7 +710,9 @@
|
|
| 710 |
{
|
| 711 |
"name": "quant-int8-prefill-h32kv8-d128-s512-pathology",
|
| 712 |
"preset": "stress",
|
| 713 |
-
"provenance": {
|
|
|
|
|
|
|
| 714 |
"attrs": {
|
| 715 |
"num_heads": 32,
|
| 716 |
"kv_num_heads": 8,
|
|
@@ -765,7 +767,9 @@
|
|
| 765 |
{
|
| 766 |
"name": "quant-int4-prefill-h32kv8-d128-s512-pathology",
|
| 767 |
"preset": "stress",
|
| 768 |
-
"provenance": {
|
|
|
|
|
|
|
| 769 |
"attrs": {
|
| 770 |
"num_heads": 32,
|
| 771 |
"kv_num_heads": 8,
|
|
@@ -1043,7 +1047,7 @@
|
|
| 1043 |
"name": "headsink-prefill-h32kv8-d128-s512-generic-pathology",
|
| 1044 |
"preset": "stress",
|
| 1045 |
"provenance": {
|
| 1046 |
-
"notes": "
|
| 1047 |
},
|
| 1048 |
"attrs": { "num_heads": 32, "kv_num_heads": 8 },
|
| 1049 |
"inputs": {
|
|
@@ -1086,7 +1090,7 @@
|
|
| 1086 |
"name": "bias-headsink-prefill-h32kv8-d128-s512-generic-pathology",
|
| 1087 |
"preset": "stress",
|
| 1088 |
"provenance": {
|
| 1089 |
-
"notes": "
|
| 1090 |
},
|
| 1091 |
"attrs": { "num_heads": 32, "kv_num_heads": 8 },
|
| 1092 |
"inputs": {
|
|
@@ -1133,7 +1137,9 @@
|
|
| 1133 |
{
|
| 1134 |
"name": "softcap-prefill-h32kv8-d128-s512-generic-pathology",
|
| 1135 |
"preset": "stress",
|
| 1136 |
-
"provenance": {
|
|
|
|
|
|
|
| 1137 |
"attrs": { "num_heads": 32, "kv_num_heads": 8, "softcap": 30 },
|
| 1138 |
"inputs": {
|
| 1139 |
"queryT": {
|
|
@@ -1460,6 +1466,74 @@
|
|
| 1460 |
]
|
| 1461 |
}
|
| 1462 |
},
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1463 |
{
|
| 1464 |
"name": "sharedkv-chunk-append-32h8kv-d128-cap1024-q128",
|
| 1465 |
"preset": "smoke",
|
|
@@ -1551,7 +1625,97 @@
|
|
| 1551 |
},
|
| 1552 |
"provenance": {
|
| 1553 |
"source": "register-geometry gate asymmetry",
|
| 1554 |
-
"notes": "
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1555 |
}
|
| 1556 |
},
|
| 1557 |
{
|
|
@@ -1576,6 +1740,439 @@
|
|
| 1576 |
{ "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
|
| 1577 |
]
|
| 1578 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1579 |
}
|
| 1580 |
]
|
| 1581 |
}
|
|
|
|
| 710 |
{
|
| 711 |
"name": "quant-int8-prefill-h32kv8-d128-s512-pathology",
|
| 712 |
"preset": "stress",
|
| 713 |
+
"provenance": {
|
| 714 |
+
"notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with an int8 quantized key/value cache."
|
| 715 |
+
},
|
| 716 |
"attrs": {
|
| 717 |
"num_heads": 32,
|
| 718 |
"kv_num_heads": 8,
|
|
|
|
| 767 |
{
|
| 768 |
"name": "quant-int4-prefill-h32kv8-d128-s512-pathology",
|
| 769 |
"preset": "stress",
|
| 770 |
+
"provenance": {
|
| 771 |
+
"notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with a packed int4 quantized key/value cache."
|
| 772 |
+
},
|
| 773 |
"attrs": {
|
| 774 |
"num_heads": 32,
|
| 775 |
"kv_num_heads": 8,
|
|
|
|
| 1047 |
"name": "headsink-prefill-h32kv8-d128-s512-generic-pathology",
|
| 1048 |
"preset": "stress",
|
| 1049 |
"provenance": {
|
| 1050 |
+
"notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with per-head attention sink logits."
|
| 1051 |
},
|
| 1052 |
"attrs": { "num_heads": 32, "kv_num_heads": 8 },
|
| 1053 |
"inputs": {
|
|
|
|
| 1090 |
"name": "bias-headsink-prefill-h32kv8-d128-s512-generic-pathology",
|
| 1091 |
"preset": "stress",
|
| 1092 |
"provenance": {
|
| 1093 |
+
"notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with an additive attention bias and per-head attention sink logits."
|
| 1094 |
},
|
| 1095 |
"attrs": { "num_heads": 32, "kv_num_heads": 8 },
|
| 1096 |
"inputs": {
|
|
|
|
| 1137 |
{
|
| 1138 |
"name": "softcap-prefill-h32kv8-d128-s512-generic-pathology",
|
| 1139 |
"preset": "stress",
|
| 1140 |
+
"provenance": {
|
| 1141 |
+
"notes": "Measures GroupQueryAttention over a 512-token Llama-sized prefill (32 query heads grouped 4:1 over 8 key/value heads, head width 128) with a logit softcap of 30."
|
| 1142 |
+
},
|
| 1143 |
"attrs": { "num_heads": 32, "kv_num_heads": 8, "softcap": 30 },
|
| 1144 |
"inputs": {
|
| 1145 |
"queryT": {
|
|
|
|
| 1466 |
]
|
| 1467 |
}
|
| 1468 |
},
|
| 1469 |
+
{
|
| 1470 |
+
"name": "window-headsink-rotary-decode-gptoss-64h8kv-d64-cap384-w128",
|
| 1471 |
+
"preset": "smoke",
|
| 1472 |
+
"attrs": {
|
| 1473 |
+
"num_heads": 64,
|
| 1474 |
+
"kv_num_heads": 8,
|
| 1475 |
+
"do_rotary": 1,
|
| 1476 |
+
"sliding_window_cache": 1,
|
| 1477 |
+
"local_window_size": 128
|
| 1478 |
+
},
|
| 1479 |
+
"inputs": {
|
| 1480 |
+
"queryT": {
|
| 1481 |
+
"dtype": "float32",
|
| 1482 |
+
"shape": [1, 1, 4096],
|
| 1483 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 1484 |
+
},
|
| 1485 |
+
"keyT": {
|
| 1486 |
+
"dtype": "float32",
|
| 1487 |
+
"shape": [1, 1, 512],
|
| 1488 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 1489 |
+
},
|
| 1490 |
+
"valueT": {
|
| 1491 |
+
"dtype": "float32",
|
| 1492 |
+
"shape": [1, 1, 512],
|
| 1493 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 1494 |
+
},
|
| 1495 |
+
"pastKeyT": {
|
| 1496 |
+
"dtype": "float32",
|
| 1497 |
+
"shape": [1, 8, 384, 64],
|
| 1498 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 1499 |
+
},
|
| 1500 |
+
"pastValueT": {
|
| 1501 |
+
"dtype": "float32",
|
| 1502 |
+
"shape": [1, 8, 384, 64],
|
| 1503 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 1504 |
+
},
|
| 1505 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [4095] } },
|
| 1506 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [4096] } },
|
| 1507 |
+
"cosCacheT": {
|
| 1508 |
+
"dtype": "float32",
|
| 1509 |
+
"shape": [4096, 32],
|
| 1510 |
+
"data": { "kind": "rotaryCos", "thetaStart": 0.0, "thetaStep": 0.017 }
|
| 1511 |
+
},
|
| 1512 |
+
"sinCacheT": {
|
| 1513 |
+
"dtype": "float32",
|
| 1514 |
+
"shape": [4096, 32],
|
| 1515 |
+
"data": { "kind": "rotarySin", "thetaStart": 0.0, "thetaStep": 0.017 }
|
| 1516 |
+
},
|
| 1517 |
+
"headSinkT": {
|
| 1518 |
+
"dtype": "float32",
|
| 1519 |
+
"shape": [64],
|
| 1520 |
+
"data": { "kind": "fillFloat32", "scale": 0.5, "sinStep": 0.041, "cosStep": 0.011 }
|
| 1521 |
+
}
|
| 1522 |
+
},
|
| 1523 |
+
"outputs": {
|
| 1524 |
+
"outputT": { "dtype": "float32", "shape": [1, 1, 4096] },
|
| 1525 |
+
"presentKeyT": { "dtype": "float32", "shape": [1, 8, 384, 64] },
|
| 1526 |
+
"presentValueT": { "dtype": "float32", "shape": [1, 8, 384, 64] }
|
| 1527 |
+
},
|
| 1528 |
+
"bench": {
|
| 1529 |
+
"metrics": [
|
| 1530 |
+
{
|
| 1531 |
+
"type": "bandwidth",
|
| 1532 |
+
"value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
|
| 1533 |
+
}
|
| 1534 |
+
]
|
| 1535 |
+
}
|
| 1536 |
+
},
|
| 1537 |
{
|
| 1538 |
"name": "sharedkv-chunk-append-32h8kv-d128-cap1024-q128",
|
| 1539 |
"preset": "smoke",
|
|
|
|
| 1625 |
},
|
| 1626 |
"provenance": {
|
| 1627 |
"source": "register-geometry gate asymmetry",
|
| 1628 |
+
"notes": "Measures cached float32 prefill that appends 128 tokens after 896 cached ones, filling a 1024-slot key/value buffer, with 8 query heads over 2 key/value heads of width 256."
|
| 1629 |
+
}
|
| 1630 |
+
},
|
| 1631 |
+
{
|
| 1632 |
+
"name": "bidirectional-share-append-32h8kv-d128-cap1024-q128",
|
| 1633 |
+
"preset": "smoke",
|
| 1634 |
+
"vars": { "batch": 1, "qSeq": 128, "kvSeq": 1024, "heads": 32, "headDim": 128 },
|
| 1635 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "causal": 0 },
|
| 1636 |
+
"inputs": {
|
| 1637 |
+
"queryT": {
|
| 1638 |
+
"dtype": "float32",
|
| 1639 |
+
"shape": [1, 128, 4096],
|
| 1640 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 1641 |
+
},
|
| 1642 |
+
"keyT": {
|
| 1643 |
+
"dtype": "float32",
|
| 1644 |
+
"shape": [1, 128, 1024],
|
| 1645 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 1646 |
+
},
|
| 1647 |
+
"valueT": {
|
| 1648 |
+
"dtype": "float32",
|
| 1649 |
+
"shape": [1, 128, 1024],
|
| 1650 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 1651 |
+
},
|
| 1652 |
+
"pastKeyT": {
|
| 1653 |
+
"dtype": "float32",
|
| 1654 |
+
"shape": [1, 8, 1024, 128],
|
| 1655 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 1656 |
+
},
|
| 1657 |
+
"pastValueT": {
|
| 1658 |
+
"dtype": "float32",
|
| 1659 |
+
"shape": [1, 8, 1024, 128],
|
| 1660 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 1661 |
+
},
|
| 1662 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [1023] } },
|
| 1663 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [1024] } }
|
| 1664 |
+
},
|
| 1665 |
+
"outputs": {
|
| 1666 |
+
"outputT": { "dtype": "float32", "shape": [1, 128, 4096] },
|
| 1667 |
+
"presentKeyT": { "dtype": "float32", "shape": [1, 8, 1024, 128] },
|
| 1668 |
+
"presentValueT": { "dtype": "float32", "shape": [1, 8, 1024, 128] }
|
| 1669 |
+
},
|
| 1670 |
+
"bench": {
|
| 1671 |
+
"metrics": [
|
| 1672 |
+
{ "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
|
| 1673 |
+
]
|
| 1674 |
+
}
|
| 1675 |
+
},
|
| 1676 |
+
{
|
| 1677 |
+
"name": "bidirectional-share-append-f16-32h8kv-d128-cap512-q64",
|
| 1678 |
+
"preset": "smoke",
|
| 1679 |
+
"vars": { "batch": 1, "qSeq": 64, "kvSeq": 512, "heads": 32, "headDim": 128 },
|
| 1680 |
+
"attrs": { "num_heads": 32, "kv_num_heads": 8, "causal": 0 },
|
| 1681 |
+
"inputs": {
|
| 1682 |
+
"queryT": {
|
| 1683 |
+
"dtype": "float16",
|
| 1684 |
+
"shape": [1, 64, 4096],
|
| 1685 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 1686 |
+
},
|
| 1687 |
+
"keyT": {
|
| 1688 |
+
"dtype": "float16",
|
| 1689 |
+
"shape": [1, 64, 1024],
|
| 1690 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 1691 |
+
},
|
| 1692 |
+
"valueT": {
|
| 1693 |
+
"dtype": "float16",
|
| 1694 |
+
"shape": [1, 64, 1024],
|
| 1695 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 1696 |
+
},
|
| 1697 |
+
"pastKeyT": {
|
| 1698 |
+
"dtype": "float16",
|
| 1699 |
+
"shape": [1, 8, 512, 128],
|
| 1700 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 1701 |
+
},
|
| 1702 |
+
"pastValueT": {
|
| 1703 |
+
"dtype": "float16",
|
| 1704 |
+
"shape": [1, 8, 512, 128],
|
| 1705 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 1706 |
+
},
|
| 1707 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [511] } },
|
| 1708 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [512] } }
|
| 1709 |
+
},
|
| 1710 |
+
"outputs": {
|
| 1711 |
+
"outputT": { "dtype": "float16", "shape": [1, 64, 4096] },
|
| 1712 |
+
"presentKeyT": { "dtype": "float16", "shape": [1, 8, 512, 128] },
|
| 1713 |
+
"presentValueT": { "dtype": "float16", "shape": [1, 8, 512, 128] }
|
| 1714 |
+
},
|
| 1715 |
+
"bench": {
|
| 1716 |
+
"metrics": [
|
| 1717 |
+
{ "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
|
| 1718 |
+
]
|
| 1719 |
}
|
| 1720 |
},
|
| 1721 |
{
|
|
|
|
| 1740 |
{ "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
|
| 1741 |
]
|
| 1742 |
}
|
| 1743 |
+
},
|
| 1744 |
+
{
|
| 1745 |
+
"name": "causal-prefill-f16-h4kv2-d64-q64",
|
| 1746 |
+
"provenance": {
|
| 1747 |
+
"source": "synthetic",
|
| 1748 |
+
"test": "Causal half-precision prompt prefill",
|
| 1749 |
+
"notes": "Each query attends only itself and earlier key positions. Verify the complete attention output and exact retention of both key/value inputs in the present caches."
|
| 1750 |
+
},
|
| 1751 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2 },
|
| 1752 |
+
"inputs": {
|
| 1753 |
+
"queryT": {
|
| 1754 |
+
"dtype": "float16",
|
| 1755 |
+
"shape": [1, 64, 256],
|
| 1756 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.0625 }
|
| 1757 |
+
},
|
| 1758 |
+
"keyT": {
|
| 1759 |
+
"dtype": "float16",
|
| 1760 |
+
"shape": [1, 64, 128],
|
| 1761 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.0625 }
|
| 1762 |
+
},
|
| 1763 |
+
"valueT": {
|
| 1764 |
+
"dtype": "float16",
|
| 1765 |
+
"shape": [1, 64, 128],
|
| 1766 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.031, "scale": 0.0625 }
|
| 1767 |
+
},
|
| 1768 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [63] } },
|
| 1769 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [64] } }
|
| 1770 |
+
},
|
| 1771 |
+
"outputs": {
|
| 1772 |
+
"outputT": { "dtype": "float16", "shape": [1, 64, 256] },
|
| 1773 |
+
"presentKeyT": { "dtype": "float16", "shape": [1, 2, 64, 64] },
|
| 1774 |
+
"presentValueT": { "dtype": "float16", "shape": [1, 2, 64, 64] }
|
| 1775 |
+
},
|
| 1776 |
+
"preset": "model",
|
| 1777 |
+
"bench": {
|
| 1778 |
+
"metrics": [
|
| 1779 |
+
{
|
| 1780 |
+
"type": "gflops",
|
| 1781 |
+
"value": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * (dim(shapes.queryT, 1) + 1) * dim(shapes.queryT, 2)"
|
| 1782 |
+
}
|
| 1783 |
+
]
|
| 1784 |
+
}
|
| 1785 |
+
},
|
| 1786 |
+
{
|
| 1787 |
+
"name": "causal-prefill-f16-h4kv2-d128-q65",
|
| 1788 |
+
"provenance": {
|
| 1789 |
+
"source": "synthetic",
|
| 1790 |
+
"test": "Causal half-precision prompt prefill",
|
| 1791 |
+
"notes": "Each query attends only itself and earlier key positions. Verify the complete attention output and exact retention of both key/value inputs in the present caches."
|
| 1792 |
+
},
|
| 1793 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2 },
|
| 1794 |
+
"inputs": {
|
| 1795 |
+
"queryT": {
|
| 1796 |
+
"dtype": "float16",
|
| 1797 |
+
"shape": [1, 65, 512],
|
| 1798 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.0625 }
|
| 1799 |
+
},
|
| 1800 |
+
"keyT": {
|
| 1801 |
+
"dtype": "float16",
|
| 1802 |
+
"shape": [1, 65, 256],
|
| 1803 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.0625 }
|
| 1804 |
+
},
|
| 1805 |
+
"valueT": {
|
| 1806 |
+
"dtype": "float16",
|
| 1807 |
+
"shape": [1, 65, 256],
|
| 1808 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.031, "scale": 0.0625 }
|
| 1809 |
+
},
|
| 1810 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [64] } },
|
| 1811 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65] } }
|
| 1812 |
+
},
|
| 1813 |
+
"outputs": {
|
| 1814 |
+
"outputT": { "dtype": "float16", "shape": [1, 65, 512] },
|
| 1815 |
+
"presentKeyT": { "dtype": "float16", "shape": [1, 2, 65, 128] },
|
| 1816 |
+
"presentValueT": { "dtype": "float16", "shape": [1, 2, 65, 128] }
|
| 1817 |
+
},
|
| 1818 |
+
"preset": "model",
|
| 1819 |
+
"bench": {
|
| 1820 |
+
"metrics": [
|
| 1821 |
+
{
|
| 1822 |
+
"type": "gflops",
|
| 1823 |
+
"value": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * (dim(shapes.queryT, 1) + 1) * dim(shapes.queryT, 2)"
|
| 1824 |
+
}
|
| 1825 |
+
]
|
| 1826 |
+
}
|
| 1827 |
+
},
|
| 1828 |
+
{
|
| 1829 |
+
"name": "causal-prefill-f16-h8kv2-d64-q512",
|
| 1830 |
+
"provenance": {
|
| 1831 |
+
"source": "synthetic",
|
| 1832 |
+
"test": "Causal half-precision prompt prefill",
|
| 1833 |
+
"notes": "Each query attends only itself and earlier key positions. Verify the complete attention output and exact retention of both key/value inputs in the present caches."
|
| 1834 |
+
},
|
| 1835 |
+
"attrs": { "num_heads": 8, "kv_num_heads": 2 },
|
| 1836 |
+
"inputs": {
|
| 1837 |
+
"queryT": {
|
| 1838 |
+
"dtype": "float16",
|
| 1839 |
+
"shape": [1, 512, 512],
|
| 1840 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.0625 }
|
| 1841 |
+
},
|
| 1842 |
+
"keyT": {
|
| 1843 |
+
"dtype": "float16",
|
| 1844 |
+
"shape": [1, 512, 128],
|
| 1845 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.0625 }
|
| 1846 |
+
},
|
| 1847 |
+
"valueT": {
|
| 1848 |
+
"dtype": "float16",
|
| 1849 |
+
"shape": [1, 512, 128],
|
| 1850 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.031, "scale": 0.0625 }
|
| 1851 |
+
},
|
| 1852 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [511] } },
|
| 1853 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [512] } }
|
| 1854 |
+
},
|
| 1855 |
+
"outputs": {
|
| 1856 |
+
"outputT": { "dtype": "float16", "shape": [1, 512, 512] },
|
| 1857 |
+
"presentKeyT": { "dtype": "float16", "shape": [1, 2, 512, 64] },
|
| 1858 |
+
"presentValueT": { "dtype": "float16", "shape": [1, 2, 512, 64] }
|
| 1859 |
+
},
|
| 1860 |
+
"preset": "model",
|
| 1861 |
+
"bench": {
|
| 1862 |
+
"metrics": [
|
| 1863 |
+
{
|
| 1864 |
+
"type": "gflops",
|
| 1865 |
+
"value": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * (dim(shapes.queryT, 1) + 1) * dim(shapes.queryT, 2)"
|
| 1866 |
+
}
|
| 1867 |
+
]
|
| 1868 |
+
}
|
| 1869 |
+
},
|
| 1870 |
+
{
|
| 1871 |
+
"name": "prefill-f32-h8kv1-d32-q33",
|
| 1872 |
+
"preset": "model",
|
| 1873 |
+
"provenance": { "notes": "Small bidirectional prefill with query tails and varying head widths." },
|
| 1874 |
+
"vars": { "batch": 1, "qSeq": 33, "kvSeq": 33, "heads": 8, "kvHeads": 1, "headDim": 32 },
|
| 1875 |
+
"attrs": { "num_heads": 8, "kv_num_heads": 1, "scale": 0.125, "causal": 0 },
|
| 1876 |
+
"inputs": {
|
| 1877 |
+
"queryT": { "shape": [1, 33, 256], "dtype": "float32", "dist": "normal", "seed": 1075, "scale": 0.2 },
|
| 1878 |
+
"keyT": { "shape": [1, 33, 32], "dtype": "float32", "dist": "normal", "seed": 1076, "scale": 0.2 },
|
| 1879 |
+
"valueT": { "shape": [1, 33, 32], "dtype": "float32", "dist": "normal", "seed": 1077, "scale": 0.2 },
|
| 1880 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [32] } },
|
| 1881 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [33] } }
|
| 1882 |
+
},
|
| 1883 |
+
"outputs": {
|
| 1884 |
+
"outputT": { "shape": [1, 33, 256], "dtype": "float32" },
|
| 1885 |
+
"presentKeyT": { "shape": [1, 1, 33, 32], "dtype": "float32" },
|
| 1886 |
+
"presentValueT": { "shape": [1, 1, 33, 32], "dtype": "float32" }
|
| 1887 |
+
},
|
| 1888 |
+
"bench": {
|
| 1889 |
+
"metrics": [
|
| 1890 |
+
{ "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
|
| 1891 |
+
]
|
| 1892 |
+
}
|
| 1893 |
+
},
|
| 1894 |
+
{
|
| 1895 |
+
"name": "prefill-f32-h8kv1-d128-q33",
|
| 1896 |
+
"preset": "model",
|
| 1897 |
+
"provenance": { "notes": "Small bidirectional prefill with query tails and varying head widths." },
|
| 1898 |
+
"vars": { "batch": 1, "qSeq": 33, "kvSeq": 33, "heads": 8, "kvHeads": 1, "headDim": 128 },
|
| 1899 |
+
"attrs": { "num_heads": 8, "kv_num_heads": 1, "scale": 0.125, "causal": 0 },
|
| 1900 |
+
"inputs": {
|
| 1901 |
+
"queryT": { "shape": [1, 33, 1024], "dtype": "float32", "dist": "normal", "seed": 1075, "scale": 0.2 },
|
| 1902 |
+
"keyT": { "shape": [1, 33, 128], "dtype": "float32", "dist": "normal", "seed": 1076, "scale": 0.2 },
|
| 1903 |
+
"valueT": { "shape": [1, 33, 128], "dtype": "float32", "dist": "normal", "seed": 1077, "scale": 0.2 },
|
| 1904 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [32] } },
|
| 1905 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [33] } }
|
| 1906 |
+
},
|
| 1907 |
+
"outputs": {
|
| 1908 |
+
"outputT": { "shape": [1, 33, 1024], "dtype": "float32" },
|
| 1909 |
+
"presentKeyT": { "shape": [1, 1, 33, 128], "dtype": "float32" },
|
| 1910 |
+
"presentValueT": { "shape": [1, 1, 33, 128], "dtype": "float32" }
|
| 1911 |
+
},
|
| 1912 |
+
"bench": {
|
| 1913 |
+
"metrics": [
|
| 1914 |
+
{ "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
|
| 1915 |
+
]
|
| 1916 |
+
}
|
| 1917 |
+
},
|
| 1918 |
+
{
|
| 1919 |
+
"name": "prefill-f32-h8kv1-d256-q32",
|
| 1920 |
+
"preset": "model",
|
| 1921 |
+
"provenance": { "notes": "Small bidirectional prefill with query tails and varying head widths." },
|
| 1922 |
+
"vars": { "batch": 1, "qSeq": 32, "kvSeq": 32, "heads": 8, "kvHeads": 1, "headDim": 256 },
|
| 1923 |
+
"attrs": { "num_heads": 8, "kv_num_heads": 1, "scale": 0.125, "causal": 0 },
|
| 1924 |
+
"inputs": {
|
| 1925 |
+
"queryT": { "shape": [1, 32, 2048], "dtype": "float32", "dist": "normal", "seed": 1075, "scale": 0.2 },
|
| 1926 |
+
"keyT": { "shape": [1, 32, 256], "dtype": "float32", "dist": "normal", "seed": 1076, "scale": 0.2 },
|
| 1927 |
+
"valueT": { "shape": [1, 32, 256], "dtype": "float32", "dist": "normal", "seed": 1077, "scale": 0.2 },
|
| 1928 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [31] } },
|
| 1929 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [32] } }
|
| 1930 |
+
},
|
| 1931 |
+
"outputs": {
|
| 1932 |
+
"outputT": { "shape": [1, 32, 2048], "dtype": "float32" },
|
| 1933 |
+
"presentKeyT": { "shape": [1, 1, 32, 256], "dtype": "float32" },
|
| 1934 |
+
"presentValueT": { "shape": [1, 1, 32, 256], "dtype": "float32" }
|
| 1935 |
+
},
|
| 1936 |
+
"bench": {
|
| 1937 |
+
"metrics": [
|
| 1938 |
+
{ "type": "gflops", "value": "4 * args.batch * args.qSeq * args.kvSeq * args.heads * args.headDim" }
|
| 1939 |
+
]
|
| 1940 |
+
}
|
| 1941 |
+
},
|
| 1942 |
+
{
|
| 1943 |
+
"name": "cooperative-window-h4kv2-d122-cap256",
|
| 1944 |
+
"preset": "smoke",
|
| 1945 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
|
| 1946 |
+
"inputs": {
|
| 1947 |
+
"queryT": {
|
| 1948 |
+
"dtype": "float32",
|
| 1949 |
+
"shape": [1, 1, 488],
|
| 1950 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 1951 |
+
},
|
| 1952 |
+
"keyT": {
|
| 1953 |
+
"dtype": "float32",
|
| 1954 |
+
"shape": [1, 1, 244],
|
| 1955 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 1956 |
+
},
|
| 1957 |
+
"valueT": {
|
| 1958 |
+
"dtype": "float32",
|
| 1959 |
+
"shape": [1, 1, 244],
|
| 1960 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 1961 |
+
},
|
| 1962 |
+
"pastKeyT": {
|
| 1963 |
+
"dtype": "float32",
|
| 1964 |
+
"shape": [1, 2, 256, 122],
|
| 1965 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 1966 |
+
},
|
| 1967 |
+
"pastValueT": {
|
| 1968 |
+
"dtype": "float32",
|
| 1969 |
+
"shape": [1, 2, 256, 122],
|
| 1970 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 1971 |
+
},
|
| 1972 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
|
| 1973 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
|
| 1974 |
+
},
|
| 1975 |
+
"outputs": {
|
| 1976 |
+
"outputT": { "dtype": "float32", "shape": [1, 1, 488] },
|
| 1977 |
+
"presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 122] },
|
| 1978 |
+
"presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 122] }
|
| 1979 |
+
},
|
| 1980 |
+
"bench": {
|
| 1981 |
+
"metrics": [
|
| 1982 |
+
{
|
| 1983 |
+
"type": "bandwidth",
|
| 1984 |
+
"value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
|
| 1985 |
+
}
|
| 1986 |
+
]
|
| 1987 |
+
}
|
| 1988 |
+
},
|
| 1989 |
+
{
|
| 1990 |
+
"name": "cooperative-window-h4kv2-d124-cap256",
|
| 1991 |
+
"preset": "smoke",
|
| 1992 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
|
| 1993 |
+
"inputs": {
|
| 1994 |
+
"queryT": {
|
| 1995 |
+
"dtype": "float32",
|
| 1996 |
+
"shape": [1, 1, 496],
|
| 1997 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 1998 |
+
},
|
| 1999 |
+
"keyT": {
|
| 2000 |
+
"dtype": "float32",
|
| 2001 |
+
"shape": [1, 1, 248],
|
| 2002 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 2003 |
+
},
|
| 2004 |
+
"valueT": {
|
| 2005 |
+
"dtype": "float32",
|
| 2006 |
+
"shape": [1, 1, 248],
|
| 2007 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 2008 |
+
},
|
| 2009 |
+
"pastKeyT": {
|
| 2010 |
+
"dtype": "float32",
|
| 2011 |
+
"shape": [1, 2, 256, 124],
|
| 2012 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 2013 |
+
},
|
| 2014 |
+
"pastValueT": {
|
| 2015 |
+
"dtype": "float32",
|
| 2016 |
+
"shape": [1, 2, 256, 124],
|
| 2017 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 2018 |
+
},
|
| 2019 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
|
| 2020 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
|
| 2021 |
+
},
|
| 2022 |
+
"outputs": {
|
| 2023 |
+
"outputT": { "dtype": "float32", "shape": [1, 1, 496] },
|
| 2024 |
+
"presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 124] },
|
| 2025 |
+
"presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 124] }
|
| 2026 |
+
},
|
| 2027 |
+
"bench": {
|
| 2028 |
+
"metrics": [
|
| 2029 |
+
{
|
| 2030 |
+
"type": "bandwidth",
|
| 2031 |
+
"value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
|
| 2032 |
+
}
|
| 2033 |
+
]
|
| 2034 |
+
}
|
| 2035 |
+
},
|
| 2036 |
+
{
|
| 2037 |
+
"name": "cooperative-window-h4kv2-d246-cap256",
|
| 2038 |
+
"preset": "smoke",
|
| 2039 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
|
| 2040 |
+
"inputs": {
|
| 2041 |
+
"queryT": {
|
| 2042 |
+
"dtype": "float32",
|
| 2043 |
+
"shape": [1, 1, 984],
|
| 2044 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 2045 |
+
},
|
| 2046 |
+
"keyT": {
|
| 2047 |
+
"dtype": "float32",
|
| 2048 |
+
"shape": [1, 1, 492],
|
| 2049 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 2050 |
+
},
|
| 2051 |
+
"valueT": {
|
| 2052 |
+
"dtype": "float32",
|
| 2053 |
+
"shape": [1, 1, 492],
|
| 2054 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 2055 |
+
},
|
| 2056 |
+
"pastKeyT": {
|
| 2057 |
+
"dtype": "float32",
|
| 2058 |
+
"shape": [1, 2, 256, 246],
|
| 2059 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 2060 |
+
},
|
| 2061 |
+
"pastValueT": {
|
| 2062 |
+
"dtype": "float32",
|
| 2063 |
+
"shape": [1, 2, 256, 246],
|
| 2064 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 2065 |
+
},
|
| 2066 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
|
| 2067 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
|
| 2068 |
+
},
|
| 2069 |
+
"outputs": {
|
| 2070 |
+
"outputT": { "dtype": "float32", "shape": [1, 1, 984] },
|
| 2071 |
+
"presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 246] },
|
| 2072 |
+
"presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 246] }
|
| 2073 |
+
},
|
| 2074 |
+
"bench": {
|
| 2075 |
+
"metrics": [
|
| 2076 |
+
{
|
| 2077 |
+
"type": "bandwidth",
|
| 2078 |
+
"value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
|
| 2079 |
+
}
|
| 2080 |
+
]
|
| 2081 |
+
}
|
| 2082 |
+
},
|
| 2083 |
+
{
|
| 2084 |
+
"name": "cooperative-window-h4kv2-d248-cap256",
|
| 2085 |
+
"preset": "smoke",
|
| 2086 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
|
| 2087 |
+
"inputs": {
|
| 2088 |
+
"queryT": {
|
| 2089 |
+
"dtype": "float32",
|
| 2090 |
+
"shape": [1, 1, 992],
|
| 2091 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 2092 |
+
},
|
| 2093 |
+
"keyT": {
|
| 2094 |
+
"dtype": "float32",
|
| 2095 |
+
"shape": [1, 1, 496],
|
| 2096 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 2097 |
+
},
|
| 2098 |
+
"valueT": {
|
| 2099 |
+
"dtype": "float32",
|
| 2100 |
+
"shape": [1, 1, 496],
|
| 2101 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 2102 |
+
},
|
| 2103 |
+
"pastKeyT": {
|
| 2104 |
+
"dtype": "float32",
|
| 2105 |
+
"shape": [1, 2, 256, 248],
|
| 2106 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 2107 |
+
},
|
| 2108 |
+
"pastValueT": {
|
| 2109 |
+
"dtype": "float32",
|
| 2110 |
+
"shape": [1, 2, 256, 248],
|
| 2111 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 2112 |
+
},
|
| 2113 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
|
| 2114 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
|
| 2115 |
+
},
|
| 2116 |
+
"outputs": {
|
| 2117 |
+
"outputT": { "dtype": "float32", "shape": [1, 1, 992] },
|
| 2118 |
+
"presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 248] },
|
| 2119 |
+
"presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 248] }
|
| 2120 |
+
},
|
| 2121 |
+
"bench": {
|
| 2122 |
+
"metrics": [
|
| 2123 |
+
{
|
| 2124 |
+
"type": "bandwidth",
|
| 2125 |
+
"value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
|
| 2126 |
+
}
|
| 2127 |
+
]
|
| 2128 |
+
}
|
| 2129 |
+
},
|
| 2130 |
+
{
|
| 2131 |
+
"name": "cooperative-window-h4kv2-d256-cap256",
|
| 2132 |
+
"preset": "smoke",
|
| 2133 |
+
"attrs": { "num_heads": 4, "kv_num_heads": 2, "sliding_window_cache": 1, "local_window_size": 256 },
|
| 2134 |
+
"inputs": {
|
| 2135 |
+
"queryT": {
|
| 2136 |
+
"dtype": "float32",
|
| 2137 |
+
"shape": [1, 1, 1024],
|
| 2138 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.011, "cosStep": 0.029 }
|
| 2139 |
+
},
|
| 2140 |
+
"keyT": {
|
| 2141 |
+
"dtype": "float32",
|
| 2142 |
+
"shape": [1, 1, 512],
|
| 2143 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.031, "cosStep": 0.007 }
|
| 2144 |
+
},
|
| 2145 |
+
"valueT": {
|
| 2146 |
+
"dtype": "float32",
|
| 2147 |
+
"shape": [1, 1, 512],
|
| 2148 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.037, "cosStep": 0.005 }
|
| 2149 |
+
},
|
| 2150 |
+
"pastKeyT": {
|
| 2151 |
+
"dtype": "float32",
|
| 2152 |
+
"shape": [1, 2, 256, 256],
|
| 2153 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.017, "cosStep": 0.019 }
|
| 2154 |
+
},
|
| 2155 |
+
"pastValueT": {
|
| 2156 |
+
"dtype": "float32",
|
| 2157 |
+
"shape": [1, 2, 256, 256],
|
| 2158 |
+
"data": { "kind": "fillFloat32", "scale": 0.2, "sinStep": 0.023, "cosStep": 0.013 }
|
| 2159 |
+
},
|
| 2160 |
+
"seqlensKT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65535] } },
|
| 2161 |
+
"totalSequenceLengthT": { "dtype": "int32", "shape": [1], "data": { "kind": "values", "values": [65536] } }
|
| 2162 |
+
},
|
| 2163 |
+
"outputs": {
|
| 2164 |
+
"outputT": { "dtype": "float32", "shape": [1, 1, 1024] },
|
| 2165 |
+
"presentKeyT": { "dtype": "float32", "shape": [1, 2, 256, 256] },
|
| 2166 |
+
"presentValueT": { "dtype": "float32", "shape": [1, 2, 256, 256] }
|
| 2167 |
+
},
|
| 2168 |
+
"bench": {
|
| 2169 |
+
"metrics": [
|
| 2170 |
+
{
|
| 2171 |
+
"type": "bandwidth",
|
| 2172 |
+
"value": "2 * 2 * numel(shapes.pastKeyT) * 4 + numel(shapes.queryT) * 4 + numel(shapes.outputT) * 4"
|
| 2173 |
+
}
|
| 2174 |
+
]
|
| 2175 |
+
}
|
| 2176 |
}
|
| 2177 |
]
|
| 2178 |
}
|
build/webgpu/gqa-attention.wgsl.jinja
CHANGED
|
@@ -2,8 +2,12 @@
|
|
| 2 |
{% set quantized = quantized | default(false) %}
|
| 3 |
{% set bits = bits | default(0) %}
|
| 4 |
{% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
|
| 5 |
-
{% if
|
| 6 |
-
{%
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
// f16 inputs widen to f32 on load; computation and shared memory remain f32,
|
| 8 |
// and results narrow only on store. The casts are identities for f32.
|
| 9 |
{% set IO = "f16" if usesF16 else "f32" %}
|
|
@@ -29,15 +33,13 @@ const Q_HIDDEN: u32 = {{ qHidden }}u;
|
|
| 29 |
const GROUP: u32 = {{ qHeads }}u / {{ kvHeads }}u;
|
| 30 |
const PACKED: u32 = {{ packed }}u;
|
| 31 |
{% if hasQNorm %}const QK_EPS: f32 = {{ qkEps }};{% endif %}
|
| 32 |
-
{% if cooperative %}const WG: u32 = {{
|
| 33 |
const NEG_INF: f32 = -3.4028234663852886e38;
|
| 34 |
-
{%
|
| 35 |
-
fn scale_value() -> f32 {
|
| 36 |
if (params.scale != 0.0) { return params.scale; }
|
| 37 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 38 |
}
|
| 39 |
|
| 40 |
-
|
| 41 |
{% if quantized %}
|
| 42 |
fn kscale(d: u32, hk: u32) -> f32 { return k_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
|
| 43 |
fn vscale(d: u32, hk: u32) -> f32 { return v_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
|
|
@@ -67,10 +69,15 @@ fn read_v(base: u32, d: u32{% if quantized %}, hk: u32{% endif %}) -> f32 {
|
|
| 67 |
}
|
| 68 |
|
| 69 |
{% if cooperative %}
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 74 |
|
| 75 |
@compute @workgroup_size(WG)
|
| 76 |
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
|
|
@@ -92,42 +99,47 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
|
|
| 92 |
{% else %}
|
| 93 |
let activeEnd = totalSeq;
|
| 94 |
{% endif %}
|
| 95 |
-
let absPos = activeEnd - params.qSeq + s;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
|
| 97 |
// Cooperative query-vector prep (post norm/rotary) into shared memory.
|
| 98 |
let qBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
|
| 99 |
-
for (var d = lane; d < HEAD_DIM; d = d + WG) {
|
| 100 |
workgroupBarrier();
|
| 101 |
{% if hasQNorm %}
|
| 102 |
var part = 0.0;
|
| 103 |
-
for (var d = lane; d < HEAD_DIM; d = d + WG) { part = part +
|
| 104 |
-
|
| 105 |
workgroupBarrier();
|
| 106 |
var ms = 0.0;
|
| 107 |
-
for (var L = 0u; L < WG; L = L + 1u) { ms = ms +
|
| 108 |
let invRms = inverseSqrt(ms / f32(HEAD_DIM) + QK_EPS);
|
| 109 |
-
for (var d = lane; d < HEAD_DIM; d = d + WG) {
|
| 110 |
workgroupBarrier();
|
| 111 |
{% endif %}
|
| 112 |
{% if hasRotary %}
|
| 113 |
// Each lane owns disjoint pairs (d, d+HALF), so the in-place rotate is safe.
|
| 114 |
for (var d = lane; d < HALF; d = d + WG) {
|
| 115 |
-
let cs = f32(cos_cache[
|
| 116 |
-
let sn = f32(sin_cache[
|
| 117 |
-
let x0 =
|
| 118 |
-
let x1 =
|
| 119 |
-
|
| 120 |
-
|
| 121 |
}
|
| 122 |
workgroupBarrier();
|
| 123 |
{% endif %}
|
| 124 |
let scale = scale_value();
|
| 125 |
-
let maxKj = absPos + 1u;
|
| 126 |
var minKj = 0u;
|
| 127 |
if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
|
| 128 |
|
| 129 |
let accBase = lane * HEAD_DIM;
|
| 130 |
-
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
| 131 |
var m = NEG_INF;
|
| 132 |
var l = 0.0;
|
| 133 |
|
|
@@ -135,34 +147,34 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
|
|
| 135 |
for (var j = minKj + lane; j < maxKj; j = j + WG) {
|
| 136 |
let base = ((b * KV_HEADS + hk) * totalSeq + j) * PACKED;
|
| 137 |
var dot = 0.0;
|
| 138 |
-
for (var d = 0u; d < HEAD_DIM; d = d + 1u) { dot = dot +
|
| 139 |
var score = dot * scale;
|
| 140 |
if (params.softcap > 0.0) { score = params.softcap * tanh(clamp(score / params.softcap, -30.0, 30.0)); }
|
| 141 |
{% if hasBias %}
|
| 142 |
let bb = select(b, 0u, params.biasBatch == 1u);
|
| 143 |
let bh = select(h, 0u, params.biasHeads == 1u);
|
| 144 |
-
score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) *
|
| 145 |
{% endif %}
|
| 146 |
let newM = max(m, score);
|
| 147 |
let corr = exp(m - newM);
|
| 148 |
let p = exp(score - newM);
|
| 149 |
l = l * corr + p;
|
| 150 |
-
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
| 151 |
m = newM;
|
| 152 |
}
|
| 153 |
|
| 154 |
// Flash merge across the WG lanes: global max, rescale, summed denom.
|
| 155 |
-
|
| 156 |
workgroupBarrier();
|
| 157 |
var gm = NEG_INF;
|
| 158 |
-
for (var L = 0u; L < WG; L = L + 1u) { gm = max(gm,
|
| 159 |
let fctr = exp(m - gm);
|
| 160 |
l = l * fctr;
|
| 161 |
-
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
| 162 |
-
|
| 163 |
workgroupBarrier();
|
| 164 |
var glsum = 0.0;
|
| 165 |
-
for (var L = 0u; L < WG; L = L + 1u) { glsum = glsum +
|
| 166 |
|
| 167 |
var sink = 0.0;
|
| 168 |
{% if hasHeadSink %}
|
|
@@ -187,7 +199,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
|
|
| 187 |
let oBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
|
| 188 |
for (var d = lane; d < HEAD_DIM; d = d + WG) {
|
| 189 |
var sumacc = 0.0;
|
| 190 |
-
for (var L = 0u; L < WG; L = L + 1u) { sumacc = sumacc +
|
| 191 |
output[oBase + d] = {{ IO }}(sumacc * invDenom);
|
| 192 |
}
|
| 193 |
}
|
|
@@ -215,7 +227,12 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
|
|
| 215 |
{% else %}
|
| 216 |
let activeEnd = totalSeq;
|
| 217 |
{% endif %}
|
| 218 |
-
let absPos = activeEnd - params.qSeq + s;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 219 |
|
| 220 |
// Private query-vector prep (post norm/rotary).
|
| 221 |
var q: array<f32, HEAD_DIM>;
|
|
@@ -229,8 +246,8 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
|
|
| 229 |
{% endif %}
|
| 230 |
{% if hasRotary %}
|
| 231 |
for (var d = 0u; d < HALF; d = d + 1u) {
|
| 232 |
-
let cs = f32(cos_cache[
|
| 233 |
-
let sn = f32(sin_cache[
|
| 234 |
let x0 = q[d];
|
| 235 |
let x1 = q[d + HALF];
|
| 236 |
q[d] = x0 * cs - x1 * sn;
|
|
@@ -238,7 +255,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
|
|
| 238 |
}
|
| 239 |
{% endif %}
|
| 240 |
let scale = scale_value();
|
| 241 |
-
let maxKj = absPos + 1u;
|
| 242 |
var minKj = 0u;
|
| 243 |
if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
|
| 244 |
|
|
@@ -256,7 +273,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid
|
|
| 256 |
{% if hasBias %}
|
| 257 |
let bb = select(b, 0u, params.biasBatch == 1u);
|
| 258 |
let bh = select(h, 0u, params.biasHeads == 1u);
|
| 259 |
-
score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) *
|
| 260 |
{% endif %}
|
| 261 |
let newM = max(m, score);
|
| 262 |
let corr = exp(m - newM);
|
|
|
|
| 2 |
{% set quantized = quantized | default(false) %}
|
| 3 |
{% set bits = bits | default(0) %}
|
| 4 |
{% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
|
| 5 |
+
{% if windowOrigin is not defined %}{% set windowOrigin = false %}{% endif %}
|
| 6 |
+
{% set needsOrigin = windowOrigin and (hasRotary or hasBias) %}
|
| 7 |
+
{% set ropePos = "(absPos + origin)" if needsOrigin else "absPos" %}
|
| 8 |
+
{% set biasStride = "params.biasCols" if needsOrigin else "totalSeq" %}
|
| 9 |
+
{% set biasCol = "(origin + j)" if needsOrigin else "j" %}
|
| 10 |
+
{{ env.wgsl.resourceDeclarations }}
|
| 11 |
// f16 inputs widen to f32 on load; computation and shared memory remain f32,
|
| 12 |
// and results narrow only on store. The casts are identities for f32.
|
| 13 |
{% set IO = "f16" if usesF16 else "f32" %}
|
|
|
|
| 33 |
const GROUP: u32 = {{ qHeads }}u / {{ kvHeads }}u;
|
| 34 |
const PACKED: u32 = {{ packed }}u;
|
| 35 |
{% if hasQNorm %}const QK_EPS: f32 = {{ qkEps }};{% endif %}
|
| 36 |
+
{% if cooperative %}const WG: u32 = {{ gqaCooperativeWorkgroupSize }}u;{% else %}const WG: u32 = {{ scalarWorkgroupSize }}u;{% endif %}
|
| 37 |
const NEG_INF: f32 = -3.4028234663852886e38;
|
| 38 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 39 |
if (params.scale != 0.0) { return params.scale; }
|
| 40 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 41 |
}
|
| 42 |
|
|
|
|
| 43 |
{% if quantized %}
|
| 44 |
fn kscale(d: u32, hk: u32) -> f32 { return k_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
|
| 45 |
fn vscale(d: u32, hk: u32) -> f32 { return v_scale[select(0u, hk * HEAD_DIM + d, params.perChannel != 0u)]; }
|
|
|
|
| 69 |
}
|
| 70 |
|
| 71 |
{% if cooperative %}
|
| 72 |
+
// One allocation avoids the per-variable 16-byte padding between arrays.
|
| 73 |
+
// The manifest budgets the structure's complete size, rounded to 16 bytes.
|
| 74 |
+
struct CooperativeState {
|
| 75 |
+
q: array<f32, HEAD_DIM>,
|
| 76 |
+
acc: array<f32, HEAD_DIM * WG>,
|
| 77 |
+
m: array<f32, WG>,
|
| 78 |
+
l: array<f32, WG>,
|
| 79 |
+
}
|
| 80 |
+
var<workgroup> state: CooperativeState;
|
| 81 |
|
| 82 |
@compute @workgroup_size(WG)
|
| 83 |
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
|
|
|
|
| 99 |
{% else %}
|
| 100 |
let activeEnd = totalSeq;
|
| 101 |
{% endif %}
|
| 102 |
+
{% if not bidirectional or hasRotary %} let absPos = activeEnd - params.qSeq + s;
|
| 103 |
+
{% endif %}{% if needsOrigin %}
|
| 104 |
+
// Absolute position of cache row 0 (windowed cache): the query end is the
|
| 105 |
+
// unclamped total while the cache holds only its last `activeEnd` rows.
|
| 106 |
+
let origin = max(params.qSeq, u32(seqlens_k[b]) + 1u) - activeEnd;
|
| 107 |
+
{% endif %}
|
| 108 |
|
| 109 |
// Cooperative query-vector prep (post norm/rotary) into shared memory.
|
| 110 |
let qBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
|
| 111 |
+
for (var d = lane; d < HEAD_DIM; d = d + WG) { state.q[d] = f32(query[qBase + d]); }
|
| 112 |
workgroupBarrier();
|
| 113 |
{% if hasQNorm %}
|
| 114 |
var part = 0.0;
|
| 115 |
+
for (var d = lane; d < HEAD_DIM; d = d + WG) { part = part + state.q[d] * state.q[d]; }
|
| 116 |
+
state.m[lane] = part;
|
| 117 |
workgroupBarrier();
|
| 118 |
var ms = 0.0;
|
| 119 |
+
for (var L = 0u; L < WG; L = L + 1u) { ms = ms + state.m[L]; }
|
| 120 |
let invRms = inverseSqrt(ms / f32(HEAD_DIM) + QK_EPS);
|
| 121 |
+
for (var d = lane; d < HEAD_DIM; d = d + WG) { state.q[d] = state.q[d] * invRms * f32(q_norm_weight[d]); }
|
| 122 |
workgroupBarrier();
|
| 123 |
{% endif %}
|
| 124 |
{% if hasRotary %}
|
| 125 |
// Each lane owns disjoint pairs (d, d+HALF), so the in-place rotate is safe.
|
| 126 |
for (var d = lane; d < HALF; d = d + WG) {
|
| 127 |
+
let cs = f32(cos_cache[{{ ropePos }} * HALF + d]);
|
| 128 |
+
let sn = f32(sin_cache[{{ ropePos }} * HALF + d]);
|
| 129 |
+
let x0 = state.q[d];
|
| 130 |
+
let x1 = state.q[d + HALF];
|
| 131 |
+
state.q[d] = x0 * cs - x1 * sn;
|
| 132 |
+
state.q[d + HALF] = x1 * cs + x0 * sn;
|
| 133 |
}
|
| 134 |
workgroupBarrier();
|
| 135 |
{% endif %}
|
| 136 |
let scale = scale_value();
|
| 137 |
+
let maxKj = {% if bidirectional %}activeEnd{% else %}absPos + 1u{% endif %};
|
| 138 |
var minKj = 0u;
|
| 139 |
if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
|
| 140 |
|
| 141 |
let accBase = lane * HEAD_DIM;
|
| 142 |
+
for (var d = 0u; d < HEAD_DIM; d = d + 1u) { state.acc[accBase + d] = 0.0; }
|
| 143 |
var m = NEG_INF;
|
| 144 |
var l = 0.0;
|
| 145 |
|
|
|
|
| 147 |
for (var j = minKj + lane; j < maxKj; j = j + WG) {
|
| 148 |
let base = ((b * KV_HEADS + hk) * totalSeq + j) * PACKED;
|
| 149 |
var dot = 0.0;
|
| 150 |
+
for (var d = 0u; d < HEAD_DIM; d = d + 1u) { dot = dot + state.q[d] * read_k(base, d{% if quantized %}, hk{% endif %}); }
|
| 151 |
var score = dot * scale;
|
| 152 |
if (params.softcap > 0.0) { score = params.softcap * tanh(clamp(score / params.softcap, -30.0, 30.0)); }
|
| 153 |
{% if hasBias %}
|
| 154 |
let bb = select(b, 0u, params.biasBatch == 1u);
|
| 155 |
let bh = select(h, 0u, params.biasHeads == 1u);
|
| 156 |
+
score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) * {{ biasStride }} + {{ biasCol }}]);
|
| 157 |
{% endif %}
|
| 158 |
let newM = max(m, score);
|
| 159 |
let corr = exp(m - newM);
|
| 160 |
let p = exp(score - newM);
|
| 161 |
l = l * corr + p;
|
| 162 |
+
for (var d = 0u; d < HEAD_DIM; d = d + 1u) { state.acc[accBase + d] = state.acc[accBase + d] * corr + p * read_v(base, d{% if quantized %}, hk{% endif %}); }
|
| 163 |
m = newM;
|
| 164 |
}
|
| 165 |
|
| 166 |
// Flash merge across the WG lanes: global max, rescale, summed denom.
|
| 167 |
+
state.m[lane] = m;
|
| 168 |
workgroupBarrier();
|
| 169 |
var gm = NEG_INF;
|
| 170 |
+
for (var L = 0u; L < WG; L = L + 1u) { gm = max(gm, state.m[L]); }
|
| 171 |
let fctr = exp(m - gm);
|
| 172 |
l = l * fctr;
|
| 173 |
+
for (var d = 0u; d < HEAD_DIM; d = d + 1u) { state.acc[accBase + d] = state.acc[accBase + d] * fctr; }
|
| 174 |
+
state.l[lane] = l;
|
| 175 |
workgroupBarrier();
|
| 176 |
var glsum = 0.0;
|
| 177 |
+
for (var L = 0u; L < WG; L = L + 1u) { glsum = glsum + state.l[L]; }
|
| 178 |
|
| 179 |
var sink = 0.0;
|
| 180 |
{% if hasHeadSink %}
|
|
|
|
| 199 |
let oBase = (b * params.qSeq + s) * Q_HIDDEN + h * HEAD_DIM;
|
| 200 |
for (var d = lane; d < HEAD_DIM; d = d + WG) {
|
| 201 |
var sumacc = 0.0;
|
| 202 |
+
for (var L = 0u; L < WG; L = L + 1u) { sumacc = sumacc + state.acc[L * HEAD_DIM + d]; }
|
| 203 |
output[oBase + d] = {{ IO }}(sumacc * invDenom);
|
| 204 |
}
|
| 205 |
}
|
|
|
|
| 227 |
{% else %}
|
| 228 |
let activeEnd = totalSeq;
|
| 229 |
{% endif %}
|
| 230 |
+
{% if not bidirectional or hasRotary %} let absPos = activeEnd - params.qSeq + s;
|
| 231 |
+
{% endif %}{% if needsOrigin %}
|
| 232 |
+
// Absolute position of cache row 0 (windowed cache): the query end is the
|
| 233 |
+
// unclamped total while the cache holds only its last `activeEnd` rows.
|
| 234 |
+
let origin = max(params.qSeq, u32(seqlens_k[b]) + 1u) - activeEnd;
|
| 235 |
+
{% endif %}
|
| 236 |
|
| 237 |
// Private query-vector prep (post norm/rotary).
|
| 238 |
var q: array<f32, HEAD_DIM>;
|
|
|
|
| 246 |
{% endif %}
|
| 247 |
{% if hasRotary %}
|
| 248 |
for (var d = 0u; d < HALF; d = d + 1u) {
|
| 249 |
+
let cs = f32(cos_cache[{{ ropePos }} * HALF + d]);
|
| 250 |
+
let sn = f32(sin_cache[{{ ropePos }} * HALF + d]);
|
| 251 |
let x0 = q[d];
|
| 252 |
let x1 = q[d + HALF];
|
| 253 |
q[d] = x0 * cs - x1 * sn;
|
|
|
|
| 255 |
}
|
| 256 |
{% endif %}
|
| 257 |
let scale = scale_value();
|
| 258 |
+
let maxKj = {% if bidirectional %}activeEnd{% else %}absPos + 1u{% endif %};
|
| 259 |
var minKj = 0u;
|
| 260 |
if (params.windowSize > 0u && maxKj > params.windowSize) { minKj = maxKj - params.windowSize; }
|
| 261 |
|
|
|
|
| 273 |
{% if hasBias %}
|
| 274 |
let bb = select(b, 0u, params.biasBatch == 1u);
|
| 275 |
let bh = select(h, 0u, params.biasHeads == 1u);
|
| 276 |
+
score = score + f32(attn_bias[((bb * params.biasHeads + bh) * params.qSeq + s) * {{ biasStride }} + {{ biasCol }}]);
|
| 277 |
{% endif %}
|
| 278 |
let newM = max(m, score);
|
| 279 |
let corr = exp(m - newM);
|
build/webgpu/gqa-present.wgsl.jinja
CHANGED
|
@@ -9,8 +9,15 @@
|
|
| 9 |
{% set bits = bits | default(0) %}
|
| 10 |
{% set qmax = qmax | default(0) %}
|
| 11 |
{% set qmin = qmin | default(0) %}
|
| 12 |
-
{%
|
| 13 |
-
{%
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 14 |
{% if mode == "transpose" %}
|
| 15 |
|
| 16 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
|
@@ -31,9 +38,7 @@ const KV_HIDDEN: u32 = {{ kvHidden }}u;
|
|
| 31 |
fn main(
|
| 32 |
@builtin(global_invocation_id) gid: vec3<u32>
|
| 33 |
) {
|
| 34 |
-
|
| 35 |
-
// flat invocation index; this reduces to gid.x when no fold is needed.
|
| 36 |
-
let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
|
| 37 |
{% if presentVec4 %}
|
| 38 |
let total = params.batchSize * KV_HEADS * params.kvSeq * HEAD_DIM_V4;
|
| 39 |
if (index >= total) {
|
|
@@ -60,9 +65,11 @@ fn main(
|
|
| 60 |
present_value[index] = {{ presentScalar }}(value[packed_index]);
|
| 61 |
{% endif %}
|
| 62 |
}
|
| 63 |
-
{%
|
| 64 |
|
| 65 |
{% set shareAppend = mode == "merge_share" and shareRegion == "append" %}
|
|
|
|
|
|
|
| 66 |
{% set cooperativeCopy = device.adapterInfo.vendor != "apple" %}
|
| 67 |
{% set cooperativeMerge = cooperativeCopy or shareAppend %}
|
| 68 |
{% set cooperativeMode = cooperativeMerge and (mode == "copy" or mode == "merge" or mode == "merge_share") %}
|
|
@@ -95,7 +102,7 @@ const HEAD_DIM: u32 = {{ headDim }}u;
|
|
| 95 |
{% if mode != "copy" %}const KV_HEADS: u32 = {{ kvHeads }}u;
|
| 96 |
{% endif %}
|
| 97 |
const WG: u32 = {{ tunables.COPY_WORKGROUP_SIZE }}u;
|
| 98 |
-
{% if mode != "copy" and (mode != "merge_share" or shareAppend) %}const KV_HIDDEN: u32 = {{ kvHeads }}u * {{ headDim }}u;
|
| 99 |
{% endif %}
|
| 100 |
{% if mode == "build_quant" or mode == "append_quant" %}const PACKED: u32 = {{ packed }}u;
|
| 101 |
{% endif %}
|
|
@@ -153,8 +160,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 153 |
{% endmacro %}
|
| 154 |
{% macro append_per_row_walk() %}
|
| 155 |
// Per-row append walk: one thread streams its row's contiguous bytes.
|
| 156 |
-
|
| 157 |
-
if (i >= params.count) { return; }
|
| 158 |
let j = i % params.keySeq;
|
| 159 |
let tmp = i / params.keySeq;
|
| 160 |
let hk = tmp % KV_HEADS;
|
|
@@ -263,8 +269,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 263 |
{% endif %}
|
| 264 |
{% elif mode == "copy" %}
|
| 265 |
// Thread i streams one contiguous (batch, kvHead, token) row.
|
| 266 |
-
|
| 267 |
-
if (i >= params.count) { return; }
|
| 268 |
let base = i * HEAD_DIM;
|
| 269 |
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
| 270 |
present_key[base + d] = src_k[base + d];
|
|
@@ -272,8 +277,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 272 |
}
|
| 273 |
{% elif mode == "merge" %}
|
| 274 |
// Per-row merge walk.
|
| 275 |
-
|
| 276 |
-
if (i >= params.count) { return; }
|
| 277 |
let t = i % params.seq;
|
| 278 |
let tmp = i / params.seq;
|
| 279 |
let hk = tmp % KV_HEADS;
|
|
@@ -296,8 +300,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 296 |
{% elif mode == "merge_share" %}
|
| 297 |
// Past and present share the full-capacity stride, so outside-window rows
|
| 298 |
// copy at the same index.
|
| 299 |
-
|
| 300 |
-
if (i >= params.count) { return; }
|
| 301 |
let t = i % params.seq;
|
| 302 |
let tmp = i / params.seq;
|
| 303 |
let hk = tmp % KV_HEADS;
|
|
@@ -316,28 +319,51 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 316 |
}
|
| 317 |
}
|
| 318 |
{% elif mode == "window_shift" %}
|
| 319 |
-
//
|
| 320 |
-
//
|
| 321 |
-
//
|
| 322 |
-
|
| 323 |
-
//
|
| 324 |
-
//
|
| 325 |
-
//
|
| 326 |
-
//
|
| 327 |
-
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 332 |
-
|
| 333 |
-
|
| 334 |
-
//
|
| 335 |
-
//
|
| 336 |
-
//
|
| 337 |
-
|
| 338 |
-
|
| 339 |
-
let
|
| 340 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 341 |
let t = i % params.seq;
|
| 342 |
let tmp = i / params.seq;
|
| 343 |
let hk = tmp % KV_HEADS;
|
|
@@ -358,6 +384,9 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 358 |
present_key[dstBase + d] = past_k[srcBase + d];
|
| 359 |
present_value[dstBase + d] = past_v[srcBase + d];
|
| 360 |
}
|
|
|
|
|
|
|
|
|
|
| 361 |
} else if (t < appendStart + params.keySeq) {
|
| 362 |
let nSrc = (b * params.keySeq + (t - appendStart)) * KV_HIDDEN + hk * HEAD_DIM;
|
| 363 |
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
|
@@ -365,6 +394,7 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 365 |
present_value[dstBase + d] = new_v[nSrc + d];
|
| 366 |
}
|
| 367 |
} else {
|
|
|
|
| 368 |
// Outside the resident window. Cleared rather than left stale so the present
|
| 369 |
// buffer does not expose stale contents from its distinct allocation.
|
| 370 |
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
|
@@ -372,11 +402,9 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 372 |
present_value[dstBase + d] = {{ zeroScalar }}(0.0);
|
| 373 |
}
|
| 374 |
}
|
|
|
|
| 375 |
{% else %}
|
| 376 |
-
|
| 377 |
-
// flat invocation index; this reduces to gid.x when no fold is needed.
|
| 378 |
-
let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
|
| 379 |
-
if (i >= params.count) { return; }
|
| 380 |
let t = i % params.seq;
|
| 381 |
let tmp = i / params.seq;
|
| 382 |
let hk = tmp % KV_HEADS;
|
|
@@ -435,4 +463,4 @@ fn main({% if needsGid %}@builtin(global_invocation_id) gid: vec3<u32>{% if coop
|
|
| 435 |
{% endif %}
|
| 436 |
{% endif %}
|
| 437 |
}
|
| 438 |
-
{%
|
|
|
|
| 9 |
{% set bits = bits | default(0) %}
|
| 10 |
{% set qmax = qmax | default(0) %}
|
| 11 |
{% set qmin = qmin | default(0) %}
|
| 12 |
+
{% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
|
| 13 |
+
{% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
|
| 14 |
+
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 15 |
+
// per-axis workgroup fold width.
|
| 16 |
+
{% if bound == "" %}
|
| 17 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% else %}
|
| 18 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
|
| 19 |
+
if ({{ name }} >= {{ bound }}) { return; }{% endif %}{% endmacro %}
|
| 20 |
+
{{ env.wgsl.resourceDeclarations }}
|
| 21 |
{% if mode == "transpose" %}
|
| 22 |
|
| 23 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
|
|
|
| 38 |
fn main(
|
| 39 |
@builtin(global_invocation_id) gid: vec3<u32>
|
| 40 |
) {
|
| 41 |
+
{{ flat_index_2d("WG", "index", "") }}
|
|
|
|
|
|
|
| 42 |
{% if presentVec4 %}
|
| 43 |
let total = params.batchSize * KV_HEADS * params.kvSeq * HEAD_DIM_V4;
|
| 44 |
if (index >= total) {
|
|
|
|
| 65 |
present_value[index] = {{ presentScalar }}(value[packed_index]);
|
| 66 |
{% endif %}
|
| 67 |
}
|
| 68 |
+
{% else %}
|
| 69 |
|
| 70 |
{% set shareAppend = mode == "merge_share" and shareRegion == "append" %}
|
| 71 |
+
{% set shiftRegion = shiftRegion | default("all") %}
|
| 72 |
+
{% set shiftRetainOnly = mode == "window_shift" and shiftRegion == "shift" %}
|
| 73 |
{% set cooperativeCopy = device.adapterInfo.vendor != "apple" %}
|
| 74 |
{% set cooperativeMerge = cooperativeCopy or shareAppend %}
|
| 75 |
{% set cooperativeMode = cooperativeMerge and (mode == "copy" or mode == "merge" or mode == "merge_share") %}
|
|
|
|
| 102 |
{% if mode != "copy" %}const KV_HEADS: u32 = {{ kvHeads }}u;
|
| 103 |
{% endif %}
|
| 104 |
const WG: u32 = {{ tunables.COPY_WORKGROUP_SIZE }}u;
|
| 105 |
+
{% if mode != "copy" and (mode != "merge_share" or shareAppend) and not shiftRetainOnly %}const KV_HIDDEN: u32 = {{ kvHeads }}u * {{ headDim }}u;
|
| 106 |
{% endif %}
|
| 107 |
{% if mode == "build_quant" or mode == "append_quant" %}const PACKED: u32 = {{ packed }}u;
|
| 108 |
{% endif %}
|
|
|
|
| 160 |
{% endmacro %}
|
| 161 |
{% macro append_per_row_walk() %}
|
| 162 |
// Per-row append walk: one thread streams its row's contiguous bytes.
|
| 163 |
+
{{ flat_index_2d("WG", guardInline=true) }}
|
|
|
|
| 164 |
let j = i % params.keySeq;
|
| 165 |
let tmp = i / params.keySeq;
|
| 166 |
let hk = tmp % KV_HEADS;
|
|
|
|
| 269 |
{% endif %}
|
| 270 |
{% elif mode == "copy" %}
|
| 271 |
// Thread i streams one contiguous (batch, kvHead, token) row.
|
| 272 |
+
{{ flat_index_2d("WG", guardInline=true) }}
|
|
|
|
| 273 |
let base = i * HEAD_DIM;
|
| 274 |
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
| 275 |
present_key[base + d] = src_k[base + d];
|
|
|
|
| 277 |
}
|
| 278 |
{% elif mode == "merge" %}
|
| 279 |
// Per-row merge walk.
|
| 280 |
+
{{ flat_index_2d("WG", guardInline=true) }}
|
|
|
|
| 281 |
let t = i % params.seq;
|
| 282 |
let tmp = i / params.seq;
|
| 283 |
let hk = tmp % KV_HEADS;
|
|
|
|
| 300 |
{% elif mode == "merge_share" %}
|
| 301 |
// Past and present share the full-capacity stride, so outside-window rows
|
| 302 |
// copy at the same index.
|
| 303 |
+
{{ flat_index_2d("WG", guardInline=true) }}
|
|
|
|
| 304 |
let t = i % params.seq;
|
| 305 |
let tmp = i / params.seq;
|
| 306 |
let hk = tmp % KV_HEADS;
|
|
|
|
| 319 |
}
|
| 320 |
}
|
| 321 |
{% elif mode == "window_shift" %}
|
| 322 |
+
// Keep the most recent min(total, capacity) tokens in contiguous cache rows.
|
| 323 |
+
// Read past and write a distinct present buffer so compaction cannot race.
|
| 324 |
+
// Rotary append uses the absolute token position; retained keys are already rotated.
|
| 325 |
+
{% if shiftRegion == "append" %}
|
| 326 |
+
// Append-only pass: the thread range spans just the appended window (batch x
|
| 327 |
+
// kvHead x keySeq rows), while `params.seq` stays the cache capacity. The shift
|
| 328 |
+
// pass owns every other row and the two ranges are disjoint, so their order is
|
| 329 |
+
// immaterial.
|
| 330 |
+
{{ flat_index_2d("WG", guardInline=true) }}
|
| 331 |
+
let j = i % params.keySeq;
|
| 332 |
+
let tmp = i / params.keySeq;
|
| 333 |
+
let hk = tmp % KV_HEADS;
|
| 334 |
+
let b = tmp / KV_HEADS;
|
| 335 |
+
let absTotal = max(params.keySeq, u32(seqlens_k[b]) + 1u);
|
| 336 |
+
let residentBefore = min(absTotal - params.keySeq, params.seq);
|
| 337 |
+
// `max(0u, a + b - c)` does not clamp in u32: the subtraction wraps first, so a
|
| 338 |
+
// step that evicts nothing reads back ~2^32 instead of 0. Compare before
|
| 339 |
+
// subtracting.
|
| 340 |
+
let filled = residentBefore + params.keySeq;
|
| 341 |
+
let evicted = select(0u, filled - params.seq, filled > params.seq);
|
| 342 |
+
let appendStart = residentBefore - evicted;
|
| 343 |
+
let t = appendStart + j;
|
| 344 |
+
let dstBase = ((b * KV_HEADS + hk) * params.seq + t) * HEAD_DIM;
|
| 345 |
+
let nSrc = (b * params.keySeq + j) * KV_HIDDEN + hk * HEAD_DIM;
|
| 346 |
+
// Row appendStart + j carries absolute position absTotal - keySeq + j: the cache
|
| 347 |
+
// holds only the last L positions, so the row index is short of the absolute one
|
| 348 |
+
// by the window origin. The reference rotates the same rows at
|
| 349 |
+
// `appendStart + cacheOrigin`, which is the same number.
|
| 350 |
+
let pos = absTotal - params.keySeq + j;
|
| 351 |
+
var k: array<f32, HEAD_DIM>;
|
| 352 |
+
for (var d = 0u; d < HEAD_DIM; d = d + 1u) { k[d] = new_k[nSrc + d]; }
|
| 353 |
+
for (var d = 0u; d < HALF; d = d + 1u) {
|
| 354 |
+
let cs = cos_cache[pos * HALF + d];
|
| 355 |
+
let sn = sin_cache[pos * HALF + d];
|
| 356 |
+
let x0 = k[d];
|
| 357 |
+
let x1 = k[d + HALF];
|
| 358 |
+
k[d] = x0 * cs - x1 * sn;
|
| 359 |
+
k[d + HALF] = x1 * cs + x0 * sn;
|
| 360 |
+
}
|
| 361 |
+
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
| 362 |
+
present_key[dstBase + d] = k[d];
|
| 363 |
+
present_value[dstBase + d] = new_v[nSrc + d];
|
| 364 |
+
}
|
| 365 |
+
{% else %}
|
| 366 |
+
{{ flat_index_2d("WG", guardInline=true) }}
|
| 367 |
let t = i % params.seq;
|
| 368 |
let tmp = i / params.seq;
|
| 369 |
let hk = tmp % KV_HEADS;
|
|
|
|
| 384 |
present_key[dstBase + d] = past_k[srcBase + d];
|
| 385 |
present_value[dstBase + d] = past_v[srcBase + d];
|
| 386 |
}
|
| 387 |
+
{% if shiftRegion == "shift" %}
|
| 388 |
+
} else if (t >= appendStart + params.keySeq) {
|
| 389 |
+
{% else %}
|
| 390 |
} else if (t < appendStart + params.keySeq) {
|
| 391 |
let nSrc = (b * params.keySeq + (t - appendStart)) * KV_HIDDEN + hk * HEAD_DIM;
|
| 392 |
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
|
|
|
| 394 |
present_value[dstBase + d] = new_v[nSrc + d];
|
| 395 |
}
|
| 396 |
} else {
|
| 397 |
+
{% endif %}
|
| 398 |
// Outside the resident window. Cleared rather than left stale so the present
|
| 399 |
// buffer does not expose stale contents from its distinct allocation.
|
| 400 |
for (var d = 0u; d < HEAD_DIM; d = d + 1u) {
|
|
|
|
| 402 |
present_value[dstBase + d] = {{ zeroScalar }}(0.0);
|
| 403 |
}
|
| 404 |
}
|
| 405 |
+
{% endif %}
|
| 406 |
{% else %}
|
| 407 |
+
{{ flat_index_2d("WG", guardInline=true) }}
|
|
|
|
|
|
|
|
|
|
| 408 |
let t = i % params.seq;
|
| 409 |
let tmp = i / params.seq;
|
| 410 |
let hk = tmp % KV_HEADS;
|
|
|
|
| 463 |
{% endif %}
|
| 464 |
{% endif %}
|
| 465 |
}
|
| 466 |
+
{% endif %}
|
build/webgpu/gqa-qprep.wgsl.jinja
CHANGED
|
@@ -1,5 +1,9 @@
|
|
| 1 |
-
{%
|
| 2 |
-
{%
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
|
| 4 |
// One invocation per (batch, query head, query token) applies optional per-head
|
| 5 |
// RMS normalization followed by NeoX half-split rotary embedding at the query's
|
|
@@ -21,9 +25,7 @@ const QK_EPS: f32 = {{ qkEps }};
|
|
| 21 |
|
| 22 |
@compute @workgroup_size(WG)
|
| 23 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 24 |
-
|
| 25 |
-
// Reduces to gid.x when the dispatch does not fold.
|
| 26 |
-
let qi = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
|
| 27 |
let total = params.batch * Q_HEADS * params.qSeq;
|
| 28 |
if (qi >= total) { return; }
|
| 29 |
let s = qi % params.qSeq;
|
|
|
|
| 1 |
+
{% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
|
| 2 |
+
{% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
|
| 3 |
+
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 4 |
+
// per-axis workgroup fold width.
|
| 5 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
|
| 6 |
+
{{ env.wgsl.resourceDeclarations }}
|
| 7 |
|
| 8 |
// One invocation per (batch, query head, query token) applies optional per-head
|
| 9 |
// RMS normalization followed by NeoX half-split rotary embedding at the query's
|
|
|
|
| 25 |
|
| 26 |
@compute @workgroup_size(WG)
|
| 27 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 28 |
+
{{ flat_index_2d("WG", "qi", "") }}
|
|
|
|
|
|
|
| 29 |
let total = params.batch * Q_HEADS * params.qSeq;
|
| 30 |
if (qi >= total) { return; }
|
| 31 |
let s = qi % params.qSeq;
|
build/webgpu/manifest.json
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,32 +1,32 @@
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.GroupQueryAttention",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
-
"attention-rank4-tiled.wgsl.jinja": "
|
| 11 |
-
"attn-flash-decode-splitk-merge.wgsl.jinja": "
|
| 12 |
-
"attn-flash-decode-splitk.wgsl.jinja": "
|
| 13 |
-
"attn-flash-online.wgsl.jinja": "
|
| 14 |
-
"attn-flash-prefill-cluster.wgsl.jinja": "
|
| 15 |
-
"attn-flash-q32-broadcast.wgsl.jinja": "
|
| 16 |
-
"attn-materialized-rowstats-combine-f32.wgsl.jinja": "
|
| 17 |
-
"attn-materialized-sgmat-f32.wgsl.jinja": "
|
| 18 |
-
"attn-online-scalar.wgsl.jinja": "
|
| 19 |
-
"bench.json": "
|
| 20 |
-
"gqa-attention.wgsl.jinja": "
|
| 21 |
-
"gqa-present.wgsl.jinja": "/
|
| 22 |
-
"gqa-qprep.wgsl.jinja": "
|
| 23 |
-
"manifest.json": "
|
| 24 |
-
"test.json": "
|
| 25 |
}
|
| 26 |
},
|
| 27 |
-
"provenance": { "kernel": { "sha": "
|
| 28 |
"webgpu": {
|
| 29 |
-
"manifestSpec": "2.
|
| 30 |
"variants": {
|
| 31 |
"qkv_present_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 32 |
"past_kv_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
@@ -40,9 +40,12 @@
|
|
| 40 |
"new_kv_share_append_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 41 |
"new_kv_share_append_headsink_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 42 |
"new_kv_share_append_rotary_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
| 43 |
"qkv_present_tiled_nosg": ["attention-rank4-tiled.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 44 |
"qkv_present_flash": ["attn-flash-online.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
| 45 |
"qkv_present": ["attn-online-scalar.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
| 46 |
"quant_int8": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 47 |
"quant_int4": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 48 |
"quant_int8_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
@@ -50,17 +53,9 @@
|
|
| 50 |
"qkv_present_flash_cluster": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 51 |
"past_kv_bias_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 52 |
"past_kv_qnorm_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja", "gqa-qprep.wgsl.jinja"],
|
| 53 |
-
"quant_int8_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 54 |
-
"qkv_present_flash_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 55 |
-
"qkv_present_flash_cluster_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 56 |
-
"past_kv_bias_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 57 |
-
"past_kv_qnorm_rotary_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja", "gqa-qprep.wgsl.jinja"],
|
| 58 |
"past_kv_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 59 |
"new_kv_past_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 60 |
"window_shift_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 61 |
-
"past_kv_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 62 |
-
"new_kv_past_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 63 |
-
"window_shift_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 64 |
"past_kv_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 65 |
"new_kv_past_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 66 |
"past_kv_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
@@ -68,32 +63,27 @@
|
|
| 68 |
"past_kv_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 69 |
"past_kv_bias_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 70 |
"window_shift_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 71 |
-
"past_kv_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 72 |
-
"new_kv_past_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 73 |
-
"past_kv_rotary_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 74 |
-
"past_kv_softcap_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 75 |
-
"past_kv_headsink_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 76 |
-
"past_kv_bias_headsink_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 77 |
-
"window_shift_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 78 |
"qkv_present_flash_q32_broadcast": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 79 |
"qkv_present_flash_q32_shared": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 80 |
"past_kv": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 81 |
"past_kv_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 82 |
"past_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
|
|
|
|
|
|
| 83 |
"new_kv_past": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 84 |
"window_shift_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
|
|
|
| 85 |
"new_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 86 |
"past_kv_bias": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 87 |
"past_kv_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 88 |
"past_kv_bias_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 89 |
"quant_int8_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 90 |
"quant_int4_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 91 |
-
"quant_int8_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 92 |
-
"quant_int4_flash_prefill_nosg": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 93 |
"share_append_split_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 94 |
-
"share_append_split_decode_splitk_nosg": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 95 |
"share_append_split_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 96 |
-
"
|
| 97 |
}
|
| 98 |
}
|
| 99 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.GroupQueryAttention",
|
| 3 |
+
"id": "_com_microsoft_groupqueryattention_webgpu_0c84e3e",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
+
"attention-rank4-tiled.wgsl.jinja": "084x5bbpxeFR95/Q4nGzLrLnk9JJrd9XA1KJHxwU9CY=",
|
| 11 |
+
"attn-flash-decode-splitk-merge.wgsl.jinja": "nKMUcaMP7mcKykLHoGsn/KnsIaJiJRHDsceE7dKVcwo=",
|
| 12 |
+
"attn-flash-decode-splitk.wgsl.jinja": "ZD6C2fn5I7xZvvpTlsHdLtai8JBK2Ar+xc0SrLn3SAE=",
|
| 13 |
+
"attn-flash-online.wgsl.jinja": "QGaUiZyIJuQ9xS00I8cuZWYQI8bGldV+UJ0Mp/5kj6w=",
|
| 14 |
+
"attn-flash-prefill-cluster.wgsl.jinja": "Q7DUmulv66HOVBMkAannNVhEtOWkD+uhDn5muWo4p+E=",
|
| 15 |
+
"attn-flash-q32-broadcast.wgsl.jinja": "vogUJggIxK8tRg2tDORhIXZLheesY3Y221iVFFyHVY4=",
|
| 16 |
+
"attn-materialized-rowstats-combine-f32.wgsl.jinja": "SG+V1o5KhJ1LvkUQL0jcq62C+ETRGpy0XV9tvvOQivk=",
|
| 17 |
+
"attn-materialized-sgmat-f32.wgsl.jinja": "R0+zNbbnfRKpA+R7lcgY5FPlkmLNzNBJwJe1PmaSOvE=",
|
| 18 |
+
"attn-online-scalar.wgsl.jinja": "ysj3CuuKeu+MPlnaTwnNvCu3sSf8egjHE/uhE/Essmo=",
|
| 19 |
+
"bench.json": "v06hV3rz9UPzelx7YbIHmyOWgbj/0Y9ci7iz9mQS6m0=",
|
| 20 |
+
"gqa-attention.wgsl.jinja": "Wy1oR/7g/EzxQZizt/zGtOh4qiD60SVN4pfLSjmfsDo=",
|
| 21 |
+
"gqa-present.wgsl.jinja": "r0ELXXscqaVYqTYzA9IYr/xl9jWAe20ZFPScyiw34UQ=",
|
| 22 |
+
"gqa-qprep.wgsl.jinja": "VI/IHqwEMTpvrhFza3kKRDuLO3Eau91C4aMZRW0nAjg=",
|
| 23 |
+
"manifest.json": "8m8RL1h3mBW9P8voV+gr/KpJiMheynqdWWhmcrOoi8E=",
|
| 24 |
+
"test.json": "I771sUbwb/9MXIc27bUlvNl3UsiiFxP3H0/wRMt/B9M="
|
| 25 |
}
|
| 26 |
},
|
| 27 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 28 |
"webgpu": {
|
| 29 |
+
"manifestSpec": "2.1",
|
| 30 |
"variants": {
|
| 31 |
"qkv_present_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 32 |
"past_kv_materialized_sgmat_f32": ["attn-materialized-rowstats-combine-f32.wgsl.jinja", "attn-materialized-sgmat-f32.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
| 40 |
"new_kv_share_append_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 41 |
"new_kv_share_append_headsink_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 42 |
"new_kv_share_append_rotary_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 43 |
+
"new_kv_share_append_bidirectional_split": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 44 |
"qkv_present_tiled_nosg": ["attention-rank4-tiled.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 45 |
"qkv_present_flash": ["attn-flash-online.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 46 |
+
"qkv_present_flash_causal": ["attn-flash-online.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 47 |
"qkv_present": ["attn-online-scalar.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 48 |
+
"qkv_present_causal": ["attn-online-scalar.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 49 |
"quant_int8": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 50 |
"quant_int4": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 51 |
"quant_int8_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
| 53 |
"qkv_present_flash_cluster": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 54 |
"past_kv_bias_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 55 |
"past_kv_qnorm_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja", "gqa-qprep.wgsl.jinja"],
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
"past_kv_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 57 |
"new_kv_past_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 58 |
"window_shift_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
|
|
|
|
|
|
| 59 |
"past_kv_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 60 |
"new_kv_past_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 61 |
"past_kv_rotary_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
| 63 |
"past_kv_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 64 |
"past_kv_bias_headsink_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 65 |
"window_shift_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 66 |
"qkv_present_flash_q32_broadcast": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 67 |
"qkv_present_flash_q32_shared": ["attn-flash-q32-broadcast.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 68 |
"past_kv": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 69 |
"past_kv_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 70 |
"past_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 71 |
+
"window_shift_rotary_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 72 |
+
"window_shift_rotary_headsink_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 73 |
+
"window_shift_rotary_bias_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 74 |
"new_kv_past": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 75 |
"window_shift_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 76 |
+
"window_shift_headsink_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 77 |
+
"window_shift_bias_append": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 78 |
"new_kv_qnorm_rotary": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 79 |
"past_kv_bias": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 80 |
"past_kv_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 81 |
"past_kv_bias_headsink": ["gqa-attention.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 82 |
"quant_int8_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 83 |
"quant_int4_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
|
|
|
| 84 |
"share_append_split_decode_splitk": ["attn-flash-decode-splitk-merge.wgsl.jinja", "attn-flash-decode-splitk.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
|
|
|
| 85 |
"share_append_split_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"],
|
| 86 |
+
"share_append_bidirectional_flash_prefill": ["attn-flash-prefill-cluster.wgsl.jinja", "gqa-present.wgsl.jinja"]
|
| 87 |
}
|
| 88 |
}
|
| 89 |
}
|
build/webgpu/test.json
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|