Xenova HF Staff commited on
Commit
eba4596
·
verified ·
1 Parent(s): 4e80716

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -82,13 +82,13 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
82
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
83
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
84
  - [`test.json`](build/webgpu/test.json) — correctness cases
85
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
86
  - [`expert-group-slots.wgsl.jinja`](build/webgpu/expert-group-slots.wgsl.jinja)
 
87
  - [`moe-ffn-gemv.wgsl.jinja`](build/webgpu/moe-ffn-gemv.wgsl.jinja)
88
  - [`moe-ffn-grouped.wgsl.jinja`](build/webgpu/moe-ffn-grouped.wgsl.jinja)
89
  - [`moe-ffn-stage.wgsl.jinja`](build/webgpu/moe-ffn-stage.wgsl.jinja)
90
  - [`moe-grouped-sgmat.wgsl.jinja`](build/webgpu/moe-grouped-sgmat.wgsl.jinja)
91
- - [`moe-mix-stage.wgsl.jinja`](build/webgpu/moe-mix-stage.wgsl.jinja)
92
  - [`moe-output-gemv.wgsl.jinja`](build/webgpu/moe-output-gemv.wgsl.jinja)
93
  - [`moe-output-grouped.wgsl.jinja`](build/webgpu/moe-output-grouped.wgsl.jinja)
94
  - [`moe-output-stage.wgsl.jinja`](build/webgpu/moe-output-stage.wgsl.jinja)
@@ -97,7 +97,7 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
97
  ## Use with `@huggingface/kernels`
98
 
99
  ```sh
100
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
101
  ```
102
 
103
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
82
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
83
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
84
  - [`test.json`](build/webgpu/test.json) — correctness cases
85
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
86
  - [`expert-group-slots.wgsl.jinja`](build/webgpu/expert-group-slots.wgsl.jinja)
87
+ - [`expert-slot-mix.wgsl.jinja`](build/webgpu/expert-slot-mix.wgsl.jinja)
88
  - [`moe-ffn-gemv.wgsl.jinja`](build/webgpu/moe-ffn-gemv.wgsl.jinja)
89
  - [`moe-ffn-grouped.wgsl.jinja`](build/webgpu/moe-ffn-grouped.wgsl.jinja)
90
  - [`moe-ffn-stage.wgsl.jinja`](build/webgpu/moe-ffn-stage.wgsl.jinja)
91
  - [`moe-grouped-sgmat.wgsl.jinja`](build/webgpu/moe-grouped-sgmat.wgsl.jinja)
 
92
  - [`moe-output-gemv.wgsl.jinja`](build/webgpu/moe-output-gemv.wgsl.jinja)
93
  - [`moe-output-grouped.wgsl.jinja`](build/webgpu/moe-output-grouped.wgsl.jinja)
94
  - [`moe-output-stage.wgsl.jinja`](build/webgpu/moe-output-stage.wgsl.jinja)
 
97
  ## Use with `@huggingface/kernels`
98
 
99
  ```sh
100
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
101
  ```
102
 
103
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/{moe-mix-stage.wgsl.jinja → expert-slot-mix.wgsl.jinja} RENAMED
@@ -1,3 +1,8 @@
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // Routed sum for the grouped schedule. The grouped FC2 stage produces one projected row per
@@ -11,9 +16,7 @@ const WG: u32 = {{ workgroupSize }}u;
11
 
12
  @compute @workgroup_size(WG, 1, 1)
13
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
14
- // 2D-folded flat index: gid.y carries the high bits past the per-axis dispatch fold width.
15
- // Reduces to gid.x when the dispatch does not fold.
16
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
17
  let total = params.tokenCount * HIDDEN;
