sync 6fdf6301e2bb
Browse files- README.md +3 -3
- build/webgpu/{moe-mix-stage.wgsl.jinja → expert-slot-mix.wgsl.jinja} +6 -3
- build/webgpu/manifest.json +0 -0
- build/webgpu/metadata.json +38 -38
- build/webgpu/moe-ffn-gemv.wgsl.jinja +10 -6
- build/webgpu/moe-ffn-grouped.wgsl.jinja +11 -8
- build/webgpu/moe-ffn-stage.wgsl.jinja +15 -6
- build/webgpu/moe-grouped-sgmat.wgsl.jinja +16 -7
- build/webgpu/moe-output-gemv.wgsl.jinja +1 -3
- build/webgpu/moe-output-grouped.wgsl.jinja +2 -14
- build/webgpu/moe-output-stage.wgsl.jinja +6 -3
- build/webgpu/moe-route-stage.wgsl.jinja +9 -6
- build/webgpu/test.json +1 -3
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
|
| 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.
|
| 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 |
-
|
| 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": "
|
| 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 |
-
"
|
| 13 |
-
"
|
| 14 |
-
"moe-ffn-
|
| 15 |
-
"moe-ffn-
|
| 16 |
-
"moe-
|
| 17 |
-
"moe-
|
| 18 |
-
"moe-output-gemv.wgsl.jinja": "
|
| 19 |
-
"moe-output-grouped.wgsl.jinja": "
|
| 20 |
-
"moe-output-stage.wgsl.jinja": "
|
| 21 |
-
"moe-route-stage.wgsl.jinja": "
|
| 22 |
-
"test.json": "
|
| 23 |
}
|
| 24 |
},
|
| 25 |
-
"provenance": { "kernel": { "sha": "
|
| 26 |
"webgpu": {
|
| 27 |
-
"manifestSpec": "2.
|
| 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 |
-
"
|
| 54 |
-
"
|
| 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" %}
|
|
|
|
|
|
|
|
|
|
| 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 |
-
}
|
|
|
|
| 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" %}
|
|
|
|
|
|
|
|
|
|
| 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 |
-
}
|
|
|
|
| 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" %}
|
|
|
|
|
|
|
|
|
|
| 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 |
-
}
|
|
|
|
| 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 |
-
|
| 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" %}
|
|
|
|
|
|
|
|
|
|
| 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 |
-
}
|
| 22 |
-
{%
|
|
|
|
|
|
|
| 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 |
-
//
|
| 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 * {{
|
| 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 %}
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
|
| 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": {
|