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

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
-
kernel
webgpu
wgsl
apache-2.0
WebGPU

Requires WebGPU support. See the compatibility table.