com.microsoft.GatedDeltaNet
com.microsoft · ONNX Runtime contrib operator · contrib since_version 1
Description
Token-major gated delta network with an explicit recurrent state: S_t = exp(g_t) * S_{t-1} + k_t * (beta_t * (v_t - exp(g_t) * S_{t-1}^T k_t))^T and o_t = scale * S_t^T q_t per value head; update_rule selects which terms survive. Query, key and value are (total_tokens, heads, head_size) or (batch, sequence, heads, head_size). num_heads_v is a positive multiple of num_heads_q, and value head h reads query head h * num_heads_q / num_heads_v. The V-major state and the gate fusions are float32. Bfloat16, per-key-dimension decay and compact capture are not implemented.
See the ONNX Runtime GatedDeltaNet contrib-operator spec for the reference semantics.
Inputs
| Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
|---|---|---|---|---|---|---|
queryT |
query |
T |
— | — | Query vectors, (total_tokens, num_heads_q, head_size_qk) or (batch_size, sequence_length, num_heads_q, head_size_qk). The two spellings share one memory layout; the rank-4 form lets an exporter keep a static reshape target. |
required |
keyT |
key |
T |
— | — | Key vectors with exactly the shape of query. The delta rules need L2-normalized keys, either from upstream or from qk_l2_norm=1; without them the per-chunk system is ill-conditioned and the recurrence diverges. |
required |
valueT |
value |
T |
— | — | Value vectors, (total_tokens, num_heads_v, head_size_v), sharing the leading token axes of query. |
required |
cuSeqlensT |
cu_seqlens |
TI |
1 |
— | Exclusive prefix sums of the per-request token counts, (batch_size + 1). Requires the rank-3 spelling. Offsets are clamped on device to [0, total_tokens], so a decreasing or out-of-range entry yields an empty request rather than an out-of-bounds read. Absent means uniform packing. |
optional |
decayT |
decay |
TS |
— | — | Log-space decay, (..., num_heads_v) over the same leading token axes as query, always float32. Present exactly for update_rule gated and gated_delta. The per-key-dimension spelling (..., num_heads_v, head_size_qk) is not implemented. |
optional |
betaT |
beta |
TS |
— | — | Update rate, (..., num_heads_v) over the same leading token axes as query, always float32. Present exactly for update_rule delta and gated_delta. |
optional |
initialStateT |
initial_state |
TS |
4 |
— | Incoming recurrent state, V-major (batch_size, num_heads_v, head_size_v, head_size_qk), always float32. Absent means a zero state. It is a separate allocation from final_state: WebGPU usage-scope validation tracks whole buffers, so binding one buffer as both read-only and read-write invalidates the command buffer. |
optional |
aLogT |
a_log |
TS |
1 |
— | Per-head A_log, (num_heads_v), always float32. Present exactly when gate_activation is qwen. |
optional |
dtBiasT |
dt_bias |
TS |
1 |
— | Per-head gate bias, (num_heads_v), always float32. Present exactly when gate_activation is qwen. |
optional |
Outputs
| Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
|---|---|---|---|---|---|---|
outputT |
output |
T |
same as queryT |
derived | Attention output, (total_tokens, num_heads_v, head_size_v) or (batch_size, sequence_length, num_heads_v, head_size_v). num_heads_v is the larger head count because it is a positive multiple of num_heads_q. |
required |
finalStateT |
final_state |
TS |
4 |
derived | State after the last token of each request, V-major (batch_size, num_heads_v, head_size_v, head_size_qk), always float32. A request with no tokens copies its incoming state through. |
required |
Attributes
Default values (overridable per request):
| Attribute | Default | Description |
|---|---|---|
beta_activation |
"none" |
none treats beta as the effective update rate; sigmoid applies a logistic sigmoid to it in float32. Default is none. |
chunk_size |
64 |
Accepted scheduling hint for a chunk-parallel prefill algorithm. Every value is accepted and none changes the result. Default is 64. |
gate_activation |
"none" |
none treats decay as the effective log-space decay. qwen computes -exp(a_log) * Softplus(decay + dt_bias) in float32 and requires both a_log and dt_bias. Default is none. |
qk_l2_norm |
0 |
Set to 1 to L2-normalize each query and key head vector, with epsilon 1e-12, before the recurrence. Default is 0. |
scale |
0 |
Output scaling factor applied to S_t^T q_t. The default 0.0 selects 1 / sqrt(head_size_qk). |
state_update_capacity |
0 |
Capacity for compact transition capture. The default and only supported value is 0; compact state_update capture is not implemented. |
update_rule |
"gated_delta" |
Which terms of the recurrence survive: linear drops the decay and the delta retrieval, gated keeps only the decay, delta keeps only the retrieval, and gated_delta keeps both. Default is gated_delta. |
Type constraints
| Variable | Allowed dtypes |
|---|---|
T |
float32, float16 |
TS |
float32 |
TI |
int32 |
Files
metadata.json— kernel metadata (id, digests, per-variant templates, provenance)manifest.json— the op contract (source of truth)test.json— correctness casesbench.json— benchmark casesgated-delta-net-gate.wgsl.jinjagated-delta-net.wgsl.jinja
Use with @huggingface/kernels
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.
import { getKernel } from "@huggingface/kernels";
const kernel = await getKernel("webgpu-kernels/com.microsoft.GatedDeltaNet", { version: 1 });
const { outputT, finalStateT } = await kernel({
queryT: { data: queryTData, shape: [4, 1, 12] },
keyT: { data: keyTData, shape: [4, 1, 12] },
valueT: { data: valueTData, shape: [4, 2, 9] },
decayT: { data: decayTData, shape: [4, 2] },
betaT: { data: betaTData, shape: [4, 2] },
initialStateT: { data: initialStateTData, shape: [1, 2, 9, 12] },
});
- Downloads last month
- -
Requires WebGPU support. See the compatibility table.