18
  if (index >= total) {
19
  return;
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
6
  {{ env.wgsl.resourceDeclarations }}
7
 
8
  // Routed sum for the grouped schedule. The grouped FC2 stage produces one projected row per
 
16
 
17
  @compute @workgroup_size(WG, 1, 1)
18
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
19
+ {{ flat_index_2d("WG", "index", "") }}
 
 
20
  let total = params.tokenCount * HIDDEN;
21
  if (index >= total) {
22
  return;
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.MoE",
3
- "id": "_com_microsoft_moe_webgpu_7f9cfff",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -9,71 +9,71 @@
9
  "files": {
10
  "bench.json": "dk7Y+0dMicwCWAn3BjOAjL0VyY0CqCabC1DuWCMDMis=",
11
  "expert-group-slots.wgsl.jinja": "Ta+3H2FA1qRRzksdkwgKLZiO+AokmuM+B9JVN8cv1EQ=",
12
- "manifest.json": "CRRrdafIg5G627yy63AxYRy9pW6V5I93qKobNmgV6vs=",
13
- "moe-ffn-gemv.wgsl.jinja": "2LCF3pr7u1cfGz0eUHwkSXr5JQbFgynU8IZS7wFnAuU=",
14
- "moe-ffn-grouped.wgsl.jinja": "s1GymR5z3Kj0z7/NOdw6jFgA6D5LM4yfYISrcqmoy+Y=",
15
- "moe-ffn-stage.wgsl.jinja": "NNU6+WY7JCjo8itC+pYuTQIwPbjX2PZb1P2WXNS3z/Y=",
16
- "moe-grouped-sgmat.wgsl.jinja": "+bkbK0RGxiTY2JlmPLkynnUctqoUcRkdu5OMEl+d9nA=",
17
- "moe-mix-stage.wgsl.jinja": "+xEc6boXboFbx+4AJ2/KWElM04LRuufmiIKfEK51/oo=",
18
- "moe-output-gemv.wgsl.jinja": "5yj+BoTsRHiHW92Q8fVvtojdkY7r1jHupMX+FPVQMBk=",
19
- "moe-output-grouped.wgsl.jinja": "qzP5G/WtwF4zOBrtySCLTCBflHuvKlU8dk4/CRIzxe8=",
20
- "moe-output-stage.wgsl.jinja": "DqLA91ZM2ELUfyDxELZMxHqb0qgmDh/V6Q9k1wVkQH8=",
21
- "moe-route-stage.wgsl.jinja": "aj6lm39GKmPH+aAkK7f4i5JpGdb47wPNURH4JQW/A1w=",
22
- "test.json": "5rXwo178i5IeETKSYI6HhcwOn6rOVmXBaiy8J12coIk="
23
  }
24
  },
25
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
26
  "webgpu": {
27
- "manifestSpec": "2.0",
28
  "variants": {
29
  "split_routed_fc1plain_fc3none_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
30
  "gemv_routed_fc1plain_fc3none_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
31
  "split_routed_fc1plain_fc3none_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
32
  "gemv_routed_fc1plain_fc3none_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
33
  "split_routed_fc1plain_fc3plain_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
34
  "gemv_routed_fc1plain_fc3plain_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
35
  "split_routed_fc1plain_fc3plain_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
36
  "gemv_routed_fc1plain_fc3plain_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
37
  "split_routed_fc1plain_fc3biased_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
38
  "gemv_routed_fc1plain_fc3biased_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
39
  "split_routed_fc1plain_fc3biased_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
40
  "gemv_routed_fc1plain_fc3biased_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
41
  "split_routed_fc1bias_fc3none_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
42
  "gemv_routed_fc1bias_fc3none_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
43
  "split_routed_fc1bias_fc3none_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
44
  "gemv_routed_fc1bias_fc3none_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
45
  "split_routed_fc1bias_fc3plain_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
46
  "gemv_routed_fc1bias_fc3plain_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
47
  "split_routed_fc1bias_fc3plain_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
48
  "gemv_routed_fc1bias_fc3plain_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
49
  "split_routed_fc1bias_fc3biased_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
50
  "gemv_routed_fc1bias_fc3biased_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
 
 
51
  "split_routed_fc1bias_fc3biased_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
52
  "gemv_routed_fc1bias_fc3biased_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
53
- "sgmat_grouped_routed_fc1plain_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
54
- "grouped_routed_fc1plain_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
55
- "sgmat_grouped_routed_fc1plain_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
56
- "grouped_routed_fc1plain_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
57
- "sgmat_grouped_routed_fc1plain_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
58
- "grouped_routed_fc1plain_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
59
- "sgmat_grouped_routed_fc1plain_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
60
- "grouped_routed_fc1plain_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
61
- "sgmat_grouped_routed_fc1plain_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
62
- "grouped_routed_fc1plain_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
63
- "sgmat_grouped_routed_fc1plain_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
64
- "grouped_routed_fc1plain_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
65
- "sgmat_grouped_routed_fc1bias_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
66
- "grouped_routed_fc1bias_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
67
- "sgmat_grouped_routed_fc1bias_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
68
- "grouped_routed_fc1bias_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
69
- "sgmat_grouped_routed_fc1bias_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
70
- "grouped_routed_fc1bias_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
71
- "sgmat_grouped_routed_fc1bias_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
72
- "grouped_routed_fc1bias_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
73
- "sgmat_grouped_routed_fc1bias_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
74
- "grouped_routed_fc1bias_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
75
- "sgmat_grouped_routed_fc1bias_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
76
- "grouped_routed_fc1bias_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-mix-stage.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"]
77
  }
78
  }
79
  }
 
