sync 6fdf6301e2bb
Browse files- README.md +2 -2
- build/webgpu/chunk-out.wgsl.jinja +7 -35
- build/webgpu/chunk-prep.wgsl.jinja +1 -6
- build/webgpu/chunk-scan.wgsl.jinja +6 -52
- build/webgpu/chunk-ut.wgsl.jinja +5 -38
- build/webgpu/linear-attention.scalar.wgsl.jinja +10 -27
- build/webgpu/linear-attention.serial.wgsl.jinja +5 -15
- build/webgpu/linear-attention.vec4.wgsl.jinja +20 -39
- build/webgpu/manifest.json +0 -0
- build/webgpu/metadata.json +29 -29
- build/webgpu/test.json +190 -32
README.md
CHANGED
|
@@ -68,7 +68,7 @@ One implementation is selected per call from the device capabilities, the reques
|
|
| 68 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 69 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 70 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 71 |
-
- [`bench.json`](build/webgpu/bench.json) — benchmark
|
| 72 |
- [`chunk-out.wgsl.jinja`](build/webgpu/chunk-out.wgsl.jinja)
|
| 73 |
- [`chunk-prep.wgsl.jinja`](build/webgpu/chunk-prep.wgsl.jinja)
|
| 74 |
- [`chunk-scan.wgsl.jinja`](build/webgpu/chunk-scan.wgsl.jinja)
|
|
@@ -80,7 +80,7 @@ One implementation is selected per call from the device capabilities, the reques
|
|
| 80 |
## Use with `@huggingface/kernels`
|
| 81 |
|
| 82 |
```sh
|
| 83 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 84 |
```
|
| 85 |
|
| 86 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
|
|
| 68 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 69 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 70 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 71 |
+
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
|
| 72 |
- [`chunk-out.wgsl.jinja`](build/webgpu/chunk-out.wgsl.jinja)
|
| 73 |
- [`chunk-prep.wgsl.jinja`](build/webgpu/chunk-prep.wgsl.jinja)
|
| 74 |
- [`chunk-scan.wgsl.jinja`](build/webgpu/chunk-scan.wgsl.jinja)
|
|
|
|
| 80 |
## Use with `@huggingface/kernels`
|
| 81 |
|
| 82 |
```sh
|
| 83 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 84 |
```
|
| 85 |
|
| 86 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
build/webgpu/chunk-out.wgsl.jinja
CHANGED
|
@@ -1,20 +1,11 @@
|
|
| 1 |
-
{% macro read_scalar(name, index, dtype) %}
|
| 2 |
-
{% if dtype == "float16" %}
|
| 3 |
-
|
| 4 |
-
{% endmacro %}
|
| 5 |
-
{% macro write_scalar(expr, dtype) %}
|
| 6 |
-
{% if dtype == "float16" %}
|
| 7 |
-
f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
| 8 |
-
{% endmacro -%}
|
| 9 |
-
{% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{%- endmacro %}
|
| 10 |
-
|
| 11 |
{% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
|
| 12 |
-
{% if needQuery %}
|
| 13 |
// The output scale defaults to 1/sqrt(head_dim_k), matching the recurrent kernels.
|
| 14 |
fn out_scale() -> f32 {
|
| 15 |
return select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
|
| 16 |
}
|
| 17 |
-
{% endif %}
|
| 18 |
fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
| 19 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 20 |
{% if usesDecay %}
|
|
@@ -23,18 +14,6 @@ fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
|
| 23 |
return raw;
|
| 24 |
{% endif %}
|
| 25 |
}
|
| 26 |
-
{% if needKtil %}
|
| 27 |
-
|
| 28 |
-
fn k_til(bt: u32, h: u32, i: u32) -> f32 {
|
| 29 |
-
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 30 |
-
{% if usesDecay %}
|
| 31 |
-
return raw * {{ decay_at("bt", "h", "i") }};
|
| 32 |
-
{% else %}
|
| 33 |
-
return raw;
|
| 34 |
-
{% endif %}
|
| 35 |
-
}
|
| 36 |
-
{% endif %}
|
| 37 |
-
{% if needQuery %}
|
| 38 |
|
| 39 |
// The output scale rides on q_til so both output terms -- q_til * S and P * delta,
|
| 40 |
// where P is itself built from q_til -- pick it up without a second pass over y.
|
|
@@ -46,18 +25,12 @@ fn q_til(bt: u32, q_head: u32, {% if usesDecay %}h: u32, {% endif %}i: u32, scal
|
|
| 46 |
return raw * scale;
|
| 47 |
{% endif %}
|
| 48 |
}
|
| 49 |
-
{%
|
| 50 |
-
{%
|
| 51 |
-
{%
|
| 52 |
let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
|
| 53 |
let packed_out = max(params.qNumHeads, params.kvNumHeads) * HEAD_DIM_V;
|
| 54 |
-
{%
|
| 55 |
-
let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;
|
| 56 |
-
{%- endmacro %}
|
| 57 |
-
|
| 58 |
-
{% if queryDtype == "float16" %}
|
| 59 |
-
enable f16;
|
| 60 |
-
{% endif %}
|
| 61 |
{{ env.wgsl.resourceDeclarations }}
|
| 62 |
|
| 63 |
// com.microsoft.LinearAttention, chunked prefill: outputs, from entry states.
|
|
@@ -81,7 +54,6 @@ const TK: u32 = {{ chunkTileK }}u;
|
|
| 81 |
const WG: u32 = {{ workgroupSize }}u;
|
| 82 |
|
| 83 |
{{ emit_chunk_operands(needQuery=true) }}
|
| 84 |
-
|
| 85 |
var<workgroup> qblock: array<f32, TOKEN_ROWS * HEAD_DIM_K>;
|
| 86 |
var<workgroup> ktile: array<f32, CHUNK * TK>;
|
| 87 |
var<workgroup> pmrow: array<f32, TOKEN_ROWS * CHUNK>;
|
|
|
|
| 1 |
+
{% macro read_scalar(name, index, dtype) %}{% if dtype == "float16" %}f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}{% endmacro %}
|
| 2 |
+
{% macro write_scalar(expr, dtype) %}{% if dtype == "float16" %}f16({{ expr }}){% else %}{{ expr }}{% endif %}{% endmacro %}
|
| 3 |
+
{% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
{% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
|
|
|
|
| 5 |
// The output scale defaults to 1/sqrt(head_dim_k), matching the recurrent kernels.
|
| 6 |
fn out_scale() -> f32 {
|
| 7 |
return select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
|
| 8 |
}
|
|
|
|
| 9 |
fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
| 10 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 11 |
{% if usesDecay %}
|
|
|
|
| 14 |
return raw;
|
| 15 |
{% endif %}
|
| 16 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
|
| 18 |
// The output scale rides on q_til so both output terms -- q_til * S and P * delta,
|
| 19 |
// where P is itself built from q_til -- pick it up without a second pass over y.
|
|
|
|
| 25 |
return raw * scale;
|
| 26 |
{% endif %}
|
| 27 |
}
|
| 28 |
+
{% endmacro %}
|
| 29 |
+
{% macro q_til_call(bt, q_head, h, i, scale) %}q_til({{ bt }}, {{ q_head }}, {{ h ~ ", " if usesDecay else "" }}{{ i }}, {{ scale }}){% endmacro %}
|
| 30 |
+
{% macro emit_chunk_head_setup(needHeads=false) %}
|
| 31 |
let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
|
| 32 |
let packed_out = max(params.qNumHeads, params.kvNumHeads) * HEAD_DIM_V;
|
| 33 |
+
let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 34 |
{{ env.wgsl.resourceDeclarations }}
|
| 35 |
|
| 36 |
// com.microsoft.LinearAttention, chunked prefill: outputs, from entry states.
|
|
|
|
| 54 |
const WG: u32 = {{ workgroupSize }}u;
|
| 55 |
|
| 56 |
{{ emit_chunk_operands(needQuery=true) }}
|
|
|
|
| 57 |
var<workgroup> qblock: array<f32, TOKEN_ROWS * HEAD_DIM_K>;
|
| 58 |
var<workgroup> ktile: array<f32, CHUNK * TK>;
|
| 59 |
var<workgroup> pmrow: array<f32, TOKEN_ROWS * CHUNK>;
|
build/webgpu/chunk-prep.wgsl.jinja
CHANGED
|
@@ -1,9 +1,4 @@
|
|
| 1 |
-
{% macro read_scalar(name, index, dtype) %}
|
| 2 |
-
{% if dtype == "float16" %}
|
| 3 |
-
f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
|
| 4 |
-
{% endmacro %}{% if queryDtype == "float16" %}
|
| 5 |
-
enable f16;
|
| 6 |
-
{% endif %}
|
| 7 |
{{ env.wgsl.resourceDeclarations }}
|
| 8 |
|
| 9 |
// com.microsoft.LinearAttention, chunked prefill: within-chunk decay prefix.
|
|
|
|
| 1 |
+
{% macro read_scalar(name, index, dtype) %}{% if dtype == "float16" %}f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2 |
{{ env.wgsl.resourceDeclarations }}
|
| 3 |
|
| 4 |
// com.microsoft.LinearAttention, chunked prefill: within-chunk decay prefix.
|
build/webgpu/chunk-scan.wgsl.jinja
CHANGED
|
@@ -1,20 +1,7 @@
|
|
| 1 |
-
{% macro read_scalar(name, index, dtype) %}
|
| 2 |
-
{% if dtype == "float16" %}
|
| 3 |
-
|
| 4 |
-
{% endmacro %}
|
| 5 |
-
{% macro write_scalar(expr, dtype) %}
|
| 6 |
-
{% if dtype == "float16" %}
|
| 7 |
-
f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
| 8 |
-
{% endmacro -%}
|
| 9 |
-
{% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{%- endmacro %}
|
| 10 |
-
|
| 11 |
{% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
|
| 12 |
-
{% if needQuery %}
|
| 13 |
-
// The output scale defaults to 1/sqrt(head_dim_k), matching the recurrent kernels.
|
| 14 |
-
fn out_scale() -> f32 {
|
| 15 |
-
return select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
|
| 16 |
-
}
|
| 17 |
-
{% endif %}
|
| 18 |
fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
| 19 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 20 |
{% if usesDecay %}
|
|
@@ -23,41 +10,9 @@ fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
|
| 23 |
return raw;
|
| 24 |
{% endif %}
|
| 25 |
}
|
| 26 |
-
{%
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 30 |
-
{% if usesDecay %}
|
| 31 |
-
return raw * {{ decay_at("bt", "h", "i") }};
|
| 32 |
-
{% else %}
|
| 33 |
-
return raw;
|
| 34 |
-
{% endif %}
|
| 35 |
-
}
|
| 36 |
-
{% endif %}
|
| 37 |
-
{% if needQuery %}
|
| 38 |
-
|
| 39 |
-
// The output scale rides on q_til so both output terms -- q_til * S and P * delta,
|
| 40 |
-
// where P is itself built from q_til -- pick it up without a second pass over y.
|
| 41 |
-
fn q_til(bt: u32, q_head: u32, {% if usesDecay %}h: u32, {% endif %}i: u32, scale: f32) -> f32 {
|
| 42 |
-
let raw = {{ read_scalar("query", "bt * params.qPackedDim + q_head * HEAD_DIM_K + i", queryDtype) }};
|
| 43 |
-
{% if usesDecay %}
|
| 44 |
-
return raw * {{ decay_at("bt", "h", "i") }} * scale;
|
| 45 |
-
{% else %}
|
| 46 |
-
return raw * scale;
|
| 47 |
-
{% endif %}
|
| 48 |
-
}
|
| 49 |
-
{% endif %}
|
| 50 |
-
{%- endmacro %}{% macro emit_chunk_head_setup(needHeads=false) %}
|
| 51 |
-
{% if needHeads %}
|
| 52 |
-
let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
|
| 53 |
-
let packed_out = max(params.qNumHeads, params.kvNumHeads) * HEAD_DIM_V;
|
| 54 |
-
{% endif %}
|
| 55 |
-
let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;
|
| 56 |
-
{%- endmacro %}
|
| 57 |
-
|
| 58 |
-
{% if queryDtype == "float16" or stateDtype == "float16" %}
|
| 59 |
-
enable f16;
|
| 60 |
-
{% endif %}
|
| 61 |
{{ env.wgsl.resourceDeclarations }}
|
| 62 |
|
| 63 |
// com.microsoft.LinearAttention, chunked prefill: the sequential state scan.
|
|
@@ -85,7 +40,6 @@ const TOKENS_PER_TILE: u32 = TOKEN_TILE / GROUPS;
|
|
| 85 |
const ROWS_PER_GROUP: u32 = HEAD_DIM_K / GROUPS;
|
| 86 |
|
| 87 |
{{ emit_chunk_operands() }}
|
| 88 |
-
|
| 89 |
var<workgroup> st: array<f32, HEAD_DIM_K * TILE_V>;
|
| 90 |
var<workgroup> dl: array<f32, CHUNK * TILE_V>;
|
| 91 |
// Staging for TOKEN_TILE rows of wk, then of k_hat: same shape, disjoint live ranges.
|
|
|
|
| 1 |
+
{% macro read_scalar(name, index, dtype) %}{% if dtype == "float16" %}f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}{% endmacro %}
|
| 2 |
+
{% macro write_scalar(expr, dtype) %}{% if dtype == "float16" %}f16({{ expr }}){% else %}{{ expr }}{% endif %}{% endmacro %}
|
| 3 |
+
{% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
{% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
| 6 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 7 |
{% if usesDecay %}
|
|
|
|
| 10 |
return raw;
|
| 11 |
{% endif %}
|
| 12 |
}
|
| 13 |
+
{% endmacro %}
|
| 14 |
+
{% macro emit_chunk_head_setup(needHeads=false) %}
|
| 15 |
+
let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
{{ env.wgsl.resourceDeclarations }}
|
| 17 |
|
| 18 |
// com.microsoft.LinearAttention, chunked prefill: the sequential state scan.
|
|
|
|
| 40 |
const ROWS_PER_GROUP: u32 = HEAD_DIM_K / GROUPS;
|
| 41 |
|
| 42 |
{{ emit_chunk_operands() }}
|
|
|
|
| 43 |
var<workgroup> st: array<f32, HEAD_DIM_K * TILE_V>;
|
| 44 |
var<workgroup> dl: array<f32, CHUNK * TILE_V>;
|
| 45 |
// Staging for TOKEN_TILE rows of wk, then of k_hat: same shape, disjoint live ranges.
|
build/webgpu/chunk-ut.wgsl.jinja
CHANGED
|
@@ -1,15 +1,6 @@
|
|
| 1 |
-
{% macro read_scalar(name, index, dtype) %}
|
| 2 |
-
{% if
|
| 3 |
-
f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
|
| 4 |
-
{% endmacro %}{% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{%- endmacro %}
|
| 5 |
-
|
| 6 |
{% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
|
| 7 |
-
{% if needQuery %}
|
| 8 |
-
// The output scale defaults to 1/sqrt(head_dim_k), matching the recurrent kernels.
|
| 9 |
-
fn out_scale() -> f32 {
|
| 10 |
-
return select(inverseSqrt(f32(HEAD_DIM_K)), params.scale, params.scale != 0.0);
|
| 11 |
-
}
|
| 12 |
-
{% endif %}
|
| 13 |
fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
| 14 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 15 |
{% if usesDecay %}
|
|
@@ -18,7 +9,6 @@ fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
|
| 18 |
return raw;
|
| 19 |
{% endif %}
|
| 20 |
}
|
| 21 |
-
{% if needKtil %}
|
| 22 |
|
| 23 |
fn k_til(bt: u32, h: u32, i: u32) -> f32 {
|
| 24 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
|
@@ -28,31 +18,9 @@ fn k_til(bt: u32, h: u32, i: u32) -> f32 {
|
|
| 28 |
return raw;
|
| 29 |
{% endif %}
|
| 30 |
}
|
| 31 |
-
{%
|
| 32 |
-
{%
|
| 33 |
-
|
| 34 |
-
// The output scale rides on q_til so both output terms -- q_til * S and P * delta,
|
| 35 |
-
// where P is itself built from q_til -- pick it up without a second pass over y.
|
| 36 |
-
fn q_til(bt: u32, q_head: u32, {% if usesDecay %}h: u32, {% endif %}i: u32, scale: f32) -> f32 {
|
| 37 |
-
let raw = {{ read_scalar("query", "bt * params.qPackedDim + q_head * HEAD_DIM_K + i", queryDtype) }};
|
| 38 |
-
{% if usesDecay %}
|
| 39 |
-
return raw * {{ decay_at("bt", "h", "i") }} * scale;
|
| 40 |
-
{% else %}
|
| 41 |
-
return raw * scale;
|
| 42 |
-
{% endif %}
|
| 43 |
-
}
|
| 44 |
-
{% endif %}
|
| 45 |
-
{%- endmacro %}{% macro emit_chunk_head_setup(needHeads=false) %}
|
| 46 |
-
{% if needHeads %}
|
| 47 |
-
let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
|
| 48 |
-
let packed_out = max(params.qNumHeads, params.kvNumHeads) * HEAD_DIM_V;
|
| 49 |
-
{% endif %}
|
| 50 |
-
let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;
|
| 51 |
-
{%- endmacro %}
|
| 52 |
-
|
| 53 |
-
{% if queryDtype == "float16" %}
|
| 54 |
-
enable f16;
|
| 55 |
-
{% endif %}
|
| 56 |
{{ env.wgsl.resourceDeclarations }}
|
| 57 |
|
| 58 |
// com.microsoft.LinearAttention, chunked prefill: the chunk-local delta transform.
|
|
@@ -75,7 +43,6 @@ const WG: u32 = {{ workgroupSize }}u;
|
|
| 75 |
const ENTRIES: u32 = (CHUNK * CHUNK) / WG;
|
| 76 |
|
| 77 |
{{ emit_chunk_operands(needKtil=true) }}
|
| 78 |
-
|
| 79 |
// `wm` holds W while the system is being built, then the inverse in place: row t of
|
| 80 |
// the inverse is produced from row t of W and the already-final rows above it.
|
| 81 |
var<workgroup> wm: array<f32, CHUNK * CHUNK>;
|
|
|
|
| 1 |
+
{% macro read_scalar(name, index, dtype) %}{% if dtype == "float16" %}f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}{% endmacro %}
|
| 2 |
+
{% macro decay_at(bt, h, i) %}{% if decayPerElement %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }}) * HEAD_DIM_K + ({{ i }})]{% else %}gexp[({{ bt }}) * params.decayPackedDim + ({{ h }})]{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
| 3 |
{% macro emit_chunk_operands(needKtil=false, needQuery=false) %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
fn k_hat(bt: u32, h: u32, i: u32) -> f32 {
|
| 5 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
| 6 |
{% if usesDecay %}
|
|
|
|
| 9 |
return raw;
|
| 10 |
{% endif %}
|
| 11 |
}
|
|
|
|
| 12 |
|
| 13 |
fn k_til(bt: u32, h: u32, i: u32) -> f32 {
|
| 14 |
let raw = {{ read_scalar("key", "bt * params.kPackedDim + (h / KV_PER_KEY_HEAD) * HEAD_DIM_K + i", queryDtype) }};
|
|
|
|
| 18 |
return raw;
|
| 19 |
{% endif %}
|
| 20 |
}
|
| 21 |
+
{% endmacro %}
|
| 22 |
+
{% macro emit_chunk_head_setup(needHeads=false) %}
|
| 23 |
+
let num_chunks = (params.seqLength + CHUNK - 1u) / CHUNK;{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 24 |
{{ env.wgsl.resourceDeclarations }}
|
| 25 |
|
| 26 |
// com.microsoft.LinearAttention, chunked prefill: the chunk-local delta transform.
|
|
|
|
| 43 |
const ENTRIES: u32 = (CHUNK * CHUNK) / WG;
|
| 44 |
|
| 45 |
{{ emit_chunk_operands(needKtil=true) }}
|
|
|
|
| 46 |
// `wm` holds W while the system is being built, then the inverse in place: row t of
|
| 47 |
// the inverse is produced from row t of W and the already-final rows above it.
|
| 48 |
var<workgroup> wm: array<f32, CHUNK * CHUNK>;
|
build/webgpu/linear-attention.scalar.wgsl.jinja
CHANGED
|
@@ -1,23 +1,12 @@
|
|
| 1 |
-
{% macro read_scalar(name, index, dtype) %}
|
| 2 |
-
{% if dtype == "float16" %}
|
| 3 |
-
f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
|
| 4 |
-
{% endmacro %}
|
| 5 |
-
{% macro write_scalar(expr, dtype) %}
|
| 6 |
-
{% if dtype == "float16" %}
|
| 7 |
-
f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
| 8 |
-
{% endmacro -%}
|
| 9 |
{% macro emit_tiled_setup(dvGroups=1) %}
|
| 10 |
let head_dim_k = params.qPackedDim / params.qNumHeads;
|
| 11 |
let head_dim_v = params.vPackedDim / params.kvNumHeads;
|
| 12 |
let n_key_heads = params.kPackedDim / head_dim_k;
|
| 13 |
let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
|
| 14 |
let kv_per_key_head = params.kvNumHeads / n_key_heads;
|
| 15 |
-
{% if dvGroups == 1 %}
|
| 16 |
let dv_tiles = (head_dim_v + TILE_V - 1u) / TILE_V;
|
| 17 |
-
{% else %}
|
| 18 |
-
// Tile slots per workgroup-index step: each step covers DV_GROUPS value tiles.
|
| 19 |
-
let dv_tiles = ((head_dim_v + TILE_V - 1u) / TILE_V + DV_GROUPS - 1u) / DV_GROUPS;
|
| 20 |
-
{% endif %}
|
| 21 |
let scale = select(inverseSqrt(f32(head_dim_k)), params.scale, params.scale != 0.0);
|
| 22 |
|
| 23 |
// 2D-folded flat (batch*head*dv_tile) index: wg.y carries the high bits past
|
|
@@ -32,14 +21,9 @@ f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
|
| 32 |
return;
|
| 33 |
}
|
| 34 |
|
| 35 |
-
{% if dvGroups == 1 %}
|
| 36 |
let dv_start = dv_tile_idx * TILE_V;
|
| 37 |
-
{% else %}
|
| 38 |
-
let dv_start = (dv_tile_idx * DV_GROUPS + dv_group) * TILE_V;
|
| 39 |
-
{% endif %}
|
| 40 |
let packed_out = max(params.qNumHeads, params.kvNumHeads) * head_dim_v;
|
| 41 |
-
let key_head_idx = head_idx / kv_per_key_head;
|
| 42 |
-
{%- endmacro -%}
|
| 43 |
{% macro emit_query_groups(first_group) %}
|
| 44 |
for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
|
| 45 |
let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
|
|
@@ -104,13 +88,9 @@ f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
|
| 104 |
}
|
| 105 |
}
|
| 106 |
{% endif %}
|
| 107 |
-
}
|
| 108 |
-
{%- endmacro %}
|
| 109 |
{% set usesDecay = updateRule == "gated" or updateRule == "gated_delta" %}
|
| 110 |
{% set usesBeta = updateRule == "delta" or updateRule == "gated_delta" %}
|
| 111 |
-
{% if queryDtype == "float16" or stateDtype == "float16" %}
|
| 112 |
-
enable f16;
|
| 113 |
-
{% endif %}
|
| 114 |
{% if useSubgroups %}
|
| 115 |
enable subgroups;
|
| 116 |
{% endif %}
|
|
@@ -119,9 +99,12 @@ enable subgroups;
|
|
| 119 |
const WG: u32 = {{ workgroupSize }}u;
|
| 120 |
const TILE_V: u32 = {{ tileV }}u;
|
| 121 |
|
| 122 |
-
{% if usesBeta %}
|
| 123 |
-
|
| 124 |
-
{%
|
|
|
|
|
|
|
|
|
|
| 125 |
var<workgroup> broadcast_delta: array<f32, TILE_V>;
|
| 126 |
|
| 127 |
{% endif %}
|
|
|
|
| 1 |
+
{% macro read_scalar(name, index, dtype) %}{% if dtype == "float16" %}f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}{% endmacro %}
|
| 2 |
+
{% macro write_scalar(expr, dtype) %}{% if dtype == "float16" %}f16({{ expr }}){% else %}{{ expr }}{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
{% macro emit_tiled_setup(dvGroups=1) %}
|
| 4 |
let head_dim_k = params.qPackedDim / params.qNumHeads;
|
| 5 |
let head_dim_v = params.vPackedDim / params.kvNumHeads;
|
| 6 |
let n_key_heads = params.kPackedDim / head_dim_k;
|
| 7 |
let heads_per_group = max(1u, params.qNumHeads / params.kvNumHeads);
|
| 8 |
let kv_per_key_head = params.kvNumHeads / n_key_heads;
|
|
|
|
| 9 |
let dv_tiles = (head_dim_v + TILE_V - 1u) / TILE_V;
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
let scale = select(inverseSqrt(f32(head_dim_k)), params.scale, params.scale != 0.0);
|
| 11 |
|
| 12 |
// 2D-folded flat (batch*head*dv_tile) index: wg.y carries the high bits past
|
|
|
|
| 21 |
return;
|
| 22 |
}
|
| 23 |
|
|
|
|
| 24 |
let dv_start = dv_tile_idx * TILE_V;
|
|
|
|
|
|
|
|
|
|
| 25 |
let packed_out = max(params.qNumHeads, params.kvNumHeads) * head_dim_v;
|
| 26 |
+
let key_head_idx = head_idx / kv_per_key_head;{% endmacro %}
|
|
|
|
| 27 |
{% macro emit_query_groups(first_group) %}
|
| 28 |
for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
|
| 29 |
let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
|
|
|
|
| 88 |
}
|
| 89 |
}
|
| 90 |
{% endif %}
|
| 91 |
+
}{% endmacro %}
|
|
|
|
| 92 |
{% set usesDecay = updateRule == "gated" or updateRule == "gated_delta" %}
|
| 93 |
{% set usesBeta = updateRule == "delta" or updateRule == "gated_delta" %}
|
|
|
|
|
|
|
|
|
|
| 94 |
{% if useSubgroups %}
|
| 95 |
enable subgroups;
|
| 96 |
{% endif %}
|
|
|
|
| 99 |
const WG: u32 = {{ workgroupSize }}u;
|
| 100 |
const TILE_V: u32 = {{ tileV }}u;
|
| 101 |
|
| 102 |
+
{% if usesBeta %}
|
| 103 |
+
var<workgroup> red_retrieved: array<f32, WG * TILE_V>;
|
| 104 |
+
{% endif %}
|
| 105 |
+
var<workgroup> red_preout: array<f32, WG * TILE_V>;
|
| 106 |
+
{% if usesBeta %}
|
| 107 |
+
var<workgroup> red_kq: array<f32, WG>;
|
| 108 |
var<workgroup> broadcast_delta: array<f32, TILE_V>;
|
| 109 |
|
| 110 |
{% endif %}
|
build/webgpu/linear-attention.serial.wgsl.jinja
CHANGED
|
@@ -1,11 +1,5 @@
|
|
| 1 |
-
{% macro read_scalar(name, index, dtype) %}
|
| 2 |
-
{% if dtype == "float16" %}
|
| 3 |
-
f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
|
| 4 |
-
{% endmacro %}
|
| 5 |
-
{% macro write_scalar(expr, dtype) %}
|
| 6 |
-
{% if dtype == "float16" %}
|
| 7 |
-
f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
| 8 |
-
{% endmacro -%}
|
| 9 |
{% macro emit_serial_query_groups(first_group) %}
|
| 10 |
for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
|
| 11 |
let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
|
|
@@ -17,14 +11,9 @@ f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
|
| 17 |
}
|
| 18 |
let out_idx = bt * packed_out + out_head * head_dim_v + dv_idx;
|
| 19 |
output[out_idx] = {{ write_scalar("q_out * scale", outputDtype) }};
|
| 20 |
-
}
|
| 21 |
-
{%- endmacro -%}
|
| 22 |
{% set gatedDeltaRule = updateRule == "gated_delta" %}
|
| 23 |
-
{% if queryDtype == "float16" or stateDtype == "float16" %}
|
| 24 |
-
enable f16;
|
| 25 |
-
{% endif %}
|
| 26 |
{{ env.wgsl.resourceDeclarations }}
|
| 27 |
-
|
| 28 |
{% if updateRule == "linear" %}
|
| 29 |
{% set geometry = [
|
| 30 |
["batchSize", "u32", serialBatchSize],
|
|
@@ -39,6 +28,7 @@ enable f16;
|
|
| 39 |
{% if hasStateWindow %}
|
| 40 |
{% set geometry = geometry + [["stateWindow", "u32", serialStateWindow], ["stateSlotStride", "u32", serialStateSlotStride]] %}
|
| 41 |
{% endif %}
|
|
|
|
| 42 |
struct SerialGeometry {
|
| 43 |
{% for field in geometry %}
|
| 44 |
{{ field[0] }}: {{ field[1] }},
|
|
@@ -46,7 +36,7 @@ struct SerialGeometry {
|
|
| 46 |
}
|
| 47 |
const params = SerialGeometry(
|
| 48 |
{% for field in geometry %}
|
| 49 |
-
{
|
| 50 |
{% endfor %}
|
| 51 |
);
|
| 52 |
{% endif %}
|
|
|
|
| 1 |
+
{% macro read_scalar(name, index, dtype) %}{% if dtype == "float16" %}f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}{% endmacro %}
|
| 2 |
+
{% macro write_scalar(expr, dtype) %}{% if dtype == "float16" %}f16({{ expr }}){% else %}{{ expr }}{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
{% macro emit_serial_query_groups(first_group) %}
|
| 4 |
for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
|
| 5 |
let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
|
|
|
|
| 11 |
}
|
| 12 |
let out_idx = bt * packed_out + out_head * head_dim_v + dv_idx;
|
| 13 |
output[out_idx] = {{ write_scalar("q_out * scale", outputDtype) }};
|
| 14 |
+
}{% endmacro %}
|
|
|
|
| 15 |
{% set gatedDeltaRule = updateRule == "gated_delta" %}
|
|
|
|
|
|
|
|
|
|
| 16 |
{{ env.wgsl.resourceDeclarations }}
|
|
|
|
| 17 |
{% if updateRule == "linear" %}
|
| 18 |
{% set geometry = [
|
| 19 |
["batchSize", "u32", serialBatchSize],
|
|
|
|
| 28 |
{% if hasStateWindow %}
|
| 29 |
{% set geometry = geometry + [["stateWindow", "u32", serialStateWindow], ["stateSlotStride", "u32", serialStateSlotStride]] %}
|
| 30 |
{% endif %}
|
| 31 |
+
|
| 32 |
struct SerialGeometry {
|
| 33 |
{% for field in geometry %}
|
| 34 |
{{ field[0] }}: {{ field[1] }},
|
|
|
|
| 36 |
}
|
| 37 |
const params = SerialGeometry(
|
| 38 |
{% for field in geometry %}
|
| 39 |
+
{{ ("f32(" ~ field[2] ~ ")") if field[1] == "f32" else (field[2] ~ "u") }},
|
| 40 |
{% endfor %}
|
| 41 |
);
|
| 42 |
{% endif %}
|
build/webgpu/linear-attention.vec4.wgsl.jinja
CHANGED
|
@@ -1,11 +1,5 @@
|
|
| 1 |
-
{% macro read_scalar(name, index, dtype) %}
|
| 2 |
-
{% if dtype == "float16" %}
|
| 3 |
-
f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}
|
| 4 |
-
{% endmacro %}
|
| 5 |
-
{% macro write_scalar(expr, dtype) %}
|
| 6 |
-
{% if dtype == "float16" %}
|
| 7 |
-
f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
| 8 |
-
{% endmacro -%}
|
| 9 |
{% macro emit_tiled_setup(dvGroups=1) %}
|
| 10 |
let head_dim_k = params.qPackedDim / params.qNumHeads;
|
| 11 |
let head_dim_v = params.vPackedDim / params.kvNumHeads;
|
|
@@ -38,8 +32,7 @@ f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
|
| 38 |
let dv_start = (dv_tile_idx * DV_GROUPS + dv_group) * TILE_V;
|
| 39 |
{% endif %}
|
| 40 |
let packed_out = max(params.qNumHeads, params.kvNumHeads) * head_dim_v;
|
| 41 |
-
let key_head_idx = head_idx / kv_per_key_head;
|
| 42 |
-
{%- endmacro -%}
|
| 43 |
{% macro emit_vec4_query_groups(first_group) %}
|
| 44 |
for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
|
| 45 |
let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
|
|
@@ -81,8 +74,7 @@ f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
|
| 81 |
}
|
| 82 |
}
|
| 83 |
}
|
| 84 |
-
}
|
| 85 |
-
{%- endmacro %}
|
| 86 |
{% set usesDecay = updateRule == "gated" or updateRule == "gated_delta" %}
|
| 87 |
{% set redWidth = 4 if tileV % 4 == 0 else (2 if tileV % 2 == 0 else 1) %}
|
| 88 |
{% set redVecs = (tileV / redWidth)|int %}
|
|
@@ -90,16 +82,11 @@ f16({{ expr }}){% else %}{{ expr }}{% endif %}
|
|
| 90 |
{% set redComps = ["x", "y", "z", "w"] %}
|
| 91 |
{% set usesBeta = updateRule == "delta" or updateRule == "gated_delta" %}
|
| 92 |
{% set redSlots = (2 * redVecs + 1) if usesBeta else redVecs %}
|
| 93 |
-
{% if queryDtype == "float16" or stateDtype == "float16" %}
|
| 94 |
-
enable f16;
|
| 95 |
-
{% endif %}
|
| 96 |
{% if useSubgroups %}
|
| 97 |
enable subgroups;
|
| 98 |
{% endif %}
|
| 99 |
{{ env.wgsl.resourceDeclarations }}
|
| 100 |
{% if useSubgroups %}
|
| 101 |
-
{% set skipLogicalLaneCount = true %}
|
| 102 |
-
{% set skipLogicalLaneCount = skipLogicalLaneCount is defined and skipLogicalLaneCount %}
|
| 103 |
var<workgroup> sgLaneClaims: atomic<u32>;
|
| 104 |
|
| 105 |
struct SubgroupLogicalLanes {
|
|
@@ -109,25 +96,17 @@ struct SubgroupLogicalLanes {
|
|
| 109 |
count: u32,
|
| 110 |
}
|
| 111 |
|
| 112 |
-
fn subgroup_logical_lanes() -> SubgroupLogicalLanes {
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
//
|
| 116 |
-
//
|
| 117 |
-
//
|
| 118 |
-
//
|
| 119 |
-
// the lowest active lane, which is the lane `subgroupBroadcastFirst` reads.
|
| 120 |
let ticket = atomicAdd(&sgLaneClaims, select(0u, count | (1u << 16u), rank == 0u));
|
| 121 |
let claim = subgroupBroadcastFirst(ticket);
|
| 122 |
return SubgroupLogicalLanes((claim & 0xffffu) + rank, claim >> 16u, rank, count);
|
| 123 |
}
|
| 124 |
-
{% if not skipLogicalLaneCount %}
|
| 125 |
-
|
| 126 |
-
fn subgroup_logical_count() -> u32 {
|
| 127 |
-
return atomicLoad(&sgLaneClaims) >> 16u;
|
| 128 |
-
}
|
| 129 |
-
{% endif %}
|
| 130 |
-
|
| 131 |
{% endif %}
|
| 132 |
|
| 133 |
// LANES threads cooperate on one value tile's reduction axis; DV_GROUPS such groups
|
|
@@ -159,7 +138,10 @@ var<workgroup> wg_fold_out: array<{{ redT }}, DV_GROUPS * {{ redSlots }}u>;
|
|
| 159 |
@compute @workgroup_size(WG, 1, 1)
|
| 160 |
fn main(
|
| 161 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 162 |
-
{% if
|
|
|
|
|
|
|
|
|
|
| 163 |
@builtin(local_invocation_id) lid: vec3<u32>,
|
| 164 |
{% endif %}
|
| 165 |
) {
|
|
@@ -170,7 +152,7 @@ fn main(
|
|
| 170 |
// the subgroup itself -- the dense claim order of the subgroup and the invocation's
|
| 171 |
// dense rank among its active lanes -- not from a partition of local_invocation_id,
|
| 172 |
// because WGSL does not promise which invocations share a subgroup.
|
| 173 |
-
let L = subgroup_logical_lanes();
|
| 174 |
let lane = L.rank;
|
| 175 |
let dv_group = L.ord;
|
| 176 |
{% else %}
|
|
@@ -178,7 +160,8 @@ fn main(
|
|
| 178 |
let lane = tid % LANES;
|
| 179 |
let dv_group = tid / LANES;
|
| 180 |
{% endif %}
|
| 181 |
-
let dk_base = lane * VEC4_LANES;
|
|
|
|
| 182 |
let lane_active = dk_base < head_dim_k;
|
| 183 |
|
| 184 |
// state[j] holds 4 consecutive dk rows (the 4 components) for dv slot j.
|
|
@@ -231,15 +214,13 @@ fn main(
|
|
| 231 |
if (lane_active) {
|
| 232 |
let k_base = ({{ token }} * params.kPackedDim + key_head_idx * head_dim_k) / VEC4_LANES + lane;
|
| 233 |
k_next = vec4<f32>(key[k_base]);
|
| 234 |
-
}
|
| 235 |
-
{%- endmacro %}
|
| 236 |
{% if usesBeta %}
|
| 237 |
{% macro load_query0(token) %}
|
| 238 |
if (lane_active) {
|
| 239 |
let q0_base = ({{ token }} * params.qPackedDim + q_head_0 * head_dim_k) / VEC4_LANES + lane;
|
| 240 |
q0_next = vec4<f32>(query[q0_base]);
|
| 241 |
-
}
|
| 242 |
-
{%- endmacro %}
|
| 243 |
let q_head_0 = (head_idx * params.qNumHeads) / params.kvNumHeads;
|
| 244 |
let out_head_0 = head_idx * heads_per_group;
|
| 245 |
var q0_next = vec4<f32>(0.0);
|
|
|
|
| 1 |
+
{% macro read_scalar(name, index, dtype) %}{% if dtype == "float16" %}f32({{ name }}[{{ index }}]){% else %}{{ name }}[{{ index }}]{% endif %}{% endmacro %}
|
| 2 |
+
{% macro write_scalar(expr, dtype) %}{% if dtype == "float16" %}f16({{ expr }}){% else %}{{ expr }}{% endif %}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
{% macro emit_tiled_setup(dvGroups=1) %}
|
| 4 |
let head_dim_k = params.qPackedDim / params.qNumHeads;
|
| 5 |
let head_dim_v = params.vPackedDim / params.kvNumHeads;
|
|
|
|
| 32 |
let dv_start = (dv_tile_idx * DV_GROUPS + dv_group) * TILE_V;
|
| 33 |
{% endif %}
|
| 34 |
let packed_out = max(params.qNumHeads, params.kvNumHeads) * head_dim_v;
|
| 35 |
+
let key_head_idx = head_idx / kv_per_key_head;{% endmacro %}
|
|
|
|
| 36 |
{% macro emit_vec4_query_groups(first_group) %}
|
| 37 |
for (var qg = {{ first_group }}u; qg < heads_per_group; qg = qg + 1u) {
|
| 38 |
let q_head = (head_idx * params.qNumHeads) / params.kvNumHeads + qg;
|
|
|
|
| 74 |
}
|
| 75 |
}
|
| 76 |
}
|
| 77 |
+
}{% endmacro %}
|
|
|
|
| 78 |
{% set usesDecay = updateRule == "gated" or updateRule == "gated_delta" %}
|
| 79 |
{% set redWidth = 4 if tileV % 4 == 0 else (2 if tileV % 2 == 0 else 1) %}
|
| 80 |
{% set redVecs = (tileV / redWidth)|int %}
|
|
|
|
| 82 |
{% set redComps = ["x", "y", "z", "w"] %}
|
| 83 |
{% set usesBeta = updateRule == "delta" or updateRule == "gated_delta" %}
|
| 84 |
{% set redSlots = (2 * redVecs + 1) if usesBeta else redVecs %}
|
|
|
|
|
|
|
|
|
|
| 85 |
{% if useSubgroups %}
|
| 86 |
enable subgroups;
|
| 87 |
{% endif %}
|
| 88 |
{{ env.wgsl.resourceDeclarations }}
|
| 89 |
{% if useSubgroups %}
|
|
|
|
|
|
|
| 90 |
var<workgroup> sgLaneClaims: atomic<u32>;
|
| 91 |
|
| 92 |
struct SubgroupLogicalLanes {
|
|
|
|
| 96 |
count: u32,
|
| 97 |
}
|
| 98 |
|
| 99 |
+
fn subgroup_logical_lanes(rank: u32, count: u32) -> SubgroupLogicalLanes {
|
| 100 |
+
// Every lane performs the atomic so no collective follows a lane guard: a
|
| 101 |
+
// subgroup op after a closed `if (rank == 0u)` block is a reconvergence
|
| 102 |
+
// hazard, since every subgroup operation must be reached in uniform control
|
| 103 |
+
// flow. Only the lane-0 lane adds its subgroup's claim, every other lane
|
| 104 |
+
// adds 0 and discards its snapshot. Lane 0 is the lowest lane, which is the
|
| 105 |
+
// lane `subgroupBroadcastFirst` reads.
|
|
|
|
| 106 |
let ticket = atomicAdd(&sgLaneClaims, select(0u, count | (1u << 16u), rank == 0u));
|
| 107 |
let claim = subgroupBroadcastFirst(ticket);
|
| 108 |
return SubgroupLogicalLanes((claim & 0xffffu) + rank, claim >> 16u, rank, count);
|
| 109 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 110 |
{% endif %}
|
| 111 |
|
| 112 |
// LANES threads cooperate on one value tile's reduction axis; DV_GROUPS such groups
|
|
|
|
| 138 |
@compute @workgroup_size(WG, 1, 1)
|
| 139 |
fn main(
|
| 140 |
@builtin(workgroup_id) wg: vec3<u32>,
|
| 141 |
+
{% if useSubgroups %}
|
| 142 |
+
@builtin(subgroup_invocation_id) sgLane: u32,
|
| 143 |
+
@builtin(subgroup_size) sgWidth: u32,
|
| 144 |
+
{% else %}
|
| 145 |
@builtin(local_invocation_id) lid: vec3<u32>,
|
| 146 |
{% endif %}
|
| 147 |
) {
|
|
|
|
| 152 |
// the subgroup itself -- the dense claim order of the subgroup and the invocation's
|
| 153 |
// dense rank among its active lanes -- not from a partition of local_invocation_id,
|
| 154 |
// because WGSL does not promise which invocations share a subgroup.
|
| 155 |
+
let L = subgroup_logical_lanes(sgLane, sgWidth);
|
| 156 |
let lane = L.rank;
|
| 157 |
let dv_group = L.ord;
|
| 158 |
{% else %}
|
|
|
|
| 160 |
let lane = tid % LANES;
|
| 161 |
let dv_group = tid / LANES;
|
| 162 |
{% endif %}
|
| 163 |
+
let dk_base = lane * VEC4_LANES;
|
| 164 |
+
{{ emit_tiled_setup(dvGroups=(2 if useSubgroups else dvGroups)) }}
|
| 165 |
let lane_active = dk_base < head_dim_k;
|
| 166 |
|
| 167 |
// state[j] holds 4 consecutive dk rows (the 4 components) for dv slot j.
|
|
|
|
| 214 |
if (lane_active) {
|
| 215 |
let k_base = ({{ token }} * params.kPackedDim + key_head_idx * head_dim_k) / VEC4_LANES + lane;
|
| 216 |
k_next = vec4<f32>(key[k_base]);
|
| 217 |
+
}{% endmacro %}
|
|
|
|
| 218 |
{% if usesBeta %}
|
| 219 |
{% macro load_query0(token) %}
|
| 220 |
if (lane_active) {
|
| 221 |
let q0_base = ({{ token }} * params.qPackedDim + q_head_0 * head_dim_k) / VEC4_LANES + lane;
|
| 222 |
q0_next = vec4<f32>(query[q0_base]);
|
| 223 |
+
}{% endmacro %}
|
|
|
|
| 224 |
let q_head_0 = (head_idx * params.qNumHeads) / params.kvNumHeads;
|
| 225 |
let out_head_0 = head_idx * heads_per_group;
|
| 226 |
var q0_next = vec4<f32>(0.0);
|
build/webgpu/manifest.json
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.LinearAttention",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
@@ -8,49 +8,49 @@
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "SADv8TDNnXFGZdNUCqGMV5D9W9IKy2+gSRRNOZvxh/4=",
|
| 11 |
-
"chunk-out.wgsl.jinja": "
|
| 12 |
-
"chunk-prep.wgsl.jinja": "
|
| 13 |
-
"chunk-scan.wgsl.jinja": "
|
| 14 |
-
"chunk-ut.wgsl.jinja": "
|
| 15 |
-
"linear-attention.scalar.wgsl.jinja": "
|
| 16 |
-
"linear-attention.serial.wgsl.jinja": "
|
| 17 |
-
"linear-attention.vec4.wgsl.jinja": "
|
| 18 |
-
"manifest.json": "
|
| 19 |
-
"test.json": "
|
| 20 |
}
|
| 21 |
},
|
| 22 |
-
"provenance": { "kernel": { "sha": "
|
| 23 |
"webgpu": {
|
| 24 |
-
"manifestSpec": "2.
|
| 25 |
"variants": {
|
| 26 |
"linear_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
|
|
|
|
|
|
| 27 |
"linear_state_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
|
|
|
|
|
|
| 28 |
"gated_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
|
|
|
|
|
|
| 29 |
"gated_state_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
|
|
|
|
|
|
| 30 |
"delta_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
|
|
|
|
|
|
| 31 |
"delta_state_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
|
|
|
|
|
|
| 32 |
"gated_delta_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
|
|
|
|
|
|
| 33 |
"gated_delta_state_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
|
|
|
|
|
|
| 34 |
"linear_zero_serial_small_dk": ["linear-attention.serial.wgsl.jinja"],
|
| 35 |
"linear_state_serial_small_dk": ["linear-attention.serial.wgsl.jinja"],
|
| 36 |
"gated_delta_zero_serial_small_dk": ["linear-attention.serial.wgsl.jinja"],
|
| 37 |
-
"gated_delta_state_serial_small_dk": ["linear-attention.serial.wgsl.jinja"]
|
| 38 |
-
"linear_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 39 |
-
"linear_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 40 |
-
"gated_delta_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 41 |
-
"gated_delta_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 42 |
-
"gated_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 43 |
-
"gated_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 44 |
-
"delta_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 45 |
-
"delta_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 46 |
-
"linear_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 47 |
-
"linear_state_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 48 |
-
"gated_delta_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 49 |
-
"gated_delta_state_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 50 |
-
"gated_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 51 |
-
"gated_state_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 52 |
-
"delta_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 53 |
-
"delta_state_vec4": ["linear-attention.vec4.wgsl.jinja"]
|
| 54 |
}
|
| 55 |
}
|
| 56 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.LinearAttention",
|
| 3 |
+
"id": "_com_microsoft_linearattention_webgpu_f649015",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "SADv8TDNnXFGZdNUCqGMV5D9W9IKy2+gSRRNOZvxh/4=",
|
| 11 |
+
"chunk-out.wgsl.jinja": "NtmWHwJETU/jVQ4S7IhzVPTMr5/EIgigBd9vXPff+X4=",
|
| 12 |
+
"chunk-prep.wgsl.jinja": "EFWTZh76y5wM+swu+IaId+L491AVXqYz2QbzwKIc+Oc=",
|
| 13 |
+
"chunk-scan.wgsl.jinja": "Jexma8+n/iicl/obGiEGFg8+et86bGJZV7CXvXemrtA=",
|
| 14 |
+
"chunk-ut.wgsl.jinja": "I5ALen2vUhua450yTy5D1/aoUk+c32yDnOYRiBxkRrM=",
|
| 15 |
+
"linear-attention.scalar.wgsl.jinja": "NVuNbrDhuK97SfPabpsEUBJhQ6PXf8cKpTa7O62BZss=",
|
| 16 |
+
"linear-attention.serial.wgsl.jinja": "LOksC9OSD5RnpMPXqtEFWi37cOPlUnEHtn5es2B0bOw=",
|
| 17 |
+
"linear-attention.vec4.wgsl.jinja": "7KHl8IX+MVeQWq9dA97PSIJb8FKm3JgyDnuVndvg7fk=",
|
| 18 |
+
"manifest.json": "VLl0IELx7QOCW+y935REgTG7VUHQF5o/Up17IFamO/8=",
|
| 19 |
+
"test.json": "2sKl1HvKpMz2KteF7M+SRa+Csr/a4gw2ZyHqwA4Ow44="
|
| 20 |
}
|
| 21 |
},
|
| 22 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 23 |
"webgpu": {
|
| 24 |
+
"manifestSpec": "2.1",
|
| 25 |
"variants": {
|
| 26 |
"linear_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
| 27 |
+
"linear_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 28 |
+
"linear_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 29 |
"linear_state_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
| 30 |
+
"linear_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 31 |
+
"linear_state_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 32 |
"gated_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
| 33 |
+
"gated_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 34 |
+
"gated_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 35 |
"gated_state_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja"],
|
| 36 |
+
"gated_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 37 |
+
"gated_state_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 38 |
"delta_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
| 39 |
+
"delta_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 40 |
+
"delta_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 41 |
"delta_state_chunked": ["chunk-out.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
| 42 |
+
"delta_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 43 |
+
"delta_state_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 44 |
"gated_delta_zero_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
| 45 |
+
"gated_delta_zero_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 46 |
+
"gated_delta_zero_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 47 |
"gated_delta_state_chunked": ["chunk-out.wgsl.jinja", "chunk-prep.wgsl.jinja", "chunk-scan.wgsl.jinja", "chunk-ut.wgsl.jinja"],
|
| 48 |
+
"gated_delta_state_scalar": ["linear-attention.scalar.wgsl.jinja"],
|
| 49 |
+
"gated_delta_state_vec4": ["linear-attention.vec4.wgsl.jinja"],
|
| 50 |
"linear_zero_serial_small_dk": ["linear-attention.serial.wgsl.jinja"],
|
| 51 |
"linear_state_serial_small_dk": ["linear-attention.serial.wgsl.jinja"],
|
| 52 |
"gated_delta_zero_serial_small_dk": ["linear-attention.serial.wgsl.jinja"],
|
| 53 |
+
"gated_delta_state_serial_small_dk": ["linear-attention.serial.wgsl.jinja"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 54 |
}
|
| 55 |
}
|
| 56 |
}
|
build/webgpu/test.json
CHANGED
|
@@ -368,9 +368,9 @@
|
|
| 368 |
}
|
| 369 |
},
|
| 370 |
{
|
| 371 |
-
"name": "
|
| 372 |
"provenance": {
|
| 373 |
-
"notes": "
|
| 374 |
},
|
| 375 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 4, "update_rule": "gated_delta", "scale": 0.25 },
|
| 376 |
"inputs": {
|
|
@@ -533,7 +533,7 @@
|
|
| 533 |
{
|
| 534 |
"name": "linear_zero_scalar_f16_seq128",
|
| 535 |
"provenance": {
|
| 536 |
-
"notes": "Float16 query and state
|
| 537 |
},
|
| 538 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
|
| 539 |
"inputs": {
|
|
@@ -559,7 +559,7 @@
|
|
| 559 |
}
|
| 560 |
},
|
| 561 |
{
|
| 562 |
-
"name": "
|
| 563 |
"provenance": {
|
| 564 |
"notes": "A 128-token float16 recurrence omits past state but does not produce zero output. Positive-offset keys and values around 0.5 keep output and present state at order-one magnitude, making multiplicative errors observable on the serial, scalar, and vec4 zero-state routes."
|
| 565 |
},
|
|
@@ -589,7 +589,7 @@
|
|
| 589 |
{
|
| 590 |
"name": "linear_state_scalar_f16_seq128",
|
| 591 |
"provenance": {
|
| 592 |
-
"notes": "A compact float16 recurrence
|
| 593 |
},
|
| 594 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
|
| 595 |
"inputs": {
|
|
@@ -620,7 +620,7 @@
|
|
| 620 |
}
|
| 621 |
},
|
| 622 |
{
|
| 623 |
-
"name": "
|
| 624 |
"provenance": {
|
| 625 |
"notes": "A supplied past state at amplitude 0.3, positive-offset keys, and values around 0.5 keep both outputs at order-one magnitude. Dropping the initial state or uniformly rescaling either output therefore exceeds tolerance."
|
| 626 |
},
|
|
@@ -760,7 +760,7 @@
|
|
| 760 |
{
|
| 761 |
"name": "gated_delta_scalar_headdimk_gt_128_partial_dv_tile",
|
| 762 |
"provenance": {
|
| 763 |
-
"notes": "
|
| 764 |
},
|
| 765 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.5 },
|
| 766 |
"inputs": {
|
|
@@ -803,7 +803,7 @@
|
|
| 803 |
{
|
| 804 |
"name": "gated_delta_default_rule_no_updateRule_arg",
|
| 805 |
"provenance": {
|
| 806 |
-
"notes": "Omitting `
|
| 807 |
},
|
| 808 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1 },
|
| 809 |
"inputs": {
|
|
@@ -892,9 +892,9 @@
|
|
| 892 |
}
|
| 893 |
},
|
| 894 |
{
|
| 895 |
-
"name": "
|
| 896 |
"provenance": {
|
| 897 |
-
"notes": "
|
| 898 |
},
|
| 899 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated_delta", "scale": 0.08838834764831845 },
|
| 900 |
"inputs": {
|
|
@@ -978,9 +978,9 @@
|
|
| 978 |
}
|
| 979 |
},
|
| 980 |
{
|
| 981 |
-
"name": "
|
| 982 |
"provenance": {
|
| 983 |
-
"notes": "
|
| 984 |
},
|
| 985 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated_delta", "scale": 0.08838834764831845 },
|
| 986 |
"inputs": {
|
|
@@ -1024,7 +1024,7 @@
|
|
| 1024 |
"name": "linear_zero_vec4_dk256_wg_gt_subgroup",
|
| 1025 |
"provenance": {
|
| 1026 |
"source": "onnxruntime/contrib_ops/webgpu/bert/linear_attention.wgsl.template",
|
| 1027 |
-
"test": "two-level subgroup reduction
|
| 1028 |
"notes": "The dk reduction must span the whole workgroup. head_dim_k/4 lanes exceed the subgroup size here, so a bare subgroupAdd only sums one subgroup's partial dot products."
|
| 1029 |
},
|
| 1030 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "linear" },
|
|
@@ -1054,7 +1054,7 @@
|
|
| 1054 |
"name": "gated_delta_vec4_dk132_wg_gt_subgroup",
|
| 1055 |
"provenance": {
|
| 1056 |
"source": "onnxruntime/contrib_ops/webgpu/bert/linear_attention.wgsl.template",
|
| 1057 |
-
"test": "two-level subgroup reduction
|
| 1058 |
"notes": "The dk reduction must span the whole workgroup. head_dim_k/4 lanes exceed the subgroup size here, so a bare subgroupAdd only sums one subgroup's partial dot products."
|
| 1059 |
},
|
| 1060 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated_delta" },
|
|
@@ -1129,7 +1129,7 @@
|
|
| 1129 |
{
|
| 1130 |
"name": "linear_zero_scalar_dk17_above_serial_cap_unaligned",
|
| 1131 |
"provenance": {
|
| 1132 |
-
"notes": "
|
| 1133 |
},
|
| 1134 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "linear", "scale": 0.5 },
|
| 1135 |
"inputs": {
|
|
@@ -1156,9 +1156,7 @@
|
|
| 1156 |
},
|
| 1157 |
{
|
| 1158 |
"name": "linear_state_scalar_dk17_above_serial_cap_unaligned",
|
| 1159 |
-
"provenance": {
|
| 1160 |
-
"notes": "Key head size 17 exceeds the serial limit and is not divisible by four, selecting the scalar recurrence with a supplied past state."
|
| 1161 |
-
},
|
| 1162 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "linear", "scale": 0.5 },
|
| 1163 |
"inputs": {
|
| 1164 |
"queryT": {
|
|
@@ -1677,7 +1675,7 @@
|
|
| 1677 |
{
|
| 1678 |
"name": "inverse_gqa_gated_delta_state_q2_kv4",
|
| 1679 |
"provenance": {
|
| 1680 |
-
"notes": "Inverse GQA: kvNumHeads exceeds qNumHeads, so the output carries one head per KV head
|
| 1681 |
},
|
| 1682 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 4, "update_rule": "gated_delta", "scale": 0.6 },
|
| 1683 |
"inputs": {
|
|
@@ -1750,9 +1748,7 @@
|
|
| 1750 |
},
|
| 1751 |
{
|
| 1752 |
"name": "inverse_gqa_gated_delta_state_dk32_tiled",
|
| 1753 |
-
"provenance": {
|
| 1754 |
-
"notes": "Inverse GQA with a head dimension of 32, above the serial kernel's ceiling of 16, so the tiled scalar and vectorized kernels are the ones selected rather than merely eligible."
|
| 1755 |
-
},
|
| 1756 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 4, "update_rule": "gated_delta", "scale": 0.35 },
|
| 1757 |
"inputs": {
|
| 1758 |
"queryT": {
|
|
@@ -1949,7 +1945,7 @@
|
|
| 1949 |
{
|
| 1950 |
"name": "gated_delta_zero_scalar_dk17_no_past",
|
| 1951 |
"provenance": {
|
| 1952 |
-
"notes": "
|
| 1953 |
},
|
| 1954 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.25 },
|
| 1955 |
"inputs": {
|
|
@@ -1979,7 +1975,7 @@
|
|
| 1979 |
{
|
| 1980 |
"name": "gated_delta_zero_vec4_dk20_no_past",
|
| 1981 |
"provenance": {
|
| 1982 |
-
"notes": "
|
| 1983 |
},
|
| 1984 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.25 },
|
| 1985 |
"inputs": {
|
|
@@ -2009,7 +2005,7 @@
|
|
| 2009 |
{
|
| 2010 |
"name": "gated_zero_chunked_seq1024",
|
| 2011 |
"provenance": {
|
| 2012 |
-
"notes": "
|
| 2013 |
},
|
| 2014 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated", "scale": 0.35 },
|
| 2015 |
"inputs": {
|
|
@@ -2042,7 +2038,7 @@
|
|
| 2042 |
{
|
| 2043 |
"name": "gated_state_chunked_seq1024",
|
| 2044 |
"provenance": {
|
| 2045 |
-
"notes": "
|
| 2046 |
},
|
| 2047 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated", "scale": 0.35 },
|
| 2048 |
"inputs": {
|
|
@@ -2080,7 +2076,7 @@
|
|
| 2080 |
{
|
| 2081 |
"name": "delta_zero_chunked_seq1024",
|
| 2082 |
"provenance": {
|
| 2083 |
-
"notes": "
|
| 2084 |
},
|
| 2085 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "delta", "scale": 0.35 },
|
| 2086 |
"inputs": {
|
|
@@ -2113,7 +2109,7 @@
|
|
| 2113 |
{
|
| 2114 |
"name": "delta_state_chunked_seq1024",
|
| 2115 |
"provenance": {
|
| 2116 |
-
"notes": "
|
| 2117 |
},
|
| 2118 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "delta", "scale": 0.35 },
|
| 2119 |
"inputs": {
|
|
@@ -2151,7 +2147,7 @@
|
|
| 2151 |
{
|
| 2152 |
"name": "gated_delta_zero_chunked_seq1024",
|
| 2153 |
"provenance": {
|
| 2154 |
-
"notes": "
|
| 2155 |
},
|
| 2156 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.35 },
|
| 2157 |
"inputs": {
|
|
@@ -2194,7 +2190,7 @@
|
|
| 2194 |
{
|
| 2195 |
"name": "linear_zero_chunked_seq1024",
|
| 2196 |
"provenance": {
|
| 2197 |
-
"notes": "
|
| 2198 |
},
|
| 2199 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "linear", "scale": 0.35 },
|
| 2200 |
"inputs": {
|
|
@@ -2227,7 +2223,7 @@
|
|
| 2227 |
{
|
| 2228 |
"name": "linear_state_chunked_seq1024",
|
| 2229 |
"provenance": {
|
| 2230 |
-
"notes": "
|
| 2231 |
},
|
| 2232 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.35 },
|
| 2233 |
"inputs": {
|
|
@@ -2260,7 +2256,7 @@
|
|
| 2260 |
{
|
| 2261 |
"name": "gated_delta_state_chunked_seq1024",
|
| 2262 |
"provenance": {
|
| 2263 |
-
"notes": "
|
| 2264 |
},
|
| 2265 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated_delta", "scale": 0.35 },
|
| 2266 |
"inputs": {
|
|
@@ -3138,6 +3134,168 @@
|
|
| 3138 |
"outputT": { "dtype": "float32", "shape": [2, 1, 27] },
|
| 3139 |
"presentStateT": { "dtype": "float16", "shape": [2, 2, 3, 16, 9], "tolerance": 6e-8, "relTolerance": 0.0005 }
|
| 3140 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3141 |
}
|
| 3142 |
]
|
| 3143 |
}
|
|
|
|
| 368 |
}
|
| 369 |
},
|
| 370 |
{
|
| 371 |
+
"name": "gated_delta_headdim6_seq128_offset_value_scale",
|
| 372 |
"provenance": {
|
| 373 |
+
"notes": "A 128-step gated-delta recurrence with 4 query and 4 key/value heads, head width 6 (key) by 12 (value). Key values scale to |k|^2 near 0.7 and V is offset to 1.0, so decay, beta, and the running update change the output measurably at each step."
|
| 374 |
},
|
| 375 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 4, "update_rule": "gated_delta", "scale": 0.25 },
|
| 376 |
"inputs": {
|
|
|
|
| 533 |
{
|
| 534 |
"name": "linear_zero_scalar_f16_seq128",
|
| 535 |
"provenance": {
|
| 536 |
+
"notes": "Float16 query and state over a 128-token zero-state linear recurrence check the output at every token."
|
| 537 |
},
|
| 538 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
|
| 539 |
"inputs": {
|
|
|
|
| 559 |
}
|
| 560 |
},
|
| 561 |
{
|
| 562 |
+
"name": "linear_zero_f16_seq128_offset_value_scale",
|
| 563 |
"provenance": {
|
| 564 |
"notes": "A 128-token float16 recurrence omits past state but does not produce zero output. Positive-offset keys and values around 0.5 keep output and present state at order-one magnitude, making multiplicative errors observable on the serial, scalar, and vec4 zero-state routes."
|
| 565 |
},
|
|
|
|
| 589 |
{
|
| 590 |
"name": "linear_state_scalar_f16_seq128",
|
| 591 |
"provenance": {
|
| 592 |
+
"notes": "A compact float16 recurrence with a supplied initial state verifies that the first output incorporates that state."
|
| 593 |
},
|
| 594 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.25 },
|
| 595 |
"inputs": {
|
|
|
|
| 620 |
}
|
| 621 |
},
|
| 622 |
{
|
| 623 |
+
"name": "linear_state_f16_seq128_offset_value_scale",
|
| 624 |
"provenance": {
|
| 625 |
"notes": "A supplied past state at amplitude 0.3, positive-offset keys, and values around 0.5 keep both outputs at order-one magnitude. Dropping the initial state or uniformly rescaling either output therefore exceeds tolerance."
|
| 626 |
},
|
|
|
|
| 760 |
{
|
| 761 |
"name": "gated_delta_scalar_headdimk_gt_128_partial_dv_tile",
|
| 762 |
"provenance": {
|
| 763 |
+
"notes": "A gated-delta step with 2 query heads sharing 1 key/value head. Key head width 130 exceeds 128 and is not a multiple of four; value head width 10 leaves a two-element remainder after grouping into fours, exercising both the key and value tails together."
|
| 764 |
},
|
| 765 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.5 },
|
| 766 |
"inputs": {
|
|
|
|
| 803 |
{
|
| 804 |
"name": "gated_delta_default_rule_no_updateRule_arg",
|
| 805 |
"provenance": {
|
| 806 |
+
"notes": "Omitting `update_rule` requests the schema-default gated-delta recurrence with decay, beta, and past state. Key head size four checks a compact reduction."
|
| 807 |
},
|
| 808 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1 },
|
| 809 |
"inputs": {
|
|
|
|
| 892 |
}
|
| 893 |
},
|
| 894 |
{
|
| 895 |
+
"name": "gated_delta_f32_dk128_dv128_offset_value_scale",
|
| 896 |
"provenance": {
|
| 897 |
+
"notes": "For float32 dK=dV=128, key magnitude near 0.7 makes the gated-delta correction visible; O(1) values and queries expose incorrect decay, beta or output scaling in both output and state."
|
| 898 |
},
|
| 899 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated_delta", "scale": 0.08838834764831845 },
|
| 900 |
"inputs": {
|
|
|
|
| 978 |
}
|
| 979 |
},
|
| 980 |
{
|
| 981 |
+
"name": "gated_delta_f16_dk128_dv128_offset_value_scale",
|
| 982 |
"provenance": {
|
| 983 |
+
"notes": "For float16 dK=dV=128, key magnitude near 0.7 makes the gated-delta correction visible; O(1) values and queries expose incorrect decay, beta or output scaling in both output and state."
|
| 984 |
},
|
| 985 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated_delta", "scale": 0.08838834764831845 },
|
| 986 |
"inputs": {
|
|
|
|
| 1024 |
"name": "linear_zero_vec4_dk256_wg_gt_subgroup",
|
| 1025 |
"provenance": {
|
| 1026 |
"source": "onnxruntime/contrib_ops/webgpu/bert/linear_attention.wgsl.template",
|
| 1027 |
+
"test": "two-level subgroup reduction",
|
| 1028 |
"notes": "The dk reduction must span the whole workgroup. head_dim_k/4 lanes exceed the subgroup size here, so a bare subgroupAdd only sums one subgroup's partial dot products."
|
| 1029 |
},
|
| 1030 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "linear" },
|
|
|
|
| 1054 |
"name": "gated_delta_vec4_dk132_wg_gt_subgroup",
|
| 1055 |
"provenance": {
|
| 1056 |
"source": "onnxruntime/contrib_ops/webgpu/bert/linear_attention.wgsl.template",
|
| 1057 |
+
"test": "two-level subgroup reduction",
|
| 1058 |
"notes": "The dk reduction must span the whole workgroup. head_dim_k/4 lanes exceed the subgroup size here, so a bare subgroupAdd only sums one subgroup's partial dot products."
|
| 1059 |
},
|
| 1060 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated_delta" },
|
|
|
|
| 1129 |
{
|
| 1130 |
"name": "linear_zero_scalar_dk17_above_serial_cap_unaligned",
|
| 1131 |
"provenance": {
|
| 1132 |
+
"notes": "A linear-attention step (no gating) from a zero initial state, with 1 query and 1 key/value head, key head width 17 (odd, not a multiple of four) and value head width 2."
|
| 1133 |
},
|
| 1134 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "linear", "scale": 0.5 },
|
| 1135 |
"inputs": {
|
|
|
|
| 1156 |
},
|
| 1157 |
{
|
| 1158 |
"name": "linear_state_scalar_dk17_above_serial_cap_unaligned",
|
| 1159 |
+
"provenance": { "notes": "Key head size 17 has a partial four-lane tail and a supplied past state." },
|
|
|
|
|
|
|
| 1160 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "linear", "scale": 0.5 },
|
| 1161 |
"inputs": {
|
| 1162 |
"queryT": {
|
|
|
|
| 1675 |
{
|
| 1676 |
"name": "inverse_gqa_gated_delta_state_q2_kv4",
|
| 1677 |
"provenance": {
|
| 1678 |
+
"notes": "Inverse GQA: kvNumHeads exceeds qNumHeads, so the output carries one head per KV head and several KV heads share a query head (KV head h reads query head floor(h * qNumHeads / kvNumHeads)), under the gated-delta update rule."
|
| 1679 |
},
|
| 1680 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 4, "update_rule": "gated_delta", "scale": 0.6 },
|
| 1681 |
"inputs": {
|
|
|
|
| 1748 |
},
|
| 1749 |
{
|
| 1750 |
"name": "inverse_gqa_gated_delta_state_dk32_tiled",
|
| 1751 |
+
"provenance": { "notes": "Inverse GQA with head dimension 32 checks the recurrent output across grouped heads." },
|
|
|
|
|
|
|
| 1752 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 4, "update_rule": "gated_delta", "scale": 0.35 },
|
| 1753 |
"inputs": {
|
| 1754 |
"queryT": {
|
|
|
|
| 1945 |
{
|
| 1946 |
"name": "gated_delta_zero_scalar_dk17_no_past",
|
| 1947 |
"provenance": {
|
| 1948 |
+
"notes": "An omitted past state and key head size 17 check zero-state recurrence with a partial four-lane tail."
|
| 1949 |
},
|
| 1950 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.25 },
|
| 1951 |
"inputs": {
|
|
|
|
| 1975 |
{
|
| 1976 |
"name": "gated_delta_zero_vec4_dk20_no_past",
|
| 1977 |
"provenance": {
|
| 1978 |
+
"notes": "An omitted past state with key head size 20 checks zero-state recurrence and a partial value-head tile."
|
| 1979 |
},
|
| 1980 |
"attrs": { "q_num_heads": 1, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.25 },
|
| 1981 |
"inputs": {
|
|
|
|
| 2005 |
{
|
| 2006 |
"name": "gated_zero_chunked_seq1024",
|
| 2007 |
"provenance": {
|
| 2008 |
+
"notes": "A 1,024-token gated recurrence with compact heads, zero entry state, and elementwise decay checks every output and final state."
|
| 2009 |
},
|
| 2010 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated", "scale": 0.35 },
|
| 2011 |
"inputs": {
|
|
|
|
| 2038 |
{
|
| 2039 |
"name": "gated_state_chunked_seq1024",
|
| 2040 |
"provenance": {
|
| 2041 |
+
"notes": "A 1,024-token gated recurrence with compact heads, supplied entry state, and per-head decay checks every output and final state."
|
| 2042 |
},
|
| 2043 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated", "scale": 0.35 },
|
| 2044 |
"inputs": {
|
|
|
|
| 2076 |
{
|
| 2077 |
"name": "delta_zero_chunked_seq1024",
|
| 2078 |
"provenance": {
|
| 2079 |
+
"notes": "A 1,024-token delta recurrence with compact heads, zero entry state, and shared beta checks every output and final state."
|
| 2080 |
},
|
| 2081 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "delta", "scale": 0.35 },
|
| 2082 |
"inputs": {
|
|
|
|
| 2109 |
{
|
| 2110 |
"name": "delta_state_chunked_seq1024",
|
| 2111 |
"provenance": {
|
| 2112 |
+
"notes": "A 1,024-token delta recurrence with compact heads, supplied entry state, and per-head beta checks every output and final state."
|
| 2113 |
},
|
| 2114 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "delta", "scale": 0.35 },
|
| 2115 |
"inputs": {
|
|
|
|
| 2147 |
{
|
| 2148 |
"name": "gated_delta_zero_chunked_seq1024",
|
| 2149 |
"provenance": {
|
| 2150 |
+
"notes": "A 1,024-token gated-delta recurrence with compact heads, zero entry state, elementwise decay, and shared beta checks every output and final state."
|
| 2151 |
},
|
| 2152 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "gated_delta", "scale": 0.35 },
|
| 2153 |
"inputs": {
|
|
|
|
| 2190 |
{
|
| 2191 |
"name": "linear_zero_chunked_seq1024",
|
| 2192 |
"provenance": {
|
| 2193 |
+
"notes": "A 1,024-token linear recurrence with compact heads and zero entry state checks every output and final state."
|
| 2194 |
},
|
| 2195 |
"attrs": { "q_num_heads": 2, "kv_num_heads": 1, "update_rule": "linear", "scale": 0.35 },
|
| 2196 |
"inputs": {
|
|
|
|
| 2223 |
{
|
| 2224 |
"name": "linear_state_chunked_seq1024",
|
| 2225 |
"provenance": {
|
| 2226 |
+
"notes": "A 1,024-token linear recurrence with compact heads, supplied entry state, and grouped query heads checks every output and final state."
|
| 2227 |
},
|
| 2228 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "linear", "scale": 0.35 },
|
| 2229 |
"inputs": {
|
|
|
|
| 2256 |
{
|
| 2257 |
"name": "gated_delta_state_chunked_seq1024",
|
| 2258 |
"provenance": {
|
| 2259 |
+
"notes": "A 1,024-token gated-delta recurrence with compact heads, supplied entry state, per-head decay, and per-head beta checks every output and final state."
|
| 2260 |
},
|
| 2261 |
"attrs": { "q_num_heads": 4, "kv_num_heads": 2, "update_rule": "gated_delta", "scale": 0.35 },
|
| 2262 |
"inputs": {
|
|
|
|
| 3134 |
"outputT": { "dtype": "float32", "shape": [2, 1, 27] },
|
| 3135 |
"presentStateT": { "dtype": "float16", "shape": [2, 2, 3, 16, 9], "tolerance": 6e-8, "relTolerance": 0.0005 }
|
| 3136 |
}
|
| 3137 |
+
},
|
| 3138 |
+
{
|
| 3139 |
+
"name": "ort_gated_delta_state_window3_dk128_dv128_decode",
|
| 3140 |
+
"attrs": {
|
| 3141 |
+
"update_rule": "gated_delta",
|
| 3142 |
+
"q_num_heads": 2,
|
| 3143 |
+
"kv_num_heads": 2,
|
| 3144 |
+
"state_window": 3,
|
| 3145 |
+
"scale": 0.08838834764831843
|
| 3146 |
+
},
|
| 3147 |
+
"inputs": {
|
| 3148 |
+
"queryT": {
|
| 3149 |
+
"dtype": "float32",
|
| 3150 |
+
"shape": [2, 4, 256],
|
| 3151 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.017, "scale": 0.1 }
|
| 3152 |
+
},
|
| 3153 |
+
"keyT": {
|
| 3154 |
+
"dtype": "float32",
|
| 3155 |
+
"shape": [2, 4, 128],
|
| 3156 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.013, "scale": 0.1 }
|
| 3157 |
+
},
|
| 3158 |
+
"valueT": {
|
| 3159 |
+
"dtype": "float32",
|
| 3160 |
+
"shape": [2, 4, 256],
|
| 3161 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.023, "cosStep": 0.019, "scale": 0.1 }
|
| 3162 |
+
},
|
| 3163 |
+
"pastStateT": {
|
| 3164 |
+
"dtype": "float32",
|
| 3165 |
+
"shape": [3, 2, 2, 128, 128],
|
| 3166 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.007, "scale": 0.1 }
|
| 3167 |
+
},
|
| 3168 |
+
"decayT": {
|
| 3169 |
+
"dtype": "float32",
|
| 3170 |
+
"shape": [2, 4, 2],
|
| 3171 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.29, "cosStep": 0.31, "scale": 0.5, "offset": -0.5 }
|
| 3172 |
+
},
|
| 3173 |
+
"betaT": {
|
| 3174 |
+
"dtype": "float32",
|
| 3175 |
+
"shape": [2, 4, 2],
|
| 3176 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.4, "offset": 0.5 }
|
| 3177 |
+
}
|
| 3178 |
+
},
|
| 3179 |
+
"outputs": {
|
| 3180 |
+
"outputT": { "dtype": "float32", "shape": [2, 4, 256], "tolerance": 0.001, "relTolerance": 0.005 },
|
| 3181 |
+
"presentStateT": { "dtype": "float32", "shape": [3, 2, 2, 128, 128], "tolerance": 0.005, "relTolerance": 0.005 }
|
| 3182 |
+
}
|
| 3183 |
+
},
|
| 3184 |
+
{
|
| 3185 |
+
"name": "ort_gated_delta_state_window5_wider_than_sequence_with_past",
|
| 3186 |
+
"attrs": {
|
| 3187 |
+
"update_rule": "gated_delta",
|
| 3188 |
+
"q_num_heads": 2,
|
| 3189 |
+
"kv_num_heads": 2,
|
| 3190 |
+
"state_window": 5,
|
| 3191 |
+
"scale": 0.08838834764831843
|
| 3192 |
+
},
|
| 3193 |
+
"inputs": {
|
| 3194 |
+
"queryT": {
|
| 3195 |
+
"dtype": "float32",
|
| 3196 |
+
"shape": [2, 2, 256],
|
| 3197 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.017, "scale": 0.1 }
|
| 3198 |
+
},
|
| 3199 |
+
"keyT": {
|
| 3200 |
+
"dtype": "float32",
|
| 3201 |
+
"shape": [2, 2, 128],
|
| 3202 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.013, "scale": 0.1 }
|
| 3203 |
+
},
|
| 3204 |
+
"valueT": {
|
| 3205 |
+
"dtype": "float32",
|
| 3206 |
+
"shape": [2, 2, 256],
|
| 3207 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.023, "cosStep": 0.019, "scale": 0.1 }
|
| 3208 |
+
},
|
| 3209 |
+
"pastStateT": {
|
| 3210 |
+
"dtype": "float32",
|
| 3211 |
+
"shape": [5, 2, 2, 128, 128],
|
| 3212 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.007, "scale": 0.1 }
|
| 3213 |
+
},
|
| 3214 |
+
"decayT": {
|
| 3215 |
+
"dtype": "float32",
|
| 3216 |
+
"shape": [2, 2, 2],
|
| 3217 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.29, "cosStep": 0.31, "scale": 0.5, "offset": -0.5 }
|
| 3218 |
+
},
|
| 3219 |
+
"betaT": {
|
| 3220 |
+
"dtype": "float32",
|
| 3221 |
+
"shape": [2, 2, 2],
|
| 3222 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.4, "offset": 0.5 }
|
| 3223 |
+
}
|
| 3224 |
+
},
|
| 3225 |
+
"outputs": {
|
| 3226 |
+
"outputT": { "dtype": "float32", "shape": [2, 2, 256], "tolerance": 0.001, "relTolerance": 0.005 },
|
| 3227 |
+
"presentStateT": { "dtype": "float32", "shape": [5, 2, 2, 128, 128], "tolerance": 0.005, "relTolerance": 0.005 }
|
| 3228 |
+
}
|
| 3229 |
+
},
|
| 3230 |
+
{
|
| 3231 |
+
"name": "ort_gated_delta_standard_gqa_n4_q8_kv2",
|
| 3232 |
+
"attrs": { "update_rule": "gated_delta", "q_num_heads": 8, "kv_num_heads": 2, "scale": 0.17677669529663687 },
|
| 3233 |
+
"inputs": {
|
| 3234 |
+
"queryT": {
|
| 3235 |
+
"dtype": "float32",
|
| 3236 |
+
"shape": [2, 10, 256],
|
| 3237 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.017, "scale": 0.1 }
|
| 3238 |
+
},
|
| 3239 |
+
"keyT": {
|
| 3240 |
+
"dtype": "float32",
|
| 3241 |
+
"shape": [2, 10, 64],
|
| 3242 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.013, "scale": 0.1 }
|
| 3243 |
+
},
|
| 3244 |
+
"valueT": {
|
| 3245 |
+
"dtype": "float32",
|
| 3246 |
+
"shape": [2, 10, 128],
|
| 3247 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.023, "cosStep": 0.019, "scale": 0.1 }
|
| 3248 |
+
},
|
| 3249 |
+
"decayT": {
|
| 3250 |
+
"dtype": "float32",
|
| 3251 |
+
"shape": [2, 10, 2],
|
| 3252 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.29, "cosStep": 0.31, "scale": 0.5, "offset": -0.5 }
|
| 3253 |
+
},
|
| 3254 |
+
"betaT": {
|
| 3255 |
+
"dtype": "float32",
|
| 3256 |
+
"shape": [2, 10, 2],
|
| 3257 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.4, "offset": 0.5 }
|
| 3258 |
+
}
|
| 3259 |
+
},
|
| 3260 |
+
"outputs": {
|
| 3261 |
+
"outputT": { "dtype": "float32", "shape": [2, 10, 512], "tolerance": 0.00056, "relTolerance": 0.005 },
|
| 3262 |
+
"presentStateT": { "dtype": "float32", "shape": [2, 2, 32, 64], "tolerance": 0.000084, "relTolerance": 0.005 }
|
| 3263 |
+
}
|
| 3264 |
+
},
|
| 3265 |
+
{
|
| 3266 |
+
"name": "ort_gated_delta_standard_gqa_n16_q16_kv1",
|
| 3267 |
+
"attrs": { "update_rule": "gated_delta", "q_num_heads": 16, "kv_num_heads": 1, "scale": 0.17677669529663687 },
|
| 3268 |
+
"inputs": {
|
| 3269 |
+
"queryT": {
|
| 3270 |
+
"dtype": "float32",
|
| 3271 |
+
"shape": [2, 10, 512],
|
| 3272 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.017, "scale": 0.1 }
|
| 3273 |
+
},
|
| 3274 |
+
"keyT": {
|
| 3275 |
+
"dtype": "float32",
|
| 3276 |
+
"shape": [2, 10, 32],
|
| 3277 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.013, "scale": 0.1 }
|
| 3278 |
+
},
|
| 3279 |
+
"valueT": {
|
| 3280 |
+
"dtype": "float32",
|
| 3281 |
+
"shape": [2, 10, 64],
|
| 3282 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.023, "cosStep": 0.019, "scale": 0.1 }
|
| 3283 |
+
},
|
| 3284 |
+
"decayT": {
|
| 3285 |
+
"dtype": "float32",
|
| 3286 |
+
"shape": [2, 10, 1],
|
| 3287 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.29, "cosStep": 0.31, "scale": 0.5, "offset": -0.5 }
|
| 3288 |
+
},
|
| 3289 |
+
"betaT": {
|
| 3290 |
+
"dtype": "float32",
|
| 3291 |
+
"shape": [2, 10, 1],
|
| 3292 |
+
"data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.23, "scale": 0.4, "offset": 0.5 }
|
| 3293 |
+
}
|
| 3294 |
+
},
|
| 3295 |
+
"outputs": {
|
| 3296 |
+
"outputT": { "dtype": "float32", "shape": [2, 10, 1024], "tolerance": 0.00075, "relTolerance": 0.005 },
|
| 3297 |
+
"presentStateT": { "dtype": "float32", "shape": [2, 1, 32, 64], "tolerance": 0.00074, "relTolerance": 0.005 }
|
| 3298 |
+
}
|
| 3299 |
}
|
| 3300 |
]
|
| 3301 |
}
|