sync 6fdf6301e2bb
Browse files- README.md +5 -2
- build/webgpu/attention-rank4-tiled.wgsl.jinja +1 -25
- build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja +2 -9
- build/webgpu/attn-flash-decode-splitk.wgsl.jinja +14 -115
- build/webgpu/attn-flash-online.wgsl.jinja +8 -44
- build/webgpu/attn-flash-prefill-cluster.wgsl.jinja +41 -194
- build/webgpu/attn-flash-q32-broadcast.wgsl.jinja +46 -11
- build/webgpu/attn-materialized-apply-f32.wgsl.jinja +3 -2
- build/webgpu/attn-materialized-rowstats-combine-f32.wgsl.jinja +9 -5
- build/webgpu/attn-materialized-score-f32.wgsl.jinja +6 -85
- build/webgpu/attn-materialized-sgmat-f32.wgsl.jinja +28 -54
- build/webgpu/attn-materialized-softmax-f32.wgsl.jinja +6 -41
- build/webgpu/attn-online-scalar.wgsl.jinja +39 -28
- build/webgpu/attn-small-head-parallel.wgsl.jinja +3 -9
- build/webgpu/attn-small-head-value.wgsl.jinja +198 -0
- build/webgpu/bench.json +0 -0
- build/webgpu/manifest.json +224 -362
- build/webgpu/metadata.json +24 -21
- build/webgpu/mha-small-seq-blocked.wgsl.jinja +5 -4
- build/webgpu/mha-small-seq.wgsl.jinja +26 -6
- build/webgpu/test.json +0 -0
README.md
CHANGED
|
@@ -52,6 +52,8 @@ Attributes and default values (overridable per request):
|
|
| 52 |
|
| 53 |
One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
|
| 54 |
|
|
|
|
|
|
|
| 55 |
- `qkv_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences with bias and optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
|
| 56 |
- `qkv_no_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences without bias and with optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
|
| 57 |
- `qkv_no_bias_small_seq` — Whole-head attention for short bidirectional float32 requests without bias. One workgroup stages a head's complete key and value planes, and each participating invocation owns one query row through scoring, normalization, and context accumulation.
|
|
@@ -100,7 +102,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
|
|
| 100 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 101 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 102 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 103 |
-
- [`bench.json`](build/webgpu/bench.json) — benchmark
|
| 104 |
- [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
|
| 105 |
- [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
|
| 106 |
- [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
|
|
@@ -114,13 +116,14 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
|
|
| 114 |
- [`attn-materialized-softmax-f32.wgsl.jinja`](build/webgpu/attn-materialized-softmax-f32.wgsl.jinja)
|
| 115 |
- [`attn-online-scalar.wgsl.jinja`](build/webgpu/attn-online-scalar.wgsl.jinja)
|
| 116 |
- [`attn-small-head-parallel.wgsl.jinja`](build/webgpu/attn-small-head-parallel.wgsl.jinja)
|
|
|
|
| 117 |
- [`mha-small-seq-blocked.wgsl.jinja`](build/webgpu/mha-small-seq-blocked.wgsl.jinja)
|
| 118 |
- [`mha-small-seq.wgsl.jinja`](build/webgpu/mha-small-seq.wgsl.jinja)
|
| 119 |
|
| 120 |
## Use with `@huggingface/kernels`
|
| 121 |
|
| 122 |
```sh
|
| 123 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 124 |
```
|
| 125 |
|
| 126 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
|
|
| 52 |
|
| 53 |
One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
|
| 54 |
|
| 55 |
+
- `qkv_no_bias_small_head_value_subgroups` — Cooperatively reduces values across key lanes and reuses keys/values for two queries when device workgroup storage fits. Uses portable-width subgroup collectives with f32 scores, probabilities and accumulation. Short key sequences use the serial value path.
|
| 56 |
+
- `qkv_no_bias_small_head_value_tree` — Cooperatively reduces values across key lanes and reuses keys/values for two queries when device workgroup storage fits. Uses shared-memory tree reductions with f32 scores, probabilities and accumulation. Short key sequences use the serial value path.
|
| 57 |
- `qkv_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences with bias and optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
|
| 58 |
- `qkv_no_bias_small_seq_blocked` — Blocked whole-head attention for short key sequences without bias and with optional causal masking. A workgroup stages one head's K and V once, then assigns each query in a small block to a group of key lanes whose register-resident online-softmax partials are merged in one pass.
|
| 59 |
- `qkv_no_bias_small_seq` — Whole-head attention for short bidirectional float32 requests without bias. One workgroup stages a head's complete key and value planes, and each participating invocation owns one query row through scoring, normalization, and context accumulation.
|
|
|
|
| 102 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 103 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 104 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 105 |
+
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
|
| 106 |
- [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
|
| 107 |
- [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
|
| 108 |
- [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
|
|
|
|
| 116 |
- [`attn-materialized-softmax-f32.wgsl.jinja`](build/webgpu/attn-materialized-softmax-f32.wgsl.jinja)
|
| 117 |
- [`attn-online-scalar.wgsl.jinja`](build/webgpu/attn-online-scalar.wgsl.jinja)
|
| 118 |
- [`attn-small-head-parallel.wgsl.jinja`](build/webgpu/attn-small-head-parallel.wgsl.jinja)
|
| 119 |
+
- [`attn-small-head-value.wgsl.jinja`](build/webgpu/attn-small-head-value.wgsl.jinja)
|
| 120 |
- [`mha-small-seq-blocked.wgsl.jinja`](build/webgpu/mha-small-seq-blocked.wgsl.jinja)
|
| 121 |
- [`mha-small-seq.wgsl.jinja`](build/webgpu/mha-small-seq.wgsl.jinja)
|
| 122 |
|
| 123 |
## Use with `@huggingface/kernels`
|
| 124 |
|
| 125 |
```sh
|
| 126 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 127 |
```
|
| 128 |
|
| 129 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
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 |
-
{% if scalar == "f16" %}
|
| 2 |
-
enable f16;
|
| 3 |
-
{% endif %}
|
| 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
|
|
@@ -49,6 +46,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 49 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 50 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 51 |
}
|
|
|
|
| 52 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 53 |
return exp(shifted_value(value, maxValue));
|
| 54 |
}
|
|
@@ -104,11 +102,6 @@ fn main(
|
|
| 104 |
// V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
|
| 105 |
let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
|
| 106 |
outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
|
| 107 |
-
{% endif %}
|
| 108 |
-
{% if hasGate %}
|
| 109 |
-
// The gated route multiplies the normalized attention output elementwise by its gate.
|
| 110 |
-
let gateV = vec4<f32>(gate[qBaseV4 + d4]);
|
| 111 |
-
outValue = outValue * (vec4<f32>(1.0) / (vec4<f32>(1.0) + exp(-gateV)));
|
| 112 |
{% endif %}
|
| 113 |
output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
|
| 114 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
{% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
|
|
|
|
| 2 |
{% set qSeq = qSeq | default(0) %}
|
| 3 |
+
{% set qHidden = qHidden | default(0) %}
|
| 4 |
{{ env.wgsl.resourceDeclarations }}
|
| 5 |
|
| 6 |
// Split-K flash decode, pass 2 of 2. Combines the per-split un-normalized online
|
|
|
|
| 46 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 47 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 48 |
}
|
| 49 |
+
|
| 50 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 51 |
return exp(shifted_value(value, maxValue));
|
| 52 |
}
|
|
|
|
| 102 |
// V bias row base: 2 * qHidden (skip the packed Q and K bias blocks) + this head.
|
| 103 |
let vBiasBase = 2u * Q_HIDDEN + h * HEAD_DIM + d4 * 4u;
|
| 104 |
outValue = outValue + vec4<f32>(bias[vBiasBase], bias[vBiasBase + 1u], bias[vBiasBase + 2u], bias[vBiasBase + 3u]);
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
{% endif %}
|
| 106 |
output[qBaseV4 + d4] = vec4<{{ scalar }}>(outValue);
|
| 107 |
}
|
build/webgpu/attn-flash-decode-splitk.wgsl.jinja
CHANGED
|
@@ -1,11 +1,6 @@
|
|
| 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 |
-
{% if maskIsBool is not defined %}{% set maskIsBool = false %}{% endif %}
|
| 7 |
-
{% set scale = scale | default("0.0") %}
|
| 8 |
-
{% if hasRotary is not defined %}{% set hasRotary = false %}{% endif %}
|
| 9 |
{% set qSeq = qSeq | default(0) %}
|
| 10 |
{% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
|
| 11 |
{% if useSubgroups %}
|
|
@@ -13,7 +8,8 @@ enable subgroups;
|
|
| 13 |
{% endif %}
|
| 14 |
{{ env.wgsl.resourceDeclarations }}
|
| 15 |
|
| 16 |
-
// Split-K flash attention
|
|
|
|
| 17 |
// handles decode and short-query, long-context prefill inputs.
|
| 18 |
//
|
| 19 |
// The non-split flash decode launches only `batch * numHeads` workgroups, each
|
|
@@ -56,6 +52,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 56 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 57 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 58 |
}
|
|
|
|
| 59 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 60 |
return exp(shifted_value(value, maxValue));
|
| 61 |
}
|
|
@@ -63,7 +60,7 @@ fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
|
| 63 |
var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
|
| 64 |
var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
|
| 65 |
var<workgroup> probs: array<f32, WG>;
|
| 66 |
-
{% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) %}
|
| 67 |
{% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
|
| 68 |
{% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
|
| 69 |
{% if coopQk %}
|
|
@@ -132,41 +129,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
|
|
| 132 |
return combinedMD;
|
| 133 |
}
|
| 134 |
{% else %}
|
| 135 |
-
{% set
|
| 136 |
-
{% set mdStreams = mdStreams if mdStreams is defined else 1 %}
|
| 137 |
-
{% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
|
| 138 |
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 139 |
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
| 140 |
-
{% if mdStreamed %}
|
| 141 |
-
|
| 142 |
-
// In-place fold of {{ mdStreams }} streams. Input partials occupy
|
| 143 |
-
// partialM/partialD; stream s returns its merged pair in slot s * WG.
|
| 144 |
-
fn combine_partials_streams(lidx: u32) {
|
| 145 |
-
workgroupBarrier();
|
| 146 |
-
var stride = WG / 2u;
|
| 147 |
-
loop {
|
| 148 |
-
if (stride == 0u) {
|
| 149 |
-
break;
|
| 150 |
-
}
|
| 151 |
-
if (lidx < stride) {
|
| 152 |
-
{% for s in range(mdStreams) %}
|
| 153 |
-
{
|
| 154 |
-
let slot = {{ s }}u * WG + lidx;
|
| 155 |
-
let m1 = partialM[slot];
|
| 156 |
-
let d1 = partialD[slot];
|
| 157 |
-
let m2 = partialM[slot + stride];
|
| 158 |
-
let d2 = partialD[slot + stride];
|
| 159 |
-
let mNew = max(m1, m2);
|
| 160 |
-
partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
|
| 161 |
-
partialM[slot] = mNew;
|
| 162 |
-
}
|
| 163 |
-
{% endfor %}
|
| 164 |
-
}
|
| 165 |
-
workgroupBarrier();
|
| 166 |
-
stride = stride / 2u;
|
| 167 |
-
}
|
| 168 |
-
}
|
| 169 |
-
{% else %}
|
| 170 |
|
| 171 |
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 172 |
partialM[lidx] = m;
|
|
@@ -196,54 +161,12 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
|
| 196 |
return merged;
|
| 197 |
}
|
| 198 |
{% endif %}
|
| 199 |
-
{% endif %}
|
| 200 |
-
|
| 201 |
|
| 202 |
-
{%
|
| 203 |
-
fn scale_value() -> f32 {
|
| 204 |
if (params.scale != 0.0) { return params.scale; }
|
| 205 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 206 |
}
|
| 207 |
|
| 208 |
-
|
| 209 |
-
{% if quantizedCache %}
|
| 210 |
-
{% macro emit_quant_scale4(kind, scaleBuffer) %}
|
| 211 |
-
fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
|
| 212 |
-
if (params.perChannel == 0u) {
|
| 213 |
-
return vec4<f32>({{ scaleBuffer }}[0]);
|
| 214 |
-
}
|
| 215 |
-
let base = hk * HEAD_DIM + d4 * 4u;
|
| 216 |
-
return vec4<f32>(
|
| 217 |
-
{{ scaleBuffer }}[base],
|
| 218 |
-
{{ scaleBuffer }}[base + 1u],
|
| 219 |
-
{{ scaleBuffer }}[base + 2u],
|
| 220 |
-
{{ scaleBuffer }}[base + 3u]
|
| 221 |
-
);
|
| 222 |
-
}
|
| 223 |
-
{%- endmacro %}
|
| 224 |
-
{%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
|
| 225 |
-
{{ emit_quant_scale4(kind, scaleBuffer) }}
|
| 226 |
-
fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
|
| 227 |
-
{%- if format == "int8" %}
|
| 228 |
-
return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
|
| 229 |
-
{%- else %}
|
| 230 |
-
// Two elements cover this vec4: each carries two +8-biased nibbles, low first.
|
| 231 |
-
let rowBase = indexV4 - d4;
|
| 232 |
-
let lo = {{ buffer }}[rowBase + d4 * 2u];
|
| 233 |
-
let hi = {{ buffer }}[rowBase + d4 * 2u + 1u];
|
| 234 |
-
let nibbles = vec4<i32>(
|
| 235 |
-
i32(lo & 0xFu), i32((lo >> 4u) & 0xFu),
|
| 236 |
-
i32(hi & 0xFu), i32((hi >> 4u) & 0xFu)
|
| 237 |
-
);
|
| 238 |
-
let signed = nibbles - vec4<i32>(8);
|
| 239 |
-
return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
|
| 240 |
-
{%- endif %}
|
| 241 |
-
}
|
| 242 |
-
{%- endmacro %}
|
| 243 |
-
|
| 244 |
-
{{ emit_quant_load4("int8", "key", "key", "k_scale") }}
|
| 245 |
-
{{ emit_quant_load4("int8", "value", "value", "v_scale") }}
|
| 246 |
-
{% else %}
|
| 247 |
fn load_key4(indexV4: u32) -> vec4<f32> {
|
| 248 |
return vec4<f32>(key[indexV4]);
|
| 249 |
}
|
|
@@ -251,16 +174,17 @@ fn load_key4(indexV4: u32) -> vec4<f32> {
|
|
| 251 |
fn load_value4(indexV4: u32) -> vec4<f32> {
|
| 252 |
return vec4<f32>(value[indexV4]);
|
| 253 |
}
|
| 254 |
-
{% endif %}
|
| 255 |
|
| 256 |
{% if hasBias %}
|
| 257 |
// Packed [Q; K; V] bias rows (token-independent). The Q bias folds into the
|
| 258 |
// query row before the Q.K dots; the K bias adds a constant to every key score
|
| 259 |
// that softmax cancels, so it is skipped; the V bias is token-independent and
|
| 260 |
// is applied once in the merge pass after the final normalize.
|
|
|
|
|
|
|
| 261 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 262 |
let offset = base + d4 * 4u;
|
| 263 |
-
return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
|
| 264 |
}
|
| 265 |
|
| 266 |
{% endif %}
|
|
@@ -290,13 +214,7 @@ fn main(
|
|
| 290 |
let tid = lid.x;
|
| 291 |
let hKv = h / (Q_HEADS / KV_HEADS);
|
| 292 |
let cacheSeq = params.kvSeq;
|
| 293 |
-
{% if cacheSeqlens %}
|
| 294 |
-
// Buffer-sharing caches retain their capacity in the physical BNSH stride;
|
| 295 |
-
// seqlens_k supplies the active end independently for each batch.
|
| 296 |
-
let kvSeq = min(cacheSeq, u32(seqlens_k[b]) + 1u);
|
| 297 |
-
{% else %}
|
| 298 |
let kvSeq = cacheSeq;
|
| 299 |
-
{% endif %}
|
| 300 |
|
| 301 |
// Query row (decode uses token zero; short-query prefill folds the token into wg.x).
|
| 302 |
{% if splitQueries %}
|
|
@@ -317,7 +235,6 @@ fn main(
|
|
| 317 |
splitEnd = kvSeq;
|
| 318 |
}
|
| 319 |
|
| 320 |
-
{% set hasBias = hasBias is defined and hasBias %}
|
| 321 |
for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
|
| 322 |
var qv = vec4<f32>(query[qBaseV4 + d4]);
|
| 323 |
{% if hasBias %}
|
|
@@ -357,7 +274,7 @@ fn main(
|
|
| 357 |
if (j < tileCount) {
|
| 358 |
let kRowV4 = kvBaseV4 + (kjBase + j) * kvTokenStrideV4;
|
| 359 |
for (var d4: u32 = lane; d4 < HEAD_DIM_V4; d4 = d4 + sgSize) {
|
| 360 |
-
accS = accS + dot(q_shared[d4], load_key4(kRowV4 + d4
|
| 361 |
}
|
| 362 |
}
|
| 363 |
let sj = subgroupAdd(accS);
|
|
@@ -376,7 +293,7 @@ fn main(
|
|
| 376 |
let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
|
| 377 |
var acc: f32 = 0.0;
|
| 378 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 379 |
-
acc = acc + dot(q_shared[d4], load_key4(kRowV4 + d4
|
| 380 |
}
|
| 381 |
score = acc * scale;
|
| 382 |
m = score;
|
|
@@ -391,21 +308,10 @@ fn main(
|
|
| 391 |
let maskQuery = 0u;
|
| 392 |
{% endif %}
|
| 393 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
|
| 394 |
-
{% if maskIsBool %}
|
| 395 |
-
// A rejected bool-mask key contributes no probability mass. The merge
|
| 396 |
-
// pass already maps a zero global denominator to an all-zero output row.
|
| 397 |
-
if (attn_mask[maskIndex] == 0u) {
|
| 398 |
-
keyAllowed = false;
|
| 399 |
-
score = -FLT_MAX;
|
| 400 |
-
dPart = 0.0;
|
| 401 |
-
}
|
| 402 |
-
{% else %}
|
| 403 |
score = score + f32(attn_mask[maskIndex]);
|
| 404 |
-
{% endif %}
|
| 405 |
m = score;
|
| 406 |
}
|
| 407 |
-
{% endif %}
|
| 408 |
-
let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
|
| 409 |
|
| 410 |
// Merge one key tile's online-softmax (maximum, denominator) partial into the
|
| 411 |
// running state, then store the per-key probabilities consumed by V accumulation.
|
|
@@ -421,7 +327,6 @@ fn main(
|
|
| 421 |
probs[tid] = prob;
|
| 422 |
workgroupBarrier();
|
| 423 |
|
| 424 |
-
|
| 425 |
{% if jSplitV %}
|
| 426 |
// j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
|
| 427 |
// J_GROUPS for dim block d4v into a register, then the groups combine
|
|
@@ -434,10 +339,7 @@ fn main(
|
|
| 434 |
loop {
|
| 435 |
if (jj >= tileCount) { break; }
|
| 436 |
vacc = vacc + probs[jj] * load_value4(
|
| 437 |
-
kvBaseV4 + (kjBase + jj) * kvTokenStrideV4 + d4v
|
| 438 |
-
d4v,
|
| 439 |
-
hKv{% endif %}
|
| 440 |
-
);
|
| 441 |
jj = jj + J_GROUPS;
|
| 442 |
}
|
| 443 |
vacc_sh[tid] = vacc;
|
|
@@ -455,10 +357,7 @@ fn main(
|
|
| 455 |
var vSum = vec4<f32>(0.0);
|
| 456 |
for (var i: u32 = 0u; i < tileCount; i = i + 1u) {
|
| 457 |
vSum = vSum + probs[i] * load_value4(
|
| 458 |
-
kvBaseV4 + (kjBase + i) * kvTokenStrideV4 + d4
|
| 459 |
-
d4,
|
| 460 |
-
hKv{% endif %}
|
| 461 |
-
);
|
| 462 |
}
|
| 463 |
running_out[d4] = running_out[d4] * correction + vSum;
|
| 464 |
}
|
|
|
|
| 1 |
{% if useSubgroups is not defined %}{% set useSubgroups = true %}{% endif %}
|
| 2 |
{% if splitQueries is not defined %}{% set splitQueries = false %}{% endif %}
|
|
|
|
|
|
|
| 3 |
{% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
|
|
|
|
|
|
|
|
|
|
| 4 |
{% set qSeq = qSeq | default(0) %}
|
| 5 |
{% set splitKWorkgroupSize = workgroupSizeSpec if workgroupSizeSpec is defined else tunables.WORKGROUP_SIZE %}
|
| 6 |
{% if useSubgroups %}
|
|
|
|
| 8 |
{% endif %}
|
| 9 |
{{ env.wgsl.resourceDeclarations }}
|
| 10 |
|
| 11 |
+
// Split-K flash attention. Single-partition direct output skips the merge pass.
|
| 12 |
+
// Otherwise this is the first of two passes. This geometry
|
| 13 |
// handles decode and short-query, long-context prefill inputs.
|
| 14 |
//
|
| 15 |
// The non-split flash decode launches only `batch * numHeads` workgroups, each
|
|
|
|
| 52 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 53 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 54 |
}
|
| 55 |
+
|
| 56 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 57 |
return exp(shifted_value(value, maxValue));
|
| 58 |
}
|
|
|
|
| 60 |
var<workgroup> q_shared: array<vec4<f32>, HEAD_DIM_V4>;
|
| 61 |
var<workgroup> running_out: array<vec4<f32>, HEAD_DIM_V4>;
|
| 62 |
var<workgroup> probs: array<f32, WG>;
|
| 63 |
+
{% set coopQk = useSubgroups and headDimV4 >= 8 and not (usesF16 and headDimV4 <= 32) and (allowCooperativeQk if allowCooperativeQk is defined else true) %}
|
| 64 |
{% set jGroups = (splitKWorkgroupSize / headDimV4)|int %}
|
| 65 |
{% set jSplitV = (splitKWorkgroupSize % headDimV4 == 0) and (jGroups >= 2) %}
|
| 66 |
{% if coopQk %}
|
|
|
|
| 129 |
return combinedMD;
|
| 130 |
}
|
| 131 |
{% else %}
|
| 132 |
+
{% set mdExtent = "WG" %}
|
|
|
|
|
|
|
| 133 |
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 134 |
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 135 |
|
| 136 |
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 137 |
partialM[lidx] = m;
|
|
|
|
| 161 |
return merged;
|
| 162 |
}
|
| 163 |
{% endif %}
|
|
|
|
|
|
|
| 164 |
|
| 165 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 166 |
if (params.scale != 0.0) { return params.scale; }
|
| 167 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 168 |
}
|
| 169 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 170 |
fn load_key4(indexV4: u32) -> vec4<f32> {
|
| 171 |
return vec4<f32>(key[indexV4]);
|
| 172 |
}
|
|
|
|
| 174 |
fn load_value4(indexV4: u32) -> vec4<f32> {
|
| 175 |
return vec4<f32>(value[indexV4]);
|
| 176 |
}
|
|
|
|
| 177 |
|
| 178 |
{% if hasBias %}
|
| 179 |
// Packed [Q; K; V] bias rows (token-independent). The Q bias folds into the
|
| 180 |
// query row before the Q.K dots; the K bias adds a constant to every key score
|
| 181 |
// that softmax cancels, so it is skipped; the V bias is token-independent and
|
| 182 |
// is applied once in the merge pass after the final normalize.
|
| 183 |
+
{% set BW = "" %}
|
| 184 |
+
{% set BC = "" %}
|
| 185 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 186 |
let offset = base + d4 * 4u;
|
| 187 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 188 |
}
|
| 189 |
|
| 190 |
{% endif %}
|
|
|
|
| 214 |
let tid = lid.x;
|
| 215 |
let hKv = h / (Q_HEADS / KV_HEADS);
|
| 216 |
let cacheSeq = params.kvSeq;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 217 |
let kvSeq = cacheSeq;
|
|
|
|
| 218 |
|
| 219 |
// Query row (decode uses token zero; short-query prefill folds the token into wg.x).
|
| 220 |
{% if splitQueries %}
|
|
|
|
| 235 |
splitEnd = kvSeq;
|
| 236 |
}
|
| 237 |
|
|
|
|
| 238 |
for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
|
| 239 |
var qv = vec4<f32>(query[qBaseV4 + d4]);
|
| 240 |
{% if hasBias %}
|
|
|
|
| 274 |
if (j < tileCount) {
|
| 275 |
let kRowV4 = kvBaseV4 + (kjBase + j) * kvTokenStrideV4;
|
| 276 |
for (var d4: u32 = lane; d4 < HEAD_DIM_V4; d4 = d4 + sgSize) {
|
| 277 |
+
accS = accS + dot(q_shared[d4], load_key4(kRowV4 + d4));
|
| 278 |
}
|
| 279 |
}
|
| 280 |
let sj = subgroupAdd(accS);
|
|
|
|
| 293 |
let kRowV4 = kvBaseV4 + kj * kvTokenStrideV4;
|
| 294 |
var acc: f32 = 0.0;
|
| 295 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 296 |
+
acc = acc + dot(q_shared[d4], load_key4(kRowV4 + d4));
|
| 297 |
}
|
| 298 |
score = acc * scale;
|
| 299 |
m = score;
|
|
|
|
| 308 |
let maskQuery = 0u;
|
| 309 |
{% endif %}
|
| 310 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQuery * params.maskSeqStride + kj;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 311 |
score = score + f32(attn_mask[maskIndex]);
|
|
|
|
| 312 |
m = score;
|
| 313 |
}
|
| 314 |
+
{% endif %} let tile = combine_partials(m, dPart, tid{% if useSubgroups %}, sgSize{% endif %});
|
|
|
|
| 315 |
|
| 316 |
// Merge one key tile's online-softmax (maximum, denominator) partial into the
|
| 317 |
// running state, then store the per-key probabilities consumed by V accumulation.
|
|
|
|
| 327 |
probs[tid] = prob;
|
| 328 |
workgroupBarrier();
|
| 329 |
|
|
|
|
| 330 |
{% if jSplitV %}
|
| 331 |
// j-split V accumulation: thread (jg, d4v) sums keys j == jg mod
|
| 332 |
// J_GROUPS for dim block d4v into a register, then the groups combine
|
|
|
|
| 339 |
loop {
|
| 340 |
if (jj >= tileCount) { break; }
|
| 341 |
vacc = vacc + probs[jj] * load_value4(
|
| 342 |
+
kvBaseV4 + (kjBase + jj) * kvTokenStrideV4 + d4v );
|
|
|
|
|
|
|
|
|
|
| 343 |
jj = jj + J_GROUPS;
|
| 344 |
}
|
| 345 |
vacc_sh[tid] = vacc;
|
|
|
|
| 357 |
var vSum = vec4<f32>(0.0);
|
| 358 |
for (var i: u32 = 0u; i < tileCount; i = i + 1u) {
|
| 359 |
vSum = vSum + probs[i] * load_value4(
|
| 360 |
+
kvBaseV4 + (kjBase + i) * kvTokenStrideV4 + d4 );
|
|
|
|
|
|
|
|
|
|
| 361 |
}
|
| 362 |
running_out[d4] = running_out[d4] * correction + vSum;
|
| 363 |
}
|
build/webgpu/attn-flash-online.wgsl.jinja
CHANGED
|
@@ -28,8 +28,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 +50,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,7 +64,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
|
|
@@ -116,41 +116,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
|
|
| 116 |
return combinedMD;
|
| 117 |
}
|
| 118 |
{% else %}
|
| 119 |
-
{% set
|
| 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;
|
|
@@ -180,22 +148,20 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
|
| 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 %}
|
|
@@ -221,7 +187,6 @@ fn main(
|
|
| 221 |
let kvTokenStrideV4 = KV_HIDDEN_V4;
|
| 222 |
|
| 223 |
// Cooperative vec4 Q-row load; init the output accumulator.
|
| 224 |
-
{% set hasBias = hasBias is defined and hasBias %}
|
| 225 |
for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
|
| 226 |
var qv = vec4<f32>(query[qBaseV4 + d4]);
|
| 227 |
{% if hasBias %}
|
|
@@ -297,7 +262,6 @@ fn main(
|
|
| 297 |
probs[tid] = prob;
|
| 298 |
workgroupBarrier();
|
| 299 |
|
| 300 |
-
|
| 301 |
// running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
|
| 302 |
// tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
|
| 303 |
let tileCount = min(WG, keyBound - kjBase);
|
|
|
|
| 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 = "Q_HEADS" %}
|
| 32 |
+
{% set kvHeads = "KV_HEADS" %}
|
| 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 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 51 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 52 |
}
|
| 53 |
+
|
| 54 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 55 |
return exp(shifted_value(value, maxValue));
|
| 56 |
}
|
|
|
|
| 64 |
// Both the subgroup and portable barrier-tree engines return the same merged
|
| 65 |
// pair to every invocation. Repeated merges require a workgroup barrier between
|
| 66 |
// calls before their shared partial storage is reused.
|
|
|
|
| 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
|
|
|
|
| 116 |
return combinedMD;
|
| 117 |
}
|
| 118 |
{% else %}
|
| 119 |
+
{% set mdExtent = "WG" %}
|
|
|
|
|
|
|
| 120 |
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 121 |
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 122 |
|
| 123 |
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 124 |
partialM[lidx] = m;
|
|
|
|
| 148 |
return merged;
|
| 149 |
}
|
| 150 |
{% endif %}
|
|
|
|
|
|
|
| 151 |
|
| 152 |
// An explicit-zero specialization bakes the scale as 0. Otherwise,
|
| 153 |
// params.scale == 0 encodes an omitted scale and selects 1/sqrt(headDim).
|
| 154 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 155 |
if (params.scale != 0.0) { return params.scale; }
|
| 156 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 157 |
}
|
|
|
|
| 158 |
{% if hasBias %}
|
| 159 |
|
| 160 |
+
{% set BW = "" %}
|
| 161 |
+
{% set BC = "" %}
|
| 162 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 163 |
let offset = base + d4 * 4u;
|
| 164 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 165 |
}
|
| 166 |
|
| 167 |
{% endif %}
|
|
|
|
| 187 |
let kvTokenStrideV4 = KV_HIDDEN_V4;
|
| 188 |
|
| 189 |
// Cooperative vec4 Q-row load; init the output accumulator.
|
|
|
|
| 190 |
for (var d4 = tid; d4 < HEAD_DIM_V4; d4 = d4 + WG) {
|
| 191 |
var qv = vec4<f32>(query[qBaseV4 + d4]);
|
| 192 |
{% if hasBias %}
|
|
|
|
| 262 |
probs[tid] = prob;
|
| 263 |
workgroupBarrier();
|
| 264 |
|
|
|
|
| 265 |
// running_out[d4] is owned by the same thread (tid ≡ d4 mod WG) across all
|
| 266 |
// tiles, so the rescale-and-accumulate below is race-free; vec4 V loads.
|
| 267 |
let tileCount = min(WG, keyBound - kjBase);
|
build/webgpu/attn-flash-prefill-cluster.wgsl.jinja
CHANGED
|
@@ -1,32 +1,27 @@
|
|
| 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 |
-
{% if maskIsKeyKeep is not defined %}{% set maskIsKeyKeep = false %}{% endif %}
|
| 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 |
{% set Q_HIDDEN = qHidden | default(0) %}
|
| 21 |
-
{% if quantCacheFormat is not defined %}{% set quantCacheFormat = "" %}{% endif %}
|
| 22 |
-
{% if useSeqlens is not defined %}{% set useSeqlens = false %}{% endif %}
|
| 23 |
// A windowed cache binds a fixed CAPACITY but keeps only the most recent
|
| 24 |
// min(total, capacity) rows resident. params.kvSeq then names the physical row
|
| 25 |
// count, which is still the right batch stride but the wrong attention bound, so
|
| 26 |
// the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
|
| 27 |
// in both roles.
|
| 28 |
-
{%
|
| 29 |
-
{%
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
{% if useSubgroups %}
|
| 31 |
enable subgroups;
|
| 32 |
{% endif %}
|
|
@@ -36,15 +31,30 @@ enable subgroups;
|
|
| 36 |
// workgroup storage; values widen to f32 when read. Other specializations stage
|
| 37 |
// them as f32. Score and weighted-value accumulation remain in f32 throughout.
|
| 38 |
{% set STAGE_MASK = hasMask and useSubgroups and stageMask is defined and stageMask %}
|
| 39 |
-
{% set MASK_IS_INT = hasMask and (maskIsKeyKeep or maskIsBool) %}
|
| 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 |
|
| 49 |
// Tiled flash prefill attention with configurable-width query clusters for
|
| 50 |
// token-major [batch, seq, heads*headDim] attention ops. Each workgroup covers
|
|
@@ -64,8 +74,7 @@ enable subgroups;
|
|
| 64 |
const HEAD_DIM: u32 = {{ headDim }}u;
|
| 65 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
| 66 |
const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
|
| 67 |
-
|
| 68 |
-
{% endif %}const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
|
| 69 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 70 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 71 |
const TILE_Q: u32 = {{ TILE_Q }}u;
|
|
@@ -79,24 +88,10 @@ const WG: u32 = (TILE_Q / QPL) * LPQ;
|
|
| 79 |
{% else %}
|
| 80 |
const WG: u32 = TILE_Q * LPQ;
|
| 81 |
{% endif %}
|
| 82 |
-
{% if MASK_IS_INT %}
|
| 83 |
-
// Key-keep masks in contrib attention use a finite low logit for a rejected
|
| 84 |
-
// key. Logical ONNX bool masks use the exclusion sentinel instead, so a fully
|
| 85 |
-
// masked row has zero mass. Keeping one declaration shape for both mask modes
|
| 86 |
-
// lets both mask modes share the score loop below.
|
| 87 |
const NEG_INF: f32 = -3.4028234663852886e38;
|
| 88 |
-
const MASK_NEG: f32 = {{ "-1e38" if maskIsKeyKeep else "-3.4028234663852886e38" }};
|
| 89 |
-
{% else %}
|
| 90 |
-
const NEG_INF: f32 = -3.4028234663852886e38;
|
| 91 |
-
{% endif %}
|
| 92 |
|
| 93 |
var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
| 94 |
var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
| 95 |
-
{% if STAGE_MASK %}
|
| 96 |
-
// Each LPQ cluster consumes one mask value per (query,key), so stage the
|
| 97 |
-
// TILE_Q x TILE_K mask tile once instead of issuing LPQ duplicate global loads.
|
| 98 |
-
var<workgroup> mask_tile: array<{{ MASK_TILE_TYPE }}, TILE_Q * TILE_K>;
|
| 99 |
-
{% endif %}
|
| 100 |
{% if not useSubgroups %}
|
| 101 |
{% if batchNoSgReduction %}
|
| 102 |
// No-subgroups cluster reduction scratch for a whole K tile. Staging every
|
|
@@ -109,62 +104,17 @@ var<workgroup> red: array<f32, WG>;
|
|
| 109 |
{% endif %}
|
| 110 |
{% endif %}
|
| 111 |
|
| 112 |
-
{%
|
| 113 |
-
fn scale_value() -> f32 {
|
| 114 |
-
{% if ATTN_SCALE_OVERRIDE is defined %}
|
| 115 |
-
return {{ ATTN_SCALE_OVERRIDE }};
|
| 116 |
-
{% else %}
|
| 117 |
if (params.scale != 0.0) { return params.scale; }
|
| 118 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 119 |
-
{% endif %}
|
| 120 |
-
}
|
| 121 |
-
|
| 122 |
-
{% if quantCacheFormat %}
|
| 123 |
-
// A quantized cache is dequantized once per key into the staged tile, then read
|
| 124 |
-
// by all TILE_Q queries in the workgroup. The unpack cost is amortized over the
|
| 125 |
-
// tile height instead of paid once per (query, key).
|
| 126 |
-
{% macro emit_quant_scale4(kind, scaleBuffer) %}
|
| 127 |
-
fn {{ kind }}scale4(d4: u32, hk: u32) -> vec4<f32> {
|
| 128 |
-
if (params.perChannel == 0u) {
|
| 129 |
-
return vec4<f32>({{ scaleBuffer }}[0]);
|
| 130 |
-
}
|
| 131 |
-
let base = hk * HEAD_DIM + d4 * 4u;
|
| 132 |
-
return vec4<f32>(
|
| 133 |
-
{{ scaleBuffer }}[base],
|
| 134 |
-
{{ scaleBuffer }}[base + 1u],
|
| 135 |
-
{{ scaleBuffer }}[base + 2u],
|
| 136 |
-
{{ scaleBuffer }}[base + 3u]
|
| 137 |
-
);
|
| 138 |
-
}
|
| 139 |
-
{%- endmacro %}
|
| 140 |
-
{%- macro emit_quant_load4(format, kind, buffer, scaleBuffer) %}
|
| 141 |
-
{{ emit_quant_scale4(kind, scaleBuffer) }}
|
| 142 |
-
fn load_{{ kind }}4(indexV4: u32, d4: u32, hk: u32) -> vec4<f32> {
|
| 143 |
-
{%- if format == "int8" %}
|
| 144 |
-
return vec4<f32>({{ buffer }}[indexV4]) * {{ kind }}scale4(d4, hk);
|
| 145 |
-
{%- else %}
|
| 146 |
-
// Two elements cover this vec4: each carries two +8-biased nibbles, low first.
|
| 147 |
-
let rowBase = indexV4 - d4;
|
| 148 |
-
let lo = {{ buffer }}[rowBase + d4 * 2u];
|
| 149 |
-
let hi = {{ buffer }}[rowBase + d4 * 2u + 1u];
|
| 150 |
-
let nibbles = vec4<i32>(
|
| 151 |
-
i32(lo & 0xFu), i32((lo >> 4u) & 0xFu),
|
| 152 |
-
i32(hi & 0xFu), i32((hi >> 4u) & 0xFu)
|
| 153 |
-
);
|
| 154 |
-
let signed = nibbles - vec4<i32>(8);
|
| 155 |
-
return vec4<f32>(signed) * {{ kind }}scale4(d4, hk);
|
| 156 |
-
{%- endif %}
|
| 157 |
}
|
| 158 |
-
{%- endmacro %}
|
| 159 |
-
|
| 160 |
-
{{ emit_quant_load4(quantCacheFormat, "key", KEY, "k_scale") }}
|
| 161 |
-
{{ emit_quant_load4(quantCacheFormat, "value", VALUE, "v_scale") }}
|
| 162 |
-
{% endif %}
|
| 163 |
|
| 164 |
{% if hasBias %}
|
|
|
|
|
|
|
| 165 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 166 |
let offset = base + d4 * 4u;
|
| 167 |
-
return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
|
| 168 |
}
|
| 169 |
|
| 170 |
{% endif %}
|
|
@@ -212,17 +162,6 @@ fn main(
|
|
| 212 |
var {{ qn("m", qi) }}: f32 = NEG_INF;
|
| 213 |
var {{ qn("l", qi) }}: f32 = 0.0;
|
| 214 |
{% endfor %}
|
| 215 |
-
{% if useSeqlens %}
|
| 216 |
-
|
| 217 |
-
// Resident rows of a windowed cache: the survivors were shifted down to [0, kvActive),
|
| 218 |
-
// so query/key DISTANCE is unchanged and every bound below reads as if the cache were
|
| 219 |
-
// exactly kvActive long. Rotary is excluded from this path (it would need the absolute
|
| 220 |
-
// position, not the cache-relative one), so pastLenForRope keeps the physical length.
|
| 221 |
-
// Clamping below by the query count matches the cache-update passes and the scalar
|
| 222 |
-
// path: a right-padded batch (seqlens_k[b]+1 < qSeq) still appends its whole chunk,
|
| 223 |
-
// so its queries score against all of it.
|
| 224 |
-
let kvActive = min({{ KVSEQ }}, max({{ QSEQ }}, u32(seqlens_k[b]) + 1u));
|
| 225 |
-
{% endif %}
|
| 226 |
// Causal ceiling per query; the key loop runs over the workgroup's union range
|
| 227 |
// (uniform trip count), masking out-of-range (query, key) pairs.
|
| 228 |
// Upper-left causal and/or sliding-window bounds. Query qIdx sits at
|
|
@@ -260,33 +199,13 @@ fn main(
|
|
| 260 |
let kj = kStart + slot;
|
| 261 |
if (kj < wgEnd) {
|
| 262 |
let base4 = kvBatch4 + kj * KV_HIDDEN_V4 + d4;
|
| 263 |
-
{% if quantCacheFormat %}
|
| 264 |
-
k_tile[i] = vec4<{{ ST }}>(load_key4(base4, d4, hKv));
|
| 265 |
-
v_tile[i] = vec4<{{ ST }}>(load_value4(base4, d4, hKv));
|
| 266 |
-
{% else %}
|
| 267 |
k_tile[i] = vec4<{{ ST }}>({{ KEY }}[base4]);
|
| 268 |
v_tile[i] = vec4<{{ ST }}>({{ VALUE }}[base4]);
|
| 269 |
-
{% endif %}
|
| 270 |
} else {
|
| 271 |
k_tile[i] = vec4<{{ ST }}>(0.0);
|
| 272 |
v_tile[i] = vec4<{{ ST }}>(0.0);
|
| 273 |
}
|
| 274 |
}
|
| 275 |
-
{% if STAGE_MASK %}
|
| 276 |
-
// The K/V load barrier also publishes this compact mask tile.
|
| 277 |
-
for (var i: u32 = tid; i < TILE_Q * TILE_K; i = i + WG) {
|
| 278 |
-
let qSlot = i / TILE_K;
|
| 279 |
-
let kSlot = i % TILE_K;
|
| 280 |
-
let maskQ = min(wg.x * TILE_Q + qSlot, {{ QSEQ }} - 1u);
|
| 281 |
-
let maskK = kStart + kSlot;
|
| 282 |
-
if (maskK < wgEnd) {
|
| 283 |
-
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + maskQ * params.maskSeqStride + maskK;
|
| 284 |
-
mask_tile[i] = {{ MASK_TILE_LOAD }};
|
| 285 |
-
} else {
|
| 286 |
-
mask_tile[i] = {{ MASK_TILE_ZERO }};
|
| 287 |
-
}
|
| 288 |
-
}
|
| 289 |
-
{% endif %}
|
| 290 |
workgroupBarrier();
|
| 291 |
// TILE_K is a small shader constant; the loop updates the named q/o slices in place.
|
| 292 |
{% if not useSubgroups and batchNoSgReduction %}
|
|
@@ -329,15 +248,7 @@ fn main(
|
|
| 329 |
if (kj >= minKj && kj < maxKj) {
|
| 330 |
{% if hasMask %}
|
| 331 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
| 332 |
-
{% if maskIsBool %}
|
| 333 |
-
if (attn_mask[maskIndex] != 0u) {
|
| 334 |
-
s[kk] = {{ score_expr("part") }};
|
| 335 |
-
} else {
|
| 336 |
-
s[kk] = MASK_NEG;
|
| 337 |
-
}
|
| 338 |
-
{% else %}
|
| 339 |
s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
|
| 340 |
-
{% endif %}
|
| 341 |
{% else %}
|
| 342 |
s[kk] = {{ score_expr("part") }};
|
| 343 |
{% endif %}
|
|
@@ -374,29 +285,11 @@ fn main(
|
|
| 374 |
let sc = partv[{{ qi }}];
|
| 375 |
if (kj >= {{ qn("minKj", qi) }} && kj < {{ qn("maxKj", qi) }}) {
|
| 376 |
{% if hasMask %}
|
| 377 |
-
{% if STAGE_MASK %}
|
| 378 |
-
let maskValue = mask_tile[(qSub + {{ qi }}u) * TILE_K + kk];
|
| 379 |
-
{% else %}
|
| 380 |
// Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
|
| 381 |
// [q, k] masks set batch/head strides to 0). {{ qn("qClamped", qi) }} keeps the seq index
|
| 382 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 383 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
|
| 384 |
-
{% endif %}
|
| 385 |
-
{% if maskIsKeyKeep %}
|
| 386 |
-
// A broadcast key mask uses 1 for a retained key and 0 for padding.
|
| 387 |
-
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
|
| 388 |
-
{% elif maskIsBool %}
|
| 389 |
-
// Logical bool: a rejected key contributes no softmax mass. Leaving
|
| 390 |
-
// the initialized NEG_INF sentinel in place makes a fully masked row
|
| 391 |
-
// land on the zero-denominator output guard below.
|
| 392 |
-
if ({{ MASK_ELEMENT }} != 0u) {
|
| 393 |
-
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
|
| 394 |
-
} else {
|
| 395 |
-
{{ qn("s", qi) }}[kk] = MASK_NEG;
|
| 396 |
-
}
|
| 397 |
-
{% else %}
|
| 398 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
|
| 399 |
-
{% endif %}
|
| 400 |
{% else %}
|
| 401 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
|
| 402 |
{% endif %}
|
|
@@ -434,29 +327,11 @@ fn main(
|
|
| 434 |
{% endif %}
|
| 435 |
if (kj >= minKj && kj < maxKj) {
|
| 436 |
{% if hasMask %}
|
| 437 |
-
{% if STAGE_MASK %}
|
| 438 |
-
let maskValue = mask_tile[qSub * TILE_K + kk];
|
| 439 |
-
{% else %}
|
| 440 |
// Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
|
| 441 |
// [q, k] masks set batch/head strides to 0). qClamped keeps the seq index
|
| 442 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 443 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
| 444 |
-
{% endif %}
|
| 445 |
-
{% if maskIsKeyKeep %}
|
| 446 |
-
// A broadcast key mask uses 1 for a retained key and 0 for padding.
|
| 447 |
-
s[kk] = {{ score_expr("part") }} + (1.0 - f32({{ MASK_ELEMENT }})) * MASK_NEG;
|
| 448 |
-
{% elif maskIsBool %}
|
| 449 |
-
// Logical bool: a rejected key contributes no softmax mass. Leaving
|
| 450 |
-
// the initialized NEG_INF sentinel in place makes a fully masked row
|
| 451 |
-
// land on the zero-denominator output guard below.
|
| 452 |
-
if ({{ MASK_ELEMENT }} != 0u) {
|
| 453 |
-
s[kk] = {{ score_expr("part") }};
|
| 454 |
-
} else {
|
| 455 |
-
s[kk] = MASK_NEG;
|
| 456 |
-
}
|
| 457 |
-
{% else %}
|
| 458 |
s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
|
| 459 |
-
{% endif %}
|
| 460 |
{% else %}
|
| 461 |
s[kk] = {{ score_expr("part") }};
|
| 462 |
{% endif %}
|
|
@@ -464,24 +339,7 @@ fn main(
|
|
| 464 |
}
|
| 465 |
{% endif %}
|
| 466 |
|
| 467 |
-
|
| 468 |
-
// exponentiated probabilities for the PV accumulation below.
|
| 469 |
-
{% for qi in range(QPL) %}
|
| 470 |
-
var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
|
| 471 |
-
for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
|
| 472 |
-
{{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
|
| 473 |
-
}
|
| 474 |
-
let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
|
| 475 |
-
let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
|
| 476 |
-
var {{ qn("pSum", qi) }}: f32 = 0.0;
|
| 477 |
-
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 478 |
-
let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
|
| 479 |
-
{{ qn("s", qi) }}[kk] = pk;
|
| 480 |
-
{{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
|
| 481 |
-
}
|
| 482 |
-
{{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
|
| 483 |
-
{{ qn("m", qi) }} = {{ qn("newMax", qi) }};
|
| 484 |
-
{% endfor %}
|
| 485 |
// A boundary tile can address V rows outside a query's attended range, and
|
| 486 |
// a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
|
| 487 |
// NaN, the guarded loop selects the V operand away for range-excluded keys.
|
|
@@ -540,22 +398,11 @@ fn main(
|
|
| 540 |
{% for qi in range(QPL) %}
|
| 541 |
if ({{ qn("qValid", qi) }}) {
|
| 542 |
let outBase4 = (b * {{ QSEQ }} + {{ qn("qIdx", qi) }}) * Q_HIDDEN_V4 + h * HEAD_DIM_V4 + lane8 * SLICE;
|
| 543 |
-
{% if hasHeadSink %}
|
| 544 |
-
// The head sink is a learned logit that competes with the keys but carries
|
| 545 |
-
// no value, so it enters the denominator only and the weighted sum above is
|
| 546 |
-
// untouched. Renormalizing against max({{ qn("m", qi) }}, sink) keeps the exponentials in
|
| 547 |
-
// range when the sink dominates a fully-masked row.
|
| 548 |
-
let sink = f32(head_sink[h]);
|
| 549 |
-
let finalM = max({{ qn("m", qi) }}, sink);
|
| 550 |
-
let accScale = exp({{ qn("m", qi) }} - finalM);
|
| 551 |
-
let inv = accScale / (exp(sink - finalM) + {{ qn("l", qi) }} * accScale);
|
| 552 |
-
{% else %}
|
| 553 |
// {{ qn("l", qi) }} == 0 means this query had no probability-bearing key: either its
|
| 554 |
// causal/window range is empty or its logical bool mask rejects every key.
|
| 555 |
// Emit 0 rather than 0/0. Other attention paths enforce the same empty-row
|
| 556 |
// contract by selecting on positive global mass.
|
| 557 |
let inv = select(0.0, 1.0 / {{ qn("l", qi) }}, {{ qn("l", qi) }} > 0.0);
|
| 558 |
-
{% endif %}
|
| 559 |
{% for c in range(SLICE_COUNT) %}
|
| 560 |
{{ OUTPUT }}[outBase4 + {{ c }}u] = vec4<{{ scalar }}>({{ attention_value(c, qi) }});
|
| 561 |
{% endfor %}
|
|
|
|
| 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 |
{% if hasMask is not defined %}{% set hasMask = false %}{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
{% set Q_HIDDEN = qHidden | default(0) %}
|
|
|
|
|
|
|
| 13 |
// A windowed cache binds a fixed CAPACITY but keeps only the most recent
|
| 14 |
// min(total, capacity) rows resident. params.kvSeq then names the physical row
|
| 15 |
// count, which is still the right batch stride but the wrong attention bound, so
|
| 16 |
// the bounds read kvActive instead. Modes without seqlens use `params.kvSeq`
|
| 17 |
// in both roles.
|
| 18 |
+
{% if hasKeyLimit is not defined %}{% set hasKeyLimit = false %}{% endif %}
|
| 19 |
+
{% if hasKeyLimit %}
|
| 20 |
+
{% set KVA = "select(params.kvSeq, min(params.keyLimit, params.kvSeq), params.keyLimit > 0u)" %}
|
| 21 |
+
{% else %}
|
| 22 |
+
{% set KVA = KVSEQ %}
|
| 23 |
+
{% endif %}
|
| 24 |
+
{% macro score_expr(part) %}{{ part }} * SCALE{% endmacro %}
|
| 25 |
{% if useSubgroups %}
|
| 26 |
enable subgroups;
|
| 27 |
{% endif %}
|
|
|
|
| 31 |
// workgroup storage; values widen to f32 when read. Other specializations stage
|
| 32 |
// them as f32. Score and weighted-value accumulation remain in f32 throughout.
|
| 33 |
{% set STAGE_MASK = hasMask and useSubgroups and stageMask is defined and stageMask %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 34 |
{% set MASK_ADDITIVE = "maskValue" if STAGE_MASK else "f32(attn_mask[maskIndex])" %}
|
| 35 |
{% set SLICE_COUNT = ((headDimV4 / LPQ) | int) %}
|
| 36 |
{% 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 %}
|
| 37 |
{% macro qn(name, qi) %}{{ name }}{% if qi > 0 %}_q{{ qi }}{% endif %}{% endmacro %}
|
| 38 |
+
{% macro emit_tile_softmax() %}
|
| 39 |
+
// Per-thread online softmax over the tile. s[kk] is reused to hold the
|
| 40 |
+
// exponentiated probabilities for the PV accumulation below.
|
| 41 |
+
{% for qi in range(QPL) %}
|
| 42 |
+
var {{ qn("tileMax", qi) }}: f32 = {{ qn("s", qi) }}[0];
|
| 43 |
+
for (var kk: u32 = 1u; kk < TILE_K; kk = kk + 1u) {
|
| 44 |
+
{{ qn("tileMax", qi) }} = max({{ qn("tileMax", qi) }}, {{ qn("s", qi) }}[kk]);
|
| 45 |
+
}
|
| 46 |
+
let {{ qn("newMax", qi) }} = max({{ qn("m", qi) }}, {{ qn("tileMax", qi) }});
|
| 47 |
+
let {{ qn("corr", qi) }} = select(exp({{ qn("m", qi) }} - {{ qn("newMax", qi) }}), 0.0, {{ qn("m", qi) }} == NEG_INF);
|
| 48 |
+
var {{ qn("pSum", qi) }}: f32 = 0.0;
|
| 49 |
+
for (var kk: u32 = 0u; kk < TILE_K; kk = kk + 1u) {
|
| 50 |
+
let pk = select(0.0, exp({{ qn("s", qi) }}[kk] - {{ qn("newMax", qi) }}), {{ qn("s", qi) }}[kk] != NEG_INF);
|
| 51 |
+
{{ qn("s", qi) }}[kk] = pk;
|
| 52 |
+
{{ qn("pSum", qi) }} = {{ qn("pSum", qi) }} + pk;
|
| 53 |
+
}
|
| 54 |
+
{{ qn("l", qi) }} = {{ qn("l", qi) }} * {{ qn("corr", qi) }} + {{ qn("pSum", qi) }};
|
| 55 |
+
{{ qn("m", qi) }} = {{ qn("newMax", qi) }};
|
| 56 |
+
{% endfor %}
|
| 57 |
+
{% endmacro %}
|
| 58 |
|
| 59 |
// Tiled flash prefill attention with configurable-width query clusters for
|
| 60 |
// token-major [batch, seq, heads*headDim] attention ops. Each workgroup covers
|
|
|
|
| 74 |
const HEAD_DIM: u32 = {{ headDim }}u;
|
| 75 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
| 76 |
const Q_HIDDEN_V4: u32 = {{ qHiddenV4 }}u; // bsh token-major hidden stride
|
| 77 |
+
const KV_HIDDEN_V4: u32 = {{ kvHiddenV4 }}u;
|
|
|
|
| 78 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 79 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 80 |
const TILE_Q: u32 = {{ TILE_Q }}u;
|
|
|
|
| 88 |
{% else %}
|
| 89 |
const WG: u32 = TILE_Q * LPQ;
|
| 90 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 91 |
const NEG_INF: f32 = -3.4028234663852886e38;
|
|
|
|
|
|
|
|
|
|
|
|
|
| 92 |
|
| 93 |
var<workgroup> k_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
| 94 |
var<workgroup> v_tile: array<vec4<{{ ST }}>, TILE_K * (HEAD_DIM / 4u)>;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
{% if not useSubgroups %}
|
| 96 |
{% if batchNoSgReduction %}
|
| 97 |
// No-subgroups cluster reduction scratch for a whole K tile. Staging every
|
|
|
|
| 104 |
{% endif %}
|
| 105 |
{% endif %}
|
| 106 |
|
| 107 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
|
|
|
|
|
|
|
|
|
| 108 |
if (params.scale != 0.0) { return params.scale; }
|
| 109 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 110 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 111 |
|
| 112 |
{% if hasBias %}
|
| 113 |
+
{% set BW = "" %}
|
| 114 |
+
{% set BC = "" %}
|
| 115 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 116 |
let offset = base + d4 * 4u;
|
| 117 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 118 |
}
|
| 119 |
|
| 120 |
{% endif %}
|
|
|
|
| 162 |
var {{ qn("m", qi) }}: f32 = NEG_INF;
|
| 163 |
var {{ qn("l", qi) }}: f32 = 0.0;
|
| 164 |
{% endfor %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 165 |
// Causal ceiling per query; the key loop runs over the workgroup's union range
|
| 166 |
// (uniform trip count), masking out-of-range (query, key) pairs.
|
| 167 |
// Upper-left causal and/or sliding-window bounds. Query qIdx sits at
|
|
|
|
| 199 |
let kj = kStart + slot;
|
| 200 |
if (kj < wgEnd) {
|
| 201 |
let base4 = kvBatch4 + kj * KV_HIDDEN_V4 + d4;
|
|
|
|
|
|
|
|
|
|
|
|
|
| 202 |
k_tile[i] = vec4<{{ ST }}>({{ KEY }}[base4]);
|
| 203 |
v_tile[i] = vec4<{{ ST }}>({{ VALUE }}[base4]);
|
|
|
|
| 204 |
} else {
|
| 205 |
k_tile[i] = vec4<{{ ST }}>(0.0);
|
| 206 |
v_tile[i] = vec4<{{ ST }}>(0.0);
|
| 207 |
}
|
| 208 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 209 |
workgroupBarrier();
|
| 210 |
// TILE_K is a small shader constant; the loop updates the named q/o slices in place.
|
| 211 |
{% if not useSubgroups and batchNoSgReduction %}
|
|
|
|
| 248 |
if (kj >= minKj && kj < maxKj) {
|
| 249 |
{% if hasMask %}
|
| 250 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 251 |
s[kk] = {{ score_expr("part") }} + f32(attn_mask[maskIndex]);
|
|
|
|
| 252 |
{% else %}
|
| 253 |
s[kk] = {{ score_expr("part") }};
|
| 254 |
{% endif %}
|
|
|
|
| 285 |
let sc = partv[{{ qi }}];
|
| 286 |
if (kj >= {{ qn("minKj", qi) }} && kj < {{ qn("maxKj", qi) }}) {
|
| 287 |
{% if hasMask %}
|
|
|
|
|
|
|
|
|
|
| 288 |
// Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
|
| 289 |
// [q, k] masks set batch/head strides to 0). {{ qn("qClamped", qi) }} keeps the seq index
|
| 290 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 291 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + {{ qn("qClamped", qi) }} * params.maskSeqStride + kj;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 292 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }} + {{ MASK_ADDITIVE }};
|
|
|
|
| 293 |
{% else %}
|
| 294 |
{{ qn("s", qi) }}[kk] = {{ score_expr("sc") }};
|
| 295 |
{% endif %}
|
|
|
|
| 327 |
{% endif %}
|
| 328 |
if (kj >= minKj && kj < maxKj) {
|
| 329 |
{% if hasMask %}
|
|
|
|
|
|
|
|
|
|
| 330 |
// Broadcast-strided mask address (a 0 stride collapses that axis; rank-2
|
| 331 |
// [q, k] masks set batch/head strides to 0). qClamped keeps the seq index
|
| 332 |
// in-bounds for padding queries in the last tile (their output is dropped).
|
| 333 |
let maskIndex = b * params.maskBatchStride + h * params.maskHeadStride + qClamped * params.maskSeqStride + kj;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 334 |
s[kk] = {{ score_expr("part") }} + {{ MASK_ADDITIVE }};
|
|
|
|
| 335 |
{% else %}
|
| 336 |
s[kk] = {{ score_expr("part") }};
|
| 337 |
{% endif %}
|
|
|
|
| 339 |
}
|
| 340 |
{% endif %}
|
| 341 |
|
| 342 |
+
{{ emit_tile_softmax() }}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 343 |
// A boundary tile can address V rows outside a query's attended range, and
|
| 344 |
// a dynamic-row producer may leave those rows unwritten. Because 0 * NaN is
|
| 345 |
// NaN, the guarded loop selects the V operand away for range-excluded keys.
|
|
|
|
| 398 |
{% for qi in range(QPL) %}
|
| 399 |
if ({{ qn("qValid", qi) }}) {
|
| 400 |
let outBase4 = (b * {{ QSEQ }} + {{ qn("qIdx", qi) }}) * Q_HIDDEN_V4 + h * HEAD_DIM_V4 + lane8 * SLICE;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 401 |
// {{ qn("l", qi) }} == 0 means this query had no probability-bearing key: either its
|
| 402 |
// causal/window range is empty or its logical bool mask rejects every key.
|
| 403 |
// Emit 0 rather than 0/0. Other attention paths enforce the same empty-row
|
| 404 |
// contract by selecting on positive global mass.
|
| 405 |
let inv = select(0.0, 1.0 / {{ qn("l", qi) }}, {{ qn("l", qi) }} > 0.0);
|
|
|
|
| 406 |
{% for c in range(SLICE_COUNT) %}
|
| 407 |
{{ OUTPUT }}[outBase4 + {{ c }}u] = vec4<{{ scalar }}>({{ attention_value(c, qi) }});
|
| 408 |
{% endfor %}
|
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,11 +11,16 @@
|
|
| 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 |
{% set CAUSAL = false if (hasCausal is defined and not hasCausal) else true %}
|
|
|
|
| 15 |
{% if USE_SUBGROUPS %}
|
| 16 |
enable subgroups;
|
| 17 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
| 18 |
{{ env.wgsl.resourceDeclarations }}
|
| 19 |
|
| 20 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
|
@@ -46,15 +52,16 @@ fn scale_value() -> f32 {
|
|
| 46 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 47 |
}
|
| 48 |
|
| 49 |
-
|
| 50 |
{% if hasBias %}
|
|
|
|
|
|
|
| 51 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 52 |
let offset = base + d4 * 4u;
|
| 53 |
-
return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
|
| 54 |
}
|
| 55 |
{% endif %}
|
| 56 |
|
| 57 |
-
@compute @workgroup_size({{ Q_STEP }}, 1, 1)
|
| 58 |
fn main(
|
| 59 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 60 |
@builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
|
|
@@ -99,9 +106,14 @@ fn main(
|
|
| 99 |
// so every lane shares a uniform trip count (subgroup ops stay reconverged);
|
| 100 |
// each lane masks its own keys past myMaxKj to NEG_INF.
|
| 101 |
{% if CAUSAL %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |
let lastQ = min(wg.x * Q_STEP + Q_STEP - 1u, params.qSeq - 1u);
|
| 103 |
-
let kvEnd =
|
| 104 |
-
let myMaxKj =
|
| 105 |
{% else %}
|
| 106 |
let kvEnd = params.kvSeq;
|
| 107 |
let myMaxKj = params.kvSeq;
|
|
@@ -193,6 +205,13 @@ fn main(
|
|
| 193 |
{% endfor %}
|
| 194 |
previous_max = new_max;
|
| 195 |
previous_denom = denom;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 196 |
|
| 197 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 198 |
{% if USE_SUBGROUPS %}
|
|
@@ -209,21 +228,37 @@ fn main(
|
|
| 209 |
}
|
| 210 |
{% endif %}
|
| 211 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 212 |
var acc: vec4<f32> = vec4<f32>(0.0);
|
|
|
|
| 213 |
{% for g in range(qkGroups) %}
|
| 214 |
{% for lane in range(4) %}
|
| 215 |
-
{% if USE_SUBGROUPS %}
|
| 216 |
-
{%
|
| 217 |
-
|
|
|
|
| 218 |
{% else %}
|
| 219 |
-
|
| 220 |
{% endif %}
|
|
|
|
|
|
|
| 221 |
{% else %}
|
| 222 |
-
acc = acc + vec4<f32>(
|
| 223 |
{% endif %}
|
| 224 |
{% endfor %}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 225 |
{% endfor %}
|
|
|
|
|
|
|
|
|
|
| 226 |
o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
|
|
|
|
| 227 |
}
|
| 228 |
{% if not USE_SUBGROUPS %}
|
| 229 |
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 |
+
{% set qHidden = qHidden | default(0) %}
|
| 18 |
{% if USE_SUBGROUPS %}
|
| 19 |
enable subgroups;
|
| 20 |
{% endif %}
|
| 21 |
+
{% if PIN_SUBGROUP_32 %}
|
| 22 |
+
enable subgroup_size_control;
|
| 23 |
+
{% endif %}
|
| 24 |
{{ env.wgsl.resourceDeclarations }}
|
| 25 |
|
| 26 |
const HEAD_DIM_V4: u32 = {{ headDimV4 }}u;
|
|
|
|
| 52 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 53 |
}
|
| 54 |
|
|
|
|
| 55 |
{% if hasBias %}
|
| 56 |
+
{% set BW = "" %}
|
| 57 |
+
{% set BC = "" %}
|
| 58 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 59 |
let offset = base + d4 * 4u;
|
| 60 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 61 |
}
|
| 62 |
{% endif %}
|
| 63 |
|
| 64 |
+
@compute @workgroup_size({{ Q_STEP }}, 1, 1){{ " @subgroup_size(32)" if PIN_SUBGROUP_32 else "" }}
|
| 65 |
fn main(
|
| 66 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 67 |
@builtin(local_invocation_id) lid: vec3<u32>{% if USE_SUBGROUPS %},
|
|
|
|
| 106 |
// so every lane shares a uniform trip count (subgroup ops stay reconverged);
|
| 107 |
// each lane masks its own keys past myMaxKj to NEG_INF.
|
| 108 |
{% if CAUSAL %}
|
| 109 |
+
{% set FIXED_CAUSAL = false %}
|
| 110 |
+
{% macro key_ceiling(query) %}
|
| 111 |
+
{% if FIXED_CAUSAL %}
|
| 112 |
+
min({{ query }} + 1u, params.kvSeq){% else %}
|
| 113 |
+
select(params.kvSeq, min({{ query }} + 1u, params.kvSeq), params.isCausal != 0u){% endif %}{% endmacro %}
|
| 114 |
let lastQ = min(wg.x * Q_STEP + Q_STEP - 1u, params.qSeq - 1u);
|
| 115 |
+
let kvEnd = {{ key_ceiling("lastQ") }};
|
| 116 |
+
let myMaxKj = {{ key_ceiling("qi") }};
|
| 117 |
{% else %}
|
| 118 |
let kvEnd = params.kvSeq;
|
| 119 |
let myMaxKj = params.kvSeq;
|
|
|
|
| 205 |
{% endfor %}
|
| 206 |
previous_max = new_max;
|
| 207 |
previous_denom = denom;
|
| 208 |
+
{% set PV_HALF = ST == "f16" and not (precisePv | default(false)) %}
|
| 209 |
+
{% set PV_FLUSH_GROUPS = 4 %}
|
| 210 |
+
{% if PV_HALF %}
|
| 211 |
+
{% for g in range(qkGroups) %}
|
| 212 |
+
let p{{ g }} = vec4<f16>(qk{{ g }});
|
| 213 |
+
{% endfor %}
|
| 214 |
+
{% endif %}
|
| 215 |
|
| 216 |
for (var d4: u32 = 0u; d4 < HEAD_DIM_V4; d4 = d4 + 1u) {
|
| 217 |
{% if USE_SUBGROUPS %}
|
|
|
|
| 228 |
}
|
| 229 |
{% endif %}
|
| 230 |
{% endif %}
|
| 231 |
+
{% if PV_HALF %}
|
| 232 |
+
var acc_f32: vec4<f32> = vec4<f32>(0.0);
|
| 233 |
+
var acc: vec4<f16> = vec4<f16>(0.0);
|
| 234 |
+
{% else %}
|
| 235 |
var acc: vec4<f32> = vec4<f32>(0.0);
|
| 236 |
+
{% endif %}
|
| 237 |
{% for g in range(qkGroups) %}
|
| 238 |
{% for lane in range(4) %}
|
| 239 |
+
{% if not USE_SUBGROUPS %}
|
| 240 |
+
{% set pvValue = "valueTile[d4 * K_STEP + " ~ (g * 4 + lane) ~ "u]" %}
|
| 241 |
+
{% elif g < 8 %}
|
| 242 |
+
{% set pvValue = "subgroupShuffle(v_local0, " ~ (g * 4 + lane) ~ "u)" %}
|
| 243 |
{% else %}
|
| 244 |
+
{% set pvValue = "subgroupShuffle(v_local1, " ~ ((g - 8) * 4 + lane) ~ "u)" %}
|
| 245 |
{% endif %}
|
| 246 |
+
{% if PV_HALF %}
|
| 247 |
+
acc = fma({{ pvValue }}, vec4<f16>(p{{ g }}.{{ components[lane] }}), acc);
|
| 248 |
{% else %}
|
| 249 |
+
acc = acc + vec4<f32>({{ pvValue }}) * qk{{ g }}.{{ components[lane] }};
|
| 250 |
{% endif %}
|
| 251 |
{% endfor %}
|
| 252 |
+
{% if PV_HALF and (g + 1) % PV_FLUSH_GROUPS == 0 %}
|
| 253 |
+
acc_f32 = acc_f32 + vec4<f32>(acc);
|
| 254 |
+
acc = vec4<f16>(0.0);
|
| 255 |
+
{% endif %}
|
| 256 |
{% endfor %}
|
| 257 |
+
{% if PV_HALF %}
|
| 258 |
+
o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc_f32;
|
| 259 |
+
{% else %}
|
| 260 |
o_tile[d4] = o_tile[d4] * vec4<f32>(o_ratio) + acc;
|
| 261 |
+
{% endif %}
|
| 262 |
}
|
| 263 |
{% if not USE_SUBGROUPS %}
|
| 264 |
workgroupBarrier();
|
build/webgpu/attn-materialized-apply-f32.wgsl.jinja
CHANGED
|
@@ -1,3 +1,4 @@
|
|
|
|
|
| 1 |
{{ env.wgsl.resourceDeclarations }}
|
| 2 |
|
| 3 |
// Register-blocked normalized-scores @ V GEMM with BHSD and BSH specializations.
|
|
@@ -39,7 +40,7 @@ var<workgroup> rowMax: array<f32, BM>;
|
|
| 39 |
|
| 40 |
{% endif %}
|
| 41 |
fn v_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
|
| 42 |
-
return value[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d];
|
| 43 |
}
|
| 44 |
|
| 45 |
fn store_y(b: u32, h: u32, q: u32, d: u32, result: f32) {
|
|
@@ -48,7 +49,7 @@ fn store_y(b: u32, h: u32, q: u32, d: u32, result: f32) {
|
|
| 48 |
{% else %}
|
| 49 |
let biased = result;
|
| 50 |
{% endif %}
|
| 51 |
-
output[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d] = biased;
|
| 52 |
}
|
| 53 |
|
| 54 |
@compute @workgroup_size(WG_DIM, WG_DIM, 1)
|
|
|
|
| 1 |
+
{% set STORAGE_F16 = false %}
|
| 2 |
{{ env.wgsl.resourceDeclarations }}
|
| 3 |
|
| 4 |
// Register-blocked normalized-scores @ V GEMM with BHSD and BSH specializations.
|
|
|
|
| 40 |
|
| 41 |
{% endif %}
|
| 42 |
fn v_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
|
| 43 |
+
return {{ "f32(" if STORAGE_F16 else "" }}value[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d]{{ ")" if STORAGE_F16 else "" }};
|
| 44 |
}
|
| 45 |
|
| 46 |
fn store_y(b: u32, h: u32, q: u32, d: u32, result: f32) {
|
|
|
|
| 49 |
{% else %}
|
| 50 |
let biased = result;
|
| 51 |
{% endif %}
|
| 52 |
+
output[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d] = {{ "f16(" if STORAGE_F16 else "" }}biased{{ ")" if STORAGE_F16 else "" }};
|
| 53 |
}
|
| 54 |
|
| 55 |
@compute @workgroup_size(WG_DIM, WG_DIM, 1)
|
build/webgpu/attn-materialized-rowstats-combine-f32.wgsl.jinja
CHANGED
|
@@ -13,6 +13,12 @@
|
|
| 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 |
{% set MAX_ONLY = maxOnly is defined and maxOnly %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
{{ env.wgsl.resourceDeclarations }}
|
| 17 |
|
| 18 |
const SLOTS: u32 = {{ statSlots }}u;
|
|
@@ -41,21 +47,19 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 41 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 42 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 43 |
}
|
| 44 |
-
{%
|
| 45 |
{% if stableUsage == "all" %}
|
| 46 |
|
| 47 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 48 |
return exp(shifted_value(value, maxValue));
|
| 49 |
}
|
| 50 |
-
{%
|
| 51 |
-
|
| 52 |
|
| 53 |
@compute @workgroup_size(WG, 1, 1)
|
| 54 |
fn main(
|
| 55 |
@builtin(global_invocation_id) gid: vec3<u32>
|
| 56 |
) {
|
| 57 |
-
|
| 58 |
-
if (row >= params.rows) { return; }
|
| 59 |
|
| 60 |
// `row` already runs over (batch, head, query) together, and the partial
|
| 61 |
// layout puts that same product one axis out from the slot, so the stride
|
|
|
|
| 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 |
{% set MAX_ONLY = maxOnly is defined and maxOnly %}
|
| 16 |
+
{% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
|
| 17 |
+
{% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
|
| 18 |
+
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 19 |
+
// per-axis workgroup fold width.
|
| 20 |
+
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
|
| 21 |
+
if ({{ name }} >= {{ bound }}) { return; }{% endmacro %}
|
| 22 |
{{ env.wgsl.resourceDeclarations }}
|
| 23 |
|
| 24 |
const SLOTS: u32 = {{ statSlots }}u;
|
|
|
|
| 47 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 48 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 49 |
}
|
| 50 |
+
{% endif %}
|
| 51 |
{% if stableUsage == "all" %}
|
| 52 |
|
| 53 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 54 |
return exp(shifted_value(value, maxValue));
|
| 55 |
}
|
| 56 |
+
{% endif %}
|
|
|
|
| 57 |
|
| 58 |
@compute @workgroup_size(WG, 1, 1)
|
| 59 |
fn main(
|
| 60 |
@builtin(global_invocation_id) gid: vec3<u32>
|
| 61 |
) {
|
| 62 |
+
{{ flat_index_2d("WG", "row", "params.rows", guardInline=true) }}
|
|
|
|
| 63 |
|
| 64 |
// `row` already runs over (batch, head, query) together, and the partial
|
| 65 |
// layout puts that same product one axis out from the slot, so the stride
|
build/webgpu/attn-materialized-score-f32.wgsl.jinja
CHANGED
|
@@ -1,3 +1,4 @@
|
|
|
|
|
| 1 |
{{ env.wgsl.resourceDeclarations }}
|
| 2 |
|
| 3 |
// f32 prefill score GEMM with BHSD and packed BSH specializations. It computes
|
|
@@ -8,22 +9,14 @@
|
|
| 8 |
// for TM * TN * 4 fused multiply-adds. K keeps a one-word shared stride
|
| 9 |
// padding, which changes the transposed workgroup-memory stride and reduces
|
| 10 |
// bank-conflict risk on banked implementations.
|
| 11 |
-
{% if scalarAccumulators is not defined %}{% set scalarAccumulators = false %}{% endif %}
|
| 12 |
-
{% if maskIsKeyKeep is not defined %}{% set maskIsKeyKeep = false %}{% endif %}
|
| 13 |
-
{% if scalarAccumulators %}
|
| 14 |
-
{% set TM_JINJA = (materializedQueryTile / materializedWorkgroupDim)|int %}
|
| 15 |
-
{% set TN_JINJA = (materializedKeyTile / materializedWorkgroupDim)|int %}
|
| 16 |
-
{% endif %}
|
| 17 |
{% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
|
|
|
|
| 18 |
|
| 19 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 20 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 21 |
const HEAD_DIM: u32 = {{ headDim }}u;
|
| 22 |
const Q_HIDDEN: u32 = Q_HEADS * HEAD_DIM;
|
| 23 |
const KV_HIDDEN: u32 = KV_HEADS * HEAD_DIM;
|
| 24 |
-
{% if maskIsKeyKeep %}
|
| 25 |
-
const NEG_INF: f32 = -3.4028234663852886e38;
|
| 26 |
-
{% endif %}
|
| 27 |
|
| 28 |
const BK: u32 = {{ materializedInnerTile }}u;
|
| 29 |
const BM: u32 = {{ materializedQueryTile }}u;
|
|
@@ -45,46 +38,20 @@ const STAT_SLOTS: u32 = {{ statSlots }}u;
|
|
| 45 |
var<workgroup> statScratch: array<f32, WG_THREADS>;
|
| 46 |
{% endif %}
|
| 47 |
|
| 48 |
-
{% if EMIT_ROW_STATS %}
|
| 49 |
-
{% set stableUsage = stableHelperUsage if stableHelperUsage is defined else "all" %}
|
| 50 |
-
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 51 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
| 52 |
// accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
|
| 53 |
// a zero final denominator according to their public semantics. Using -inf
|
| 54 |
// here changes +inf-row behavior.
|
| 55 |
const FLT_MAX: f32 = 3.4028234663852886e38;
|
| 56 |
-
{% if stableUsage != "constant" %}
|
| 57 |
-
|
| 58 |
-
fn is_finite_f32(value: f32) -> bool {
|
| 59 |
-
return select(false, value <= FLT_MAX, value >= -FLT_MAX);
|
| 60 |
-
}
|
| 61 |
-
|
| 62 |
-
// x - m that is exactly 0 when x equals a finite m, so exp(shifted) == 1
|
| 63 |
-
// exactly at the row max. `x - x` on an infinite max is a legal fast-math
|
| 64 |
-
// fold to 0, which would silently turn +inf rows finite — the explicit
|
| 65 |
-
// equality test keeps the NaN propagation of the serial kernels.
|
| 66 |
-
fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
| 67 |
-
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 68 |
-
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 69 |
-
}
|
| 70 |
-
{%- endif %}
|
| 71 |
-
{% if stableUsage == "all" %}
|
| 72 |
-
|
| 73 |
-
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 74 |
-
return exp(shifted_value(value, maxValue));
|
| 75 |
-
}
|
| 76 |
-
{%- endif %}
|
| 77 |
-
|
| 78 |
{% endif %}
|
| 79 |
-
{%
|
| 80 |
-
fn scale_value() -> f32 {
|
| 81 |
if (params.scale != 0.0) { return params.scale; }
|
| 82 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 83 |
}
|
| 84 |
|
| 85 |
-
|
| 86 |
fn q_at(b: u32, h: u32, q: u32, d: u32) -> f32 {
|
| 87 |
-
var value = query[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d];
|
| 88 |
{% if hasBias %}
|
| 89 |
value = value + bias[h * HEAD_DIM + d];
|
| 90 |
{% endif %}
|
|
@@ -92,7 +59,7 @@ fn q_at(b: u32, h: u32, q: u32, d: u32) -> f32 {
|
|
| 92 |
}
|
| 93 |
|
| 94 |
fn k_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
|
| 95 |
-
return key[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d];
|
| 96 |
}
|
| 97 |
|
| 98 |
@compute @workgroup_size(WG_DIM, WG_DIM, 1)
|
|
@@ -107,17 +74,8 @@ fn main(
|
|
| 107 |
let kBase = wg.x * BN;
|
| 108 |
let li = lid.y * WG_DIM + lid.x;
|
| 109 |
|
| 110 |
-
{% if scalarAccumulators %}
|
| 111 |
-
// Keep every register-tile cell statically addressable.
|
| 112 |
-
{% for i in range(TM_JINJA) %}
|
| 113 |
-
{% for j in range(TN_JINJA) %}
|
| 114 |
-
var acc{{ i }}_{{ j }}: f32 = 0.0;
|
| 115 |
-
{% endfor %}
|
| 116 |
-
{% endfor %}
|
| 117 |
-
{% else %}
|
| 118 |
var acc: array<f32, TM * TN>;
|
| 119 |
for (var i = 0u; i < TM * TN; i = i + 1u) { acc[i] = 0.0; }
|
| 120 |
-
{% endif %}
|
| 121 |
|
| 122 |
for (var dBase = 0u; dBase < HEAD_DIM; dBase = dBase + BK) {
|
| 123 |
// Both source tiles are loaded in their physical row-major direction.
|
|
@@ -152,19 +110,6 @@ fn main(
|
|
| 152 |
let qr0 = lid.y * TM;
|
| 153 |
let kc0 = lid.x * TN;
|
| 154 |
for (var dv = 0u; dv < D_VECS; dv = dv + 1u) {
|
| 155 |
-
{% if scalarAccumulators %}
|
| 156 |
-
{% for i in range(TM_JINJA) %}
|
| 157 |
-
let qv{{ i }} = tileQ[qr0 + {{ i }}u][dv];
|
| 158 |
-
{% endfor %}
|
| 159 |
-
{% for j in range(TN_JINJA) %}
|
| 160 |
-
let kv{{ j }} = tileK[kc0 + {{ j }}u][dv];
|
| 161 |
-
{% endfor %}
|
| 162 |
-
{% for i in range(TM_JINJA) %}
|
| 163 |
-
{% for j in range(TN_JINJA) %}
|
| 164 |
-
acc{{ i }}_{{ j }} = acc{{ i }}_{{ j }} + dot(qv{{ i }}, kv{{ j }});
|
| 165 |
-
{% endfor %}
|
| 166 |
-
{% endfor %}
|
| 167 |
-
{% else %}
|
| 168 |
var qv: array<vec4<f32>, TM>;
|
| 169 |
var kv: array<vec4<f32>, TN>;
|
| 170 |
for (var i = 0u; i < TM; i = i + 1u) { qv[i] = tileQ[qr0 + i][dv]; }
|
|
@@ -174,7 +119,6 @@ fn main(
|
|
| 174 |
acc[i * TN + j] = acc[i * TN + j] + dot(qv[i], kv[j]);
|
| 175 |
}
|
| 176 |
}
|
| 177 |
-
{% endif %}
|
| 178 |
}
|
| 179 |
workgroupBarrier();
|
| 180 |
}
|
|
@@ -189,28 +133,6 @@ fn main(
|
|
| 189 |
var rowM: array<f32, TM>;
|
| 190 |
for (var i = 0u; i < TM; i = i + 1u) { rowM[i] = -FLT_MAX; }
|
| 191 |
{% endif %}
|
| 192 |
-
{% if scalarAccumulators %}
|
| 193 |
-
{% for i in range(TM_JINJA) %}
|
| 194 |
-
{
|
| 195 |
-
let qi = qi0 + {{ i }}u;
|
| 196 |
-
if (qi < params.qSeq) {
|
| 197 |
-
{% for j in range(TN_JINJA) %}
|
| 198 |
-
let ki{{ j }} = ki0 + {{ j }}u;
|
| 199 |
-
if (ki{{ j }} < params.kvSeq) {
|
| 200 |
-
var score = acc{{ i }}_{{ j }} * scale;
|
| 201 |
-
{% if maskIsKeyKeep %}
|
| 202 |
-
score = score + (1.0 - f32(attn_mask[b * params.kvSeq + ki{{ j }}])) * NEG_INF;
|
| 203 |
-
{% endif %}
|
| 204 |
-
scores[scoreBase + qi * params.kvSeq + ki{{ j }}] = score;
|
| 205 |
-
{% if EMIT_ROW_STATS %}
|
| 206 |
-
rowM[{{ i }}u] = max(rowM[{{ i }}u], score);
|
| 207 |
-
{% endif %}
|
| 208 |
-
}
|
| 209 |
-
{% endfor %}
|
| 210 |
-
}
|
| 211 |
-
}
|
| 212 |
-
{% endfor %}
|
| 213 |
-
{% else %}
|
| 214 |
for (var i = 0u; i < TM; i = i + 1u) {
|
| 215 |
let qi = qi0 + i;
|
| 216 |
if (qi >= params.qSeq) { continue; }
|
|
@@ -225,7 +147,6 @@ fn main(
|
|
| 225 |
}
|
| 226 |
}
|
| 227 |
}
|
| 228 |
-
{% endif %}
|
| 229 |
|
| 230 |
{% if EMIT_ROW_STATS %}
|
| 231 |
// Fold each row's partial (m, d) across the WG_DIM threads that own its columns. Every lane
|
|
|
|
| 1 |
+
{% set STORAGE_F16 = false %}
|
| 2 |
{{ env.wgsl.resourceDeclarations }}
|
| 3 |
|
| 4 |
// f32 prefill score GEMM with BHSD and packed BSH specializations. It computes
|
|
|
|
| 9 |
// for TM * TN * 4 fused multiply-adds. K keeps a one-word shared stride
|
| 10 |
// padding, which changes the transposed workgroup-memory stride and reduces
|
| 11 |
// bank-conflict risk on banked implementations.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
{% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
|
| 13 |
+
{% set statSlots = statSlots | default(0) %}
|
| 14 |
|
| 15 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 16 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 17 |
const HEAD_DIM: u32 = {{ headDim }}u;
|
| 18 |
const Q_HIDDEN: u32 = Q_HEADS * HEAD_DIM;
|
| 19 |
const KV_HIDDEN: u32 = KV_HEADS * HEAD_DIM;
|
|
|
|
|
|
|
|
|
|
| 20 |
|
| 21 |
const BK: u32 = {{ materializedInnerTile }}u;
|
| 22 |
const BM: u32 = {{ materializedQueryTile }}u;
|
|
|
|
| 38 |
var<workgroup> statScratch: array<f32, WG_THREADS>;
|
| 39 |
{% endif %}
|
| 40 |
|
| 41 |
+
{% if EMIT_ROW_STATS %}// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
|
|
|
|
|
|
| 42 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
| 43 |
// accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
|
| 44 |
// a zero final denominator according to their public semantics. Using -inf
|
| 45 |
// here changes +inf-row behavior.
|
| 46 |
const FLT_MAX: f32 = 3.4028234663852886e38;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 47 |
{% endif %}
|
| 48 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 49 |
if (params.scale != 0.0) { return params.scale; }
|
| 50 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 51 |
}
|
| 52 |
|
|
|
|
| 53 |
fn q_at(b: u32, h: u32, q: u32, d: u32) -> f32 {
|
| 54 |
+
var value = {{ "f32(" if STORAGE_F16 else "" }}query[(b * params.qSeq + q) * Q_HIDDEN + h * HEAD_DIM + d]{{ ")" if STORAGE_F16 else "" }};
|
| 55 |
{% if hasBias %}
|
| 56 |
value = value + bias[h * HEAD_DIM + d];
|
| 57 |
{% endif %}
|
|
|
|
| 59 |
}
|
| 60 |
|
| 61 |
fn k_at(b: u32, h: u32, k: u32, d: u32) -> f32 {
|
| 62 |
+
return {{ "f32(" if STORAGE_F16 else "" }}key[(b * params.kvSeq + k) * KV_HIDDEN + h * HEAD_DIM + d]{{ ")" if STORAGE_F16 else "" }};
|
| 63 |
}
|
| 64 |
|
| 65 |
@compute @workgroup_size(WG_DIM, WG_DIM, 1)
|
|
|
|
| 74 |
let kBase = wg.x * BN;
|
| 75 |
let li = lid.y * WG_DIM + lid.x;
|
| 76 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 77 |
var acc: array<f32, TM * TN>;
|
| 78 |
for (var i = 0u; i < TM * TN; i = i + 1u) { acc[i] = 0.0; }
|
|
|
|
| 79 |
|
| 80 |
for (var dBase = 0u; dBase < HEAD_DIM; dBase = dBase + BK) {
|
| 81 |
// Both source tiles are loaded in their physical row-major direction.
|
|
|
|
| 110 |
let qr0 = lid.y * TM;
|
| 111 |
let kc0 = lid.x * TN;
|
| 112 |
for (var dv = 0u; dv < D_VECS; dv = dv + 1u) {
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 113 |
var qv: array<vec4<f32>, TM>;
|
| 114 |
var kv: array<vec4<f32>, TN>;
|
| 115 |
for (var i = 0u; i < TM; i = i + 1u) { qv[i] = tileQ[qr0 + i][dv]; }
|
|
|
|
| 119 |
acc[i * TN + j] = acc[i * TN + j] + dot(qv[i], kv[j]);
|
| 120 |
}
|
| 121 |
}
|
|
|
|
| 122 |
}
|
| 123 |
workgroupBarrier();
|
| 124 |
}
|
|
|
|
| 133 |
var rowM: array<f32, TM>;
|
| 134 |
for (var i = 0u; i < TM; i = i + 1u) { rowM[i] = -FLT_MAX; }
|
| 135 |
{% endif %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 136 |
for (var i = 0u; i < TM; i = i + 1u) {
|
| 137 |
let qi = qi0 + i;
|
| 138 |
if (qi >= params.qSeq) { continue; }
|
|
|
|
| 147 |
}
|
| 148 |
}
|
| 149 |
}
|
|
|
|
| 150 |
|
| 151 |
{% if EMIT_ROW_STATS %}
|
| 152 |
// Fold each row's partial (m, d) across the WG_DIM threads that own its columns. Every lane
|
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,7 +7,6 @@ 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") %}
|
|
@@ -19,27 +16,21 @@ diagnostic(off, chromium.subgroup_matrix_uniformity);
|
|
| 19 |
{% 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 %}
|
| 20 |
{% set OUT_ROW_STRIDE = "HEAD_DIM" if headMajor else "HIDDEN" %}
|
| 21 |
{% set KV_ROW_STRIDE = "HEAD_DIM" if kvHeadMajor else "KV_HIDDEN" %}
|
|
|
|
| 22 |
{% set scorePhase = phase == "score" %}
|
| 23 |
{% set FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
|
| 24 |
{% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
|
| 25 |
{% macro score_value(index, guard) %}
|
| 26 |
-
{% if
|
| 27 |
-
{% if PRIVATE_ROW_STATS %}
|
| 28 |
-
select(0.0, exp_shift(scores[
|
| 29 |
-
{%- else %}
|
| 30 |
-
select(0.0, exp_shift(scores[{{ index }}], softmax_m[{{ guard[0] }}]) / softmax_d[{{ guard[0] }}], {{ guard[1] }})
|
| 31 |
-
{%- endif %}
|
| 32 |
-
{% else %}
|
| 33 |
-
select(0.0, scores[{{ index }}], {{ guard[1] }})
|
| 34 |
-
{%- endif %}
|
| 35 |
-
{% endmacro %}
|
| 36 |
{% set TILE_M_VALUE = materializedSgmatQueryTile %}
|
| 37 |
{% set TILE_N_VALUE = materializedSgmatKeyTile %}
|
| 38 |
{% set TILE_K_VALUE = materializedSgmatInnerTile %}
|
| 39 |
-
{% set SUB_ROWS_VALUE =
|
| 40 |
-
{% set SUB_COLS_VALUE =
|
| 41 |
-
{% set ROW_BLOCKS =
|
| 42 |
-
{% set COL_BLOCKS =
|
| 43 |
{% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
|
| 44 |
{% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
|
| 45 |
{% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
|
|
@@ -47,15 +38,16 @@ select(0.0, scores[{{ index }}], {{ guard[1] }})
|
|
| 47 |
{% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
|
| 48 |
{% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
|
| 49 |
{% set DIRECT_SCORE_STORE = materializedSgmatDirectScoreStore and not EMIT_ROW_STATS %}
|
| 50 |
-
{% set DIRECT_APPLY_STORE = materializedSgmatDirectApplyStore and not hasBias and MT == "f32" %}
|
| 51 |
{% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
|
| 52 |
{% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
|
| 53 |
and materializedSgmatRuntimeDirectStore
|
| 54 |
-
and (scorePhase or (not hasBias and MT == "f32"))
|
| 55 |
and not EMIT_ROW_STATS %}
|
| 56 |
{% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
|
| 57 |
{% set SCALE_IN_Q = DIRECT_SCORE_STORE or (RUNTIME_DIRECT_STORE and scorePhase) %}
|
| 58 |
-
{% macro
|
|
|
|
| 59 |
|
| 60 |
const HEADS: u32 = {{ qNumHeads }}u;
|
| 61 |
{% if kvNumHeads is defined %}
|
|
@@ -108,14 +100,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 +119,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 "" }}
|
|
@@ -173,30 +159,18 @@ fn main(
|
|
| 173 |
let inner = params.kvSeq;
|
| 174 |
{% endif %}
|
| 175 |
{% if FUSED_SOFTMAX %}
|
| 176 |
-
{% if PRIVATE_ROW_STATS %}
|
| 177 |
-
// In the admitted BM64/BN64/BK32/WG256 loader, four adjacent lanes own the
|
| 178 |
-
// same query row for every reduction tile. Keep that row's constants private:
|
| 179 |
-
// this removes both 512 bytes of workgroup storage and the initialization
|
| 180 |
-
// barrier while preserving the exact exp/divide sequence of the shared-memory
|
| 181 |
-
// row-stats arm.
|
| 182 |
-
let private_stat_row =
|
| 183 |
-
(b * HEADS + h) * params.qSeq + min(m_base + li / 4u, params.qSeq - 1u);
|
| 184 |
-
let private_softmax_m = rowStats[private_stat_row * 2u];
|
| 185 |
-
let private_softmax_d = rowStats[private_stat_row * 2u + 1u];
|
| 186 |
-
{% else %}
|
| 187 |
// One row-stats pair per query row of the tile. A query tail clamps to the last
|
| 188 |
// real row rather than reading past the buffer; those lanes are discarded by the
|
| 189 |
// staging guard anyway, and the clamp keeps the denominator non-zero.
|
| 190 |
for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
|
| 191 |
-
let stat_row =
|
| 192 |
softmax_m[r] = rowStats[stat_row * 2u];
|
| 193 |
softmax_d[r] = rowStats[stat_row * 2u + 1u];
|
| 194 |
}
|
| 195 |
workgroupBarrier();
|
| 196 |
-
{% endif %}
|
| 197 |
{% endif %}
|
| 198 |
for (var k_base = 0u; k_base < inner; k_base += TILE_K) {
|
| 199 |
-
{% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" %}
|
| 200 |
// Full interior PV tiles can be loaded directly from storage. Query,
|
| 201 |
// reduction, and output-dimension tails use the guarded shared path below.
|
| 202 |
if (
|
|
@@ -206,7 +180,7 @@ fn main(
|
|
| 206 |
) {
|
| 207 |
for (var step = 0u; step < TILE_K; step += 8u) {
|
| 208 |
{% for row_block in range(ROW_BLOCKS) %}
|
| 209 |
-
let score_offset{{ row_block }} =
|
| 210 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
|
| 211 |
var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
|
| 212 |
subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
|
|
@@ -256,7 +230,7 @@ fn main(
|
|
| 256 |
tile_A[a_row * TILE_K + a_col + i] = loaded;
|
| 257 |
{% endif %}
|
| 258 |
{% else %}
|
| 259 |
-
let score_base =
|
| 260 |
tile_A[a_row * TILE_K + a_col + i] =
|
| 261 |
{{ "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 "" }};
|
| 262 |
{% endif %}
|
|
@@ -271,20 +245,20 @@ fn main(
|
|
| 271 |
{% if headDim % 32 == 0 %}
|
| 272 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 273 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 274 |
-
|
| 275 |
col < params.kvSeq
|
| 276 |
);
|
| 277 |
{% else %}
|
| 278 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 279 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 280 |
-
loaded =
|
| 281 |
}
|
| 282 |
tile_B[b_row * TILE_K + b_col + i] = loaded;
|
| 283 |
{% endif %}
|
| 284 |
{% else %}
|
| 285 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 286 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 287 |
-
|
| 288 |
k < params.kvSeq && col < HEAD_DIM
|
| 289 |
);
|
| 290 |
{% endif %}
|
|
@@ -303,13 +277,13 @@ fn main(
|
|
| 303 |
loaded = {{ q_tile_value(q_index("row", "k")) }};
|
| 304 |
}
|
| 305 |
{% elif FUSED_SOFTMAX %}
|
| 306 |
-
let score_base =
|
| 307 |
let loaded =
|
| 308 |
{{ "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 "" }};
|
| 309 |
{% else %}
|
| 310 |
var loaded = 0.0;
|
| 311 |
if (row < params.qSeq && k < params.kvSeq) {
|
| 312 |
-
let score_base =
|
| 313 |
loaded = scores[score_base + row * params.kvSeq + k];
|
| 314 |
}
|
| 315 |
{% endif %}
|
|
@@ -324,12 +298,12 @@ fn main(
|
|
| 324 |
{% if scorePhase %}
|
| 325 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 326 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 327 |
-
loaded =
|
| 328 |
}
|
| 329 |
{% else %}
|
| 330 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 331 |
if (k < params.kvSeq && col < HEAD_DIM) {
|
| 332 |
-
loaded =
|
| 333 |
}
|
| 334 |
{% endif %}
|
| 335 |
tile_B[idx] = loaded;
|
|
@@ -382,7 +356,7 @@ fn main(
|
|
| 382 |
{% for col_block in range(COL_BLOCKS) %}
|
| 383 |
{% if scorePhase %}
|
| 384 |
let output_offset{{ row_block }}{{ col_block }} =
|
| 385 |
-
|
| 386 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
|
| 387 |
+ n_base + base_B + {{ col_block * 8 }}u;
|
| 388 |
subgroupMatrixStore<row_major>(
|
|
@@ -448,7 +422,7 @@ fn main(
|
|
| 448 |
{% endif %}
|
| 449 |
let scored = result{% if not SCALE_IN_Q %} * scale{% endif %};
|
| 450 |
scores[
|
| 451 |
-
|
| 452 |
] = scored;
|
| 453 |
{% if EMIT_ROW_STATS %}
|
| 454 |
// Softmax sees the STORED value, so the statistics have to be taken on it
|
|
@@ -462,7 +436,7 @@ fn main(
|
|
| 462 |
{% if hasBias %}
|
| 463 |
// V bias row base: skip the packed Q and K blocks, then index this head.
|
| 464 |
{% endif %}
|
| 465 |
-
output[{{ q_index("row", "col") }}] = {{ "f16(" if
|
| 466 |
{% endif %}
|
| 467 |
}
|
| 468 |
}
|
|
@@ -490,7 +464,7 @@ fn main(
|
|
| 490 |
// eight consecutive pairs, and the combine pass reads a slot's whole column
|
| 491 |
// of rows contiguously.
|
| 492 |
let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
|
| 493 |
-
let out_index = ((
|
| 494 |
scorePartials[out_index] = stat_m{{ row_block }};
|
| 495 |
scorePartials[out_index + 1u] = stat_d{{ row_block }};
|
| 496 |
}
|
|
|
|
| 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 layout = layout | default("bsh") %}
|
|
|
|
| 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 = "HEAD_DIM" if headMajor else "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 FUSED_SOFTMAX = fusedSoftmax is defined and fusedSoftmax %}
|
| 22 |
{% set PRIVATE_ROW_STATS = FUSED_SOFTMAX and (materializedSgmatPrivateRowStats is defined and materializedSgmatPrivateRowStats) %}
|
| 23 |
{% macro score_value(index, guard) %}
|
| 24 |
+
{% set rowMax = "private_softmax_m" if PRIVATE_ROW_STATS else "softmax_m[" ~ guard[0] ~ "]" %}
|
| 25 |
+
{% set rowDenom = "private_softmax_d" if PRIVATE_ROW_STATS else "softmax_d[" ~ guard[0] ~ "]" %}
|
| 26 |
+
select(0.0, {{ "exp_shift(scores[" ~ index ~ "], " ~ rowMax ~ ") / " ~ rowDenom if FUSED_SOFTMAX else "scores[" ~ index ~ "]" }}, {{ guard[1] }}){% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
{% set TILE_M_VALUE = materializedSgmatQueryTile %}
|
| 28 |
{% set TILE_N_VALUE = materializedSgmatKeyTile %}
|
| 29 |
{% set TILE_K_VALUE = materializedSgmatInnerTile %}
|
| 30 |
+
{% set SUB_ROWS_VALUE = 16 %}
|
| 31 |
+
{% set SUB_COLS_VALUE = 32 %}
|
| 32 |
+
{% set ROW_BLOCKS = 2 %}
|
| 33 |
+
{% set COL_BLOCKS = 4 %}
|
| 34 |
{% set SUBGROUP_ROWS = (TILE_M_VALUE / SUB_ROWS_VALUE)|int %}
|
| 35 |
{% set SUBGROUP_COLS = (TILE_N_VALUE / SUB_COLS_VALUE)|int %}
|
| 36 |
{% set SUBGROUP_COUNT = SUBGROUP_ROWS * SUBGROUP_COLS %}
|
|
|
|
| 38 |
{% if hasBias is not defined %}{% set hasBias = false %}{% endif %}
|
| 39 |
{% set EMIT_ROW_STATS = emitRowStats is defined and emitRowStats %}
|
| 40 |
{% set DIRECT_SCORE_STORE = materializedSgmatDirectScoreStore and not EMIT_ROW_STATS %}
|
| 41 |
+
{% set DIRECT_APPLY_STORE = materializedSgmatDirectApplyStore and not hasBias and MT == "f32" and ST == "f32" %}
|
| 42 |
{% set DIRECT_OUTPUT_STORE = DIRECT_SCORE_STORE if scorePhase else DIRECT_APPLY_STORE %}
|
| 43 |
{% set RUNTIME_DIRECT_STORE = (not DIRECT_OUTPUT_STORE)
|
| 44 |
and materializedSgmatRuntimeDirectStore
|
| 45 |
+
and (scorePhase or (not hasBias and MT == "f32" and ST == "f32"))
|
| 46 |
and not EMIT_ROW_STATS %}
|
| 47 |
{% set ANY_DIRECT_STORE = DIRECT_OUTPUT_STORE or RUNTIME_DIRECT_STORE %}
|
| 48 |
{% set SCALE_IN_Q = DIRECT_SCORE_STORE or (RUNTIME_DIRECT_STORE and scorePhase) %}
|
| 49 |
+
{% macro operand_load(name, index) %}{% if ST != MT %}{{ MT }}({% endif %}{{ name }}[{{ index }}]{% if ST != MT %}){% endif %}{% endmacro %}
|
| 50 |
+
{% 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 %}
|
| 51 |
|
| 52 |
const HEADS: u32 = {{ qNumHeads }}u;
|
| 53 |
{% if kvNumHeads is defined %}
|
|
|
|
| 100 |
var<workgroup> softmax_d: array<f32, {{ TILE_M_VALUE }}>;
|
| 101 |
{% endif %}
|
| 102 |
{% if FUSED_SOFTMAX or EMIT_ROW_STATS %}
|
|
|
|
| 103 |
// FLT_MAX, not -inf, as the online (m, d) accumulator init: merges must keep
|
| 104 |
// `m - m` finite so an empty lane / all--inf row contributes the exact
|
| 105 |
// accumulator identity (m, d) = (-FLT_MAX, 0). Operator epilogues interpret
|
| 106 |
// a zero final denominator according to their public semantics. Using -inf
|
| 107 |
// here changes +inf-row behavior.
|
| 108 |
const FLT_MAX: f32 = 3.4028234663852886e38;
|
|
|
|
| 109 |
|
| 110 |
fn is_finite_f32(value: f32) -> bool {
|
| 111 |
return select(false, value <= FLT_MAX, value >= -FLT_MAX);
|
|
|
|
| 119 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 120 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 121 |
}
|
|
|
|
|
|
|
| 122 |
|
| 123 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 124 |
return exp(shifted_value(value, maxValue));
|
| 125 |
}
|
|
|
|
|
|
|
| 126 |
{% endif %}
|
| 127 |
|
| 128 |
@compute @workgroup_size({{ WORKGROUP_THREADS }}, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
|
|
|
|
| 159 |
let inner = params.kvSeq;
|
| 160 |
{% endif %}
|
| 161 |
{% if FUSED_SOFTMAX %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 162 |
// One row-stats pair per query row of the tile. A query tail clamps to the last
|
| 163 |
// real row rather than reading past the buffer; those lanes are discarded by the
|
| 164 |
// staging guard anyway, and the clamp keeps the denominator non-zero.
|
| 165 |
for (var r = li; r < TILE_M; r += WORKGROUP_THREADS) {
|
| 166 |
+
let stat_row = {{ SCRATCH_HEAD }} * params.qSeq + min(m_base + r, params.qSeq - 1u);
|
| 167 |
softmax_m[r] = rowStats[stat_row * 2u];
|
| 168 |
softmax_d[r] = rowStats[stat_row * 2u + 1u];
|
| 169 |
}
|
| 170 |
workgroupBarrier();
|
|
|
|
| 171 |
{% endif %}
|
| 172 |
for (var k_base = 0u; k_base < inner; k_base += TILE_K) {
|
| 173 |
+
{% if phase == "apply" and not FUSED_SOFTMAX and MT == "f32" and ST == "f32" %}
|
| 174 |
// Full interior PV tiles can be loaded directly from storage. Query,
|
| 175 |
// reduction, and output-dimension tails use the guarded shared path below.
|
| 176 |
if (
|
|
|
|
| 180 |
) {
|
| 181 |
for (var step = 0u; step < TILE_K; step += 8u) {
|
| 182 |
{% for row_block in range(ROW_BLOCKS) %}
|
| 183 |
+
let score_offset{{ row_block }} = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
|
| 184 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq + k_base + step;
|
| 185 |
var matA{{ row_block }}: subgroup_matrix_left<f32, 8, 8> =
|
| 186 |
subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 8>, row_major>(
|
|
|
|
| 230 |
tile_A[a_row * TILE_K + a_col + i] = loaded;
|
| 231 |
{% endif %}
|
| 232 |
{% else %}
|
| 233 |
+
let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
|
| 234 |
tile_A[a_row * TILE_K + a_col + i] =
|
| 235 |
{{ "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 "" }};
|
| 236 |
{% endif %}
|
|
|
|
| 245 |
{% if headDim % 32 == 0 %}
|
| 246 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 247 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 248 |
+
{{ operand_load("key", kv_index("col", "k")) }},
|
| 249 |
col < params.kvSeq
|
| 250 |
);
|
| 251 |
{% else %}
|
| 252 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 253 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 254 |
+
loaded = {{ operand_load("key", kv_index("col", "k")) }};
|
| 255 |
}
|
| 256 |
tile_B[b_row * TILE_K + b_col + i] = loaded;
|
| 257 |
{% endif %}
|
| 258 |
{% else %}
|
| 259 |
tile_B[b_row * TILE_K + b_col + i] = select(
|
| 260 |
{{ "0.0h" if MT == "f16" else "0.0" }},
|
| 261 |
+
{{ operand_load("value", kv_index("k", "col")) }},
|
| 262 |
k < params.kvSeq && col < HEAD_DIM
|
| 263 |
);
|
| 264 |
{% endif %}
|
|
|
|
| 277 |
loaded = {{ q_tile_value(q_index("row", "k")) }};
|
| 278 |
}
|
| 279 |
{% elif FUSED_SOFTMAX %}
|
| 280 |
+
let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
|
| 281 |
let loaded =
|
| 282 |
{{ "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 "" }};
|
| 283 |
{% else %}
|
| 284 |
var loaded = 0.0;
|
| 285 |
if (row < params.qSeq && k < params.kvSeq) {
|
| 286 |
+
let score_base = {{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq;
|
| 287 |
loaded = scores[score_base + row * params.kvSeq + k];
|
| 288 |
}
|
| 289 |
{% endif %}
|
|
|
|
| 298 |
{% if scorePhase %}
|
| 299 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 300 |
if (col < params.kvSeq && k < HEAD_DIM) {
|
| 301 |
+
loaded = {{ operand_load("key", kv_index("col", "k")) }};
|
| 302 |
}
|
| 303 |
{% else %}
|
| 304 |
var loaded = {{ "0.0h" if MT == "f16" else "0.0" }};
|
| 305 |
if (k < params.kvSeq && col < HEAD_DIM) {
|
| 306 |
+
loaded = {{ operand_load("value", kv_index("k", "col")) }};
|
| 307 |
}
|
| 308 |
{% endif %}
|
| 309 |
tile_B[idx] = loaded;
|
|
|
|
| 356 |
{% for col_block in range(COL_BLOCKS) %}
|
| 357 |
{% if scorePhase %}
|
| 358 |
let output_offset{{ row_block }}{{ col_block }} =
|
| 359 |
+
{{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq
|
| 360 |
+ (m_base + base_A + {{ row_block * 8 }}u) * params.kvSeq
|
| 361 |
+ n_base + base_B + {{ col_block * 8 }}u;
|
| 362 |
subgroupMatrixStore<row_major>(
|
|
|
|
| 422 |
{% endif %}
|
| 423 |
let scored = result{% if not SCALE_IN_Q %} * scale{% endif %};
|
| 424 |
scores[
|
| 425 |
+
{{ SCRATCH_HEAD }} * params.qSeq * params.kvSeq + row * params.kvSeq + col
|
| 426 |
] = scored;
|
| 427 |
{% if EMIT_ROW_STATS %}
|
| 428 |
// Softmax sees the STORED value, so the statistics have to be taken on it
|
|
|
|
| 436 |
{% if hasBias %}
|
| 437 |
// V bias row base: skip the packed Q and K blocks, then index this head.
|
| 438 |
{% endif %}
|
| 439 |
+
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 %};
|
| 440 |
{% endif %}
|
| 441 |
}
|
| 442 |
}
|
|
|
|
| 464 |
// eight consecutive pairs, and the combine pass reads a slot's whole column
|
| 465 |
// of rows contiguously.
|
| 466 |
let slot = wg.x * {{ SUBGROUP_COLS }}u + subtile_idx;
|
| 467 |
+
let out_index = (({{ SCRATCH_HEAD }} * STAT_SLOTS + slot) * params.qSeq + stat_row) * 2u;
|
| 468 |
scorePartials[out_index] = stat_m{{ row_block }};
|
| 469 |
scorePartials[out_index + 1u] = stat_d{{ row_block }};
|
| 470 |
}
|
build/webgpu/attn-materialized-softmax-f32.wgsl.jinja
CHANGED
|
@@ -1,9 +1,7 @@
|
|
| 1 |
{% set CACHE_VEC4 = cacheVec4 if cacheVec4 is defined else false %}
|
| 2 |
-
{% set CAUSAL_ROWS =
|
| 3 |
-
{% set
|
| 4 |
-
{% set
|
| 5 |
-
{% set COLS_EXPR = "params.keyLen" if RUNTIME_COLS else "COLS" %}
|
| 6 |
-
{% set ROWS_EXPR = rowsExpr if rowsExpr is defined else "params.rows" %}
|
| 7 |
{% if useSubgroups %}
|
| 8 |
enable subgroups;
|
| 9 |
{% endif %}
|
|
@@ -12,7 +10,7 @@ enable subgroups;
|
|
| 12 |
// In-place row softmax for materialized f32 attention. The default path reads
|
| 13 |
// each scalar twice; the register-cached vec4 path eliminates the second read.
|
| 14 |
const WG: u32 = {{ materializedSoftmaxWg }}u;
|
| 15 |
-
{% set SCALED_SCORE = "scores[base + c]
|
| 16 |
{% if CACHE_VEC4 %}
|
| 17 |
const COLS4: u32 = {{ materializedSoftmaxCols4 }}u;
|
| 18 |
const CACHE_VECS: u32 = (COLS4 + WG - 1u) / WG;
|
|
@@ -38,6 +36,7 @@ fn shifted_value(value: f32, maxValue: f32) -> f32 {
|
|
| 38 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 39 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 40 |
}
|
|
|
|
| 41 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 42 |
return exp(shifted_value(value, maxValue));
|
| 43 |
}
|
|
@@ -101,41 +100,9 @@ fn combine_partials(m: f32, d: f32, lidx: u32, sgSize: u32) -> vec2<f32> {
|
|
| 101 |
return combinedMD;
|
| 102 |
}
|
| 103 |
{% else %}
|
| 104 |
-
{% set
|
| 105 |
-
{% set mdStreams = mdStreams if mdStreams is defined else 1 %}
|
| 106 |
-
{% set mdExtent = "WG" if mdStreams == 1 else "WG * " ~ mdStreams ~ "u" %}
|
| 107 |
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 108 |
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
| 109 |
-
{% if mdStreamed %}
|
| 110 |
-
|
| 111 |
-
// In-place fold of {{ mdStreams }} streams. Input partials occupy
|
| 112 |
-
// partialM/partialD; stream s returns its merged pair in slot s * WG.
|
| 113 |
-
fn combine_partials_streams(lidx: u32) {
|
| 114 |
-
workgroupBarrier();
|
| 115 |
-
var stride = WG / 2u;
|
| 116 |
-
loop {
|
| 117 |
-
if (stride == 0u) {
|
| 118 |
-
break;
|
| 119 |
-
}
|
| 120 |
-
if (lidx < stride) {
|
| 121 |
-
{% for s in range(mdStreams) %}
|
| 122 |
-
{
|
| 123 |
-
let slot = {{ s }}u * WG + lidx;
|
| 124 |
-
let m1 = partialM[slot];
|
| 125 |
-
let d1 = partialD[slot];
|
| 126 |
-
let m2 = partialM[slot + stride];
|
| 127 |
-
let d2 = partialD[slot + stride];
|
| 128 |
-
let mNew = max(m1, m2);
|
| 129 |
-
partialD[slot] = d1 * exp_shift(m1, mNew) + d2 * exp_shift(m2, mNew);
|
| 130 |
-
partialM[slot] = mNew;
|
| 131 |
-
}
|
| 132 |
-
{% endfor %}
|
| 133 |
-
}
|
| 134 |
-
workgroupBarrier();
|
| 135 |
-
stride = stride / 2u;
|
| 136 |
-
}
|
| 137 |
-
}
|
| 138 |
-
{% else %}
|
| 139 |
|
| 140 |
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 141 |
partialM[lidx] = m;
|
|
@@ -165,8 +132,6 @@ fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
|
| 165 |
return merged;
|
| 166 |
}
|
| 167 |
{% endif %}
|
| 168 |
-
{% endif %}
|
| 169 |
-
|
| 170 |
|
| 171 |
@compute @workgroup_size(WG, 1, 1)
|
| 172 |
fn main(
|
|
|
|
| 1 |
{% set CACHE_VEC4 = cacheVec4 if cacheVec4 is defined else false %}
|
| 2 |
+
{% set CAUSAL_ROWS = false %}
|
| 3 |
+
{% set COLS_EXPR = "COLS" %}
|
| 4 |
+
{% set ROWS_EXPR = "params.rows" %}
|
|
|
|
|
|
|
| 5 |
{% if useSubgroups %}
|
| 6 |
enable subgroups;
|
| 7 |
{% endif %}
|
|
|
|
| 10 |
// In-place row softmax for materialized f32 attention. The default path reads
|
| 11 |
// each scalar twice; the register-cached vec4 path eliminates the second read.
|
| 12 |
const WG: u32 = {{ materializedSoftmaxWg }}u;
|
| 13 |
+
{% set SCALED_SCORE = "scores[base + c]" %}
|
| 14 |
{% if CACHE_VEC4 %}
|
| 15 |
const COLS4: u32 = {{ materializedSoftmaxCols4 }}u;
|
| 16 |
const CACHE_VECS: u32 = (COLS4 + WG - 1u) / WG;
|
|
|
|
| 36 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 37 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 38 |
}
|
| 39 |
+
|
| 40 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 41 |
return exp(shifted_value(value, maxValue));
|
| 42 |
}
|
|
|
|
| 100 |
return combinedMD;
|
| 101 |
}
|
| 102 |
{% else %}
|
| 103 |
+
{% set mdExtent = "WG" %}
|
|
|
|
|
|
|
| 104 |
var<workgroup> partialM: array<f32, {{ mdExtent }}>;
|
| 105 |
var<workgroup> partialD: array<f32, {{ mdExtent }}>;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 106 |
|
| 107 |
fn combine_partials(m: f32, d: f32, lidx: u32) -> vec2<f32> {
|
| 108 |
partialM[lidx] = m;
|
|
|
|
| 132 |
return merged;
|
| 133 |
}
|
| 134 |
{% endif %}
|
|
|
|
|
|
|
| 135 |
|
| 136 |
@compute @workgroup_size(WG, 1, 1)
|
| 137 |
fn main(
|
build/webgpu/attn-online-scalar.wgsl.jinja
CHANGED
|
@@ -1,11 +1,11 @@
|
|
| 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.
|
|
@@ -21,18 +21,39 @@ const Q_HIDDEN: u32 = {{ qHidden }}u;
|
|
| 21 |
const KV_HIDDEN: u32 = {{ kvHidden }}u;
|
| 22 |
const Q_HEADS: u32 = {{ qNumHeads }}u;
|
| 23 |
const KV_HEADS: u32 = {{ kvNumHeads }}u;
|
| 24 |
-
{% set qHeads = "
|
| 25 |
-
{% set kvHeads = "
|
| 26 |
-
{% set scale = scale | default("0.0") %}
|
| 27 |
const WG: u32 = {{ workgroupSize }}u;
|
| 28 |
|
| 29 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
var<workgroup> running_max: f32;
|
| 31 |
var<workgroup> running_denom: f32;
|
| 32 |
var<workgroup> running_out: array<f32, HEAD_DIM>;
|
| 33 |
var<workgroup> previous_scale: f32;
|
| 34 |
-
{% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
|
| 35 |
-
fn {{ name }}(value:
|
| 36 |
{{ buffer }}[tid] = value;
|
| 37 |
workgroupBarrier();
|
| 38 |
// Ceil-halving keeps every lane when the workgroup size is not a power of
|
|
@@ -42,11 +63,7 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
|
|
| 42 |
loop {
|
| 43 |
let half = (n + 1u) / 2u;
|
| 44 |
if (tid < n - half) {
|
| 45 |
-
{
|
| 46 |
-
{{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
|
| 47 |
-
{% else %}
|
| 48 |
-
{{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
|
| 49 |
-
{% endif %}
|
| 50 |
}
|
| 51 |
workgroupBarrier();
|
| 52 |
n = half;
|
|
@@ -58,22 +75,16 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
|
|
| 58 |
// slot 0 here, so the next call's first store must not run until all lanes have read it.
|
| 59 |
// `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
|
| 60 |
let reduced = {{ buffer }}[0];
|
| 61 |
-
{% if trailingBarrier %}
|
| 62 |
workgroupBarrier();
|
| 63 |
-
{% endif %}
|
| 64 |
return reduced;
|
| 65 |
}
|
| 66 |
{% endmacro %}
|
| 67 |
-
|
| 68 |
-
{
|
| 69 |
-
|
| 70 |
-
{% if ATTN_SCALE_DIM is not defined %}{% set ATTN_SCALE_DIM = "HEAD_DIM" %}{% endif %}
|
| 71 |
-
fn scale_value() -> f32 {
|
| 72 |
if (params.scale != 0.0) { return params.scale; }
|
| 73 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 74 |
}
|
| 75 |
|
| 76 |
-
|
| 77 |
@compute @workgroup_size(WG, 1, 1)
|
| 78 |
fn main(
|
| 79 |
@builtin(workgroup_id) wg: vec3<u32>,
|
|
@@ -117,7 +128,7 @@ fn main(
|
|
| 117 |
|
| 118 |
for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
|
| 119 |
let kRow = kvBase + key_token * kvTokenStride;
|
| 120 |
-
var partial_dot = 0.0;
|
| 121 |
for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
|
| 122 |
var q_value = f32(query[qBase + d]);
|
| 123 |
var k_value = f32(key[kRow + d]);
|
|
@@ -126,17 +137,17 @@ fn main(
|
|
| 126 |
q_value = q_value + f32(bias[channel]);
|
| 127 |
k_value = k_value + f32(bias[Q_HIDDEN + channel]);
|
| 128 |
{% endif %}
|
| 129 |
-
partial_dot = partial_dot
|
| 130 |
}
|
| 131 |
|
| 132 |
-
//
|
| 133 |
// workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
|
| 134 |
// through shared memory by the barrier below).
|
| 135 |
{% if hasMask %}
|
| 136 |
let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
|
| 137 |
-
let score =
|
| 138 |
{% else %}
|
| 139 |
-
let score =
|
| 140 |
{% endif %}
|
| 141 |
if (tid == 0u) {
|
| 142 |
let next_max = max(running_max, score);
|
|
|
|
| 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.
|
|
|
|
| 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>,
|
|
|
|
| 128 |
|
| 129 |
for (var key_token: u32 = minKj; key_token < maxKj; key_token = key_token + 1u) {
|
| 130 |
let kRow = kvBase + key_token * kvTokenStride;
|
| 131 |
+
var partial_dot = DotAccumulator(0.0, 0.0);
|
| 132 |
for (var d: u32 = tid; d < HEAD_DIM; d = d + WG) {
|
| 133 |
var q_value = f32(query[qBase + d]);
|
| 134 |
var k_value = f32(key[kRow + d]);
|
|
|
|
| 137 |
q_value = q_value + f32(bias[channel]);
|
| 138 |
k_value = k_value + f32(bias[Q_HIDDEN + channel]);
|
| 139 |
{% endif %}
|
| 140 |
+
partial_dot = dot_accumulate(partial_dot, q_value, k_value);
|
| 141 |
}
|
| 142 |
|
| 143 |
+
// reduce_dot returns the same partial[0] to every lane, so `score` is already
|
| 144 |
// workgroup-uniform; tid==0 advances the running max/denom/scale (broadcast
|
| 145 |
// through shared memory by the barrier below).
|
| 146 |
{% if hasMask %}
|
| 147 |
let maskIndex = {{ MASK_BATCH }}h * params.maskHeadStride + query_token * params.maskSeqStride + key_token;
|
| 148 |
+
let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value() + f32(attn_mask[maskIndex]);
|
| 149 |
{% else %}
|
| 150 |
+
let score = dot_value(reduce_dot(partial_dot, tid)) * scale_value();
|
| 151 |
{% endif %}
|
| 152 |
if (tid == 0u) {
|
| 153 |
let next_max = max(running_max, score);
|
build/webgpu/attn-small-head-parallel.wgsl.jinja
CHANGED
|
@@ -9,17 +9,15 @@ const HIDDEN: u32 = {{ qHidden }}u;
|
|
| 9 |
const KV_SEQ: u32 = {{ kvSeq }}u;
|
| 10 |
const WG: u32 = 64u;
|
| 11 |
const NEG_MAX: f32 = -3.4028234663852886e38;
|
| 12 |
-
{%
|
| 13 |
-
fn scale_value() -> f32 {
|
| 14 |
if (params.scale != 0.0) { return params.scale; }
|
| 15 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 16 |
}
|
| 17 |
|
| 18 |
-
|
| 19 |
var<workgroup> scores: array<f32, KV_SEQ>;
|
| 20 |
var<workgroup> partial: array<f32, WG>;
|
| 21 |
-
{% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
|
| 22 |
-
fn {{ name }}(value:
|
| 23 |
{{ buffer }}[tid] = value;
|
| 24 |
workgroupBarrier();
|
| 25 |
// Ceil-halving keeps every lane when the workgroup size is not a power of
|
|
@@ -45,16 +43,12 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
|
|
| 45 |
// slot 0 here, so the next call's first store must not run until all lanes have read it.
|
| 46 |
// `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
|
| 47 |
let reduced = {{ buffer }}[0];
|
| 48 |
-
{% if trailingBarrier %}
|
| 49 |
workgroupBarrier();
|
| 50 |
-
{% endif %}
|
| 51 |
return reduced;
|
| 52 |
}
|
| 53 |
{% endmacro %}
|
| 54 |
-
|
| 55 |
{{ wgsl_tree_reduce_f32("reduce_max", "max", "partial", "WG") }}
|
| 56 |
{{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
|
| 57 |
-
|
| 58 |
@compute @workgroup_size(WG, 1, 1)
|
| 59 |
fn main(
|
| 60 |
@builtin(workgroup_id) wg: vec3<u32>,
|
|
|
|
| 9 |
const KV_SEQ: u32 = {{ kvSeq }}u;
|
| 10 |
const WG: u32 = 64u;
|
| 11 |
const NEG_MAX: f32 = -3.4028234663852886e38;
|
| 12 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 13 |
if (params.scale != 0.0) { return params.scale; }
|
| 14 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 15 |
}
|
| 16 |
|
|
|
|
| 17 |
var<workgroup> scores: array<f32, KV_SEQ>;
|
| 18 |
var<workgroup> partial: array<f32, WG>;
|
| 19 |
+
{% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
|
| 20 |
+
fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
|
| 21 |
{{ buffer }}[tid] = value;
|
| 22 |
workgroupBarrier();
|
| 23 |
// Ceil-halving keeps every lane when the workgroup size is not a power of
|
|
|
|
| 43 |
// slot 0 here, so the next call's first store must not run until all lanes have read it.
|
| 44 |
// `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
|
| 45 |
let reduced = {{ buffer }}[0];
|
|
|
|
| 46 |
workgroupBarrier();
|
|
|
|
| 47 |
return reduced;
|
| 48 |
}
|
| 49 |
{% endmacro %}
|
|
|
|
| 50 |
{{ wgsl_tree_reduce_f32("reduce_max", "max", "partial", "WG") }}
|
| 51 |
{{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
|
|
|
|
| 52 |
@compute @workgroup_size(WG, 1, 1)
|
| 53 |
fn main(
|
| 54 |
@builtin(workgroup_id) wg: vec3<u32>,
|
build/webgpu/attn-small-head-value.wgsl.jinja
ADDED
|
@@ -0,0 +1,198 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{% if useValueSubgroups %}
|
| 2 |
+
enable subgroups;
|
| 3 |
+
{% endif %}
|
| 4 |
+
{{ env.wgsl.resourceDeclarations }}
|
| 5 |
+
{% set rowType = "f32" if queryBlock == 1 else "vec" ~ queryBlock ~ "<f32>" %}
|
| 6 |
+
{% set dotValueType = rowType %}
|
| 7 |
+
{% set dotType = dotValueType | default("f32") %}
|
| 8 |
+
// Retain product and addition residuals across a dot product. Explicit fma
|
| 9 |
+
// boundaries preserve the addition error transform under reassociation.
|
| 10 |
+
struct DotAccumulator {
|
| 11 |
+
hi: {{ dotType }},
|
| 12 |
+
lo: {{ dotType }},
|
| 13 |
+
}
|
| 14 |
+
fn dot_accumulate(acc: DotAccumulator, a: {{ dotType }}, b: {{ dotType }}) -> DotAccumulator {
|
| 15 |
+
let product = fma(a, b, {{ dotType }}(0.0));
|
| 16 |
+
let productError = fma(a, b, -product);
|
| 17 |
+
let sum = fma(acc.hi, {{ dotType }}(1.0), product);
|
| 18 |
+
let bv = fma({{ dotType }}(-1.0), acc.hi, sum);
|
| 19 |
+
let av = fma({{ dotType }}(-1.0), bv, sum);
|
| 20 |
+
let ae = fma({{ dotType }}(-1.0), av, acc.hi);
|
| 21 |
+
let be = fma({{ dotType }}(-1.0), bv, product);
|
| 22 |
+
let error = fma(fma(fma(ae, {{ dotType }}(1.0), be), {{ dotType }}(1.0), acc.lo), {{ dotType }}(1.0), productError);
|
| 23 |
+
let hi = fma(sum, {{ dotType }}(1.0), error);
|
| 24 |
+
return DotAccumulator(hi, fma({{ dotType }}(-1.0), fma({{ dotType }}(-1.0), sum, hi), error));
|
| 25 |
+
}
|
| 26 |
+
fn dot_value(acc: DotAccumulator) -> {{ dotType }} {
|
| 27 |
+
return fma(acc.hi, {{ dotType }}(1.0), acc.lo);
|
| 28 |
+
}
|
| 29 |
+
const HEAD_DIM: u32 = {{ headDim }}u;
|
| 30 |
+
const HIDDEN: u32 = {{ qHidden }}u;
|
| 31 |
+
const KV_SEQ: u32 = {{ kvSeq }}u;
|
| 32 |
+
const WG: u32 = {{ valueWorkgroupSize }}u;
|
| 33 |
+
const QUERY_BLOCK: u32 = {{ queryBlock }}u;
|
| 34 |
+
const NEG_MAX: f32 = -3.4028234663852886e38;
|
| 35 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
| 36 |
+
if (params.scale != 0.0) { return params.scale; }
|
| 37 |
+
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 38 |
+
}
|
| 39 |
+
var<workgroup> scores: array<{{ rowType }}, {{ scoreStorageElements }}u>;
|
| 40 |
+
var<workgroup> partial: array<{{ rowType }}, WG>;
|
| 41 |
+
|
| 42 |
+
{% if useValueSubgroups %}
|
| 43 |
+
{% macro reduce(name, collective, mode, identity) %}
|
| 44 |
+
fn {{ name }}(value: {{ rowType }}, sgLane: u32, sgId: u32, numSg: u32) -> {{ rowType }} {
|
| 45 |
+
let sub = {{ collective }}(value);
|
| 46 |
+
if (sgLane == 0u) { partial[sgId] = sub; }
|
| 47 |
+
workgroupBarrier();
|
| 48 |
+
var total = {{ rowType }}({{ identity }});
|
| 49 |
+
for (var i = 0u; i < numSg; i = i + 1u) {
|
| 50 |
+
{% if mode == "max" %}
|
| 51 |
+
total = max(total, partial[i]);
|
| 52 |
+
{% else %}
|
| 53 |
+
total = total + partial[i];
|
| 54 |
+
{% endif %}
|
| 55 |
+
}
|
| 56 |
+
workgroupBarrier();
|
| 57 |
+
return total;
|
| 58 |
+
}
|
| 59 |
+
{% endmacro %}
|
| 60 |
+
{{ reduce("reduce_max", "subgroupMax", "max", "NEG_MAX") }}
|
| 61 |
+
{{ reduce("reduce_sum", "subgroupAdd", "add", "0.0") }}
|
| 62 |
+
{% else %}
|
| 63 |
+
{% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
|
| 64 |
+
fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
|
| 65 |
+
{{ buffer }}[tid] = value;
|
| 66 |
+
workgroupBarrier();
|
| 67 |
+
// Ceil-halving keeps every lane when the workgroup size is not a power of
|
| 68 |
+
// two. For even n this matches the power-of-two tree order; for odd n, lanes
|
| 69 |
+
// [0, n-half) fold the upper tail while the middle lane carries forward.
|
| 70 |
+
var n: u32 = {{ wg }};
|
| 71 |
+
loop {
|
| 72 |
+
let half = (n + 1u) / 2u;
|
| 73 |
+
if (tid < n - half) {
|
| 74 |
+
{% if mode == "max" %}
|
| 75 |
+
{{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
|
| 76 |
+
{% else %}
|
| 77 |
+
{{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
|
| 78 |
+
{% endif %}
|
| 79 |
+
}
|
| 80 |
+
workgroupBarrier();
|
| 81 |
+
n = half;
|
| 82 |
+
if (n == 1u) {
|
| 83 |
+
break;
|
| 84 |
+
}
|
| 85 |
+
}
|
| 86 |
+
// The default trailing barrier makes this helper safe for back-to-back calls: every lane reads
|
| 87 |
+
// slot 0 here, so the next call's first store must not run until all lanes have read it.
|
| 88 |
+
// `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
|
| 89 |
+
let reduced = {{ buffer }}[0];
|
| 90 |
+
workgroupBarrier();
|
| 91 |
+
return reduced;
|
| 92 |
+
}
|
| 93 |
+
{% endmacro %}
|
| 94 |
+
{{ wgsl_tree_reduce_f32("reduce_max", "max", valueType=rowType) }}
|
| 95 |
+
{{ wgsl_tree_reduce_f32("reduce_sum", "add", valueType=rowType) }}
|
| 96 |
+
{% endif %}
|
| 97 |
+
@compute @workgroup_size(WG, 1, 1)
|
| 98 |
+
fn main(
|
| 99 |
+
@builtin(workgroup_id) wg: vec3<u32>,
|
| 100 |
+
@builtin(num_workgroups) nwg: vec3<u32>,
|
| 101 |
+
@builtin(local_invocation_id) lid: vec3<u32>
|
| 102 |
+
{% if useValueSubgroups %}
|
| 103 |
+
, @builtin(subgroup_invocation_id) sgLane: u32
|
| 104 |
+
, @builtin(subgroup_id) sgId: u32
|
| 105 |
+
, @builtin(num_subgroups) numSg: u32
|
| 106 |
+
{% endif %}
|
| 107 |
+
) {
|
| 108 |
+
let h = wg.y;
|
| 109 |
+
let batch = wg.z;
|
| 110 |
+
let tid = lid.x;
|
| 111 |
+
let kvBase = batch * KV_SEQ * HIDDEN + h * HEAD_DIM;
|
| 112 |
+
for (var query_token = wg.x * QUERY_BLOCK; query_token < params.qSeq; query_token = query_token + nwg.x * QUERY_BLOCK) {
|
| 113 |
+
{% for row in range(queryBlock) %}
|
| 114 |
+
let qBase_{{ row }} = (batch * params.qSeq + {% if queryTail and row > 0 %}min(query_token + {{ row }}u, params.qSeq - 1u){% else %}query_token{% if row > 0 %} + {{ row }}u{% endif %}{% endif %}) * HIDDEN + h * HEAD_DIM;
|
| 115 |
+
{% endfor %}
|
| 116 |
+
// Query components stay private throughout the key scan. Their lifetime
|
| 117 |
+
// ends before the value accumulators need the same registers.
|
| 118 |
+
{
|
| 119 |
+
{% for d in range(headDim) %}
|
| 120 |
+
let q_{{ d }} = {% if queryBlock > 1 %}{{ rowType }}({% endif %}{% for row in range(queryBlock) %}{% if not loop.first %}, {% endif %}f32(query[qBase_{{ row }} + {{ d }}u]){% endfor %}{% if queryBlock > 1 %}){% endif %};
|
| 121 |
+
{% endfor %}
|
| 122 |
+
for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
|
| 123 |
+
let kBase = kvBase + key_token * HIDDEN;
|
| 124 |
+
var score = DotAccumulator({{ rowType }}(0.0), {{ rowType }}(0.0));
|
| 125 |
+
{% for d in range(headDim) %}
|
| 126 |
+
score = dot_accumulate(score, q_{{ d }}, {{ rowType }}(f32(key[kBase + {{ d }}u])));
|
| 127 |
+
{% endfor %}
|
| 128 |
+
scores[key_token] = dot_value(score) * scale_value();
|
| 129 |
+
}
|
| 130 |
+
}
|
| 131 |
+
workgroupBarrier();
|
| 132 |
+
var laneMax = {{ rowType }}(NEG_MAX);
|
| 133 |
+
for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
|
| 134 |
+
laneMax = max(laneMax, scores[key_token]);
|
| 135 |
+
}
|
| 136 |
+
let rowMax = reduce_max(laneMax, {% if useValueSubgroups %}sgLane, sgId, numSg{% else %}tid{% endif %});
|
| 137 |
+
var laneSum = {{ rowType }}(0.0);
|
| 138 |
+
for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
|
| 139 |
+
let probability = exp(scores[key_token] - rowMax);
|
| 140 |
+
scores[key_token] = probability;
|
| 141 |
+
laneSum = laneSum + probability;
|
| 142 |
+
}
|
| 143 |
+
let invDenom = {{ rowType }}(1.0) / reduce_sum(laneSum, {% if useValueSubgroups %}sgLane, sgId, numSg{% else %}tid{% endif %});
|
| 144 |
+
{% for d in range(headDim) %}
|
| 145 |
+
var value_acc_{{ d }} = {{ rowType }}(0.0);
|
| 146 |
+
{% endfor %}
|
| 147 |
+
for (var key_token = tid; key_token < KV_SEQ; key_token = key_token + WG) {
|
| 148 |
+
let probability = scores[key_token];
|
| 149 |
+
let vBase = kvBase + key_token * HIDDEN;
|
| 150 |
+
{% for d in range(headDim) %}
|
| 151 |
+
value_acc_{{ d }} = value_acc_{{ d }} + probability * f32(value[vBase + {{ d }}u]);
|
| 152 |
+
{% endfor %}
|
| 153 |
+
}
|
| 154 |
+
{% if useValueSubgroups %}
|
| 155 |
+
{% for d in range(headDim) %}
|
| 156 |
+
let value_sum_{{ d }} = subgroupAdd(value_acc_{{ d }});
|
| 157 |
+
{% endfor %}
|
| 158 |
+
workgroupBarrier();
|
| 159 |
+
if (sgLane == 0u) {
|
| 160 |
+
{% for d in range(headDim) %}
|
| 161 |
+
scores[sgId * HEAD_DIM + {{ d }}u] = value_sum_{{ d }};
|
| 162 |
+
{% endfor %}
|
| 163 |
+
}
|
| 164 |
+
workgroupBarrier();
|
| 165 |
+
if (tid < HEAD_DIM) {
|
| 166 |
+
var total = {{ rowType }}(0.0);
|
| 167 |
+
for (var i = 0u; i < numSg; i = i + 1u) { total = total + scores[i * HEAD_DIM + tid]; }
|
| 168 |
+
let result = total * invDenom;
|
| 169 |
+
{% else %}
|
| 170 |
+
workgroupBarrier();
|
| 171 |
+
{% for d in range(headDim) %}
|
| 172 |
+
scores[{{ d }}u * WG + tid] = value_acc_{{ d }};
|
| 173 |
+
{% endfor %}
|
| 174 |
+
workgroupBarrier();
|
| 175 |
+
for (var stride = WG / 2u; stride > 0u; stride = stride / 2u) {
|
| 176 |
+
if (tid < stride) {
|
| 177 |
+
{% for d in range(headDim) %}
|
| 178 |
+
scores[{{ d }}u * WG + tid] = scores[{{ d }}u * WG + tid] + scores[{{ d }}u * WG + tid + stride];
|
| 179 |
+
{% endfor %}
|
| 180 |
+
}
|
| 181 |
+
workgroupBarrier();
|
| 182 |
+
}
|
| 183 |
+
if (tid < HEAD_DIM) {
|
| 184 |
+
let result = scores[tid * WG] * invDenom;
|
| 185 |
+
{% endif %}
|
| 186 |
+
{% for row in range(queryBlock) %}
|
| 187 |
+
{% if queryTail and row > 0 %}
|
| 188 |
+
if (query_token + {{ row }}u < params.qSeq) {
|
| 189 |
+
{% endif %}
|
| 190 |
+
output[qBase_{{ row }} + tid] = {{ outputScalar }}(result{% if queryBlock > 1 %}[{{ row }}]{% endif %});
|
| 191 |
+
{% if queryTail and row > 0 %}
|
| 192 |
+
}
|
| 193 |
+
{% endif %}
|
| 194 |
+
{% endfor %}
|
| 195 |
+
}
|
| 196 |
+
workgroupBarrier();
|
| 197 |
+
}
|
| 198 |
+
}
|
build/webgpu/bench.json
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
build/webgpu/manifest.json
CHANGED
|
@@ -57,7 +57,8 @@
|
|
| 57 |
"canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32",
|
| 58 |
"pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
|
| 59 |
"wave32Effective": "wave32Adapter or pinSubgroupSize32",
|
| 60 |
-
"
|
|
|
|
| 61 |
"qkvDtypesOk": "tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT",
|
| 62 |
"floatDtypeOk": "(tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(tensorDtypes.queryT)",
|
| 63 |
"qkvShapeOk": "ranks.queryT == 3 and ranks.keyT == 3 and ranks.valueT == 3 and ranks.outputT == 3 and attrs.num_heads > 0 and dim(shapes.queryT, 2) % attrs.num_heads == 0 and dim(shapes.keyT, 2) == dim(shapes.queryT, 2) and dim(shapes.valueT, 2) == dim(shapes.queryT, 2) and dim(shapes.keyT, 1) == dim(shapes.valueT, 1) and dim(shapes.queryT, 0) == dim(shapes.keyT, 0) and dim(shapes.queryT, 0) == dim(shapes.valueT, 0) and dim(shapes.outputT, 0) == dim(shapes.queryT, 0) and dim(shapes.outputT, 1) == dim(shapes.queryT, 1) and dim(shapes.outputT, 2) == dim(shapes.valueT, 2)",
|
|
@@ -66,40 +67,42 @@
|
|
| 66 |
"qkvContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and noAttnBias",
|
| 67 |
"qkvMaskContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and attnBiasOk",
|
| 68 |
"biasOk": "present.biasT and ranks.biasT == 1 and tensorDtypes.biasT == tensorDtypes.queryT and dim(shapes.biasT, 0) == 3 * dim(shapes.queryT, 2)",
|
| 69 |
-
"
|
| 70 |
-
"
|
| 71 |
-
"
|
|
|
|
|
|
|
| 72 |
"subgroupCluster4": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 4 and device.adapterInfo.subgroupMinSize % 4 == 0 and device.adapterInfo.subgroupMaxSize % 4 == 0",
|
| 73 |
"subgroupCluster8": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 8 and device.adapterInfo.subgroupMinSize % 8 == 0 and device.adapterInfo.subgroupMaxSize % 8 == 0",
|
| 74 |
"attentionDispatchFits": "dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 1) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 75 |
-
"flashHeadOk": "
|
| 76 |
"flashSizeOk": "flashHeadOk and (dim(shapes.queryT, 1) * attrs.num_heads >= tunables.FLASH_MIN_QUERY_HEADS or (dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512) or (dim(shapes.queryT, 1) > 1 and dim(shapes.keyT, 1) >= 2048)) and attentionDispatchFits",
|
| 77 |
"flashShapeOk": "qkvContractOk and flashSizeOk",
|
| 78 |
"flashMaskShapeOk": "qkvMaskContractOk and flashSizeOk",
|
| 79 |
"noBiasSplitKCount": "min(tunables.DECODE_MAX_SPLITS if dim(shapes.queryT, 1) == 1 else max(1, ceilDiv(tunables.SPLITK_TARGET_WORKGROUPS, dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads)), ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
|
| 80 |
"biasSplitKCount": "min(tunables.DECODE_MAX_SPLITS, ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
|
| 81 |
-
"noBiasPartialOutBytes": "dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount *
|
| 82 |
"noBiasStatsBytes": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * 4",
|
| 83 |
"noBiasSplitScratchFits": "noBiasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and noBiasPartialOutBytes <= device.limits.maxBufferSize and noBiasStatsBytes <= device.limits.maxStorageBufferBindingSize and noBiasStatsBytes <= device.limits.maxBufferSize",
|
| 84 |
"noBiasSplitDispatchFits": "dim(shapes.queryT, 1) * noBiasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 85 |
-
"biasPartialOutBytes": "dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount *
|
| 86 |
"biasStatsBytes": "2 * dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * 4",
|
| 87 |
"biasSplitScratchFits": "biasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and biasPartialOutBytes <= device.limits.maxBufferSize and biasStatsBytes <= device.limits.maxStorageBufferBindingSize and biasStatsBytes <= device.limits.maxBufferSize",
|
| 88 |
"biasSplitDispatchFits": "biasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 89 |
"decodeSplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
|
| 90 |
"shortQuerySplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) >= 2 and dim(shapes.queryT, 1) <= 16 and dim(shapes.keyT, 1) >= 2048 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
|
| 91 |
"decodeSplitKBiasOk": "biasOk and qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and biasSplitDispatchFits and biasSplitScratchFits",
|
| 92 |
-
"decodeSplitKPortablePreferred": "tensorDtypes.queryT == \"float32\" and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and
|
| 93 |
"materializedScoreBytes": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1) * 4",
|
| 94 |
"materializedScoreFits": "materializedScoreBytes <= device.limits.maxStorageBufferBindingSize and materializedScoreBytes <= device.limits.maxBufferSize",
|
| 95 |
"materializedWorkgroupSize": "tunables.MATERIALIZED_WORKGROUP_DIM * tunables.MATERIALIZED_WORKGROUP_DIM",
|
| 96 |
"materializedScoreStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_KEY_TILE * (tunables.MATERIALIZED_INNER_TILE + 4)) * 4",
|
| 97 |
-
"materializedApplyTileN": "tunables.MATERIALIZED_VALUE_TILE_D128 if
|
| 98 |
"materializedApplyStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) * 4",
|
| 99 |
"materializedTileGeometryOk": "tunables.MATERIALIZED_INNER_TILE % 4 == 0 and (materializedApplyTileN / tunables.MATERIALIZED_WORKGROUP_DIM) % 4 == 0 and tunables.MATERIALIZED_INNER_TILE > 0 and tunables.MATERIALIZED_WORKGROUP_DIM > 0 and tunables.MATERIALIZED_QUERY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_KEY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and materializedApplyTileN % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN >= materializedWorkgroupSize and (tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) % materializedWorkgroupSize == 0",
|
| 100 |
"materializedDeviceOk": "materializedTileGeometryOk and materializedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeX and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeY and materializedScoreStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and materializedApplyStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 101 |
"materializedWideSimdOk": "device.features.has(\"subgroups\") or (has(device.adapterInfo, \"subgroupMinSize\") and device.adapterInfo.subgroupMinSize >= 16)",
|
| 102 |
-
"materializedF32CoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and
|
| 103 |
"materializedSoftmaxStorageBytes": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE * 8 + 8",
|
| 104 |
"materializedSoftmaxResourcesFit": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE > 0 and pow2ceil(tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE) == tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE and tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE <= deviceWorkgroupCap and materializedSoftmaxStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 105 |
"materializedSgmatQueryTile": "tunables.MATERIALIZED_SGMAT_QUERY_TILE",
|
|
@@ -123,9 +126,9 @@
|
|
| 123 |
"materializedSgmatResourcesFit": "materializedSgmatGeometryOk and materializedSgmatWorkgroupSize <= deviceWorkgroupCap and materializedSgmatCompactStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 124 |
"materializedSgmatDispatchFits": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 125 |
"materializedSgmatDirectScoreStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and dim(shapes.keyT, 1) % materializedSgmatKeyTile == 0",
|
| 126 |
-
"materializedSgmatDirectApplyStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and
|
| 127 |
"materializedSgmatRuntimeDirectStore": "dim(shapes.queryT, 1) >= 2 * materializedSgmatQueryTile and dim(shapes.keyT, 1) >= 2 * materializedSgmatKeyTile",
|
| 128 |
-
"materializedSgmatCoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and
|
| 129 |
"materializedCachedSoftmaxWg": "min(tunables.MATERIALIZED_CACHED_SOFTMAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
|
| 130 |
"materializedCachedSoftmaxVecsPerLane": "ceilDiv(ceilDiv(dim(shapes.keyT, 1), 4), max(1, materializedCachedSoftmaxWg))",
|
| 131 |
"materializedCachedSoftmaxStorageBytes": "materializedCachedSoftmaxWg * 8 + 8",
|
|
@@ -135,41 +138,47 @@
|
|
| 135 |
"materializedSgmatOk": "materializedSgmatCoreOk and materializedAdaptiveSoftmaxOk",
|
| 136 |
"materializedSgmatFusedOk": "materializedSgmatCoreOk",
|
| 137 |
"materializedF32Ok": "materializedF32CoreOk and materializedAdaptiveSoftmaxOk",
|
| 138 |
-
"clusterTileKWg64": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (
|
| 139 |
-
"clusterTileKWg128": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (
|
| 140 |
-
"
|
| 141 |
-
"
|
|
|
|
| 142 |
"prefillTiledDeviceOk": "tunables.PREFILL_QUERY_TILE <= deviceWorkgroupCap and prefillTiledStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 143 |
-
"portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, pow2ceil(max(1,
|
| 144 |
-
"portableWorkgroupStorageBytes": "portableWorkgroupSize *
|
| 145 |
"portableWorkgroupOk": "tunables.WORKGROUP_SIZE > 0 and pow2ceil(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and portableWorkgroupSize <= deviceWorkgroupCap and portableWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 146 |
"fallbackShapeOk": "qkvContractOk and portableWorkgroupOk",
|
| 147 |
"fallbackMaskShapeOk": "qkvMaskContractOk and portableWorkgroupOk",
|
| 148 |
-
"smallSeqShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and dim(shapes.queryT, 1) >= 1 and dim(shapes.queryT, 1) <= tunables.SMALL_SEQ_MAX and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_MAX and
|
| 149 |
-
"smallSeqPrivateFloats": "dim(shapes.keyT, 1) +
|
| 150 |
"smallSeqWorkgroupSize": "max(32, pow2ceil(dim(shapes.queryT, 1)))",
|
| 151 |
-
"smallSeqSharedBytes": "dim(shapes.keyT, 1) *
|
| 152 |
"smallSeqResourcesFit": "smallSeqPrivateFloats <= tunables.SMALL_SEQ_MAX_PRIVATE_FLOATS and smallSeqWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqSharedBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 153 |
-
"smallSeqBlockedKvBytes": "dim(shapes.keyT, 1) *
|
| 154 |
-
"smallSeqBlockedLaneBytes": "8 +
|
| 155 |
"smallSeqBlockedKeyLanes": "min(pow2ceil(dim(shapes.keyT, 1)), 16 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 16 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else (8 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 8 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else 4))",
|
| 156 |
"smallSeqBlockedWorkgroupSize": "tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * smallSeqBlockedKeyLanes",
|
| 157 |
"smallSeqBlockedSharedBytes": "smallSeqBlockedKvBytes + smallSeqBlockedWorkgroupSize * smallSeqBlockedLaneBytes",
|
| 158 |
-
"smallSeqBlockedShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and
|
| 159 |
"smallSeqBlockedFits": "smallSeqBlockedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqBlockedWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqBlockedSharedBytes <= device.limits.maxComputeWorkgroupStorageSize and ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 160 |
"smallSeqDispatchFits": "attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 161 |
-
"materializedSgmatCoreF16Ok": "qkvContractOk and tensorDtypes.queryT == \"float16\" and tensorDtypes.keyT == \"float16\" and tensorDtypes.valueT == \"float16\" and device.features.has(\"shader-f16\") and attrs.unidirectional == 0 and
|
| 162 |
"materializedSgmatFusedF16Ok": "materializedSgmatCoreF16Ok",
|
| 163 |
-
"attentionScaleExpression": "\"select(inverseSqrt(f32(HEAD_DIM)), params.scale, params.scale != 0.0)\""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 164 |
},
|
| 165 |
"bindings": {
|
| 166 |
-
"query": { "arg": "queryT", "
|
| 167 |
-
"key": { "arg": "keyT", "
|
| 168 |
-
"value": { "arg": "valueT", "
|
| 169 |
-
"bias": { "arg": "biasT", "
|
| 170 |
-
"output": { "arg": "outputT", "
|
| 171 |
"params": {
|
| 172 |
-
"buffer": "uniform",
|
| 173 |
"struct": [
|
| 174 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 175 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
|
@@ -177,77 +186,67 @@
|
|
| 177 |
{ "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" }
|
| 178 |
]
|
| 179 |
},
|
| 180 |
-
"
|
| 181 |
"name": "params",
|
| 182 |
-
"buffer": "uniform",
|
| 183 |
"struct": [
|
| 184 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 185 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 186 |
]
|
| 187 |
},
|
| 188 |
-
"q": { "arg": "queryT", "
|
| 189 |
-
"k": { "arg": "keyT", "
|
| 190 |
-
"v": { "arg": "valueT", "
|
| 191 |
-
"y": { "arg": "outputT", "
|
| 192 |
-
"
|
| 193 |
"name": "params",
|
| 194 |
-
"buffer": "uniform",
|
| 195 |
"struct": [
|
| 196 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 197 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
| 198 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 199 |
]
|
| 200 |
},
|
| 201 |
-
"
|
| 202 |
-
"
|
| 203 |
-
"scores": { "scratch": "materializedScores", "
|
| 204 |
-
"scorePartials": { "scratch": "materializedScorePartials", "
|
| 205 |
-
"
|
| 206 |
-
"
|
| 207 |
"scratch": "materializedScorePartials",
|
| 208 |
"name": "scorePartials",
|
| 209 |
"buffer": "read-only-storage",
|
| 210 |
"elementType": "f32"
|
| 211 |
},
|
| 212 |
-
"rowStats": { "scratch": "materializedRowStats", "
|
| 213 |
-
"
|
| 214 |
"name": "params",
|
| 215 |
-
"buffer": "uniform",
|
| 216 |
"struct": [
|
| 217 |
{ "name": "rows", "type": "u32", "value": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)" }
|
| 218 |
]
|
| 219 |
},
|
| 220 |
-
"
|
| 221 |
"scratch": "materializedScores",
|
| 222 |
"name": "scores",
|
| 223 |
"buffer": "read-only-storage",
|
| 224 |
"elementType": "f32"
|
| 225 |
},
|
| 226 |
-
"
|
| 227 |
-
"
|
| 228 |
-
"
|
| 229 |
"scratch": "materializedRowStats",
|
| 230 |
"name": "rowStats",
|
| 231 |
"buffer": "read-only-storage",
|
| 232 |
"elementType": "f32"
|
| 233 |
},
|
| 234 |
-
"
|
| 235 |
"name": "params",
|
| 236 |
-
"buffer": "uniform",
|
| 237 |
"struct": [
|
| 238 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 239 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }
|
| 240 |
]
|
| 241 |
},
|
| 242 |
-
"
|
| 243 |
-
|
| 244 |
-
"name": "attn_mask",
|
| 245 |
-
"buffer": "read-only-storage",
|
| 246 |
-
"elementType": "$maskElement"
|
| 247 |
-
},
|
| 248 |
-
"params_8": {
|
| 249 |
"name": "params",
|
| 250 |
-
"buffer": "uniform",
|
| 251 |
"struct": [
|
| 252 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 253 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
|
@@ -266,53 +265,109 @@
|
|
| 266 |
{ "name": "maskSeqStride", "type": "u32", "value": "dim(shapes.attentionBiasT, 3)" }
|
| 267 |
]
|
| 268 |
},
|
| 269 |
-
"
|
| 270 |
-
"
|
| 271 |
-
"
|
| 272 |
-
"partial_out": { "scratch": "partialOut", "
|
| 273 |
-
"partial_stats": { "scratch": "partialStats", "
|
| 274 |
-
"
|
| 275 |
"name": "params",
|
| 276 |
-
"buffer": "uniform",
|
| 277 |
"struct": [
|
| 278 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
| 279 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 280 |
]
|
| 281 |
},
|
| 282 |
-
"
|
| 283 |
"scratch": "partialOut",
|
| 284 |
"name": "partial_out",
|
| 285 |
"buffer": "read-only-storage",
|
| 286 |
"elementType": "vec4<f32>"
|
| 287 |
},
|
| 288 |
-
"
|
| 289 |
"scratch": "partialStats",
|
| 290 |
"name": "partial_stats",
|
| 291 |
"buffer": "read-only-storage",
|
| 292 |
"elementType": "vec2<f32>"
|
| 293 |
},
|
| 294 |
-
"
|
| 295 |
-
"
|
| 296 |
-
"scratch": "materializedScores",
|
| 297 |
-
"name": "scores",
|
| 298 |
-
"buffer": "storage",
|
| 299 |
-
"elementType": "$softmaxElementType"
|
| 300 |
-
}
|
| 301 |
},
|
| 302 |
"variants": [
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 303 |
{
|
| 304 |
"id": "qkv_bias_small_seq_blocked",
|
| 305 |
"priority": 40,
|
| 306 |
"when": ["biasOk", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
|
| 307 |
"derive": {
|
| 308 |
-
"usesF16": false,
|
| 309 |
-
"scalar": "\"f32\"",
|
| 310 |
"inputElement": "\"vec4<f32>\"",
|
| 311 |
"outputElement": "\"vec4<f32>\"",
|
| 312 |
"inputScalar": "\"f32\"",
|
| 313 |
-
"
|
| 314 |
-
"headDim": "headDimPlan",
|
| 315 |
-
"headDimV4": "headDimPlan / 4",
|
| 316 |
"hidden": "dim(shapes.queryT, 2)",
|
| 317 |
"hiddenV4": "dim(shapes.queryT, 2) / 4",
|
| 318 |
"kvSeq": "dim(shapes.keyT, 1)",
|
|
@@ -329,12 +384,6 @@
|
|
| 329 |
"x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
|
| 330 |
"y": "attrs.num_heads",
|
| 331 |
"z": "dim(shapes.queryT, 0)"
|
| 332 |
-
},
|
| 333 |
-
"profile": {
|
| 334 |
-
"op": "\"multi_head_attention\"",
|
| 335 |
-
"variant": "\"qkv_bias_small_seq_blocked\"",
|
| 336 |
-
"numHeads": "attrs.num_heads",
|
| 337 |
-
"headDim": "headDim"
|
| 338 |
}
|
| 339 |
}
|
| 340 |
]
|
|
@@ -344,14 +393,9 @@
|
|
| 344 |
"priority": 40,
|
| 345 |
"when": ["not present.biasT", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
|
| 346 |
"derive": {
|
| 347 |
-
"usesF16": false,
|
| 348 |
-
"scalar": "\"f32\"",
|
| 349 |
"inputElement": "\"vec4<f32>\"",
|
| 350 |
"outputElement": "\"vec4<f32>\"",
|
| 351 |
-
"
|
| 352 |
-
"hasBias": false,
|
| 353 |
-
"headDim": "headDimPlan",
|
| 354 |
-
"headDimV4": "headDimPlan / 4",
|
| 355 |
"hidden": "dim(shapes.queryT, 2)",
|
| 356 |
"hiddenV4": "dim(shapes.queryT, 2) / 4",
|
| 357 |
"kvSeq": "dim(shapes.keyT, 1)",
|
|
@@ -368,12 +412,6 @@
|
|
| 368 |
"x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
|
| 369 |
"y": "attrs.num_heads",
|
| 370 |
"z": "dim(shapes.queryT, 0)"
|
| 371 |
-
},
|
| 372 |
-
"profile": {
|
| 373 |
-
"op": "\"multi_head_attention\"",
|
| 374 |
-
"variant": "\"qkv_no_bias_small_seq_blocked\"",
|
| 375 |
-
"numHeads": "attrs.num_heads",
|
| 376 |
-
"headDim": "headDim"
|
| 377 |
}
|
| 378 |
}
|
| 379 |
]
|
|
@@ -383,7 +421,6 @@
|
|
| 383 |
"priority": 60,
|
| 384 |
"when": ["not present.biasT", "smallSeqShapeOk", "not flashShapeOk", "smallSeqResourcesFit", "smallSeqDispatchFits"],
|
| 385 |
"derive": {
|
| 386 |
-
"usesF16": false,
|
| 387 |
"inputElement": "\"f32\"",
|
| 388 |
"outputElement": "\"f32\"",
|
| 389 |
"headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
|
|
@@ -396,25 +433,18 @@
|
|
| 396 |
"id": "main",
|
| 397 |
"name": "MultiHeadAttention",
|
| 398 |
"shader": "mha-small-seq.wgsl.jinja",
|
| 399 |
-
"bindings": ["query", "key", "value", "output", "
|
| 400 |
-
"dispatch": { "x": "attrs.num_heads", "y": "dim(shapes.queryT, 0)" }
|
| 401 |
-
"profile": {
|
| 402 |
-
"op": "\"multi_head_attention\"",
|
| 403 |
-
"variant": "\"qkv_no_bias_small_seq\"",
|
| 404 |
-
"numHeads": "attrs.num_heads",
|
| 405 |
-
"headDim": "dim(shapes.queryT, 2) / attrs.num_heads"
|
| 406 |
-
}
|
| 407 |
}
|
| 408 |
]
|
| 409 |
},
|
| 410 |
{
|
| 411 |
"id": "qkv_no_bias_tiled_nosg",
|
| 412 |
"priority": 19,
|
| 413 |
-
"when": ["not present.biasT", "qkvContractOk", "not smallHeadParallelOk", "
|
| 414 |
"supersededBy": ["qkv_no_bias_flash_cluster_nosg", "qkv_no_bias_flash_cluster_lpq4_nosg"],
|
| 415 |
"derive": {
|
| 416 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 417 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 418 |
"blockM": "tunables.PREFILL_QUERY_TILE",
|
| 419 |
"vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
|
| 420 |
},
|
|
@@ -449,8 +479,8 @@
|
|
| 449 |
}
|
| 450 |
],
|
| 451 |
"dispatch": {
|
| 452 |
-
"x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) /
|
| 453 |
-
"y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) /
|
| 454 |
"z": 1
|
| 455 |
}
|
| 456 |
}
|
|
@@ -459,14 +489,12 @@
|
|
| 459 |
{
|
| 460 |
"id": "qkv_bias_flash_q32_broadcast_f32_d128",
|
| 461 |
"priority": 30,
|
| 462 |
-
"when": ["tensorDtypes.queryT == \"float32\"", "biasOk", "attrs.unidirectional == 0", "flashShapeOk", "q32BroadcastF32RegisterGeometry", "dim(shapes.queryT, 1) >= 31", "ceilDiv(dim(shapes.queryT, 1), 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "
|
| 463 |
"requires": { "features": ["subgroups"] },
|
| 464 |
"derive": {
|
| 465 |
-
"hasBias": true,
|
| 466 |
"hasCausal": false,
|
| 467 |
"usesF16": false,
|
| 468 |
"scalar": "\"f32\"",
|
| 469 |
-
"inputVec4": "\"vec4<f32>\"",
|
| 470 |
"inputElement": "\"vec4<f32>\"",
|
| 471 |
"outputElement": "\"vec4<f32>\"",
|
| 472 |
"inputScalar": "\"f32\"",
|
|
@@ -484,13 +512,12 @@
|
|
| 484 |
"name": "MultiHeadAttention.FlashQ32BroadcastF32Bias",
|
| 485 |
"shader": "attn-flash-q32-broadcast.wgsl.jinja",
|
| 486 |
"derive": { "layout": "\"bsh\"" },
|
| 487 |
-
"bindings": ["query", "key", "value", "bias", "output", "
|
| 488 |
"dispatch": {
|
| 489 |
"x": "ceilDiv(dim(shapes.queryT, 1), 32)",
|
| 490 |
"y": "attrs.num_heads",
|
| 491 |
"z": "dim(shapes.queryT, 0)"
|
| 492 |
-
}
|
| 493 |
-
"subgroupCollectivesWidth": 32
|
| 494 |
}
|
| 495 |
]
|
| 496 |
},
|
|
@@ -499,8 +526,6 @@
|
|
| 499 |
"priority": 10,
|
| 500 |
"when": ["not present.biasT", "smallHeadParallelOk"],
|
| 501 |
"derive": {
|
| 502 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 503 |
-
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 504 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 505 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 506 |
"outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -511,9 +536,9 @@
|
|
| 511 |
"passes": [
|
| 512 |
{
|
| 513 |
"id": "main",
|
| 514 |
-
"name": "
|
| 515 |
"shader": "attn-small-head-parallel.wgsl.jinja",
|
| 516 |
-
"bindings": ["query", "key", "value", "output", "
|
| 517 |
"dispatch": {
|
| 518 |
"x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 519 |
"y": "attrs.num_heads",
|
|
@@ -525,12 +550,9 @@
|
|
| 525 |
{
|
| 526 |
"id": "qkv_no_bias_tiled_attn_bias_nosg",
|
| 527 |
"priority": 17,
|
| 528 |
-
"when": ["not present.biasT", "qkvMaskContractOk", "not smallHeadParallelOk", "
|
| 529 |
"derive": {
|
| 530 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 531 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 532 |
-
"hasMask": true,
|
| 533 |
-
"maskIsBool": false,
|
| 534 |
"blockM": "tunables.PREFILL_QUERY_TILE",
|
| 535 |
"vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
|
| 536 |
},
|
|
@@ -577,8 +599,8 @@
|
|
| 577 |
}
|
| 578 |
],
|
| 579 |
"dispatch": {
|
| 580 |
-
"x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) /
|
| 581 |
-
"y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) /
|
| 582 |
"z": 1
|
| 583 |
}
|
| 584 |
}
|
|
@@ -594,10 +616,7 @@
|
|
| 594 |
},
|
| 595 |
"derive": {
|
| 596 |
"qNumHeads": "attrs.num_heads",
|
| 597 |
-
"headDim": "headDimPlan",
|
| 598 |
"qHidden": "dim(shapes.queryT, 2)",
|
| 599 |
-
"hasBias": false,
|
| 600 |
-
"useSubgroups": true,
|
| 601 |
"statSlots": "materializedSgmatStatSlots",
|
| 602 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 603 |
},
|
|
@@ -616,19 +635,18 @@
|
|
| 616 |
"name": "MultiHeadAttention.MaterializedScoresSgmat",
|
| 617 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 618 |
"derive": { "phase": "\"score\"", "emitRowStats": true },
|
| 619 |
-
"bindings": ["
|
| 620 |
"dispatch": {
|
| 621 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 622 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
| 623 |
"z": "dim(shapes.queryT, 0) * attrs.num_heads"
|
| 624 |
-
}
|
| 625 |
-
"subgroupCollectivesWidth": 32
|
| 626 |
},
|
| 627 |
{
|
| 628 |
"id": "rowstats",
|
| 629 |
"name": "MultiHeadAttention.MaterializedRowStatsCombine",
|
| 630 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 631 |
-
"bindings": ["
|
| 632 |
"dispatch": {
|
| 633 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 634 |
"y": 1,
|
|
@@ -640,7 +658,7 @@
|
|
| 640 |
"name": "MultiHeadAttention.MaterializedApplySgmat",
|
| 641 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 642 |
"derive": { "phase": "\"apply\"", "fusedSoftmax": true },
|
| 643 |
-
"bindings": ["
|
| 644 |
"dispatch": {
|
| 645 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 646 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
@@ -659,10 +677,7 @@
|
|
| 659 |
},
|
| 660 |
"derive": {
|
| 661 |
"qNumHeads": "attrs.num_heads",
|
| 662 |
-
"headDim": "headDimPlan",
|
| 663 |
"qHidden": "dim(shapes.queryT, 2)",
|
| 664 |
-
"hasBias": true,
|
| 665 |
-
"useSubgroups": true,
|
| 666 |
"statSlots": "materializedSgmatStatSlots",
|
| 667 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 668 |
},
|
|
@@ -681,19 +696,18 @@
|
|
| 681 |
"name": "MultiHeadAttention.MaterializedScoresSgmatBias",
|
| 682 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 683 |
"derive": { "phase": "\"score\"", "emitRowStats": true },
|
| 684 |
-
"bindings": ["
|
| 685 |
"dispatch": {
|
| 686 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 687 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
| 688 |
"z": "dim(shapes.queryT, 0) * attrs.num_heads"
|
| 689 |
-
}
|
| 690 |
-
"subgroupCollectivesWidth": 32
|
| 691 |
},
|
| 692 |
{
|
| 693 |
"id": "rowstats",
|
| 694 |
"name": "MultiHeadAttention.MaterializedRowStatsCombineBias",
|
| 695 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 696 |
-
"bindings": ["
|
| 697 |
"dispatch": {
|
| 698 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 699 |
"y": 1,
|
|
@@ -705,7 +719,7 @@
|
|
| 705 |
"name": "MultiHeadAttention.MaterializedApplySgmatBias",
|
| 706 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 707 |
"derive": { "phase": "\"apply\"", "fusedSoftmax": true },
|
| 708 |
-
"bindings": ["
|
| 709 |
"dispatch": {
|
| 710 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 711 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
@@ -724,10 +738,7 @@
|
|
| 724 |
},
|
| 725 |
"derive": {
|
| 726 |
"qNumHeads": "attrs.num_heads",
|
| 727 |
-
"headDim": "headDimPlan",
|
| 728 |
"qHidden": "dim(shapes.queryT, 2)",
|
| 729 |
-
"hasBias": false,
|
| 730 |
-
"useSubgroups": true,
|
| 731 |
"statSlots": "materializedSgmatStatSlots",
|
| 732 |
"statQuerySeq": "dim(shapes.queryT, 1)",
|
| 733 |
"operandF16": true,
|
|
@@ -749,19 +760,18 @@
|
|
| 749 |
"name": "MultiHeadAttention.MaterializedScoresSgmatF16",
|
| 750 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 751 |
"derive": { "phase": "\"score\"", "emitRowStats": true },
|
| 752 |
-
"bindings": ["query", "key", "scores", "scorePartials", "
|
| 753 |
"dispatch": {
|
| 754 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 755 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
| 756 |
"z": "dim(shapes.queryT, 0) * attrs.num_heads"
|
| 757 |
-
}
|
| 758 |
-
"subgroupCollectivesWidth": 32
|
| 759 |
},
|
| 760 |
{
|
| 761 |
"id": "rowstats",
|
| 762 |
"name": "MultiHeadAttention.MaterializedRowStatsCombine",
|
| 763 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 764 |
-
"bindings": ["
|
| 765 |
"dispatch": {
|
| 766 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 767 |
"y": 1,
|
|
@@ -773,7 +783,7 @@
|
|
| 773 |
"name": "MultiHeadAttention.MaterializedApplySgmatF16",
|
| 774 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 775 |
"derive": { "phase": "\"apply\"", "fusedSoftmax": true },
|
| 776 |
-
"bindings": ["
|
| 777 |
"dispatch": {
|
| 778 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 779 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
@@ -785,15 +795,12 @@
|
|
| 785 |
{
|
| 786 |
"id": "qkv_no_bias_flash_cluster_lpq4_nosg",
|
| 787 |
"priority": 20,
|
| 788 |
-
"when": ["not present.biasT", "flashShapeOk", "
|
| 789 |
"requires": {},
|
| 790 |
"derive": {
|
| 791 |
-
"hasBias": false,
|
| 792 |
"hasCausal": true,
|
| 793 |
-
"hasWindow": false,
|
| 794 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 795 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 796 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 797 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 798 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 799 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -826,15 +833,12 @@
|
|
| 826 |
{
|
| 827 |
"id": "qkv_no_bias_flash_cluster_nosg",
|
| 828 |
"priority": 20,
|
| 829 |
-
"when": ["not present.biasT", "flashShapeOk", "
|
| 830 |
"requires": {},
|
| 831 |
"derive": {
|
| 832 |
-
"hasBias": false,
|
| 833 |
"hasCausal": true,
|
| 834 |
-
"hasWindow": false,
|
| 835 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 836 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 837 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 838 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 839 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 840 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -867,15 +871,12 @@
|
|
| 867 |
{
|
| 868 |
"id": "qkv_bias_flash_cluster_nosg",
|
| 869 |
"priority": 19,
|
| 870 |
-
"when": ["biasOk", "flashShapeOk", "
|
| 871 |
"requires": {},
|
| 872 |
"derive": {
|
| 873 |
-
"hasBias": true,
|
| 874 |
"hasCausal": true,
|
| 875 |
-
"hasWindow": false,
|
| 876 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 877 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 878 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 879 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 880 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 881 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -910,15 +911,12 @@
|
|
| 910 |
{
|
| 911 |
"id": "qkv_no_bias_flash_cluster_lpq4",
|
| 912 |
"priority": 22,
|
| 913 |
-
"when": ["not present.biasT", "flashShapeOk", "
|
| 914 |
"requires": { "features": ["subgroups"] },
|
| 915 |
"derive": {
|
| 916 |
-
"hasBias": false,
|
| 917 |
"hasCausal": true,
|
| 918 |
-
"hasWindow": false,
|
| 919 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 920 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 921 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 922 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 923 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 924 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -942,23 +940,19 @@
|
|
| 942 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 943 |
"y": "attrs.num_heads",
|
| 944 |
"z": "dim(shapes.queryT, 0)"
|
| 945 |
-
}
|
| 946 |
-
"subgroupCollectivesWidth": "portable"
|
| 947 |
}
|
| 948 |
]
|
| 949 |
},
|
| 950 |
{
|
| 951 |
"id": "qkv_no_bias_flash_cluster",
|
| 952 |
"priority": 22,
|
| 953 |
-
"when": ["not present.biasT", "flashShapeOk", "
|
| 954 |
"requires": { "features": ["subgroups"] },
|
| 955 |
"derive": {
|
| 956 |
-
"hasBias": false,
|
| 957 |
"hasCausal": true,
|
| 958 |
-
"hasWindow": false,
|
| 959 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 960 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 961 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 962 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 963 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 964 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -982,23 +976,19 @@
|
|
| 982 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 983 |
"y": "attrs.num_heads",
|
| 984 |
"z": "dim(shapes.queryT, 0)"
|
| 985 |
-
}
|
| 986 |
-
"subgroupCollectivesWidth": "portable"
|
| 987 |
}
|
| 988 |
]
|
| 989 |
},
|
| 990 |
{
|
| 991 |
"id": "qkv_bias_flash_cluster",
|
| 992 |
"priority": 21,
|
| 993 |
-
"when": ["biasOk", "flashShapeOk", "
|
| 994 |
"requires": { "features": ["subgroups"] },
|
| 995 |
"derive": {
|
| 996 |
-
"hasBias": true,
|
| 997 |
"hasCausal": true,
|
| 998 |
-
"hasWindow": false,
|
| 999 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1000 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1001 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1002 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1003 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1004 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -1024,25 +1014,19 @@
|
|
| 1024 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 1025 |
"y": "attrs.num_heads",
|
| 1026 |
"z": "dim(shapes.queryT, 0)"
|
| 1027 |
-
}
|
| 1028 |
-
"subgroupCollectivesWidth": "portable"
|
| 1029 |
}
|
| 1030 |
]
|
| 1031 |
},
|
| 1032 |
{
|
| 1033 |
"id": "qkv_no_bias_flash_cluster_attn_bias",
|
| 1034 |
"priority": 22,
|
| 1035 |
-
"when": ["not present.biasT", "flashMaskShapeOk", "
|
| 1036 |
"requires": { "features": ["subgroups"] },
|
| 1037 |
"derive": {
|
| 1038 |
-
"hasBias": false,
|
| 1039 |
"hasCausal": true,
|
| 1040 |
-
"hasWindow": false,
|
| 1041 |
-
"hasMask": true,
|
| 1042 |
-
"maskIsBool": false,
|
| 1043 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1044 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1045 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1046 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1047 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1048 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1062,30 +1046,24 @@
|
|
| 1062 |
"name": "MultiHeadAttention.Flash",
|
| 1063 |
"shader": "attn-flash-prefill-cluster.wgsl.jinja",
|
| 1064 |
"derive": { "layout": "\"bsh\"" },
|
| 1065 |
-
"bindings": ["query", "key", "value", "
|
| 1066 |
"dispatch": {
|
| 1067 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 1068 |
"y": "attrs.num_heads",
|
| 1069 |
"z": "dim(shapes.queryT, 0)"
|
| 1070 |
-
}
|
| 1071 |
-
"subgroupCollectivesWidth": "portable"
|
| 1072 |
}
|
| 1073 |
]
|
| 1074 |
},
|
| 1075 |
{
|
| 1076 |
"id": "qkv_bias_flash_cluster_attn_bias",
|
| 1077 |
"priority": 21,
|
| 1078 |
-
"when": ["biasOk", "flashMaskShapeOk", "
|
| 1079 |
"requires": { "features": ["subgroups"] },
|
| 1080 |
"derive": {
|
| 1081 |
-
"hasBias": true,
|
| 1082 |
"hasCausal": true,
|
| 1083 |
-
"hasWindow": false,
|
| 1084 |
-
"hasMask": true,
|
| 1085 |
-
"maskIsBool": false,
|
| 1086 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1087 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1088 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1089 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1090 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1091 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1107,13 +1085,12 @@
|
|
| 1107 |
"name": "MultiHeadAttention.Flash",
|
| 1108 |
"shader": "attn-flash-prefill-cluster.wgsl.jinja",
|
| 1109 |
"derive": { "layout": "\"bsh\"" },
|
| 1110 |
-
"bindings": ["query", "key", "value", "
|
| 1111 |
"dispatch": {
|
| 1112 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 1113 |
"y": "attrs.num_heads",
|
| 1114 |
"z": "dim(shapes.queryT, 0)"
|
| 1115 |
-
}
|
| 1116 |
-
"subgroupCollectivesWidth": "portable"
|
| 1117 |
}
|
| 1118 |
]
|
| 1119 |
},
|
|
@@ -1123,9 +1100,7 @@
|
|
| 1123 |
"when": ["not present.biasT", "decodeSplitKNoBiasOk or shortQuerySplitKNoBiasOk"],
|
| 1124 |
"requires": {},
|
| 1125 |
"derive": {
|
| 1126 |
-
"combineSubgroups": false,
|
| 1127 |
"useSubgroups": false,
|
| 1128 |
-
"hasWindow": false,
|
| 1129 |
"splitQueries": true,
|
| 1130 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1131 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
|
@@ -1157,7 +1132,7 @@
|
|
| 1157 |
"name": "MultiHeadAttention.DecodeSplitKNoSg",
|
| 1158 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1159 |
"derive": { "layout": "\"bsh\"" },
|
| 1160 |
-
"bindings": ["
|
| 1161 |
"dispatch": {
|
| 1162 |
"x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
|
| 1163 |
"y": "attrs.num_heads",
|
|
@@ -1169,7 +1144,7 @@
|
|
| 1169 |
"name": "MultiHeadAttention.DecodeSplitKMergeNoSg",
|
| 1170 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1171 |
"derive": { "layout": "\"bsh\"" },
|
| 1172 |
-
"bindings": ["
|
| 1173 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1174 |
}
|
| 1175 |
]
|
|
@@ -1180,10 +1155,7 @@
|
|
| 1180 |
"when": ["biasOk", "decodeSplitKBiasOk"],
|
| 1181 |
"requires": {},
|
| 1182 |
"derive": {
|
| 1183 |
-
"combineSubgroups": false,
|
| 1184 |
"useSubgroups": false,
|
| 1185 |
-
"hasBias": true,
|
| 1186 |
-
"hasWindow": false,
|
| 1187 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1188 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1189 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1215,7 +1187,7 @@
|
|
| 1215 |
"name": "MultiHeadAttention.DecodeSplitKBiasNoSg",
|
| 1216 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1217 |
"derive": { "layout": "\"bsh\"" },
|
| 1218 |
-
"bindings": ["
|
| 1219 |
"dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1220 |
},
|
| 1221 |
{
|
|
@@ -1223,7 +1195,7 @@
|
|
| 1223 |
"name": "MultiHeadAttention.DecodeSplitKMergeBiasNoSg",
|
| 1224 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1225 |
"derive": { "layout": "\"bsh\"" },
|
| 1226 |
-
"bindings": ["
|
| 1227 |
"dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1228 |
}
|
| 1229 |
]
|
|
@@ -1235,8 +1207,6 @@
|
|
| 1235 |
"demoteWhen": ["decodeSplitKPortablePreferred"],
|
| 1236 |
"requires": { "features": ["subgroups"] },
|
| 1237 |
"derive": {
|
| 1238 |
-
"combineSubgroups": true,
|
| 1239 |
-
"hasWindow": false,
|
| 1240 |
"splitQueries": true,
|
| 1241 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1242 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
|
@@ -1268,20 +1238,19 @@
|
|
| 1268 |
"name": "MultiHeadAttention.DecodeSplitK",
|
| 1269 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1270 |
"derive": { "layout": "\"bsh\"" },
|
| 1271 |
-
"bindings": ["
|
| 1272 |
"dispatch": {
|
| 1273 |
"x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
|
| 1274 |
"y": "attrs.num_heads",
|
| 1275 |
"z": "dim(shapes.queryT, 0)"
|
| 1276 |
-
}
|
| 1277 |
-
"subgroupCollectivesWidth": "portable"
|
| 1278 |
},
|
| 1279 |
{
|
| 1280 |
"id": "merge",
|
| 1281 |
"name": "MultiHeadAttention.DecodeSplitKMerge",
|
| 1282 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1283 |
"derive": { "layout": "\"bsh\"" },
|
| 1284 |
-
"bindings": ["
|
| 1285 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1286 |
}
|
| 1287 |
]
|
|
@@ -1292,9 +1261,6 @@
|
|
| 1292 |
"when": ["biasOk", "decodeSplitKBiasOk"],
|
| 1293 |
"requires": { "features": ["subgroups"] },
|
| 1294 |
"derive": {
|
| 1295 |
-
"combineSubgroups": true,
|
| 1296 |
-
"hasBias": true,
|
| 1297 |
-
"hasWindow": false,
|
| 1298 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1299 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1300 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1326,16 +1292,15 @@
|
|
| 1326 |
"name": "MultiHeadAttention.DecodeSplitKBias",
|
| 1327 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1328 |
"derive": { "layout": "\"bsh\"" },
|
| 1329 |
-
"bindings": ["
|
| 1330 |
-
"dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1331 |
-
"subgroupCollectivesWidth": "portable"
|
| 1332 |
},
|
| 1333 |
{
|
| 1334 |
"id": "merge",
|
| 1335 |
"name": "MultiHeadAttention.DecodeSplitKMergeBias",
|
| 1336 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1337 |
"derive": { "layout": "\"bsh\"" },
|
| 1338 |
-
"bindings": ["
|
| 1339 |
"dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1340 |
}
|
| 1341 |
]
|
|
@@ -1348,7 +1313,6 @@
|
|
| 1348 |
"derive": {
|
| 1349 |
"qNumHeads": "attrs.num_heads",
|
| 1350 |
"kvNumHeads": "attrs.num_heads",
|
| 1351 |
-
"headDim": "headDimPlan",
|
| 1352 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1353 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1354 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
|
@@ -1358,7 +1322,6 @@
|
|
| 1358 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 1359 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
| 1360 |
"applyTileN": "materializedApplyTileN",
|
| 1361 |
-
"hasBias": false,
|
| 1362 |
"useSubgroups": "device.features.has(\"subgroups\")"
|
| 1363 |
},
|
| 1364 |
"intermediates": [
|
|
@@ -1374,7 +1337,7 @@
|
|
| 1374 |
"name": "MultiHeadAttention.MaterializedScores",
|
| 1375 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1376 |
"derive": { "layout": "\"bsh\"" },
|
| 1377 |
-
"bindings": ["
|
| 1378 |
"dispatch": {
|
| 1379 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1380 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1386,20 +1349,19 @@
|
|
| 1386 |
"name": "MultiHeadAttention.MaterializedSoftmax",
|
| 1387 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 1388 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 1389 |
-
"bindings": ["
|
| 1390 |
"dispatch": {
|
| 1391 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1392 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1393 |
"z": 1
|
| 1394 |
-
}
|
| 1395 |
-
"subgroupCollectivesWidth": "portable"
|
| 1396 |
},
|
| 1397 |
{
|
| 1398 |
"id": "apply",
|
| 1399 |
"name": "MultiHeadAttention.MaterializedApply",
|
| 1400 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1401 |
"derive": { "layout": "\"bsh\"" },
|
| 1402 |
-
"bindings": ["
|
| 1403 |
"dispatch": {
|
| 1404 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1405 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1415,7 +1377,6 @@
|
|
| 1415 |
"derive": {
|
| 1416 |
"qNumHeads": "attrs.num_heads",
|
| 1417 |
"kvNumHeads": "attrs.num_heads",
|
| 1418 |
-
"headDim": "headDimPlan",
|
| 1419 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1420 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1421 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
|
@@ -1425,7 +1386,6 @@
|
|
| 1425 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 1426 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
| 1427 |
"applyTileN": "materializedApplyTileN",
|
| 1428 |
-
"hasBias": true,
|
| 1429 |
"useSubgroups": "device.features.has(\"subgroups\")"
|
| 1430 |
},
|
| 1431 |
"intermediates": [
|
|
@@ -1441,7 +1401,7 @@
|
|
| 1441 |
"name": "MultiHeadAttention.MaterializedScoresBias",
|
| 1442 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1443 |
"derive": { "layout": "\"bsh\"" },
|
| 1444 |
-
"bindings": ["
|
| 1445 |
"dispatch": {
|
| 1446 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1447 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1453,20 +1413,19 @@
|
|
| 1453 |
"name": "MultiHeadAttention.MaterializedSoftmaxBias",
|
| 1454 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 1455 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 1456 |
-
"bindings": ["
|
| 1457 |
"dispatch": {
|
| 1458 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1459 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1460 |
"z": 1
|
| 1461 |
-
}
|
| 1462 |
-
"subgroupCollectivesWidth": "portable"
|
| 1463 |
},
|
| 1464 |
{
|
| 1465 |
"id": "apply",
|
| 1466 |
"name": "MultiHeadAttention.MaterializedApplyBias",
|
| 1467 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1468 |
"derive": { "layout": "\"bsh\"" },
|
| 1469 |
-
"bindings": ["
|
| 1470 |
"dispatch": {
|
| 1471 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1472 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1483,14 +1442,11 @@
|
|
| 1483 |
"derive": {
|
| 1484 |
"qNumHeads": "attrs.num_heads",
|
| 1485 |
"kvNumHeads": "attrs.num_heads",
|
| 1486 |
-
"headDim": "headDimPlan",
|
| 1487 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1488 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1489 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
| 1490 |
"materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
|
| 1491 |
"applyTileN": "materializedApplyTileN",
|
| 1492 |
-
"hasBias": false,
|
| 1493 |
-
"useSubgroups": "device.features.has(\"subgroups\")",
|
| 1494 |
"statSlots": "materializedGemmStatSlots",
|
| 1495 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 1496 |
},
|
|
@@ -1509,7 +1465,7 @@
|
|
| 1509 |
"name": "MultiHeadAttention.MaterializedScoresFused",
|
| 1510 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1511 |
"derive": { "layout": "\"bsh\"", "emitRowStats": true },
|
| 1512 |
-
"bindings": ["
|
| 1513 |
"dispatch": {
|
| 1514 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1515 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1521,7 +1477,7 @@
|
|
| 1521 |
"name": "MultiHeadAttention.MaterializedGemmRowStatsCombine",
|
| 1522 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 1523 |
"derive": { "maxOnly": true },
|
| 1524 |
-
"bindings": ["
|
| 1525 |
"dispatch": {
|
| 1526 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1527 |
"y": 1,
|
|
@@ -1533,7 +1489,7 @@
|
|
| 1533 |
"name": "MultiHeadAttention.MaterializedApplyFused",
|
| 1534 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1535 |
"derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
|
| 1536 |
-
"bindings": ["
|
| 1537 |
"dispatch": {
|
| 1538 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1539 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1549,14 +1505,11 @@
|
|
| 1549 |
"derive": {
|
| 1550 |
"qNumHeads": "attrs.num_heads",
|
| 1551 |
"kvNumHeads": "attrs.num_heads",
|
| 1552 |
-
"headDim": "headDimPlan",
|
| 1553 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1554 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1555 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
| 1556 |
"materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
|
| 1557 |
"applyTileN": "materializedApplyTileN",
|
| 1558 |
-
"hasBias": true,
|
| 1559 |
-
"useSubgroups": "device.features.has(\"subgroups\")",
|
| 1560 |
"statSlots": "materializedGemmStatSlots",
|
| 1561 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 1562 |
},
|
|
@@ -1575,7 +1528,7 @@
|
|
| 1575 |
"name": "MultiHeadAttention.MaterializedScoresBiasFused",
|
| 1576 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1577 |
"derive": { "layout": "\"bsh\"", "emitRowStats": true },
|
| 1578 |
-
"bindings": ["
|
| 1579 |
"dispatch": {
|
| 1580 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1581 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1587,7 +1540,7 @@
|
|
| 1587 |
"name": "MultiHeadAttention.MaterializedGemmRowStatsCombineBias",
|
| 1588 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 1589 |
"derive": { "maxOnly": true },
|
| 1590 |
-
"bindings": ["
|
| 1591 |
"dispatch": {
|
| 1592 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1593 |
"y": 1,
|
|
@@ -1599,7 +1552,7 @@
|
|
| 1599 |
"name": "MultiHeadAttention.MaterializedApplyBiasFused",
|
| 1600 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1601 |
"derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
|
| 1602 |
-
"bindings": ["
|
| 1603 |
"dispatch": {
|
| 1604 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1605 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
@@ -1611,12 +1564,11 @@
|
|
| 1611 |
{
|
| 1612 |
"id": "qkv_no_bias_flash_q32_broadcast",
|
| 1613 |
"priority": 30,
|
| 1614 |
-
"when": ["tensorDtypes.queryT == \"float16\"", "not present.biasT", "flashShapeOk", "(dim(shapes.queryT, 2) / attrs.num_heads) % 4 == 0", "(dim(shapes.queryT, 2) / attrs.num_heads) >= 64", "(dim(shapes.queryT, 2) / attrs.num_heads) <= 256", "dim(shapes.queryT, 1) >= 31", "
|
| 1615 |
"requires": { "features": ["subgroups", "shader-f16"] },
|
| 1616 |
"derive": {
|
| 1617 |
"usesF16": true,
|
| 1618 |
"scalar": "\"f16\"",
|
| 1619 |
-
"inputVec4": "\"vec4<f16>\"",
|
| 1620 |
"inputElement": "\"vec4<f16>\"",
|
| 1621 |
"outputElement": "\"vec4<f16>\"",
|
| 1622 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -1637,8 +1589,7 @@
|
|
| 1637 |
"x": "ceilDiv(dim(shapes.queryT, 1), 32)",
|
| 1638 |
"y": "attrs.num_heads",
|
| 1639 |
"z": "dim(shapes.queryT, 0)"
|
| 1640 |
-
}
|
| 1641 |
-
"subgroupCollectivesWidth": 32
|
| 1642 |
}
|
| 1643 |
]
|
| 1644 |
},
|
|
@@ -1650,7 +1601,6 @@
|
|
| 1650 |
"derive": {
|
| 1651 |
"usesF16": true,
|
| 1652 |
"scalar": "\"f16\"",
|
| 1653 |
-
"inputVec4": "\"vec4<f16>\"",
|
| 1654 |
"inputElement": "\"vec4<f16>\"",
|
| 1655 |
"outputElement": "\"vec4<f16>\"",
|
| 1656 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -1681,14 +1631,8 @@
|
|
| 1681 |
"priority": 0,
|
| 1682 |
"when": ["not present.biasT and fallbackMaskShapeOk"],
|
| 1683 |
"derive": {
|
| 1684 |
-
"headsFromParams": false,
|
| 1685 |
-
"hasBias": false,
|
| 1686 |
"hasCausal": true,
|
| 1687 |
-
"hasWindow": false,
|
| 1688 |
"hasKeyLimit": false,
|
| 1689 |
-
"hasMask": true,
|
| 1690 |
-
"scaleFallbackRsqrt": true,
|
| 1691 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1692 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1693 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1694 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1707,7 +1651,7 @@
|
|
| 1707 |
"name": "MultiHeadAttention",
|
| 1708 |
"shader": "attn-online-scalar.wgsl.jinja",
|
| 1709 |
"derive": { "layout": "\"bsh\"" },
|
| 1710 |
-
"bindings": ["query", "key", "value", "
|
| 1711 |
"dispatch": {
|
| 1712 |
"x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1713 |
"y": "attrs.num_heads",
|
|
@@ -1721,14 +1665,8 @@
|
|
| 1721 |
"priority": 0,
|
| 1722 |
"when": ["biasOk and fallbackMaskShapeOk"],
|
| 1723 |
"derive": {
|
| 1724 |
-
"headsFromParams": false,
|
| 1725 |
-
"hasBias": true,
|
| 1726 |
"hasCausal": true,
|
| 1727 |
-
"hasWindow": false,
|
| 1728 |
"hasKeyLimit": false,
|
| 1729 |
-
"hasMask": true,
|
| 1730 |
-
"scaleFallbackRsqrt": true,
|
| 1731 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1732 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1733 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1734 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1747,7 +1685,7 @@
|
|
| 1747 |
"name": "MultiHeadAttention",
|
| 1748 |
"shader": "attn-online-scalar.wgsl.jinja",
|
| 1749 |
"derive": { "layout": "\"bsh\"" },
|
| 1750 |
-
"bindings": ["query", "key", "value", "
|
| 1751 |
"dispatch": {
|
| 1752 |
"x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1753 |
"y": "attrs.num_heads",
|
|
@@ -1761,13 +1699,8 @@
|
|
| 1761 |
"priority": 0,
|
| 1762 |
"when": ["not present.biasT and fallbackShapeOk and not flashShapeOk"],
|
| 1763 |
"derive": {
|
| 1764 |
-
"headsFromParams": false,
|
| 1765 |
-
"hasBias": false,
|
| 1766 |
"hasCausal": true,
|
| 1767 |
-
"hasWindow": false,
|
| 1768 |
"hasKeyLimit": false,
|
| 1769 |
-
"scaleFallbackRsqrt": true,
|
| 1770 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1771 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1772 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1773 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1799,13 +1732,8 @@
|
|
| 1799 |
"priority": 0,
|
| 1800 |
"when": ["biasOk and fallbackShapeOk and not flashShapeOk"],
|
| 1801 |
"derive": {
|
| 1802 |
-
"headsFromParams": false,
|
| 1803 |
-
"hasBias": true,
|
| 1804 |
"hasCausal": true,
|
| 1805 |
-
"hasWindow": false,
|
| 1806 |
"hasKeyLimit": false,
|
| 1807 |
-
"scaleFallbackRsqrt": true,
|
| 1808 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1809 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1810 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1811 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1837,17 +1765,10 @@
|
|
| 1837 |
"priority": 20,
|
| 1838 |
"when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashShapeOk"],
|
| 1839 |
"derive": {
|
| 1840 |
-
"headsFromParams": false,
|
| 1841 |
-
"hasBias": false,
|
| 1842 |
"hasCausal": true,
|
| 1843 |
-
"hasWindow": false,
|
| 1844 |
"combineSubgroups": true,
|
| 1845 |
-
"hasMask": false,
|
| 1846 |
-
"maskIsBool": false,
|
| 1847 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1848 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1849 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1850 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1851 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1852 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1853 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -1864,8 +1785,7 @@
|
|
| 1864 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1865 |
"derive": { "layout": "\"bsh\"" },
|
| 1866 |
"bindings": ["query", "key", "value", "output", "params"],
|
| 1867 |
-
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1868 |
-
"subgroupCollectivesWidth": "portable"
|
| 1869 |
}
|
| 1870 |
]
|
| 1871 |
},
|
|
@@ -1874,17 +1794,10 @@
|
|
| 1874 |
"priority": 20,
|
| 1875 |
"when": ["device.features.has(\"subgroups\")", "biasOk", "flashShapeOk"],
|
| 1876 |
"derive": {
|
| 1877 |
-
"headsFromParams": false,
|
| 1878 |
-
"hasBias": true,
|
| 1879 |
"hasCausal": true,
|
| 1880 |
-
"hasWindow": false,
|
| 1881 |
"combineSubgroups": true,
|
| 1882 |
-
"hasMask": false,
|
| 1883 |
-
"maskIsBool": false,
|
| 1884 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1885 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1886 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1887 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1888 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1889 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1890 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1903,8 +1816,7 @@
|
|
| 1903 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1904 |
"derive": { "layout": "\"bsh\"" },
|
| 1905 |
"bindings": ["query", "key", "value", "bias", "output", "params"],
|
| 1906 |
-
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1907 |
-
"subgroupCollectivesWidth": "portable"
|
| 1908 |
}
|
| 1909 |
]
|
| 1910 |
},
|
|
@@ -1913,17 +1825,10 @@
|
|
| 1913 |
"priority": 18,
|
| 1914 |
"when": ["true", "not present.biasT", "flashShapeOk"],
|
| 1915 |
"derive": {
|
| 1916 |
-
"headsFromParams": false,
|
| 1917 |
-
"hasBias": false,
|
| 1918 |
"hasCausal": true,
|
| 1919 |
-
"hasWindow": false,
|
| 1920 |
"combineSubgroups": false,
|
| 1921 |
-
"hasMask": false,
|
| 1922 |
-
"maskIsBool": false,
|
| 1923 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1924 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1925 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1926 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1927 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1928 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1929 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -1949,17 +1854,10 @@
|
|
| 1949 |
"priority": 17,
|
| 1950 |
"when": ["true", "biasOk", "flashShapeOk"],
|
| 1951 |
"derive": {
|
| 1952 |
-
"headsFromParams": false,
|
| 1953 |
-
"hasBias": true,
|
| 1954 |
"hasCausal": true,
|
| 1955 |
-
"hasWindow": false,
|
| 1956 |
"combineSubgroups": false,
|
| 1957 |
-
"hasMask": false,
|
| 1958 |
-
"maskIsBool": false,
|
| 1959 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1960 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1961 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1962 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1963 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1964 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1965 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -1987,17 +1885,10 @@
|
|
| 1987 |
"priority": 20,
|
| 1988 |
"when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashMaskShapeOk"],
|
| 1989 |
"derive": {
|
| 1990 |
-
"headsFromParams": false,
|
| 1991 |
-
"hasBias": false,
|
| 1992 |
"hasCausal": true,
|
| 1993 |
-
"hasWindow": false,
|
| 1994 |
"combineSubgroups": true,
|
| 1995 |
-
"hasMask": true,
|
| 1996 |
-
"maskIsBool": false,
|
| 1997 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1998 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1999 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 2000 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2001 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2002 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2003 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -2013,9 +1904,8 @@
|
|
| 2013 |
"name": "MultiHeadAttention.Flash",
|
| 2014 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 2015 |
"derive": { "layout": "\"bsh\"" },
|
| 2016 |
-
"bindings": ["query", "key", "value", "
|
| 2017 |
-
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 2018 |
-
"subgroupCollectivesWidth": "portable"
|
| 2019 |
}
|
| 2020 |
]
|
| 2021 |
},
|
|
@@ -2024,17 +1914,10 @@
|
|
| 2024 |
"priority": 20,
|
| 2025 |
"when": ["device.features.has(\"subgroups\")", "biasOk", "flashMaskShapeOk"],
|
| 2026 |
"derive": {
|
| 2027 |
-
"headsFromParams": false,
|
| 2028 |
-
"hasBias": true,
|
| 2029 |
"hasCausal": true,
|
| 2030 |
-
"hasWindow": false,
|
| 2031 |
"combineSubgroups": true,
|
| 2032 |
-
"hasMask": true,
|
| 2033 |
-
"maskIsBool": false,
|
| 2034 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 2035 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 2036 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 2037 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2038 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2039 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2040 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -2052,9 +1935,8 @@
|
|
| 2052 |
"name": "MultiHeadAttention.Flash",
|
| 2053 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 2054 |
"derive": { "layout": "\"bsh\"" },
|
| 2055 |
-
"bindings": ["query", "key", "value", "
|
| 2056 |
-
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 2057 |
-
"subgroupCollectivesWidth": "portable"
|
| 2058 |
}
|
| 2059 |
]
|
| 2060 |
},
|
|
@@ -2063,17 +1945,10 @@
|
|
| 2063 |
"priority": 18,
|
| 2064 |
"when": ["true", "not present.biasT", "flashMaskShapeOk"],
|
| 2065 |
"derive": {
|
| 2066 |
-
"headsFromParams": false,
|
| 2067 |
-
"hasBias": false,
|
| 2068 |
"hasCausal": true,
|
| 2069 |
-
"hasWindow": false,
|
| 2070 |
"combineSubgroups": false,
|
| 2071 |
-
"hasMask": true,
|
| 2072 |
-
"maskIsBool": false,
|
| 2073 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 2074 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 2075 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 2076 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2077 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2078 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2079 |
"qNumHeads": "attrs.num_heads",
|
|
@@ -2089,7 +1964,7 @@
|
|
| 2089 |
"name": "MultiHeadAttention.NoBiasOnlineFlashNoSg",
|
| 2090 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 2091 |
"derive": { "layout": "\"bsh\"" },
|
| 2092 |
-
"bindings": ["query", "key", "value", "
|
| 2093 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 2094 |
}
|
| 2095 |
]
|
|
@@ -2099,17 +1974,10 @@
|
|
| 2099 |
"priority": 17,
|
| 2100 |
"when": ["true", "biasOk", "flashMaskShapeOk"],
|
| 2101 |
"derive": {
|
| 2102 |
-
"headsFromParams": false,
|
| 2103 |
-
"hasBias": true,
|
| 2104 |
"hasCausal": true,
|
| 2105 |
-
"hasWindow": false,
|
| 2106 |
"combineSubgroups": false,
|
| 2107 |
-
"hasMask": true,
|
| 2108 |
-
"maskIsBool": false,
|
| 2109 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 2110 |
-
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 2111 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 2112 |
-
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2113 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2114 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 2115 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
@@ -2127,7 +1995,7 @@
|
|
| 2127 |
"name": "MultiHeadAttention.BiasOnlineFlashNoSg",
|
| 2128 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 2129 |
"derive": { "layout": "\"bsh\"" },
|
| 2130 |
-
"bindings": ["query", "key", "value", "
|
| 2131 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 2132 |
}
|
| 2133 |
]
|
|
@@ -2142,13 +2010,11 @@
|
|
| 2142 |
},
|
| 2143 |
"derive": {
|
| 2144 |
"qNumHeads": "attrs.num_heads",
|
| 2145 |
-
"headDim": "headDimPlan",
|
| 2146 |
"qHidden": "dim(shapes.queryT, 2)",
|
| 2147 |
"materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
|
| 2148 |
"materializedSoftmaxCols": "dim(shapes.keyT, 1)",
|
| 2149 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 2150 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
| 2151 |
-
"hasBias": false,
|
| 2152 |
"useSubgroups": true
|
| 2153 |
},
|
| 2154 |
"intermediates": [
|
|
@@ -2164,7 +2030,7 @@
|
|
| 2164 |
"name": "MultiHeadAttention.MaterializedScoresSgmat",
|
| 2165 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2166 |
"derive": { "phase": "\"score\"" },
|
| 2167 |
-
"bindings": ["
|
| 2168 |
"dispatch": {
|
| 2169 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 2170 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
@@ -2176,20 +2042,19 @@
|
|
| 2176 |
"name": "MultiHeadAttention.MaterializedSoftmaxSgmat",
|
| 2177 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 2178 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 2179 |
-
"bindings": ["
|
| 2180 |
"dispatch": {
|
| 2181 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2182 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2183 |
"z": 1
|
| 2184 |
-
}
|
| 2185 |
-
"subgroupCollectivesWidth": "portable"
|
| 2186 |
},
|
| 2187 |
{
|
| 2188 |
"id": "apply",
|
| 2189 |
"name": "MultiHeadAttention.MaterializedApplySgmat",
|
| 2190 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2191 |
"derive": { "phase": "\"apply\"" },
|
| 2192 |
-
"bindings": ["
|
| 2193 |
"dispatch": {
|
| 2194 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 2195 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
@@ -2208,13 +2073,11 @@
|
|
| 2208 |
},
|
| 2209 |
"derive": {
|
| 2210 |
"qNumHeads": "attrs.num_heads",
|
| 2211 |
-
"headDim": "headDimPlan",
|
| 2212 |
"qHidden": "dim(shapes.queryT, 2)",
|
| 2213 |
"materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
|
| 2214 |
"materializedSoftmaxCols": "dim(shapes.keyT, 1)",
|
| 2215 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 2216 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
| 2217 |
-
"hasBias": true,
|
| 2218 |
"useSubgroups": true
|
| 2219 |
},
|
| 2220 |
"intermediates": [
|
|
@@ -2230,7 +2093,7 @@
|
|
| 2230 |
"name": "MultiHeadAttention.MaterializedScoresSgmatBias",
|
| 2231 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2232 |
"derive": { "phase": "\"score\"" },
|
| 2233 |
-
"bindings": ["
|
| 2234 |
"dispatch": {
|
| 2235 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 2236 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
@@ -2242,20 +2105,19 @@
|
|
| 2242 |
"name": "MultiHeadAttention.MaterializedSoftmaxSgmatBias",
|
| 2243 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 2244 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 2245 |
-
"bindings": ["
|
| 2246 |
"dispatch": {
|
| 2247 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2248 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2249 |
"z": 1
|
| 2250 |
-
}
|
| 2251 |
-
"subgroupCollectivesWidth": "portable"
|
| 2252 |
},
|
| 2253 |
{
|
| 2254 |
"id": "apply",
|
| 2255 |
"name": "MultiHeadAttention.MaterializedApplySgmatBias",
|
| 2256 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2257 |
"derive": { "phase": "\"apply\"" },
|
| 2258 |
-
"bindings": ["
|
| 2259 |
"dispatch": {
|
| 2260 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 2261 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
|
|
| 57 |
"canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32",
|
| 58 |
"pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
|
| 59 |
"wave32Effective": "wave32Adapter or pinSubgroupSize32",
|
| 60 |
+
"wave32SubgroupsUsable": "subgroupsWave32 or pinSubgroupSize32",
|
| 61 |
+
"headDim": "dim(shapes.queryT, 2) / attrs.num_heads if (ranks.queryT == 3 and attrs.num_heads > 0) else 0",
|
| 62 |
"qkvDtypesOk": "tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT",
|
| 63 |
"floatDtypeOk": "(tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(tensorDtypes.queryT)",
|
| 64 |
"qkvShapeOk": "ranks.queryT == 3 and ranks.keyT == 3 and ranks.valueT == 3 and ranks.outputT == 3 and attrs.num_heads > 0 and dim(shapes.queryT, 2) % attrs.num_heads == 0 and dim(shapes.keyT, 2) == dim(shapes.queryT, 2) and dim(shapes.valueT, 2) == dim(shapes.queryT, 2) and dim(shapes.keyT, 1) == dim(shapes.valueT, 1) and dim(shapes.queryT, 0) == dim(shapes.keyT, 0) and dim(shapes.queryT, 0) == dim(shapes.valueT, 0) and dim(shapes.outputT, 0) == dim(shapes.queryT, 0) and dim(shapes.outputT, 1) == dim(shapes.queryT, 1) and dim(shapes.outputT, 2) == dim(shapes.valueT, 2)",
|
|
|
|
| 67 |
"qkvContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and noAttnBias",
|
| 68 |
"qkvMaskContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and attnBiasOk",
|
| 69 |
"biasOk": "present.biasT and ranks.biasT == 1 and tensorDtypes.biasT == tensorDtypes.queryT and dim(shapes.biasT, 0) == 3 * dim(shapes.queryT, 2)",
|
| 70 |
+
"hasBias": "present.biasT",
|
| 71 |
+
"hasMask": "present.attentionBiasT",
|
| 72 |
+
"q32BroadcastSubgroupLanes": "32 if wave32SubgroupsUsable else 0",
|
| 73 |
+
"q32BroadcastF32HeadVectors": "headDim / 4 if headDim % 4 == 0 else 0",
|
| 74 |
+
"q32BroadcastF32RegisterGeometry": "wave32SubgroupsUsable and q32BroadcastF32HeadVectors == q32BroadcastSubgroupLanes",
|
| 75 |
"subgroupCluster4": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 4 and device.adapterInfo.subgroupMinSize % 4 == 0 and device.adapterInfo.subgroupMaxSize % 4 == 0",
|
| 76 |
"subgroupCluster8": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 8 and device.adapterInfo.subgroupMinSize % 8 == 0 and device.adapterInfo.subgroupMaxSize % 8 == 0",
|
| 77 |
"attentionDispatchFits": "dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 1) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 78 |
+
"flashHeadOk": "headDim % 4 == 0 and headDim >= 32 and headDim <= 256",
|
| 79 |
"flashSizeOk": "flashHeadOk and (dim(shapes.queryT, 1) * attrs.num_heads >= tunables.FLASH_MIN_QUERY_HEADS or (dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512) or (dim(shapes.queryT, 1) > 1 and dim(shapes.keyT, 1) >= 2048)) and attentionDispatchFits",
|
| 80 |
"flashShapeOk": "qkvContractOk and flashSizeOk",
|
| 81 |
"flashMaskShapeOk": "qkvMaskContractOk and flashSizeOk",
|
| 82 |
"noBiasSplitKCount": "min(tunables.DECODE_MAX_SPLITS if dim(shapes.queryT, 1) == 1 else max(1, ceilDiv(tunables.SPLITK_TARGET_WORKGROUPS, dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads)), ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
|
| 83 |
"biasSplitKCount": "min(tunables.DECODE_MAX_SPLITS, ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))",
|
| 84 |
+
"noBiasPartialOutBytes": "dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * headDim * 4",
|
| 85 |
"noBiasStatsBytes": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * 4",
|
| 86 |
"noBiasSplitScratchFits": "noBiasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and noBiasPartialOutBytes <= device.limits.maxBufferSize and noBiasStatsBytes <= device.limits.maxStorageBufferBindingSize and noBiasStatsBytes <= device.limits.maxBufferSize",
|
| 87 |
"noBiasSplitDispatchFits": "dim(shapes.queryT, 1) * noBiasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 88 |
+
"biasPartialOutBytes": "dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * headDim * 4",
|
| 89 |
"biasStatsBytes": "2 * dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * 4",
|
| 90 |
"biasSplitScratchFits": "biasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and biasPartialOutBytes <= device.limits.maxBufferSize and biasStatsBytes <= device.limits.maxStorageBufferBindingSize and biasStatsBytes <= device.limits.maxBufferSize",
|
| 91 |
"biasSplitDispatchFits": "biasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 92 |
"decodeSplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
|
| 93 |
"shortQuerySplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) >= 2 and dim(shapes.queryT, 1) <= 16 and dim(shapes.keyT, 1) >= 2048 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits",
|
| 94 |
"decodeSplitKBiasOk": "biasOk and qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and biasSplitDispatchFits and biasSplitScratchFits",
|
| 95 |
+
"decodeSplitKPortablePreferred": "tensorDtypes.queryT == \"float32\" and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and headDim / 4 < device.adapterInfo.subgroupMinSize",
|
| 96 |
"materializedScoreBytes": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1) * 4",
|
| 97 |
"materializedScoreFits": "materializedScoreBytes <= device.limits.maxStorageBufferBindingSize and materializedScoreBytes <= device.limits.maxBufferSize",
|
| 98 |
"materializedWorkgroupSize": "tunables.MATERIALIZED_WORKGROUP_DIM * tunables.MATERIALIZED_WORKGROUP_DIM",
|
| 99 |
"materializedScoreStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_KEY_TILE * (tunables.MATERIALIZED_INNER_TILE + 4)) * 4",
|
| 100 |
+
"materializedApplyTileN": "tunables.MATERIALIZED_VALUE_TILE_D128 if headDim == 128 else tunables.MATERIALIZED_VALUE_TILE",
|
| 101 |
"materializedApplyStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) * 4",
|
| 102 |
"materializedTileGeometryOk": "tunables.MATERIALIZED_INNER_TILE % 4 == 0 and (materializedApplyTileN / tunables.MATERIALIZED_WORKGROUP_DIM) % 4 == 0 and tunables.MATERIALIZED_INNER_TILE > 0 and tunables.MATERIALIZED_WORKGROUP_DIM > 0 and tunables.MATERIALIZED_QUERY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_KEY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and materializedApplyTileN % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN >= materializedWorkgroupSize and (tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) % materializedWorkgroupSize == 0",
|
| 103 |
"materializedDeviceOk": "materializedTileGeometryOk and materializedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeX and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeY and materializedScoreStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and materializedApplyStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 104 |
"materializedWideSimdOk": "device.features.has(\"subgroups\") or (has(device.adapterInfo, \"subgroupMinSize\") and device.adapterInfo.subgroupMinSize >= 16)",
|
| 105 |
+
"materializedF32CoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDim >= 64 and headDim <= 128 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and materializedScoreFits and materializedDeviceOk",
|
| 106 |
"materializedSoftmaxStorageBytes": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE * 8 + 8",
|
| 107 |
"materializedSoftmaxResourcesFit": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE > 0 and pow2ceil(tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE) == tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE and tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE <= deviceWorkgroupCap and materializedSoftmaxStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 108 |
"materializedSgmatQueryTile": "tunables.MATERIALIZED_SGMAT_QUERY_TILE",
|
|
|
|
| 126 |
"materializedSgmatResourcesFit": "materializedSgmatGeometryOk and materializedSgmatWorkgroupSize <= deviceWorkgroupCap and materializedSgmatCompactStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 127 |
"materializedSgmatDispatchFits": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 128 |
"materializedSgmatDirectScoreStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and dim(shapes.keyT, 1) % materializedSgmatKeyTile == 0",
|
| 129 |
+
"materializedSgmatDirectApplyStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and headDim % materializedSgmatKeyTile == 0",
|
| 130 |
"materializedSgmatRuntimeDirectStore": "dim(shapes.queryT, 1) >= 2 * materializedSgmatQueryTile and dim(shapes.keyT, 1) >= 2 * materializedSgmatKeyTile",
|
| 131 |
+
"materializedSgmatCoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDim >= 32 and headDim <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits",
|
| 132 |
"materializedCachedSoftmaxWg": "min(tunables.MATERIALIZED_CACHED_SOFTMAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
|
| 133 |
"materializedCachedSoftmaxVecsPerLane": "ceilDiv(ceilDiv(dim(shapes.keyT, 1), 4), max(1, materializedCachedSoftmaxWg))",
|
| 134 |
"materializedCachedSoftmaxStorageBytes": "materializedCachedSoftmaxWg * 8 + 8",
|
|
|
|
| 138 |
"materializedSgmatOk": "materializedSgmatCoreOk and materializedAdaptiveSoftmaxOk",
|
| 139 |
"materializedSgmatFusedOk": "materializedSgmatCoreOk",
|
| 140 |
"materializedF32Ok": "materializedF32CoreOk and materializedAdaptiveSoftmaxOk",
|
| 141 |
+
"clusterTileKWg64": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDim * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_SMALL * 4))))",
|
| 142 |
+
"clusterTileKWg128": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDim * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_LARGE * 4))))",
|
| 143 |
+
"smallHeadShapeOk": "qkvContractOk and headDim < 32 and dim(shapes.keyT, 1) >= 64 and attrs.unidirectional == 0 and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 144 |
+
"smallHeadParallelOk": "smallHeadShapeOk and dim(shapes.keyT, 1) <= 2048",
|
| 145 |
+
"prefillTiledStorageBytes": "headDim * tunables.PREFILL_QUERY_TILE * 4",
|
| 146 |
"prefillTiledDeviceOk": "tunables.PREFILL_QUERY_TILE <= deviceWorkgroupCap and prefillTiledStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 147 |
+
"portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, pow2ceil(max(1, headDim)))",
|
| 148 |
+
"portableWorkgroupStorageBytes": "portableWorkgroupSize * 8 + max(1, headDim) * 4 + 16",
|
| 149 |
"portableWorkgroupOk": "tunables.WORKGROUP_SIZE > 0 and pow2ceil(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and portableWorkgroupSize <= deviceWorkgroupCap and portableWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 150 |
"fallbackShapeOk": "qkvContractOk and portableWorkgroupOk",
|
| 151 |
"fallbackMaskShapeOk": "qkvMaskContractOk and portableWorkgroupOk",
|
| 152 |
+
"smallSeqShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and dim(shapes.queryT, 1) >= 1 and dim(shapes.queryT, 1) <= tunables.SMALL_SEQ_MAX and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_MAX and headDim >= 1 and attrs.unidirectional == 0",
|
| 153 |
+
"smallSeqPrivateFloats": "dim(shapes.keyT, 1) + headDim",
|
| 154 |
"smallSeqWorkgroupSize": "max(32, pow2ceil(dim(shapes.queryT, 1)))",
|
| 155 |
+
"smallSeqSharedBytes": "dim(shapes.keyT, 1) * headDim * 8",
|
| 156 |
"smallSeqResourcesFit": "smallSeqPrivateFloats <= tunables.SMALL_SEQ_MAX_PRIVATE_FLOATS and smallSeqWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqSharedBytes <= device.limits.maxComputeWorkgroupStorageSize",
|
| 157 |
+
"smallSeqBlockedKvBytes": "dim(shapes.keyT, 1) * headDim * 8",
|
| 158 |
+
"smallSeqBlockedLaneBytes": "8 + headDim * 4",
|
| 159 |
"smallSeqBlockedKeyLanes": "min(pow2ceil(dim(shapes.keyT, 1)), 16 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 16 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else (8 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 8 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else 4))",
|
| 160 |
"smallSeqBlockedWorkgroupSize": "tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * smallSeqBlockedKeyLanes",
|
| 161 |
"smallSeqBlockedSharedBytes": "smallSeqBlockedKvBytes + smallSeqBlockedWorkgroupSize * smallSeqBlockedLaneBytes",
|
| 162 |
+
"smallSeqBlockedShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and headDim % 4 == 0 and headDim >= 4 and headDim <= tunables.SMALL_SEQ_BLOCKED_MAX_HEAD_DIM and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_BLOCKED_MAX_KV and dim(shapes.queryT, 1) >= 1",
|
| 163 |
"smallSeqBlockedFits": "smallSeqBlockedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqBlockedWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqBlockedSharedBytes <= device.limits.maxComputeWorkgroupStorageSize and ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 164 |
"smallSeqDispatchFits": "attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
|
| 165 |
+
"materializedSgmatCoreF16Ok": "qkvContractOk and tensorDtypes.queryT == \"float16\" and tensorDtypes.keyT == \"float16\" and tensorDtypes.valueT == \"float16\" and device.features.has(\"shader-f16\") and attrs.unidirectional == 0 and headDim >= 32 and headDim <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits",
|
| 166 |
"materializedSgmatFusedF16Ok": "materializedSgmatCoreF16Ok",
|
| 167 |
+
"attentionScaleExpression": "\"select(inverseSqrt(f32(HEAD_DIM)), params.scale, params.scale != 0.0)\"",
|
| 168 |
+
"smallHeadValueWg": "pow(2, log2ceil(min(64, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX) + 1) - 1)",
|
| 169 |
+
"smallHeadValueStorage": "max(dim(shapes.keyT, 1), smallHeadValueWg * headDim)",
|
| 170 |
+
"smallHeadValueQueryBlock": "2 if dim(shapes.queryT, 1) >= 2 and (smallHeadValueStorage + smallHeadValueWg) * 8 <= device.limits.maxComputeWorkgroupStorageSize else 1",
|
| 171 |
+
"smallHeadValueEstimatedSteps": "headDim * (ceilDiv(dim(shapes.keyT, 1), smallHeadValueWg) + 2 * log2ceil(smallHeadValueWg) * (2 if tensorDtypes.queryT == \"float16\" else 1) + 2)",
|
| 172 |
+
"smallHeadValueSgReductionSteps": "ceilDiv(smallHeadValueWg, max(1, device.adapterInfo.subgroupMinSize)) if has(device.adapterInfo, \"subgroupMinSize\") else log2ceil(smallHeadValueWg)",
|
| 173 |
+
"smallHeadValueSgEstimatedSteps": "headDim * (ceilDiv(dim(shapes.keyT, 1), smallHeadValueWg) + 2 * smallHeadValueSgReductionSteps * (2 if tensorDtypes.queryT == \"float16\" else 1) + 2)"
|
| 174 |
},
|
| 175 |
"bindings": {
|
| 176 |
+
"query": { "arg": "queryT", "elementType": "$inputElement" },
|
| 177 |
+
"key": { "arg": "keyT", "elementType": "$inputElement" },
|
| 178 |
+
"value": { "arg": "valueT", "elementType": "$inputElement" },
|
| 179 |
+
"bias": { "arg": "biasT", "elementType": "$inputScalar" },
|
| 180 |
+
"output": { "arg": "outputT", "elementType": "$outputElement" },
|
| 181 |
"params": {
|
|
|
|
| 182 |
"struct": [
|
| 183 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 184 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
|
|
|
| 186 |
{ "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" }
|
| 187 |
]
|
| 188 |
},
|
| 189 |
+
"params_main": {
|
| 190 |
"name": "params",
|
|
|
|
| 191 |
"struct": [
|
| 192 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 193 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 194 |
]
|
| 195 |
},
|
| 196 |
+
"q": { "arg": "queryT", "elementType": "$scalar" },
|
| 197 |
+
"k": { "arg": "keyT", "elementType": "$scalar" },
|
| 198 |
+
"v": { "arg": "valueT", "elementType": "$scalar" },
|
| 199 |
+
"y": { "arg": "outputT", "elementType": "$scalar" },
|
| 200 |
+
"params_scores": {
|
| 201 |
"name": "params",
|
|
|
|
| 202 |
"struct": [
|
| 203 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 204 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
| 205 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 206 |
]
|
| 207 |
},
|
| 208 |
+
"query_f32": { "arg": "queryT", "name": "query", "elementType": "f32" },
|
| 209 |
+
"key_f32": { "arg": "keyT", "name": "key", "elementType": "f32" },
|
| 210 |
+
"scores": { "scratch": "materializedScores", "elementType": "f32" },
|
| 211 |
+
"scorePartials": { "scratch": "materializedScorePartials", "elementType": "f32" },
|
| 212 |
+
"bias_f32": { "arg": "biasT", "name": "bias", "elementType": "f32" },
|
| 213 |
+
"scorePartials_f32": {
|
| 214 |
"scratch": "materializedScorePartials",
|
| 215 |
"name": "scorePartials",
|
| 216 |
"buffer": "read-only-storage",
|
| 217 |
"elementType": "f32"
|
| 218 |
},
|
| 219 |
+
"rowStats": { "scratch": "materializedRowStats", "elementType": "f32" },
|
| 220 |
+
"params_rows": {
|
| 221 |
"name": "params",
|
|
|
|
| 222 |
"struct": [
|
| 223 |
{ "name": "rows", "type": "u32", "value": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)" }
|
| 224 |
]
|
| 225 |
},
|
| 226 |
+
"scores_f32": {
|
| 227 |
"scratch": "materializedScores",
|
| 228 |
"name": "scores",
|
| 229 |
"buffer": "read-only-storage",
|
| 230 |
"elementType": "f32"
|
| 231 |
},
|
| 232 |
+
"value_f32": { "arg": "valueT", "name": "value", "elementType": "f32" },
|
| 233 |
+
"output_f32": { "arg": "outputT", "name": "output", "elementType": "f32" },
|
| 234 |
+
"rowStats_f32": {
|
| 235 |
"scratch": "materializedRowStats",
|
| 236 |
"name": "rowStats",
|
| 237 |
"buffer": "read-only-storage",
|
| 238 |
"elementType": "f32"
|
| 239 |
},
|
| 240 |
+
"params_apply": {
|
| 241 |
"name": "params",
|
|
|
|
| 242 |
"struct": [
|
| 243 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 244 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }
|
| 245 |
]
|
| 246 |
},
|
| 247 |
+
"attn_mask_main": { "arg": "attentionBiasT", "name": "attn_mask", "elementType": "$maskElement" },
|
| 248 |
+
"params__uniform": {
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 249 |
"name": "params",
|
|
|
|
| 250 |
"struct": [
|
| 251 |
{ "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" },
|
| 252 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
|
|
|
| 265 |
{ "name": "maskSeqStride", "type": "u32", "value": "dim(shapes.attentionBiasT, 3)" }
|
| 266 |
]
|
| 267 |
},
|
| 268 |
+
"query_query_t": { "arg": "queryT", "name": "query", "elementType": "$inputVec4" },
|
| 269 |
+
"key_key_t": { "arg": "keyT", "name": "key", "elementType": "$inputVec4" },
|
| 270 |
+
"value_value_t": { "arg": "valueT", "name": "value", "elementType": "$inputVec4" },
|
| 271 |
+
"partial_out": { "scratch": "partialOut", "elementType": "vec4<f32>" },
|
| 272 |
+
"partial_stats": { "scratch": "partialStats", "elementType": "vec2<f32>" },
|
| 273 |
+
"params_kv_seq_scale": {
|
| 274 |
"name": "params",
|
|
|
|
| 275 |
"struct": [
|
| 276 |
{ "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" },
|
| 277 |
{ "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }
|
| 278 |
]
|
| 279 |
},
|
| 280 |
+
"partial_out_merge": {
|
| 281 |
"scratch": "partialOut",
|
| 282 |
"name": "partial_out",
|
| 283 |
"buffer": "read-only-storage",
|
| 284 |
"elementType": "vec4<f32>"
|
| 285 |
},
|
| 286 |
+
"partial_stats_merge": {
|
| 287 |
"scratch": "partialStats",
|
| 288 |
"name": "partial_stats",
|
| 289 |
"buffer": "read-only-storage",
|
| 290 |
"elementType": "vec2<f32>"
|
| 291 |
},
|
| 292 |
+
"output_merge": { "arg": "outputT", "name": "output", "elementType": "$inputVec4" },
|
| 293 |
+
"scores_softmax": { "scratch": "materializedScores", "name": "scores", "elementType": "$softmaxElementType" }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 294 |
},
|
| 295 |
"variants": [
|
| 296 |
+
{
|
| 297 |
+
"id": "qkv_no_bias_small_head_value_subgroups",
|
| 298 |
+
"priority": 12,
|
| 299 |
+
"when": ["not present.biasT", "smallHeadShapeOk", "headDim > 0", "headDim <= smallHeadValueWg", "(smallHeadValueStorage + smallHeadValueWg) * 4 * smallHeadValueQueryBlock <= device.limits.maxComputeWorkgroupStorageSize", "device.wgslLanguageFeatures.has(\"subgroup_id\")"],
|
| 300 |
+
"derive": {
|
| 301 |
+
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 302 |
+
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 303 |
+
"outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 304 |
+
"qHidden": "dim(shapes.queryT, 2)",
|
| 305 |
+
"headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
|
| 306 |
+
"kvSeq": "dim(shapes.keyT, 1)",
|
| 307 |
+
"valueWorkgroupSize": "smallHeadValueWg",
|
| 308 |
+
"scoreStorageElements": "smallHeadValueStorage",
|
| 309 |
+
"queryBlock": "smallHeadValueQueryBlock",
|
| 310 |
+
"queryTail": "dim(shapes.queryT, 1) % smallHeadValueQueryBlock != 0",
|
| 311 |
+
"useValueSubgroups": true
|
| 312 |
+
},
|
| 313 |
+
"passes": [
|
| 314 |
+
{
|
| 315 |
+
"id": "main",
|
| 316 |
+
"name": "MultiHeadAttention.SmallHeadValueSubgroups",
|
| 317 |
+
"shader": "attn-small-head-value.wgsl.jinja",
|
| 318 |
+
"bindings": ["query", "key", "value", "output", "params_main"],
|
| 319 |
+
"dispatch": {
|
| 320 |
+
"x": "min(ceilDiv(dim(shapes.queryT, 1), smallHeadValueQueryBlock), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 321 |
+
"y": "attrs.num_heads",
|
| 322 |
+
"z": "dim(shapes.queryT, 0)"
|
| 323 |
+
}
|
| 324 |
+
}
|
| 325 |
+
],
|
| 326 |
+
"requires": { "features": ["subgroups"] },
|
| 327 |
+
"demoteWhen": ["dim(shapes.keyT, 1) <= smallHeadValueSgEstimatedSteps"]
|
| 328 |
+
},
|
| 329 |
+
{
|
| 330 |
+
"id": "qkv_no_bias_small_head_value_tree",
|
| 331 |
+
"priority": 11,
|
| 332 |
+
"when": ["not present.biasT", "smallHeadShapeOk", "headDim > 0", "headDim <= smallHeadValueWg", "(smallHeadValueStorage + smallHeadValueWg) * 4 * smallHeadValueQueryBlock <= device.limits.maxComputeWorkgroupStorageSize", "true"],
|
| 333 |
+
"derive": {
|
| 334 |
+
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 335 |
+
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 336 |
+
"outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 337 |
+
"qHidden": "dim(shapes.queryT, 2)",
|
| 338 |
+
"headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
|
| 339 |
+
"kvSeq": "dim(shapes.keyT, 1)",
|
| 340 |
+
"valueWorkgroupSize": "smallHeadValueWg",
|
| 341 |
+
"scoreStorageElements": "smallHeadValueStorage",
|
| 342 |
+
"queryBlock": "smallHeadValueQueryBlock",
|
| 343 |
+
"queryTail": "dim(shapes.queryT, 1) % smallHeadValueQueryBlock != 0",
|
| 344 |
+
"useValueSubgroups": false
|
| 345 |
+
},
|
| 346 |
+
"passes": [
|
| 347 |
+
{
|
| 348 |
+
"id": "main",
|
| 349 |
+
"name": "MultiHeadAttention.SmallHeadValueTree",
|
| 350 |
+
"shader": "attn-small-head-value.wgsl.jinja",
|
| 351 |
+
"bindings": ["query", "key", "value", "output", "params_main"],
|
| 352 |
+
"dispatch": {
|
| 353 |
+
"x": "min(ceilDiv(dim(shapes.queryT, 1), smallHeadValueQueryBlock), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 354 |
+
"y": "attrs.num_heads",
|
| 355 |
+
"z": "dim(shapes.queryT, 0)"
|
| 356 |
+
}
|
| 357 |
+
}
|
| 358 |
+
],
|
| 359 |
+
"requires": { "features": [] },
|
| 360 |
+
"demoteWhen": ["dim(shapes.keyT, 1) <= smallHeadValueEstimatedSteps"]
|
| 361 |
+
},
|
| 362 |
{
|
| 363 |
"id": "qkv_bias_small_seq_blocked",
|
| 364 |
"priority": 40,
|
| 365 |
"when": ["biasOk", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
|
| 366 |
"derive": {
|
|
|
|
|
|
|
| 367 |
"inputElement": "\"vec4<f32>\"",
|
| 368 |
"outputElement": "\"vec4<f32>\"",
|
| 369 |
"inputScalar": "\"f32\"",
|
| 370 |
+
"headDimV4": "headDim / 4",
|
|
|
|
|
|
|
| 371 |
"hidden": "dim(shapes.queryT, 2)",
|
| 372 |
"hiddenV4": "dim(shapes.queryT, 2) / 4",
|
| 373 |
"kvSeq": "dim(shapes.keyT, 1)",
|
|
|
|
| 384 |
"x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
|
| 385 |
"y": "attrs.num_heads",
|
| 386 |
"z": "dim(shapes.queryT, 0)"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 387 |
}
|
| 388 |
}
|
| 389 |
]
|
|
|
|
| 393 |
"priority": 40,
|
| 394 |
"when": ["not present.biasT", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"],
|
| 395 |
"derive": {
|
|
|
|
|
|
|
| 396 |
"inputElement": "\"vec4<f32>\"",
|
| 397 |
"outputElement": "\"vec4<f32>\"",
|
| 398 |
+
"headDimV4": "headDim / 4",
|
|
|
|
|
|
|
|
|
|
| 399 |
"hidden": "dim(shapes.queryT, 2)",
|
| 400 |
"hiddenV4": "dim(shapes.queryT, 2) / 4",
|
| 401 |
"kvSeq": "dim(shapes.keyT, 1)",
|
|
|
|
| 412 |
"x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)",
|
| 413 |
"y": "attrs.num_heads",
|
| 414 |
"z": "dim(shapes.queryT, 0)"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 415 |
}
|
| 416 |
}
|
| 417 |
]
|
|
|
|
| 421 |
"priority": 60,
|
| 422 |
"when": ["not present.biasT", "smallSeqShapeOk", "not flashShapeOk", "smallSeqResourcesFit", "smallSeqDispatchFits"],
|
| 423 |
"derive": {
|
|
|
|
| 424 |
"inputElement": "\"f32\"",
|
| 425 |
"outputElement": "\"f32\"",
|
| 426 |
"headDim": "dim(shapes.queryT, 2) / attrs.num_heads",
|
|
|
|
| 433 |
"id": "main",
|
| 434 |
"name": "MultiHeadAttention",
|
| 435 |
"shader": "mha-small-seq.wgsl.jinja",
|
| 436 |
+
"bindings": ["query", "key", "value", "output", "params_main"],
|
| 437 |
+
"dispatch": { "x": "attrs.num_heads", "y": "dim(shapes.queryT, 0)" }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 438 |
}
|
| 439 |
]
|
| 440 |
},
|
| 441 |
{
|
| 442 |
"id": "qkv_no_bias_tiled_nosg",
|
| 443 |
"priority": 19,
|
| 444 |
+
"when": ["not present.biasT", "qkvContractOk", "not smallHeadParallelOk", "headDim % 4 == 0", "headDim <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"],
|
| 445 |
"supersededBy": ["qkv_no_bias_flash_cluster_nosg", "qkv_no_bias_flash_cluster_lpq4_nosg"],
|
| 446 |
"derive": {
|
| 447 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 448 |
"blockM": "tunables.PREFILL_QUERY_TILE",
|
| 449 |
"vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
|
| 450 |
},
|
|
|
|
| 479 |
}
|
| 480 |
],
|
| 481 |
"dispatch": {
|
| 482 |
+
"x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
|
| 483 |
+
"y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
|
| 484 |
"z": 1
|
| 485 |
}
|
| 486 |
}
|
|
|
|
| 489 |
{
|
| 490 |
"id": "qkv_bias_flash_q32_broadcast_f32_d128",
|
| 491 |
"priority": 30,
|
| 492 |
+
"when": ["tensorDtypes.queryT == \"float32\"", "biasOk", "attrs.unidirectional == 0", "flashShapeOk", "q32BroadcastF32RegisterGeometry", "dim(shapes.queryT, 1) >= 31", "ceilDiv(dim(shapes.queryT, 1), 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32SubgroupsUsable"],
|
| 493 |
"requires": { "features": ["subgroups"] },
|
| 494 |
"derive": {
|
|
|
|
| 495 |
"hasCausal": false,
|
| 496 |
"usesF16": false,
|
| 497 |
"scalar": "\"f32\"",
|
|
|
|
| 498 |
"inputElement": "\"vec4<f32>\"",
|
| 499 |
"outputElement": "\"vec4<f32>\"",
|
| 500 |
"inputScalar": "\"f32\"",
|
|
|
|
| 512 |
"name": "MultiHeadAttention.FlashQ32BroadcastF32Bias",
|
| 513 |
"shader": "attn-flash-q32-broadcast.wgsl.jinja",
|
| 514 |
"derive": { "layout": "\"bsh\"" },
|
| 515 |
+
"bindings": ["query", "key", "value", "bias", "output", "params_scores"],
|
| 516 |
"dispatch": {
|
| 517 |
"x": "ceilDiv(dim(shapes.queryT, 1), 32)",
|
| 518 |
"y": "attrs.num_heads",
|
| 519 |
"z": "dim(shapes.queryT, 0)"
|
| 520 |
+
}
|
|
|
|
| 521 |
}
|
| 522 |
]
|
| 523 |
},
|
|
|
|
| 526 |
"priority": 10,
|
| 527 |
"when": ["not present.biasT", "smallHeadParallelOk"],
|
| 528 |
"derive": {
|
|
|
|
|
|
|
| 529 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 530 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 531 |
"outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 536 |
"passes": [
|
| 537 |
{
|
| 538 |
"id": "main",
|
| 539 |
+
"name": "MultiHeadAttention.SmallHeadParallel",
|
| 540 |
"shader": "attn-small-head-parallel.wgsl.jinja",
|
| 541 |
+
"bindings": ["query", "key", "value", "output", "params_main"],
|
| 542 |
"dispatch": {
|
| 543 |
"x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 544 |
"y": "attrs.num_heads",
|
|
|
|
| 550 |
{
|
| 551 |
"id": "qkv_no_bias_tiled_attn_bias_nosg",
|
| 552 |
"priority": 17,
|
| 553 |
+
"when": ["not present.biasT", "qkvMaskContractOk", "not smallHeadParallelOk", "headDim % 4 == 0", "headDim <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"],
|
| 554 |
"derive": {
|
| 555 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
|
|
|
|
|
|
| 556 |
"blockM": "tunables.PREFILL_QUERY_TILE",
|
| 557 |
"vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads"
|
| 558 |
},
|
|
|
|
| 599 |
}
|
| 600 |
],
|
| 601 |
"dispatch": {
|
| 602 |
+
"x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
|
| 603 |
+
"y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)",
|
| 604 |
"z": 1
|
| 605 |
}
|
| 606 |
}
|
|
|
|
| 616 |
},
|
| 617 |
"derive": {
|
| 618 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 619 |
"qHidden": "dim(shapes.queryT, 2)",
|
|
|
|
|
|
|
| 620 |
"statSlots": "materializedSgmatStatSlots",
|
| 621 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 622 |
},
|
|
|
|
| 635 |
"name": "MultiHeadAttention.MaterializedScoresSgmat",
|
| 636 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 637 |
"derive": { "phase": "\"score\"", "emitRowStats": true },
|
| 638 |
+
"bindings": ["query_f32", "key_f32", "scores", "scorePartials", "params_scores"],
|
| 639 |
"dispatch": {
|
| 640 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 641 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
| 642 |
"z": "dim(shapes.queryT, 0) * attrs.num_heads"
|
| 643 |
+
}
|
|
|
|
| 644 |
},
|
| 645 |
{
|
| 646 |
"id": "rowstats",
|
| 647 |
"name": "MultiHeadAttention.MaterializedRowStatsCombine",
|
| 648 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 649 |
+
"bindings": ["scorePartials_f32", "rowStats", "params_rows"],
|
| 650 |
"dispatch": {
|
| 651 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 652 |
"y": 1,
|
|
|
|
| 658 |
"name": "MultiHeadAttention.MaterializedApplySgmat",
|
| 659 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 660 |
"derive": { "phase": "\"apply\"", "fusedSoftmax": true },
|
| 661 |
+
"bindings": ["scores_f32", "value_f32", "output_f32", "rowStats_f32", "params_apply"],
|
| 662 |
"dispatch": {
|
| 663 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 664 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
|
|
| 677 |
},
|
| 678 |
"derive": {
|
| 679 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 680 |
"qHidden": "dim(shapes.queryT, 2)",
|
|
|
|
|
|
|
| 681 |
"statSlots": "materializedSgmatStatSlots",
|
| 682 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 683 |
},
|
|
|
|
| 696 |
"name": "MultiHeadAttention.MaterializedScoresSgmatBias",
|
| 697 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 698 |
"derive": { "phase": "\"score\"", "emitRowStats": true },
|
| 699 |
+
"bindings": ["query_f32", "key_f32", "bias_f32", "scores", "scorePartials", "params_scores"],
|
| 700 |
"dispatch": {
|
| 701 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 702 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
| 703 |
"z": "dim(shapes.queryT, 0) * attrs.num_heads"
|
| 704 |
+
}
|
|
|
|
| 705 |
},
|
| 706 |
{
|
| 707 |
"id": "rowstats",
|
| 708 |
"name": "MultiHeadAttention.MaterializedRowStatsCombineBias",
|
| 709 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 710 |
+
"bindings": ["scorePartials_f32", "rowStats", "params_rows"],
|
| 711 |
"dispatch": {
|
| 712 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 713 |
"y": 1,
|
|
|
|
| 719 |
"name": "MultiHeadAttention.MaterializedApplySgmatBias",
|
| 720 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 721 |
"derive": { "phase": "\"apply\"", "fusedSoftmax": true },
|
| 722 |
+
"bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "rowStats_f32", "params_apply"],
|
| 723 |
"dispatch": {
|
| 724 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 725 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
|
|
| 738 |
},
|
| 739 |
"derive": {
|
| 740 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 741 |
"qHidden": "dim(shapes.queryT, 2)",
|
|
|
|
|
|
|
| 742 |
"statSlots": "materializedSgmatStatSlots",
|
| 743 |
"statQuerySeq": "dim(shapes.queryT, 1)",
|
| 744 |
"operandF16": true,
|
|
|
|
| 760 |
"name": "MultiHeadAttention.MaterializedScoresSgmatF16",
|
| 761 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 762 |
"derive": { "phase": "\"score\"", "emitRowStats": true },
|
| 763 |
+
"bindings": ["query", "key", "scores", "scorePartials", "params_scores"],
|
| 764 |
"dispatch": {
|
| 765 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 766 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
| 767 |
"z": "dim(shapes.queryT, 0) * attrs.num_heads"
|
| 768 |
+
}
|
|
|
|
| 769 |
},
|
| 770 |
{
|
| 771 |
"id": "rowstats",
|
| 772 |
"name": "MultiHeadAttention.MaterializedRowStatsCombine",
|
| 773 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 774 |
+
"bindings": ["scorePartials_f32", "rowStats", "params_rows"],
|
| 775 |
"dispatch": {
|
| 776 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 777 |
"y": 1,
|
|
|
|
| 783 |
"name": "MultiHeadAttention.MaterializedApplySgmatF16",
|
| 784 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 785 |
"derive": { "phase": "\"apply\"", "fusedSoftmax": true },
|
| 786 |
+
"bindings": ["scores_f32", "value", "output", "rowStats_f32", "params_apply"],
|
| 787 |
"dispatch": {
|
| 788 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 789 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
|
|
| 795 |
{
|
| 796 |
"id": "qkv_no_bias_flash_cluster_lpq4_nosg",
|
| 797 |
"priority": 20,
|
| 798 |
+
"when": ["not present.biasT", "flashShapeOk", "headDim % 16 == 0", "headDim % 32 != 0", "dim(shapes.queryT, 1) >= 31"],
|
| 799 |
"requires": {},
|
| 800 |
"derive": {
|
|
|
|
| 801 |
"hasCausal": true,
|
|
|
|
| 802 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 803 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 804 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 805 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 806 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 833 |
{
|
| 834 |
"id": "qkv_no_bias_flash_cluster_nosg",
|
| 835 |
"priority": 20,
|
| 836 |
+
"when": ["not present.biasT", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31"],
|
| 837 |
"requires": {},
|
| 838 |
"derive": {
|
|
|
|
| 839 |
"hasCausal": true,
|
|
|
|
| 840 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 841 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 842 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 843 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 844 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 871 |
{
|
| 872 |
"id": "qkv_bias_flash_cluster_nosg",
|
| 873 |
"priority": 19,
|
| 874 |
+
"when": ["biasOk", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31"],
|
| 875 |
"requires": {},
|
| 876 |
"derive": {
|
|
|
|
| 877 |
"hasCausal": true,
|
|
|
|
| 878 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 879 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 880 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 881 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 882 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 911 |
{
|
| 912 |
"id": "qkv_no_bias_flash_cluster_lpq4",
|
| 913 |
"priority": 22,
|
| 914 |
+
"when": ["not present.biasT", "flashShapeOk", "headDim % 16 == 0", "headDim % 32 != 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster4"],
|
| 915 |
"requires": { "features": ["subgroups"] },
|
| 916 |
"derive": {
|
|
|
|
| 917 |
"hasCausal": true,
|
|
|
|
| 918 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 919 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 920 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 921 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 922 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 940 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 941 |
"y": "attrs.num_heads",
|
| 942 |
"z": "dim(shapes.queryT, 0)"
|
| 943 |
+
}
|
|
|
|
| 944 |
}
|
| 945 |
]
|
| 946 |
},
|
| 947 |
{
|
| 948 |
"id": "qkv_no_bias_flash_cluster",
|
| 949 |
"priority": 22,
|
| 950 |
+
"when": ["not present.biasT", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
|
| 951 |
"requires": { "features": ["subgroups"] },
|
| 952 |
"derive": {
|
|
|
|
| 953 |
"hasCausal": true,
|
|
|
|
| 954 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 955 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 956 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 957 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 958 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 976 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 977 |
"y": "attrs.num_heads",
|
| 978 |
"z": "dim(shapes.queryT, 0)"
|
| 979 |
+
}
|
|
|
|
| 980 |
}
|
| 981 |
]
|
| 982 |
},
|
| 983 |
{
|
| 984 |
"id": "qkv_bias_flash_cluster",
|
| 985 |
"priority": 21,
|
| 986 |
+
"when": ["biasOk", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
|
| 987 |
"requires": { "features": ["subgroups"] },
|
| 988 |
"derive": {
|
|
|
|
| 989 |
"hasCausal": true,
|
|
|
|
| 990 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 991 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 992 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 993 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 994 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 1014 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 1015 |
"y": "attrs.num_heads",
|
| 1016 |
"z": "dim(shapes.queryT, 0)"
|
| 1017 |
+
}
|
|
|
|
| 1018 |
}
|
| 1019 |
]
|
| 1020 |
},
|
| 1021 |
{
|
| 1022 |
"id": "qkv_no_bias_flash_cluster_attn_bias",
|
| 1023 |
"priority": 22,
|
| 1024 |
+
"when": ["not present.biasT", "flashMaskShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
|
| 1025 |
"requires": { "features": ["subgroups"] },
|
| 1026 |
"derive": {
|
|
|
|
| 1027 |
"hasCausal": true,
|
|
|
|
|
|
|
|
|
|
| 1028 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1029 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1030 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1031 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1032 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1046 |
"name": "MultiHeadAttention.Flash",
|
| 1047 |
"shader": "attn-flash-prefill-cluster.wgsl.jinja",
|
| 1048 |
"derive": { "layout": "\"bsh\"" },
|
| 1049 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
|
| 1050 |
"dispatch": {
|
| 1051 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 1052 |
"y": "attrs.num_heads",
|
| 1053 |
"z": "dim(shapes.queryT, 0)"
|
| 1054 |
+
}
|
|
|
|
| 1055 |
}
|
| 1056 |
]
|
| 1057 |
},
|
| 1058 |
{
|
| 1059 |
"id": "qkv_bias_flash_cluster_attn_bias",
|
| 1060 |
"priority": 21,
|
| 1061 |
+
"when": ["biasOk", "flashMaskShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"],
|
| 1062 |
"requires": { "features": ["subgroups"] },
|
| 1063 |
"derive": {
|
|
|
|
| 1064 |
"hasCausal": true,
|
|
|
|
|
|
|
|
|
|
| 1065 |
"usesF16": "tensorDtypes.queryT == \"float16\"",
|
| 1066 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1067 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1068 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1069 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1085 |
"name": "MultiHeadAttention.Flash",
|
| 1086 |
"shader": "attn-flash-prefill-cluster.wgsl.jinja",
|
| 1087 |
"derive": { "layout": "\"bsh\"" },
|
| 1088 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
|
| 1089 |
"dispatch": {
|
| 1090 |
"x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)",
|
| 1091 |
"y": "attrs.num_heads",
|
| 1092 |
"z": "dim(shapes.queryT, 0)"
|
| 1093 |
+
}
|
|
|
|
| 1094 |
}
|
| 1095 |
]
|
| 1096 |
},
|
|
|
|
| 1100 |
"when": ["not present.biasT", "decodeSplitKNoBiasOk or shortQuerySplitKNoBiasOk"],
|
| 1101 |
"requires": {},
|
| 1102 |
"derive": {
|
|
|
|
| 1103 |
"useSubgroups": false,
|
|
|
|
| 1104 |
"splitQueries": true,
|
| 1105 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1106 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
|
|
|
| 1132 |
"name": "MultiHeadAttention.DecodeSplitKNoSg",
|
| 1133 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1134 |
"derive": { "layout": "\"bsh\"" },
|
| 1135 |
+
"bindings": ["query_query_t", "key_key_t", "value_value_t", "partial_out", "partial_stats", "params_kv_seq_scale"],
|
| 1136 |
"dispatch": {
|
| 1137 |
"x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
|
| 1138 |
"y": "attrs.num_heads",
|
|
|
|
| 1144 |
"name": "MultiHeadAttention.DecodeSplitKMergeNoSg",
|
| 1145 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1146 |
"derive": { "layout": "\"bsh\"" },
|
| 1147 |
+
"bindings": ["partial_out_merge", "partial_stats_merge", "output_merge"],
|
| 1148 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1149 |
}
|
| 1150 |
]
|
|
|
|
| 1155 |
"when": ["biasOk", "decodeSplitKBiasOk"],
|
| 1156 |
"requires": {},
|
| 1157 |
"derive": {
|
|
|
|
| 1158 |
"useSubgroups": false,
|
|
|
|
|
|
|
| 1159 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1160 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1161 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1187 |
"name": "MultiHeadAttention.DecodeSplitKBiasNoSg",
|
| 1188 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1189 |
"derive": { "layout": "\"bsh\"" },
|
| 1190 |
+
"bindings": ["query_query_t", "key_key_t", "value_value_t", "bias", "partial_out", "partial_stats", "params_kv_seq_scale"],
|
| 1191 |
"dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1192 |
},
|
| 1193 |
{
|
|
|
|
| 1195 |
"name": "MultiHeadAttention.DecodeSplitKMergeBiasNoSg",
|
| 1196 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1197 |
"derive": { "layout": "\"bsh\"" },
|
| 1198 |
+
"bindings": ["partial_out_merge", "partial_stats_merge", "bias", "output_merge"],
|
| 1199 |
"dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1200 |
}
|
| 1201 |
]
|
|
|
|
| 1207 |
"demoteWhen": ["decodeSplitKPortablePreferred"],
|
| 1208 |
"requires": { "features": ["subgroups"] },
|
| 1209 |
"derive": {
|
|
|
|
|
|
|
| 1210 |
"splitQueries": true,
|
| 1211 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1212 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
|
|
|
| 1238 |
"name": "MultiHeadAttention.DecodeSplitK",
|
| 1239 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1240 |
"derive": { "layout": "\"bsh\"" },
|
| 1241 |
+
"bindings": ["query_query_t", "key_key_t", "value_value_t", "partial_out", "partial_stats", "params_kv_seq_scale"],
|
| 1242 |
"dispatch": {
|
| 1243 |
"x": "dim(shapes.queryT, 1) * noBiasSplitKCount",
|
| 1244 |
"y": "attrs.num_heads",
|
| 1245 |
"z": "dim(shapes.queryT, 0)"
|
| 1246 |
+
}
|
|
|
|
| 1247 |
},
|
| 1248 |
{
|
| 1249 |
"id": "merge",
|
| 1250 |
"name": "MultiHeadAttention.DecodeSplitKMerge",
|
| 1251 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1252 |
"derive": { "layout": "\"bsh\"" },
|
| 1253 |
+
"bindings": ["partial_out_merge", "partial_stats_merge", "output_merge"],
|
| 1254 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1255 |
}
|
| 1256 |
]
|
|
|
|
| 1261 |
"when": ["biasOk", "decodeSplitKBiasOk"],
|
| 1262 |
"requires": { "features": ["subgroups"] },
|
| 1263 |
"derive": {
|
|
|
|
|
|
|
|
|
|
| 1264 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1265 |
"inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1266 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1292 |
"name": "MultiHeadAttention.DecodeSplitKBias",
|
| 1293 |
"shader": "attn-flash-decode-splitk.wgsl.jinja",
|
| 1294 |
"derive": { "layout": "\"bsh\"" },
|
| 1295 |
+
"bindings": ["query_query_t", "key_key_t", "value_value_t", "bias", "partial_out", "partial_stats", "params_kv_seq_scale"],
|
| 1296 |
+
"dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
|
|
|
| 1297 |
},
|
| 1298 |
{
|
| 1299 |
"id": "merge",
|
| 1300 |
"name": "MultiHeadAttention.DecodeSplitKMergeBias",
|
| 1301 |
"shader": "attn-flash-decode-splitk-merge.wgsl.jinja",
|
| 1302 |
"derive": { "layout": "\"bsh\"" },
|
| 1303 |
+
"bindings": ["partial_out_merge", "partial_stats_merge", "bias", "output_merge"],
|
| 1304 |
"dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1305 |
}
|
| 1306 |
]
|
|
|
|
| 1313 |
"derive": {
|
| 1314 |
"qNumHeads": "attrs.num_heads",
|
| 1315 |
"kvNumHeads": "attrs.num_heads",
|
|
|
|
| 1316 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1317 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1318 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
|
|
|
| 1322 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 1323 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
| 1324 |
"applyTileN": "materializedApplyTileN",
|
|
|
|
| 1325 |
"useSubgroups": "device.features.has(\"subgroups\")"
|
| 1326 |
},
|
| 1327 |
"intermediates": [
|
|
|
|
| 1337 |
"name": "MultiHeadAttention.MaterializedScores",
|
| 1338 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1339 |
"derive": { "layout": "\"bsh\"" },
|
| 1340 |
+
"bindings": ["query_f32", "key_f32", "scores", "params_scores"],
|
| 1341 |
"dispatch": {
|
| 1342 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1343 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1349 |
"name": "MultiHeadAttention.MaterializedSoftmax",
|
| 1350 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 1351 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 1352 |
+
"bindings": ["scores_softmax", "params_rows"],
|
| 1353 |
"dispatch": {
|
| 1354 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1355 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1356 |
"z": 1
|
| 1357 |
+
}
|
|
|
|
| 1358 |
},
|
| 1359 |
{
|
| 1360 |
"id": "apply",
|
| 1361 |
"name": "MultiHeadAttention.MaterializedApply",
|
| 1362 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1363 |
"derive": { "layout": "\"bsh\"" },
|
| 1364 |
+
"bindings": ["scores_f32", "value_f32", "output_f32", "params_apply"],
|
| 1365 |
"dispatch": {
|
| 1366 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1367 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1377 |
"derive": {
|
| 1378 |
"qNumHeads": "attrs.num_heads",
|
| 1379 |
"kvNumHeads": "attrs.num_heads",
|
|
|
|
| 1380 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1381 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1382 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
|
|
|
| 1386 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 1387 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
| 1388 |
"applyTileN": "materializedApplyTileN",
|
|
|
|
| 1389 |
"useSubgroups": "device.features.has(\"subgroups\")"
|
| 1390 |
},
|
| 1391 |
"intermediates": [
|
|
|
|
| 1401 |
"name": "MultiHeadAttention.MaterializedScoresBias",
|
| 1402 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1403 |
"derive": { "layout": "\"bsh\"" },
|
| 1404 |
+
"bindings": ["query_f32", "key_f32", "bias_f32", "scores", "params_scores"],
|
| 1405 |
"dispatch": {
|
| 1406 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1407 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1413 |
"name": "MultiHeadAttention.MaterializedSoftmaxBias",
|
| 1414 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 1415 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 1416 |
+
"bindings": ["scores_softmax", "params_rows"],
|
| 1417 |
"dispatch": {
|
| 1418 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1419 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 1420 |
"z": 1
|
| 1421 |
+
}
|
|
|
|
| 1422 |
},
|
| 1423 |
{
|
| 1424 |
"id": "apply",
|
| 1425 |
"name": "MultiHeadAttention.MaterializedApplyBias",
|
| 1426 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1427 |
"derive": { "layout": "\"bsh\"" },
|
| 1428 |
+
"bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "params_apply"],
|
| 1429 |
"dispatch": {
|
| 1430 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1431 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1442 |
"derive": {
|
| 1443 |
"qNumHeads": "attrs.num_heads",
|
| 1444 |
"kvNumHeads": "attrs.num_heads",
|
|
|
|
| 1445 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1446 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1447 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
| 1448 |
"materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
|
| 1449 |
"applyTileN": "materializedApplyTileN",
|
|
|
|
|
|
|
| 1450 |
"statSlots": "materializedGemmStatSlots",
|
| 1451 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 1452 |
},
|
|
|
|
| 1465 |
"name": "MultiHeadAttention.MaterializedScoresFused",
|
| 1466 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1467 |
"derive": { "layout": "\"bsh\"", "emitRowStats": true },
|
| 1468 |
+
"bindings": ["query_f32", "key_f32", "scores", "scorePartials", "params_scores"],
|
| 1469 |
"dispatch": {
|
| 1470 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1471 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1477 |
"name": "MultiHeadAttention.MaterializedGemmRowStatsCombine",
|
| 1478 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 1479 |
"derive": { "maxOnly": true },
|
| 1480 |
+
"bindings": ["scorePartials_f32", "rowStats", "params_rows"],
|
| 1481 |
"dispatch": {
|
| 1482 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1483 |
"y": 1,
|
|
|
|
| 1489 |
"name": "MultiHeadAttention.MaterializedApplyFused",
|
| 1490 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1491 |
"derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
|
| 1492 |
+
"bindings": ["scores_f32", "value_f32", "output_f32", "rowStats_f32", "params_apply"],
|
| 1493 |
"dispatch": {
|
| 1494 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1495 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1505 |
"derive": {
|
| 1506 |
"qNumHeads": "attrs.num_heads",
|
| 1507 |
"kvNumHeads": "attrs.num_heads",
|
|
|
|
| 1508 |
"materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE",
|
| 1509 |
"materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE",
|
| 1510 |
"materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE",
|
| 1511 |
"materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM",
|
| 1512 |
"applyTileN": "materializedApplyTileN",
|
|
|
|
|
|
|
| 1513 |
"statSlots": "materializedGemmStatSlots",
|
| 1514 |
"statQuerySeq": "dim(shapes.queryT, 1)"
|
| 1515 |
},
|
|
|
|
| 1528 |
"name": "MultiHeadAttention.MaterializedScoresBiasFused",
|
| 1529 |
"shader": "attn-materialized-score-f32.wgsl.jinja",
|
| 1530 |
"derive": { "layout": "\"bsh\"", "emitRowStats": true },
|
| 1531 |
+
"bindings": ["query_f32", "key_f32", "bias_f32", "scores", "scorePartials", "params_scores"],
|
| 1532 |
"dispatch": {
|
| 1533 |
"x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)",
|
| 1534 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1540 |
"name": "MultiHeadAttention.MaterializedGemmRowStatsCombineBias",
|
| 1541 |
"shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja",
|
| 1542 |
"derive": { "maxOnly": true },
|
| 1543 |
+
"bindings": ["scorePartials_f32", "rowStats", "params_rows"],
|
| 1544 |
"dispatch": {
|
| 1545 |
"x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1546 |
"y": 1,
|
|
|
|
| 1552 |
"name": "MultiHeadAttention.MaterializedApplyBiasFused",
|
| 1553 |
"shader": "attn-materialized-apply-f32.wgsl.jinja",
|
| 1554 |
"derive": { "layout": "\"bsh\"", "fusedSoftmax": true },
|
| 1555 |
+
"bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "rowStats_f32", "params_apply"],
|
| 1556 |
"dispatch": {
|
| 1557 |
"x": "ceilDiv(headDim, applyTileN)",
|
| 1558 |
"y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)",
|
|
|
|
| 1564 |
{
|
| 1565 |
"id": "qkv_no_bias_flash_q32_broadcast",
|
| 1566 |
"priority": 30,
|
| 1567 |
+
"when": ["tensorDtypes.queryT == \"float16\"", "not present.biasT", "flashShapeOk", "(dim(shapes.queryT, 2) / attrs.num_heads) % 4 == 0", "(dim(shapes.queryT, 2) / attrs.num_heads) >= 64", "(dim(shapes.queryT, 2) / attrs.num_heads) <= 256", "dim(shapes.queryT, 1) >= 31", "wave32SubgroupsUsable"],
|
| 1568 |
"requires": { "features": ["subgroups", "shader-f16"] },
|
| 1569 |
"derive": {
|
| 1570 |
"usesF16": true,
|
| 1571 |
"scalar": "\"f16\"",
|
|
|
|
| 1572 |
"inputElement": "\"vec4<f16>\"",
|
| 1573 |
"outputElement": "\"vec4<f16>\"",
|
| 1574 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 1589 |
"x": "ceilDiv(dim(shapes.queryT, 1), 32)",
|
| 1590 |
"y": "attrs.num_heads",
|
| 1591 |
"z": "dim(shapes.queryT, 0)"
|
| 1592 |
+
}
|
|
|
|
| 1593 |
}
|
| 1594 |
]
|
| 1595 |
},
|
|
|
|
| 1601 |
"derive": {
|
| 1602 |
"usesF16": true,
|
| 1603 |
"scalar": "\"f16\"",
|
|
|
|
| 1604 |
"inputElement": "\"vec4<f16>\"",
|
| 1605 |
"outputElement": "\"vec4<f16>\"",
|
| 1606 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 1631 |
"priority": 0,
|
| 1632 |
"when": ["not present.biasT and fallbackMaskShapeOk"],
|
| 1633 |
"derive": {
|
|
|
|
|
|
|
| 1634 |
"hasCausal": true,
|
|
|
|
| 1635 |
"hasKeyLimit": false,
|
|
|
|
|
|
|
|
|
|
| 1636 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1637 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1638 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1651 |
"name": "MultiHeadAttention",
|
| 1652 |
"shader": "attn-online-scalar.wgsl.jinja",
|
| 1653 |
"derive": { "layout": "\"bsh\"" },
|
| 1654 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
|
| 1655 |
"dispatch": {
|
| 1656 |
"x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1657 |
"y": "attrs.num_heads",
|
|
|
|
| 1665 |
"priority": 0,
|
| 1666 |
"when": ["biasOk and fallbackMaskShapeOk"],
|
| 1667 |
"derive": {
|
|
|
|
|
|
|
| 1668 |
"hasCausal": true,
|
|
|
|
| 1669 |
"hasKeyLimit": false,
|
|
|
|
|
|
|
|
|
|
| 1670 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1671 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1672 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1685 |
"name": "MultiHeadAttention",
|
| 1686 |
"shader": "attn-online-scalar.wgsl.jinja",
|
| 1687 |
"derive": { "layout": "\"bsh\"" },
|
| 1688 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
|
| 1689 |
"dispatch": {
|
| 1690 |
"x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
|
| 1691 |
"y": "attrs.num_heads",
|
|
|
|
| 1699 |
"priority": 0,
|
| 1700 |
"when": ["not present.biasT and fallbackShapeOk and not flashShapeOk"],
|
| 1701 |
"derive": {
|
|
|
|
|
|
|
| 1702 |
"hasCausal": true,
|
|
|
|
| 1703 |
"hasKeyLimit": false,
|
|
|
|
|
|
|
| 1704 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1705 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1706 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1732 |
"priority": 0,
|
| 1733 |
"when": ["biasOk and fallbackShapeOk and not flashShapeOk"],
|
| 1734 |
"derive": {
|
|
|
|
|
|
|
| 1735 |
"hasCausal": true,
|
|
|
|
| 1736 |
"hasKeyLimit": false,
|
|
|
|
|
|
|
| 1737 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1738 |
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
| 1739 |
"inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1765 |
"priority": 20,
|
| 1766 |
"when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashShapeOk"],
|
| 1767 |
"derive": {
|
|
|
|
|
|
|
| 1768 |
"hasCausal": true,
|
|
|
|
| 1769 |
"combineSubgroups": true,
|
|
|
|
|
|
|
| 1770 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1771 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1772 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1773 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1774 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 1785 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1786 |
"derive": { "layout": "\"bsh\"" },
|
| 1787 |
"bindings": ["query", "key", "value", "output", "params"],
|
| 1788 |
+
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
|
|
|
| 1789 |
}
|
| 1790 |
]
|
| 1791 |
},
|
|
|
|
| 1794 |
"priority": 20,
|
| 1795 |
"when": ["device.features.has(\"subgroups\")", "biasOk", "flashShapeOk"],
|
| 1796 |
"derive": {
|
|
|
|
|
|
|
| 1797 |
"hasCausal": true,
|
|
|
|
| 1798 |
"combineSubgroups": true,
|
|
|
|
|
|
|
| 1799 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1800 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1801 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1802 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1803 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1816 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1817 |
"derive": { "layout": "\"bsh\"" },
|
| 1818 |
"bindings": ["query", "key", "value", "bias", "output", "params"],
|
| 1819 |
+
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
|
|
|
| 1820 |
}
|
| 1821 |
]
|
| 1822 |
},
|
|
|
|
| 1825 |
"priority": 18,
|
| 1826 |
"when": ["true", "not present.biasT", "flashShapeOk"],
|
| 1827 |
"derive": {
|
|
|
|
|
|
|
| 1828 |
"hasCausal": true,
|
|
|
|
| 1829 |
"combineSubgroups": false,
|
|
|
|
|
|
|
| 1830 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1831 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1832 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1833 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1834 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 1854 |
"priority": 17,
|
| 1855 |
"when": ["true", "biasOk", "flashShapeOk"],
|
| 1856 |
"derive": {
|
|
|
|
|
|
|
| 1857 |
"hasCausal": true,
|
|
|
|
| 1858 |
"combineSubgroups": false,
|
|
|
|
|
|
|
| 1859 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1860 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1861 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1862 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1863 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1885 |
"priority": 20,
|
| 1886 |
"when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashMaskShapeOk"],
|
| 1887 |
"derive": {
|
|
|
|
|
|
|
| 1888 |
"hasCausal": true,
|
|
|
|
| 1889 |
"combineSubgroups": true,
|
|
|
|
|
|
|
| 1890 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1891 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1892 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1893 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1894 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 1904 |
"name": "MultiHeadAttention.Flash",
|
| 1905 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1906 |
"derive": { "layout": "\"bsh\"" },
|
| 1907 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
|
| 1908 |
+
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
|
|
|
| 1909 |
}
|
| 1910 |
]
|
| 1911 |
},
|
|
|
|
| 1914 |
"priority": 20,
|
| 1915 |
"when": ["device.features.has(\"subgroups\")", "biasOk", "flashMaskShapeOk"],
|
| 1916 |
"derive": {
|
|
|
|
|
|
|
| 1917 |
"hasCausal": true,
|
|
|
|
| 1918 |
"combineSubgroups": true,
|
|
|
|
|
|
|
| 1919 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1920 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1921 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1922 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1923 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1935 |
"name": "MultiHeadAttention.Flash",
|
| 1936 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1937 |
"derive": { "layout": "\"bsh\"" },
|
| 1938 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
|
| 1939 |
+
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
|
|
|
| 1940 |
}
|
| 1941 |
]
|
| 1942 |
},
|
|
|
|
| 1945 |
"priority": 18,
|
| 1946 |
"when": ["true", "not present.biasT", "flashMaskShapeOk"],
|
| 1947 |
"derive": {
|
|
|
|
|
|
|
| 1948 |
"hasCausal": true,
|
|
|
|
| 1949 |
"combineSubgroups": false,
|
|
|
|
|
|
|
| 1950 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1951 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1952 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1953 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1954 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 1964 |
"name": "MultiHeadAttention.NoBiasOnlineFlashNoSg",
|
| 1965 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1966 |
"derive": { "layout": "\"bsh\"" },
|
| 1967 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"],
|
| 1968 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 1969 |
}
|
| 1970 |
]
|
|
|
|
| 1974 |
"priority": 17,
|
| 1975 |
"when": ["true", "biasOk", "flashMaskShapeOk"],
|
| 1976 |
"derive": {
|
|
|
|
|
|
|
| 1977 |
"hasCausal": true,
|
|
|
|
| 1978 |
"combineSubgroups": false,
|
|
|
|
|
|
|
| 1979 |
"maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1980 |
"scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1981 |
"inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1982 |
"outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
|
| 1983 |
"inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
|
|
|
|
| 1995 |
"name": "MultiHeadAttention.BiasOnlineFlashNoSg",
|
| 1996 |
"shader": "attn-flash-online.wgsl.jinja",
|
| 1997 |
"derive": { "layout": "\"bsh\"" },
|
| 1998 |
+
"bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"],
|
| 1999 |
"dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" }
|
| 2000 |
}
|
| 2001 |
]
|
|
|
|
| 2010 |
},
|
| 2011 |
"derive": {
|
| 2012 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 2013 |
"qHidden": "dim(shapes.queryT, 2)",
|
| 2014 |
"materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
|
| 2015 |
"materializedSoftmaxCols": "dim(shapes.keyT, 1)",
|
| 2016 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 2017 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
|
|
|
| 2018 |
"useSubgroups": true
|
| 2019 |
},
|
| 2020 |
"intermediates": [
|
|
|
|
| 2030 |
"name": "MultiHeadAttention.MaterializedScoresSgmat",
|
| 2031 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2032 |
"derive": { "phase": "\"score\"" },
|
| 2033 |
+
"bindings": ["query_f32", "key_f32", "scores", "params_scores"],
|
| 2034 |
"dispatch": {
|
| 2035 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 2036 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
|
|
| 2042 |
"name": "MultiHeadAttention.MaterializedSoftmaxSgmat",
|
| 2043 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 2044 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 2045 |
+
"bindings": ["scores_softmax", "params_rows"],
|
| 2046 |
"dispatch": {
|
| 2047 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2048 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2049 |
"z": 1
|
| 2050 |
+
}
|
|
|
|
| 2051 |
},
|
| 2052 |
{
|
| 2053 |
"id": "apply",
|
| 2054 |
"name": "MultiHeadAttention.MaterializedApplySgmat",
|
| 2055 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2056 |
"derive": { "phase": "\"apply\"" },
|
| 2057 |
+
"bindings": ["scores_f32", "value_f32", "output_f32", "params_apply"],
|
| 2058 |
"dispatch": {
|
| 2059 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 2060 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
|
|
| 2073 |
},
|
| 2074 |
"derive": {
|
| 2075 |
"qNumHeads": "attrs.num_heads",
|
|
|
|
| 2076 |
"qHidden": "dim(shapes.queryT, 2)",
|
| 2077 |
"materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE",
|
| 2078 |
"materializedSoftmaxCols": "dim(shapes.keyT, 1)",
|
| 2079 |
"materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4",
|
| 2080 |
"softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"",
|
|
|
|
| 2081 |
"useSubgroups": true
|
| 2082 |
},
|
| 2083 |
"intermediates": [
|
|
|
|
| 2093 |
"name": "MultiHeadAttention.MaterializedScoresSgmatBias",
|
| 2094 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2095 |
"derive": { "phase": "\"score\"" },
|
| 2096 |
+
"bindings": ["query_f32", "key_f32", "bias_f32", "scores", "params_scores"],
|
| 2097 |
"dispatch": {
|
| 2098 |
"x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)",
|
| 2099 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
|
|
|
| 2105 |
"name": "MultiHeadAttention.MaterializedSoftmaxSgmatBias",
|
| 2106 |
"shader": "attn-materialized-softmax-f32.wgsl.jinja",
|
| 2107 |
"derive": { "cacheVec4": "materializedCachedSoftmaxOk" },
|
| 2108 |
+
"bindings": ["scores_softmax", "params_rows"],
|
| 2109 |
"dispatch": {
|
| 2110 |
"x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2111 |
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)",
|
| 2112 |
"z": 1
|
| 2113 |
+
}
|
|
|
|
| 2114 |
},
|
| 2115 |
{
|
| 2116 |
"id": "apply",
|
| 2117 |
"name": "MultiHeadAttention.MaterializedApplySgmatBias",
|
| 2118 |
"shader": "attn-materialized-sgmat-f32.wgsl.jinja",
|
| 2119 |
"derive": { "phase": "\"apply\"" },
|
| 2120 |
+
"bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "params_apply"],
|
| 2121 |
"dispatch": {
|
| 2122 |
"x": "ceilDiv(headDim, materializedSgmatKeyTile)",
|
| 2123 |
"y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)",
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,36 +1,39 @@
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.MultiHeadAttention",
|
| 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-apply-f32.wgsl.jinja": "
|
| 17 |
-
"attn-materialized-rowstats-combine-f32.wgsl.jinja": "
|
| 18 |
-
"attn-materialized-score-f32.wgsl.jinja": "
|
| 19 |
-
"attn-materialized-sgmat-f32.wgsl.jinja": "
|
| 20 |
-
"attn-materialized-softmax-f32.wgsl.jinja": "
|
| 21 |
-
"attn-online-scalar.wgsl.jinja": "
|
| 22 |
-
"attn-small-head-parallel.wgsl.jinja": "
|
| 23 |
-
"
|
| 24 |
-
"
|
| 25 |
-
"
|
| 26 |
-
"mha-small-seq.wgsl.jinja": "
|
| 27 |
-
"
|
|
|
|
| 28 |
}
|
| 29 |
},
|
| 30 |
-
"provenance": { "kernel": { "sha": "
|
| 31 |
"webgpu": {
|
| 32 |
-
"manifestSpec": "2.
|
| 33 |
"variants": {
|
|
|
|
|
|
|
| 34 |
"qkv_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
|
| 35 |
"qkv_no_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
|
| 36 |
"qkv_no_bias_small_seq": ["mha-small-seq.wgsl.jinja"],
|
|
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.MultiHeadAttention",
|
| 3 |
+
"id": "_com_microsoft_multiheadattention_webgpu_e942a19",
|
| 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": "JIhoYhGkuat/Wrvix6ddD19+bmzHnc0yJdwWbW9o+nw=",
|
| 12 |
+
"attn-flash-decode-splitk.wgsl.jinja": "tDj5VlltkM0N99vjpj6z5wtlIv00benlaAboa4e5vWI=",
|
| 13 |
+
"attn-flash-online.wgsl.jinja": "1DGpp4kzd/v4gzAviJKnYGCT7BwDomzXK3He6CIT+1E=",
|
| 14 |
+
"attn-flash-prefill-cluster.wgsl.jinja": "pjDvKLZ3CSluSAPNfb9j3U7yKEy+wVOfOAyZ1jdrjv8=",
|
| 15 |
+
"attn-flash-q32-broadcast.wgsl.jinja": "puOnD1vGGeVg6HmaUi8k4JAZJviJM1dXWJyMZ2xuURM=",
|
| 16 |
+
"attn-materialized-apply-f32.wgsl.jinja": "E2Vas8y/zcjHQb4Ks5TPnC4/lAhVRawXR5S1OwJB1Tc=",
|
| 17 |
+
"attn-materialized-rowstats-combine-f32.wgsl.jinja": "CBltqk0CN9vvscb6D29o1T43vHWZYlU9fFpO9av2Sa8=",
|
| 18 |
+
"attn-materialized-score-f32.wgsl.jinja": "bxW7HdJcF47NTbyXVHcaWrtCO0xCJ8pW5GW9/FTuGGg=",
|
| 19 |
+
"attn-materialized-sgmat-f32.wgsl.jinja": "3e4ZOgVD5XSnYjVldRVq5QL72KiloCkr91OZj/VU0ZU=",
|
| 20 |
+
"attn-materialized-softmax-f32.wgsl.jinja": "bpHIAAuUa+nrniTKxXEScXT5rJr5nObXY/sJ81hby5A=",
|
| 21 |
+
"attn-online-scalar.wgsl.jinja": "AbpGvmzka15nPEj3QZPyMaW48Lo/9TqtndP0SdyqV/U=",
|
| 22 |
+
"attn-small-head-parallel.wgsl.jinja": "Y0otPUGeAa4DzQUi+LnWLRTATc1i/awpIipRqbGDoQ0=",
|
| 23 |
+
"attn-small-head-value.wgsl.jinja": "+6NLq7AJT3lRkAwsW2AvjGE0nhkD7dkBLqh6SpOFjy4=",
|
| 24 |
+
"bench.json": "ateXBEoOVVMHNfiuXPmDW6sKzU1NRVRA9ndzslgsD2A=",
|
| 25 |
+
"manifest.json": "VPWToEdcoP2Mr6k2oQuf/BhwUZdbcPl9J5YLZsMBMaE=",
|
| 26 |
+
"mha-small-seq-blocked.wgsl.jinja": "WudxiAnGnY7dSEnlhHtwvQKvm4AN3ngQ9cbBfI6OUBM=",
|
| 27 |
+
"mha-small-seq.wgsl.jinja": "1MR+JuEbffrS+QkuAZzFzWCLPX4VX9RJ5lmNIOm0/o0=",
|
| 28 |
+
"test.json": "IEMAo5OqrMxRC7xc5Y/8xmjWObXuWw1SJhJeGGzDhwc="
|
| 29 |
}
|
| 30 |
},
|
| 31 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 32 |
"webgpu": {
|
| 33 |
+
"manifestSpec": "2.1",
|
| 34 |
"variants": {
|
| 35 |
+
"qkv_no_bias_small_head_value_subgroups": ["attn-small-head-value.wgsl.jinja"],
|
| 36 |
+
"qkv_no_bias_small_head_value_tree": ["attn-small-head-value.wgsl.jinja"],
|
| 37 |
"qkv_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
|
| 38 |
"qkv_no_bias_small_seq_blocked": ["mha-small-seq-blocked.wgsl.jinja"],
|
| 39 |
"qkv_no_bias_small_seq": ["mha-small-seq.wgsl.jinja"],
|
build/webgpu/mha-small-seq-blocked.wgsl.jinja
CHANGED
|
@@ -37,19 +37,20 @@ 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 |
}
|
| 43 |
-
{%
|
| 44 |
-
fn scale_value() -> f32 {
|
| 45 |
if (params.scale != 0.0) { return params.scale; }
|
| 46 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 47 |
}
|
| 48 |
-
|
| 49 |
{% if hasBias %}
|
|
|
|
|
|
|
| 50 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 51 |
let offset = base + d4 * 4u;
|
| 52 |
-
return vec4<f32>(bias[offset], bias[offset + 1u], bias[offset + 2u], bias[offset + 3u]);
|
| 53 |
}
|
| 54 |
{% endif %}
|
| 55 |
|
|
|
|
| 37 |
let equalFiniteMax = select(false, value == maxValue, is_finite_f32(maxValue));
|
| 38 |
return select(value - maxValue, 0.0, equalFiniteMax);
|
| 39 |
}
|
| 40 |
+
|
| 41 |
fn exp_shift(value: f32, maxValue: f32) -> f32 {
|
| 42 |
return exp(shifted_value(value, maxValue));
|
| 43 |
}
|
| 44 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 45 |
if (params.scale != 0.0) { return params.scale; }
|
| 46 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 47 |
}
|
|
|
|
| 48 |
{% if hasBias %}
|
| 49 |
+
{% set BW = "" %}
|
| 50 |
+
{% set BC = "" %}
|
| 51 |
fn load_bias4(base: u32, d4: u32) -> vec4<f32> {
|
| 52 |
let offset = base + d4 * 4u;
|
| 53 |
+
return vec4<f32>({{ BW }}bias[offset]{{ BC }}, {{ BW }}bias[offset + 1u]{{ BC }}, {{ BW }}bias[offset + 2u]{{ BC }}, {{ BW }}bias[offset + 3u]{{ BC }});
|
| 54 |
}
|
| 55 |
{% endif %}
|
| 56 |
|
build/webgpu/mha-small-seq.wgsl.jinja
CHANGED
|
@@ -1,4 +1,26 @@
|
|
| 1 |
{{ env.wgsl.resourceDeclarations }}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2 |
|
| 3 |
// Whole-head attention for tiny bidirectional sequences without bias or mask.
|
| 4 |
// One workgroup owns one (head, batch) and stages the complete K and V planes in
|
|
@@ -9,13 +31,11 @@ const HEAD_DIM: u32 = {{ headDim }}u;
|
|
| 9 |
const KV_SEQ: u32 = {{ kvSeq }}u;
|
| 10 |
const HIDDEN: u32 = {{ hidden }}u;
|
| 11 |
const WG: u32 = {{ workgroupSize }}u;
|
| 12 |
-
{%
|
| 13 |
-
fn scale_value() -> f32 {
|
| 14 |
if (params.scale != 0.0) { return params.scale; }
|
| 15 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 16 |
}
|
| 17 |
|
| 18 |
-
|
| 19 |
var<workgroup> kShared: array<f32, KV_SEQ * HEAD_DIM>;
|
| 20 |
var<workgroup> vShared: array<f32, KV_SEQ * HEAD_DIM>;
|
| 21 |
|
|
@@ -56,11 +76,11 @@ fn main(
|
|
| 56 |
var scores: array<f32, KV_SEQ>;
|
| 57 |
var maxScore = -3.4028234663852886e38;
|
| 58 |
for (var k = 0u; k < KV_SEQ; k++) {
|
| 59 |
-
var dot = 0.0;
|
| 60 |
for (var d = 0u; d < HEAD_DIM; d++) {
|
| 61 |
-
dot
|
| 62 |
}
|
| 63 |
-
let s = dot * scale;
|
| 64 |
scores[k] = s;
|
| 65 |
maxScore = max(maxScore, s);
|
| 66 |
}
|
|
|
|
| 1 |
{{ env.wgsl.resourceDeclarations }}
|
| 2 |
+
{% set dotType = "f32" %}
|
| 3 |
+
// Retain product and addition residuals across a dot product. Explicit fma
|
| 4 |
+
// boundaries preserve the addition error transform under reassociation.
|
| 5 |
+
struct DotAccumulator {
|
| 6 |
+
hi: {{ dotType }},
|
| 7 |
+
lo: {{ dotType }},
|
| 8 |
+
}
|
| 9 |
+
fn dot_accumulate(acc: DotAccumulator, a: {{ dotType }}, b: {{ dotType }}) -> DotAccumulator {
|
| 10 |
+
let product = fma(a, b, {{ dotType }}(0.0));
|
| 11 |
+
let productError = fma(a, b, -product);
|
| 12 |
+
let sum = fma(acc.hi, {{ dotType }}(1.0), product);
|
| 13 |
+
let bv = fma({{ dotType }}(-1.0), acc.hi, sum);
|
| 14 |
+
let av = fma({{ dotType }}(-1.0), bv, sum);
|
| 15 |
+
let ae = fma({{ dotType }}(-1.0), av, acc.hi);
|
| 16 |
+
let be = fma({{ dotType }}(-1.0), bv, product);
|
| 17 |
+
let error = fma(fma(fma(ae, {{ dotType }}(1.0), be), {{ dotType }}(1.0), acc.lo), {{ dotType }}(1.0), productError);
|
| 18 |
+
let hi = fma(sum, {{ dotType }}(1.0), error);
|
| 19 |
+
return DotAccumulator(hi, fma({{ dotType }}(-1.0), fma({{ dotType }}(-1.0), sum, hi), error));
|
| 20 |
+
}
|
| 21 |
+
fn dot_value(acc: DotAccumulator) -> {{ dotType }} {
|
| 22 |
+
return fma(acc.hi, {{ dotType }}(1.0), acc.lo);
|
| 23 |
+
}
|
| 24 |
|
| 25 |
// Whole-head attention for tiny bidirectional sequences without bias or mask.
|
| 26 |
// One workgroup owns one (head, batch) and stages the complete K and V planes in
|
|
|
|
| 31 |
const KV_SEQ: u32 = {{ kvSeq }}u;
|
| 32 |
const HIDDEN: u32 = {{ hidden }}u;
|
| 33 |
const WG: u32 = {{ workgroupSize }}u;
|
| 34 |
+
{% set ATTN_SCALE_DIM = "HEAD_DIM" %}fn scale_value() -> f32 {
|
|
|
|
| 35 |
if (params.scale != 0.0) { return params.scale; }
|
| 36 |
return inverseSqrt(f32({{ ATTN_SCALE_DIM }}));
|
| 37 |
}
|
| 38 |
|
|
|
|
| 39 |
var<workgroup> kShared: array<f32, KV_SEQ * HEAD_DIM>;
|
| 40 |
var<workgroup> vShared: array<f32, KV_SEQ * HEAD_DIM>;
|
| 41 |
|
|
|
|
| 76 |
var scores: array<f32, KV_SEQ>;
|
| 77 |
var maxScore = -3.4028234663852886e38;
|
| 78 |
for (var k = 0u; k < KV_SEQ; k++) {
|
| 79 |
+
var dot = DotAccumulator(0.0, 0.0);
|
| 80 |
for (var d = 0u; d < HEAD_DIM; d++) {
|
| 81 |
+
dot = dot_accumulate(dot, qRow[d], kShared[k * HEAD_DIM + d]);
|
| 82 |
}
|
| 83 |
+
let s = dot_value(dot) * scale;
|
| 84 |
scores[k] = s;
|
| 85 |
maxScore = max(maxScore, s);
|
| 86 |
}
|
build/webgpu/test.json
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|