1
  {
2
  "name": "com.microsoft.MoE",
3
+ "id": "_com_microsoft_moe_webgpu_81aa200",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
9
  "files": {
10
  "bench.json": "dk7Y+0dMicwCWAn3BjOAjL0VyY0CqCabC1DuWCMDMis=",
11
  "expert-group-slots.wgsl.jinja": "Ta+3H2FA1qRRzksdkwgKLZiO+AokmuM+B9JVN8cv1EQ=",
12
+ "expert-slot-mix.wgsl.jinja": "VwX7A1f2r+lxMzG/I5r0KUpGr0Y5gW/CrI/wdwx0Ah4=",
13
+ "manifest.json": "wBjvzcYwg6LKo5NBDt5ygh3hYrahgWaBsqIBLN2bfXY=",
14
+ "moe-ffn-gemv.wgsl.jinja": "jeIxUFT5qPPAB55dEtF1kV897NLEOxKBMGMwrpYuOgQ=",
15
+ "moe-ffn-grouped.wgsl.jinja": "20KqFJ5ukAEOmtA+p3FuttT6GHFWuoyq3bMIT+O3vyQ=",
16
+ "moe-ffn-stage.wgsl.jinja": "Abb7HoT+GLy160X3+CX2pOLyQPIF3rNxq/0xRPkwHM0=",
17
+ "moe-grouped-sgmat.wgsl.jinja": "blIiUWR5EYpFXm2Gdgs8gw3G/Ttq8yuaqiRJNSoeQgo=",
18
+ "moe-output-gemv.wgsl.jinja": "s3IrdnP/Os57QXA1w6f//dMa4pK5uarFK0ur3tiJipw=",
19
+ "moe-output-grouped.wgsl.jinja": "9m4bid355i/L86ZqBjNMBS8YhU6Ta+RaIy8cg7t+rgk=",
20
+ "moe-output-stage.wgsl.jinja": "lII4acbshM29uYlxAlJU+CcR+vRTREMc6CrJoYEnl5M=",
21
+ "moe-route-stage.wgsl.jinja": "uenMx8qmkRe13OPiwga5rGndXTGuqIJKrurm1Shd8xM=",
22
+ "test.json": "ififQJvtLtWJR4hKuaZHpgwomm3UdMab9oorgFBZN74="
23
  }
24
  },
25
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
26
  "webgpu": {
27
+ "manifestSpec": "2.1",
28
  "variants": {
29
  "split_routed_fc1plain_fc3none_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
30
  "gemv_routed_fc1plain_fc3none_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
31
+ "sgmat_grouped_routed_fc1plain_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
32
+ "grouped_routed_fc1plain_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
33
  "split_routed_fc1plain_fc3none_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
34
  "gemv_routed_fc1plain_fc3none_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
35
+ "sgmat_grouped_routed_fc1plain_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
36
+ "grouped_routed_fc1plain_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
37
  "split_routed_fc1plain_fc3plain_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
38
  "gemv_routed_fc1plain_fc3plain_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
39
+ "sgmat_grouped_routed_fc1plain_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
40
+ "grouped_routed_fc1plain_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
41
  "split_routed_fc1plain_fc3plain_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
42
  "gemv_routed_fc1plain_fc3plain_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
43
+ "sgmat_grouped_routed_fc1plain_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
44
+ "grouped_routed_fc1plain_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
45
  "split_routed_fc1plain_fc3biased_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
46
  "gemv_routed_fc1plain_fc3biased_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
47
+ "sgmat_grouped_routed_fc1plain_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
48
+ "grouped_routed_fc1plain_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
49
  "split_routed_fc1plain_fc3biased_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
50
  "gemv_routed_fc1plain_fc3biased_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
51
+ "sgmat_grouped_routed_fc1plain_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
52
+ "grouped_routed_fc1plain_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
53
  "split_routed_fc1bias_fc3none_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
54
  "gemv_routed_fc1bias_fc3none_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
55
+ "sgmat_grouped_routed_fc1bias_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
56
+ "grouped_routed_fc1bias_fc3none_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
57
  "split_routed_fc1bias_fc3none_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
58
  "gemv_routed_fc1bias_fc3none_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
59
+ "sgmat_grouped_routed_fc1bias_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
60
+ "grouped_routed_fc1bias_fc3none_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
61
  "split_routed_fc1bias_fc3plain_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
62
  "gemv_routed_fc1bias_fc3plain_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
63
+ "sgmat_grouped_routed_fc1bias_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
64
+ "grouped_routed_fc1bias_fc3plain_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
65
  "split_routed_fc1bias_fc3plain_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
66
  "gemv_routed_fc1bias_fc3plain_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
67
+ "sgmat_grouped_routed_fc1bias_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
68
+ "grouped_routed_fc1bias_fc3plain_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
69
  "split_routed_fc1bias_fc3biased_fc2plain": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
70
  "gemv_routed_fc1bias_fc3biased_fc2plain": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
71
+ "sgmat_grouped_routed_fc1bias_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
72
+ "grouped_routed_fc1bias_fc3biased_fc2plain": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
73
  "split_routed_fc1bias_fc3biased_fc2bias": ["moe-ffn-stage.wgsl.jinja", "moe-output-stage.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
74
  "gemv_routed_fc1bias_fc3biased_fc2bias": ["moe-ffn-gemv.wgsl.jinja", "moe-output-gemv.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
75
+ "sgmat_grouped_routed_fc1bias_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-grouped-sgmat.wgsl.jinja", "moe-route-stage.wgsl.jinja"],
76
+ "grouped_routed_fc1bias_fc3biased_fc2bias": ["expert-group-slots.wgsl.jinja", "expert-slot-mix.wgsl.jinja", "moe-ffn-grouped.wgsl.jinja", "moe-output-grouped.wgsl.jinja", "moe-route-stage.wgsl.jinja"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77
  }
78
  }
79
  }
build/webgpu/moe-ffn-gemv.wgsl.jinja CHANGED
@@ -12,15 +12,22 @@ const FC1_ROWS: u32 = {{ fc1Rows }}u;
12
  const TOP_K: u32 = {{ topK }}u;
13
  const LANES: u32 = {{ decodeLanes }}u;
14
  const ROWS: u32 = {{ decodeRows }}u;
15
- {% if activation == "gelu" %}fn tanh_safe(x: f32) -> f32 {
 
 
 
16
  if (x > 10.0) { return 1.0; }
17
  if (x < -10.0) { return -1.0; }
 
 
 
18
  return tanh(x);
19
  }
20
 
21
  fn gelu_tanh(v: f32) -> f32 {
22
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
23
- }{% endif %}
 
24
  {% set isSwiglu = activation == "swiglu" %}
25
  {% set secondFromFc3 = hasFc3 and (not isSwiglu or swigluFusion == 0) %}
26
  {% set hasSecond = isSwiglu or hasFc3 %}
@@ -45,7 +52,6 @@ fn swiglu(gate_in: f32, up_in: f32) -> f32 {
45
  }
46
  {% endif %}
47
 
48
-
49
  {% macro rowlane_fold(arrays, lanes="LANES", lane="lane", slot="slot") %}
50
  var n = {{ lanes }};
51
  while (n > 1u) {
@@ -57,9 +63,7 @@ fn swiglu(gate_in: f32, up_in: f32) -> f32 {
57
  }
58
  workgroupBarrier();
59
  n = half;
60
- }
61
- {%- endmacro %}
62
-
63
  var<workgroup> primary_partial: array<f32, {{ decodeLanes * decodeRows }}>;
64
  {% if hasSecond %}
65
  var<workgroup> secondary_partial: array<f32, {{ decodeLanes * decodeRows }}>;
 
12
  const TOP_K: u32 = {{ topK }}u;
13
  const LANES: u32 = {{ decodeLanes }}u;
14
  const ROWS: u32 = {{ decodeRows }}u;
15
+ {% if activation == "gelu" %}
16
+ fn tanh_safe(x: f32) -> f32 {
17
+ // tanh rounds to its saturated value for these tails in f32. Return that
18
+ // value directly, including for infinite input, before invoking the builtin.
19
  if (x > 10.0) { return 1.0; }
20
  if (x < -10.0) { return -1.0; }
21
+ // For tiny |x|, return x directly to preserve its sign and magnitude without
22
+ // relying on backend-specific builtin behavior near zero.
23
+ if (x > -1.0e-4 && x < 1.0e-4) { return x; }
24
  return tanh(x);
25
  }
26
 
27
  fn gelu_tanh(v: f32) -> f32 {
28
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
29
+ }
30
+ {% endif %}
31
  {% set isSwiglu = activation == "swiglu" %}
32
  {% set secondFromFc3 = hasFc3 and (not isSwiglu or swigluFusion == 0) %}
33
  {% set hasSecond = isSwiglu or hasFc3 %}
 
52
  }
