File size: 10,656 Bytes
3d17c9b
3ab8080
3d17c9b
3ab8080
 
 
 
3d17c9b
3ab8080
 
 
 
 
 
 
 
 
 
 
 
4e80716
3ab8080
4e80716
 
 
 
 
 
 
 
3ab8080
 
 
4e80716
3ab8080
4e80716
3ab8080
 
 
 
 
 
 
 
 
 
 
 
 
 
4e80716
3ab8080
 
 
 
 
 
 
4e80716
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3ab8080
 
4e80716
3ab8080
 
eba4596
3ab8080
eba4596
3ab8080
 
 
4e80716
3ab8080
 
 
 
 
 
 
4e80716
eba4596
4e80716
 
 
3ab8080
 
4e80716
3ab8080
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
---
library_name: kernels
license: apache-2.0
tags:
- kernel
- webgpu
- wgsl
---
# com.microsoft.MoE

`com.microsoft`  ·  ONNX Runtime contrib operator  ·  contrib since_version 1

## Description

Mixture of Experts: applies softmax to `router_probs`, routes each token to the top-`k` experts, applies FC1 and `activation_type`, projects through FC2, then sums the selected outputs using their routing probabilities. SwiGLU takes its operands from a separate FC3 (`swiglu_fusion` 0) or a fused FC1 in interleaved (1) or concatenated (2) order; SiLU may also use FC3 as its multiplicative linear projection. This inference package supports float32 and dense routing (`use_sparse_mixer = 0`); float16, bfloat16, and sparse mixing are not implemented. Quantized weights use `com.microsoft.QMoE`.

See the [ONNX Runtime `MoE` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.MoE) for the reference semantics.

## Inputs

| Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- | --- |
| `inputT` | `input` | `T` | — | — | Token activations, either 2D `(num_tokens, hidden_size)` or 3D `(batch_size, sequence_length, hidden_size)`. | required |
| `routerT` | `router_probs` | `T` | `2` | — | 2D router logits of shape `(num_tokens, num_experts)`, where `num_tokens` is the product of every leading dimension of `input`. A full softmax is applied before top-k selection. | required |
| `fc1T` | `fc1_experts_weights` | `T` | `3` | — | 3D first-layer expert weights of shape `(num_experts, fusion_size * inter_size, hidden_size)`, where `fusion_size` is 2 for fused SwiGLU (`swiglu_fusion` 1 or 2) and 1 otherwise. | required |
| `fc1BiasT` | `fc1_experts_bias` | `T` | `2` | — | Optional 2D FC1 bias of shape `(num_experts, fusion_size * inter_size)`. | optional |
| `fc2T` | `fc2_experts_weights` | `T` | `3` | — | 3D second-layer expert weights of shape `(num_experts, hidden_size, inter_size)`. | required |
| `fc2BiasT` | `fc2_experts_bias` | `T` | `2` | — | Optional 2D FC2 bias of shape `(num_experts, hidden_size)`, added per expert before that expert's routing weight is applied. | optional |
| `fc3T` | `fc3_experts_weights` | `T` | `3` | — | Optional 3D third-layer expert weights of shape `(num_experts, inter_size, hidden_size)`. It supplies the separate linear operand for SwiGLU when `swiglu_fusion` is 0, or the multiplicative linear projection for SiLU gating. Other activations do not consume FC3. | optional |
| `fc3BiasT` | `fc3_experts_bias` | `T` | `2` | — | Optional 2D FC3 bias of shape `(num_experts, inter_size)`. | optional |

## Outputs

| Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- | --- |
| `outputT` | `output` | `T` | same as `inputT` | same as `inputT` | Routed expert output with the same shape as `input`. | required |

## Attributes

Attributes and default values (overridable per request):

