Xenova HF Staff commited on
Commit
08b6b14
·
verified ·
1 Parent(s): 59216c2

sync 6fdf6301e2bb

Browse files
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 + tuning 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,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.2
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
- 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 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
- {% endif %}
50
- {%- endmacro %}{% macro q_til_call(bt, q_head, h, i, scale) %}q_til({{ bt }}, {{ q_head }}, {% if usesDecay %}{{ h }}, {% endif %}{{ i }}, {{ scale }}){%- 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" %}
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
- 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 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
- {% 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.
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 dtype == "float16" %}
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
- {% endif %}
32
- {% if needQuery %}
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 %}var<workgroup> red_retrieved: array<f32, WG * TILE_V>;
123
- {% endif %}var<workgroup> red_preout: array<f32, WG * TILE_V>;
124
- {% if usesBeta %}var<workgroup> red_kq: array<f32, WG>;
 
 
 
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
- {% if field[1] == "f32" %}f32({{ field[2] }}){% else %}{{ field[2] }}u{% endif %},
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
- let rank = subgroupExclusiveAdd(1u);
114
- let count = subgroupAdd(1u);
115
- // Every lane performs the atomic so no collective follows a lane guard (a
116
- // subgroup op after a closed `if (rank == 0u)` block is the reconvergence
117
- // hazard the render gate flags): only the rank-0 lane adds its subgroup's
118
- // claim, every other lane adds 0 and discards its snapshot. The rank-0 lane is
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 not useSubgroups %}
 
 
 
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;{{ emit_tiled_setup(dvGroups=(2 if useSubgroups else dvGroups)) }}
 
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": "_com_microsoft_linearattention_webgpu_13a2aaf",
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": "VyrRYpaK7TsAfpKBMZpBiusQfnbDGhS5QIvnA/Vk6Tg=",
12
- "chunk-prep.wgsl.jinja": "Iq5ZLzNK/CtE7/i7Tm4pZi5rxMKq6dZ4wEpBqrapUdQ=",
13
- "chunk-scan.wgsl.jinja": "aPOMJxL7BhXVa05PwX9Uxl86dRbnuc9DkEwoFZq1Vvo=",
14
- "chunk-ut.wgsl.jinja": "MfZk82OOr4oJCJWi4vU/54BLerVSyS31uRsWeGDakw0=",
15
- "linear-attention.scalar.wgsl.jinja": "O8NRGtK2wpemLnSk96xiGRl50mqbqHQ3TiEkZpCGlOc=",
16
- "linear-attention.serial.wgsl.jinja": "gWOSDZOYAlaSi4wPdGqEq1UWfGxpVDxYTFeUdMDp6lw=",
17
- "linear-attention.vec4.wgsl.jinja": "bUlUYYjiumu6uCY/V7+EnbuhmBljoZXUnskKNYS0Hac=",
18
- "manifest.json": "7SA3QNshHZfuljHXCA9dVbCwtV4wiiH+4EO78NGq9mU=",
19
- "test.json": "NQ2VDXoy+rsZMH/agLHNEpOXHOydLQq+C2veCcWdpyk="
20
  }
21
  },
22
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
23
  "webgpu": {
24
- "manifestSpec": "2.0",
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": "gated_delta_headdim6_seq128_offset_value_scale_lock",
372
  "provenance": {
373
- "notes": "Long-recurrence scale lock at headDimK=6, the non-vec4 width served by the serial-small and scalar gated-delta routes. Scaling the key to |k|^2 about 0.7 and offsetting V to 1.0 makes the 128-step recurrence materially update and converge. Missing or duplicated decay, incorrect beta, and uniform output scaling therefore produce clear output errors."
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 select the scalar linear-rule implementation over a model-shaped 128-token zero-state recurrence."
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": "linear_zero_f16_seq128_offset_value_scale_lock",
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 selects the serial small-key-dimension route and verifies that its supplied initial state is incorporated."
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": "linear_state_f16_seq128_offset_value_scale_lock",
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": "headDimK=130 (not %4 -> scalar variant; >128 so WG=256 with 130 active lanes -> tid<head_dim_k lane masking in tree reduce). headDimV=10, TILE_V=4 -> dv_tiles=3 with a partial last tile (2/4 valid) exercising dv_start+j<head_dim_v guards on output AND present_state. gated_delta rule, per-head decay + per-head beta."
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 `updateRule` selects the schema-default gated-delta recurrence with decay, beta, and past state. Key head size 4 exercises the vec4 route."
807
  },
808
  "attrs": { "q_num_heads": 2, "kv_num_heads": 1 },
809
  "inputs": {
@@ -892,9 +892,9 @@
892
  }
893
  },
894
  {
895
- "name": "gated_delta_f32_dk128_dv128_offset_value_scale_lock",
896
  "provenance": {
897
- "notes": "Scale lock for the f32 dK=dV=128 gated-delta route. Scaling the key to |k|^2 about 0.7 makes the delta correction material, while V around 1.0 and an O(1) query keep outputT well-conditioned. This makes a missing or duplicated decay, incorrect beta, or uniform output scaling observable on outputT rather than relying only on presentStateT."
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": "gated_delta_f16_dk128_dv128_offset_value_scale_lock",
982
  "provenance": {
983
- "notes": "Float16 dK=dV=128 gated-delta scale lock covering both vec4 and scalar state routes. Scaling the key to |k|^2 about 0.7 makes the delta correction material, while V around 1.0 and an O(1) query keep outputs well-conditioned. Missing or duplicated decay, incorrect beta, and uniform output scaling therefore exceed tolerance on 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,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 (PR #28412)",
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 (PR #28412)",
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": "headDimK=17 is above serialHeadDimFits (dk<=16) and not %4, so neither linear_zero_serial_small_dk nor linear_zero_vec4 is eligible and the scalar zero-state kernel must run. WG=pow2ceil(17)=32 with 17 active lanes exercises the masked dk reduction; TILE_V=2 covers head_dim_v in one tile."
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 rather than per query head and several KV heads share a query head (KV head h reads query head floor(h * qNumHeads / kvNumHeads)). The gated-delta rule, which is the regime the inverse layout exists for."
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": "With past state omitted, key head size 17 exceeds the serial limit and is not divisible by four, selecting the scalar zero-state recurrence."
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": "Omitted past_state with headDimK=20 crosses the serial-kernel ceiling while retaining vec4 alignment, selecting the vector zero-state recurrence and a partial dV tile."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the gated rule with a zero entry state and an elementwise decay gate."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the gated rule with a supplied entry state and a per-head decay gate."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the delta rule with a zero entry state and a shared beta column."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the delta rule with a supplied entry state and per-head beta."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the gated-delta rule with a zero entry state, elementwise decay and a shared beta column."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the linear rule with a zero entry state."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the linear rule with a supplied entry state and grouped query heads."
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": "Selects the chunked prefill decomposition, which needs seqLength >= 1024. Compact head dims keep the fixture small while covering the gated-delta rule with a supplied entry state, per-head decay and per-head beta."
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
  }