53
  {% endif %}
54
 
 
55
  {% macro rowlane_fold(arrays, lanes="LANES", lane="lane", slot="slot") %}
56
  var n = {{ lanes }};
57
  while (n > 1u) {
 
63
  }
64
  workgroupBarrier();
65
  n = half;
66
+ }{% endmacro %}
 
 
67
  var<workgroup> primary_partial: array<f32, {{ decodeLanes * decodeRows }}>;
68
  {% if hasSecond %}
69
  var<workgroup> secondary_partial: array<f32, {{ decodeLanes * decodeRows }}>;
build/webgpu/moe-ffn-grouped.wgsl.jinja CHANGED
@@ -12,15 +12,22 @@ const NTILE: u32 = {{ groupTileN }}u;
12
  const KTILE: u32 = {{ groupTileK }}u;
13
  const THREADS_SIDE: u32 = {{ groupThreads }}u;
14
  const THREADS: u32 = THREADS_SIDE * THREADS_SIDE;
15
- {% if activation == "gelu" %}fn tanh_safe(x: f32) -> f32 {
 
 
 
16
  if (x > 10.0) { return 1.0; }
17
  if (x < -10.0) { return -1.0; }
 
 
 
18
  return tanh(x);
19
  }
20
 
21
  fn gelu_tanh(v: f32) -> f32 {
22
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
23
- }{% endif %}
 
24
  {% set isSwiglu = activation == "swiglu" %}
25
  {% set secondFromFc3 = hasFc3 and (not isSwiglu or swigluFusion == 0) %}
26
  {% set hasSecond = isSwiglu or hasFc3 %}
@@ -51,7 +58,6 @@ fn swiglu(gate_in: f32, up_in: f32) -> f32 {
51
  }
52
  {% endif %}
53
 
54
-
55
  const KVEC: u32 = {{ groupTileKVec }}u;
56
  {% macro stage_group_tiles(aLoad, bLoad, b2Load, kExtent, nExtent, guarded, bLoadVec4="", b2LoadVec4="", aLoadVec4="") %}
57
  for (var idx = tid; idx < MTILE * KVEC; idx = idx + THREADS) {
@@ -104,8 +110,7 @@ const KVEC: u32 = {{ groupTileKVec }}u;
104
  {% if b2Load %}
105
  b2_tile[idx] = b2_vec;
106
  {% endif %}
107
- }
108
- {%- endmacro %}
109
 