| Attribute | Default | Description |
| --- | --- | --- |
| `activation_alpha` | `1` | Alpha parameter used by the activation; the schema default is 1. |
| `activation_beta` | `0` | Beta parameter used by the activation; the schema default is 0. |
| `activation_type` | `"relu"` | Activation applied to the FC1 projection: `relu`, `gelu`, `silu`, `swiglu`, or `identity`. The schema default is `relu`. |
| `k` | `1` | Number of experts selected per token; the schema default is 1. |
| `normalize_routing_weights` | `0` | Whether to normalize the selected routing weights; the schema default is 0. |
| `swiglu_fusion` | `0` | 0 keeps the SwiGLU operands in separate FC1/FC3 GEMMs, 1 interleaves them in one FC1 row, and 2 concatenates them. The schema default is 0. |
| `swiglu_limit` | — | Optional SwiGLU clamp limit; omission means no clamp. |
| `use_sparse_mixer` | `0` | Whether to use sparse-mixer routing. The standard default and only supported value is 0. |

## Type constraints

| Variable | Allowed dtypes |
| --- | --- |
| `T` | `float32` |

## Implementation variants

One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.

- `sgmat_grouped_routed_fc1plain_fc3none_fc2plain` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1plain_fc3none_fc2bias` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1plain_fc3plain_fc2plain` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1plain_fc3plain_fc2bias` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1plain_fc3biased_fc2plain` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1plain_fc3biased_fc2bias` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1bias_fc3none_fc2plain` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1bias_fc3none_fc2bias` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1bias_fc3plain_fc2plain` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1bias_fc3plain_fc2bias` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1bias_fc3biased_fc2plain` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.
- `sgmat_grouped_routed_fc1bias_fc3biased_fc2bias` — Expert-grouped f32 matrix projections load public weight layouts directly. Two statically interleaved accumulation chains limit rounding growth; the input tile is reused for result publication. Requires compatible subgroups, f32 fragments, workgroup limits, and complete weight tiles.

## Device requirements

Some implementation variants require `subgroup-matrix` and `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.

## Files

- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
- [`test.json`](build/webgpu/test.json) — correctness cases
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
- [`expert-group-slots.wgsl.jinja`](build/webgpu/expert-group-slots.wgsl.jinja)
- [`expert-slot-mix.wgsl.jinja`](build/webgpu/expert-slot-mix.wgsl.jinja)
- [`moe-ffn-gemv.wgsl.jinja`](build/webgpu/moe-ffn-gemv.wgsl.jinja)
- [`moe-ffn-grouped.wgsl.jinja`](build/webgpu/moe-ffn-grouped.wgsl.jinja)
- [`moe-ffn-stage.wgsl.jinja`](build/webgpu/moe-ffn-stage.wgsl.jinja)
- [`moe-grouped-sgmat.wgsl.jinja`](build/webgpu/moe-grouped-sgmat.wgsl.jinja)
- [`moe-output-gemv.wgsl.jinja`](build/webgpu/moe-output-gemv.wgsl.jinja)
- [`moe-output-grouped.wgsl.jinja`](build/webgpu/moe-output-grouped.wgsl.jinja)
- [`moe-output-stage.wgsl.jinja`](build/webgpu/moe-output-stage.wgsl.jinja)
- [`moe-route-stage.wgsl.jinja`](build/webgpu/moe-route-stage.wgsl.jinja)

## Use with `@huggingface/kernels`

```sh
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
```

Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.

The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
It follows the `v1` branch as fixes land. To pin exact artifact bytes, pass a 40-character commit `revision` instead of `version`.

Replace each `*Data` placeholder with a typed array containing the corresponding input data.

```js
import { getKernel } from "@huggingface/kernels";

const kernel = await getKernel("webgpu-kernels/com.microsoft.MoE", { version: 1 });
const { outputT } = await kernel({
  inputT: { data: inputTData, shape: [1, 1] },
  routerT: { data: routerTData, shape: [1, 2] },
  fc1T: { data: fc1TData, shape: [2, 1, 1] },
  fc2T: { data: fc2TData, shape: [2, 1, 1] },
});
```