110
  {% macro group_tile_loop(aLoad, bLoad, b2Load, kExtent, nExtent, regM, regN, bLoadVec4="", b2LoadVec4="", aLoadVec4="") %}
111
  {% for r in range(regM) %}
@@ -158,9 +163,7 @@ const KVEC: u32 = {{ groupTileKVec }}u;
158
  // Orders this step's tile reads before the next step overwrites them.
159
  workgroupBarrier();
160
  k_base = k_base + KTILE;
161
- }
162
- {%- endmacro %}
163
-
164
  var<workgroup> row_slot: array<u32, {{ groupTileM }}>;
165
  var<workgroup> a_tile: array<vec4<f32>, {{ groupTileM * groupTileKVec }}>;
166
  var<workgroup> b_tile: array<vec4<f32>, {{ groupTileN * groupTileKVec }}>;
 
12
  const KTILE: u32 = {{ groupTileK }}u;
13
  const THREADS_SIDE: u32 = {{ groupThreads }}u;
14
  const THREADS: u32 = THREADS_SIDE * THREADS_SIDE;
15
+ {% if activation == "gelu" %}
16
+ fn tanh_safe(x: f32) -> f32 {
17
+ // tanh rounds to its saturated value for these tails in f32. Return that
18
+ // value directly, including for infinite input, before invoking the builtin.
19
  if (x > 10.0) { return 1.0; }
20
  if (x < -10.0) { return -1.0; }
21
+ // For tiny |x|, return x directly to preserve its sign and magnitude without
22
+ // relying on backend-specific builtin behavior near zero.
23
+ if (x > -1.0e-4 && x < 1.0e-4) { return x; }
24
  return tanh(x);
25
  }
26
 
27
  fn gelu_tanh(v: f32) -> f32 {
28
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
29
+ }
30
+ {% endif %}
31
  {% set isSwiglu = activation == "swiglu" %}
32
  {% set secondFromFc3 = hasFc3 and (not isSwiglu or swigluFusion == 0) %}
33
  {% set hasSecond = isSwiglu or hasFc3 %}
 
58
  }
59
  {% endif %}
60
 
 
61
  const KVEC: u32 = {{ groupTileKVec }}u;
62
  {% macro stage_group_tiles(aLoad, bLoad, b2Load, kExtent, nExtent, guarded, bLoadVec4="", b2LoadVec4="", aLoadVec4="") %}
63
  for (var idx = tid; idx < MTILE * KVEC; idx = idx + THREADS) {
 
110
  {% if b2Load %}
111
  b2_tile[idx] = b2_vec;
112
  {% endif %}
113
+ }{% endmacro %}
 
114
 
115
  {% macro group_tile_loop(aLoad, bLoad, b2Load, kExtent, nExtent, regM, regN, bLoadVec4="", b2LoadVec4="", aLoadVec4="") %}
116
  {% for r in range(regM) %}
 
163
  // Orders this step's tile reads before the next step overwrites them.
164
  workgroupBarrier();
165
  k_base = k_base + KTILE;
166
+ }{% endmacro %}
 
 
167
  var<workgroup> row_slot: array<u32, {{ groupTileM }}>;
168
  var<workgroup> a_tile: array<vec4<f32>, {{ groupTileM * groupTileKVec }}>;
169
  var<workgroup> b_tile: array<vec4<f32>, {{ groupTileN * groupTileKVec }}>;
build/webgpu/moe-ffn-stage.wgsl.jinja CHANGED
@@ -1,3 +1,8 @@
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // FC1 (and, where the schema splits them, FC3) projection plus the activation. One thread owns
@@ -11,15 +16,22 @@ const INTER_DIV: u32 = max(1u, INTER);
11
  const FC1_ROWS: u32 = {{ fc1Rows }}u;
12
  const TOP_K: u32 = {{ topK }}u;
13
  const WG: u32 = {{ workgroupSize }}u;
14
- {% if activation == "gelu" %}fn tanh_safe(x: f32) -> f32 {
 
 
 
15
  if (x > 10.0) { return 1.0; }
16
  if (x < -10.0) { return -1.0; }
 
 
 
17
  return tanh(x);
18
  }
19
 
20
  fn gelu_tanh(v: f32) -> f32 {
21
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
22
- }{% endif %}
 
23
 
24
  fn fc1_row(expert: u32, row: u32) -> u32 {
25
  return (expert * FC1_ROWS + row) * HIDDEN;
@@ -68,12 +80,9 @@ fn swiglu(gate_in: f32, up_in: f32) -> f32 {
68
  }
69
  {% endif %}
70
 
71
-
72
  @compute @workgroup_size(WG, 1, 1)
73
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
74
- // 2D-folded flat index: gid.y carries the high bits past the per-axis dispatch fold width.
75
- // Reduces to gid.x when the dispatch does not fold.
76
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
77
  let total = params.tokenCount * TOP_K * INTER;
78
  if (index >= total) {
79
  return;
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
6
  {{ env.wgsl.resourceDeclarations }}
7
 
8
  // FC1 (and, where the schema splits them, FC3) projection plus the activation. One thread owns
 
16
  const FC1_ROWS: u32 = {{ fc1Rows }}u;
17
  const TOP_K: u32 = {{ topK }}u;
18
  const WG: u32 = {{ workgroupSize }}u;
19
+ {% if activation == "gelu" %}
20
+ fn tanh_safe(x: f32) -> f32 {
21
+ // tanh rounds to its saturated value for these tails in f32. Return that
22
+ // value directly, including for infinite input, before invoking the builtin.
23
  if (x > 10.0) { return 1.0; }
24
  if (x < -10.0) { return -1.0; }
25
+ // For tiny |x|, return x directly to preserve its sign and magnitude without
26
+ // relying on backend-specific builtin behavior near zero.
27
+ if (x > -1.0e-4 && x < 1.0e-4) { return x; }
28
  return tanh(x);
29
  }
30
 
31
  fn gelu_tanh(v: f32) -> f32 {
32
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
33
+ }
34
+ {% endif %}
35
 
36
  fn fc1_row(expert: u32, row: u32) -> u32 {
37
  return (expert * FC1_ROWS + row) * HIDDEN;
 
80
  }
81
  {% endif %}
82
 
 
83
  @compute @workgroup_size(WG, 1, 1)
84
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
85
+ {{ flat_index_2d("WG", "index", "") }}
 
 
86
  let total = params.tokenCount * TOP_K * INTER;
87
  if (index >= total) {
88
  return;
build/webgpu/moe-grouped-sgmat.wgsl.jinja CHANGED
@@ -4,22 +4,29 @@ enable subgroup_size_control;
4
  {% endif %}
5
  enable chromium_experimental_subgroup_matrix;
6
  diagnostic(off, chromium.subgroup_matrix_uniformity);
7
-
8
  {{ env.wgsl.resourceDeclarations }}
9
  {% set ffn = matrixStage == "ffn" %}
10
  {% set second = ffn and (hasFc3 or activation == "swiglu") %}
11
  {% set reduction = hidden if ffn else inter %}
12
  {% set columns = inter if ffn else hidden %}
13
- {% if ffn and activation == "gelu" %}fn tanh_safe(x: f32) -> f32 {
 
 
 
14
  if (x > 10.0) { return 1.0; }
15
  if (x < -10.0) { return -1.0; }
 
 
 
16
  return tanh(x);
17
  }
18
 
19
  fn gelu_tanh(v: f32) -> f32 {
20
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
21
- }{% endif %}
22
- {% if ffn %}{% if activation == "swiglu" %}
 
 
23
 
24
  fn swiglu(gate_in: f32, up_in: f32) -> f32 {
25
  {% if hasSwigluLimit %}
@@ -45,7 +52,7 @@ const TILE_K: u32 = 32u;
45
  const SUB_ROWS: u32 = 16u;
46
  const SUB_COLS: u32 = {{ 16 if second else 32 }}u;
47
  // Two independent f32 chains reduce rounding growth over long reductions.
48
- // Jinja assigns alternating 8-wide steps to the chains at compile time.
49
  // Each subgroup publishes four result banks per chain after A is dead.
50
  // The two row groups reuse those banks with barriers on both sides.
51
  var<workgroup> tile_A: array<f32, {{ (groupedSgmatSharedBytes / 4) | int }}>;
@@ -76,7 +83,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
76
  let subtile_id = local_idx / sg_size;
77
  let subtile_idy = subtile_id % 2u;
78
  let subtile_idx = subtile_id / 2u;
79
- let n_base = wid.y * {{ 32 if second else 64 }}u;
80
  {% for r in range(2) %}{% for c in range(4) %}{% for chain in range(2) %}
81
  var mat{{ ["C","D","E","F"][chain] }}{{ r }}{{ c }}: subgroup_matrix_result<f32, 8, 8>;
82
  {% endfor %}{% endfor %}{% endfor %}
@@ -161,6 +168,8 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
161
  {% endfor %}{% endfor %}
162
  }
163
  }
164
- {% if r == 0 %} workgroupBarrier();{% endif %}
 
 
165
  {% endfor %}
166
  }
 
4
  {% endif %}
5
  enable chromium_experimental_subgroup_matrix;
6
  diagnostic(off, chromium.subgroup_matrix_uniformity);
 
7
  {{ env.wgsl.resourceDeclarations }}
8
  {% set ffn = matrixStage == "ffn" %}
9
  {% set second = ffn and (hasFc3 or activation == "swiglu") %}
10
  {% set reduction = hidden if ffn else inter %}
11
  {% set columns = inter if ffn else hidden %}
12
+ {% if ffn and activation == "gelu" %}
13
+ fn tanh_safe(x: f32) -> f32 {
14
+ // tanh rounds to its saturated value for these tails in f32. Return that
15
+ // value directly, including for infinite input, before invoking the builtin.
16
  if (x > 10.0) { return 1.0; }
17
  if (x < -10.0) { return -1.0; }
18
+ // For tiny |x|, return x directly to preserve its sign and magnitude without
19
+ // relying on backend-specific builtin behavior near zero.
20
+ if (x > -1.0e-4 && x < 1.0e-4) { return x; }
21
  return tanh(x);
22
  }
23
 
24
  fn gelu_tanh(v: f32) -> f32 {
25
  return 0.5 * v * (1.0 + tanh_safe(0.7978845608028654 * (v + 0.044715 * v * v * v)));
26
+ }
27
+ {% endif %}
28
+ {% if ffn %}
29
+ {% if activation == "swiglu" %}
30
 
31
  fn swiglu(gate_in: f32, up_in: f32) -> f32 {
32
  {% if hasSwigluLimit %}
 
52
  const SUB_ROWS: u32 = 16u;
53
  const SUB_COLS: u32 = {{ 16 if second else 32 }}u;
54
  // Two independent f32 chains reduce rounding growth over long reductions.
55
+ // Alternating 8-wide steps are assigned to the chains at compile time.
56
  // Each subgroup publishes four result banks per chain after A is dead.
57
  // The two row groups reuse those banks with barriers on both sides.
58
  var<workgroup> tile_A: array<f32, {{ (groupedSgmatSharedBytes / 4) | int }}>;
 
83
  let subtile_id = local_idx / sg_size;
84
  let subtile_idy = subtile_id % 2u;
85
  let subtile_idx = subtile_id / 2u;
86
+ let n_base = wid.y * {{ tileCols }}u;
87
  {% for r in range(2) %}{% for c in range(4) %}{% for chain in range(2) %}
88
  var mat{{ ["C","D","E","F"][chain] }}{{ r }}{{ c }}: subgroup_matrix_result<f32, 8, 8>;
89
  {% endfor %}{% endfor %}{% endfor %}
 
168
  {% endfor %}{% endfor %}
169
  }
170
  }
171
+ {% if r == 0 %}
172
+ workgroupBarrier();
173
+ {% endif %}
174
  {% endfor %}
175
  }
build/webgpu/moe-output-gemv.wgsl.jinja CHANGED
@@ -19,9 +19,7 @@ const ROWS: u32 = {{ decodeRows }}u;
19
  }
20
  workgroupBarrier();
21
  n = half;
22
- }
23
- {%- endmacro %}
24
-
25
  var<workgroup> partial: array<f32, {{ decodeLanes * decodeRows }}>;
26
 
27
  @compute @workgroup_size(LANES, ROWS, 1)
 
19
  }
20
  workgroupBarrier();
21
  n = half;
22
+ }{% endmacro %}
 
 
23
  var<workgroup> partial: array<f32, {{ decodeLanes * decodeRows }}>;
24
 
25
  @compute @workgroup_size(LANES, ROWS, 1)
build/webgpu/moe-output-grouped.wgsl.jinja CHANGED
@@ -67,16 +67,12 @@ const KVEC: u32 = {{ groupTileKVec }}u;
67
  {% if b2Load %}
68
  b2_tile[idx] = b2_vec;
69
  {% endif %}
70
- }
71
- {%- endmacro %}
72
 
73
  {% macro group_tile_loop(aLoad, bLoad, b2Load, kExtent, nExtent, regM, regN, bLoadVec4="", b2LoadVec4="", aLoadVec4="") %}
74
  {% for r in range(regM) %}
75
  {% for c in range(regN) %}
76
  var acc_{{ r }}_{{ c }} = 0.0;
77
- {% if b2Load %}
78
- var acc2_{{ r }}_{{ c }} = 0.0;
79
- {% endif %}
80
  {% endfor %}
81
  {% endfor %}
82
 
@@ -105,25 +101,17 @@ const KVEC: u32 = {{ groupTileKVec }}u;
105
  {% endfor %}
106
  {% for c in range(regN) %}
107
  let b{{ c }} = b_tile[(lid.x * {{ regN }}u + {{ c }}u) * KVEC + kv];
108
- {% if b2Load %}
109
- let s{{ c }} = b2_tile[(lid.x * {{ regN }}u + {{ c }}u) * KVEC + kv];
110
- {% endif %}
111
  {% endfor %}
112
  {% for r in range(regM) %}
113
  {% for c in range(regN) %}
114
  acc_{{ r }}_{{ c }} = acc_{{ r }}_{{ c }} + dot(a{{ r }}, b{{ c }});
115
- {% if b2Load %}
116
- acc2_{{ r }}_{{ c }} = acc2_{{ r }}_{{ c }} + dot(a{{ r }}, s{{ c }});
117
- {% endif %}
118
  {% endfor %}
119
  {% endfor %}
120
  }
121
  // Orders this step's tile reads before the next step overwrites them.
122
  workgroupBarrier();
123
  k_base = k_base + KTILE;
124
- }
125
- {%- endmacro %}
126
-
127
  var<workgroup> row_slot: array<u32, {{ groupTileM }}>;
128
  var<workgroup> a_tile: array<vec4<f32>, {{ groupTileM * groupTileKVec }}>;
129
  var<workgroup> b_tile: array<vec4<f32>, {{ groupTileN * groupTileKVec }}>;
 
67
  {% if b2Load %}
68
  b2_tile[idx] = b2_vec;
69
  {% endif %}
70
+ }{% endmacro %}
 
71
 
72
  {% macro group_tile_loop(aLoad, bLoad, b2Load, kExtent, nExtent, regM, regN, bLoadVec4="", b2LoadVec4="", aLoadVec4="") %}
73
  {% for r in range(regM) %}
74
  {% for c in range(regN) %}
75
  var acc_{{ r }}_{{ c }} = 0.0;
 
 
 
76
  {% endfor %}
77
  {% endfor %}
78
 
 
101
  {% endfor %}
102
  {% for c in range(regN) %}
103
  let b{{ c }} = b_tile[(lid.x * {{ regN }}u + {{ c }}u) * KVEC + kv];
 
 
 
104
  {% endfor %}
105
  {% for r in range(regM) %}
106
  {% for c in range(regN) %}
107
  acc_{{ r }}_{{ c }} = acc_{{ r }}_{{ c }} + dot(a{{ r }}, b{{ c }});
 
 
 
108
  {% endfor %}
109
  {% endfor %}
110
  }
111
  // Orders this step's tile reads before the next step overwrites them.
112
  workgroupBarrier();
113
  k_base = k_base + KTILE;
114
+ }{% endmacro %}
 
 
115
  var<workgroup> row_slot: array<u32, {{ groupTileM }}>;
116
  var<workgroup> a_tile: array<vec4<f32>, {{ groupTileM * groupTileKVec }}>;
117
  var<workgroup> b_tile: array<vec4<f32>, {{ groupTileN * groupTileKVec }}>;
build/webgpu/moe-output-stage.wgsl.jinja CHANGED
@@ -1,3 +1,8 @@
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // FC2 projection and the routed sum. One thread owns one output column of one token and walks
@@ -9,9 +14,7 @@ const WG: u32 = {{ workgroupSize }}u;
9
 
10
  @compute @workgroup_size(WG, 1, 1)
11
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
- // 2D-folded flat index: gid.y carries the high bits past the per-axis dispatch fold width.
13
- // Reduces to gid.x when the dispatch does not fold.
14
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
15
  let total = params.tokenCount * HIDDEN;
16
  if (index >= total) {
17
  return;
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
6
  {{ env.wgsl.resourceDeclarations }}
7
 
8
  // FC2 projection and the routed sum. One thread owns one output column of one token and walks
 
14
 
15
  @compute @workgroup_size(WG, 1, 1)
16
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
17
+ {{ flat_index_2d("WG", "index", "") }}
 
 
18
  let total = params.tokenCount * HIDDEN;
19
  if (index >= total) {
20
  return;
build/webgpu/moe-route-stage.wgsl.jinja CHANGED
@@ -1,21 +1,24 @@
 
 
 
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // One thread per token applies softmax to the router logits, selects TOP_K
4
  // experts, and writes their optionally renormalized probabilities. Selection is
5
  // O(TOP_K * EXPERTS) with no scratch. Equal probabilities choose the higher
6
  // expert index.
7
- const TOKENS: u32 = {{ tokens }}u;
8
  const EXPERTS: u32 = {{ experts }}u;
9
  const TOP_K: u32 = {{ topK }}u;
10
  const WG: u32 = {{ workgroupSize }}u;
11
 
12
  @compute @workgroup_size(WG, 1, 1)
13
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
14
- // gid.y carries the high bits past the per-dimension dispatch limit.
15
- let token = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
16
- if (token >= TOKENS) {
17
- return;
18
- }
19
 
20
  let router_base = token * EXPERTS;
21
  var max_logit = router_probs[router_base];
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
6
+ if ({{ name }} >= {{ bound }}) {
7
+ return;
8
+ }{% endmacro %}
9
  {{ env.wgsl.resourceDeclarations }}
10
 
11
  // One thread per token applies softmax to the router logits, selects TOP_K
12
  // experts, and writes their optionally renormalized probabilities. Selection is
13
  // O(TOP_K * EXPERTS) with no scratch. Equal probabilities choose the higher
14
  // expert index.
 
15
  const EXPERTS: u32 = {{ experts }}u;
16
  const TOP_K: u32 = {{ topK }}u;
17
  const WG: u32 = {{ workgroupSize }}u;
18
 
19
  @compute @workgroup_size(WG, 1, 1)
20
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
21
+ {{ flat_index_2d("WG", "token", "params.tokenCount") }}
 
 
 
 
22
 
23
  let router_base = token * EXPERTS;
24
  var max_logit = router_probs[router_base];
build/webgpu/test.json CHANGED
@@ -1918,9 +1918,7 @@
1918
  },
1919
  {
1920
  "name": "grouped_prefill_identity_no_fc3",
1921
- "provenance": {
1922
- "notes": "The smallest default-tile prefill that meets the grouped schedule's routed-slot threshold, and exercises identity activation without FC3 after expert grouping."
1923
- },
1924
  "attrs": { "k": 2, "activation_type": "identity", "normalize_routing_weights": 1 },
1925
  "inputs": {
1926
  "inputT": {
 
1918
  },
1919
  {
1920
  "name": "grouped_prefill_identity_no_fc3",
1921
+ "provenance": { "notes": "Compact prefill with grouped experts, identity activation and no FC3." },
 
 
1922
  "attrs": { "k": 2, "activation_type": "identity", "normalize_routing_weights": 1 },
1923
  "inputs": {
1924
  "inputT": {