The gated delta rule from
llama.cpp as a single kernel
(
gated_delta_net) — the linear-attention recurrence behind Qwen3-Next and Qwen3.5, which eager torch
spells out in ~200 ops per layer.
l2_norm is exposed alongside it.
1import torch
2from kernels import get_kernel
3
4gdn = get_kernel("marcsun13/ggml-gated-delta-net", version=1)
5
6n_seqs, n_tokens, n_heads, head_dim = 1, 1, 32, 128
7q = torch.randn(n_seqs, n_tokens, n_heads, head_dim, device="mps")
8k = torch.randn_like(q) # one head per value head, not expanded
9v = torch.randn_like(q)
10g = torch.randn(n_seqs, n_tokens, n_heads, device="mps") # log-domain gate
11beta = torch.rand(n_seqs, n_tokens, n_heads, device="mps")
12state = torch.zeros(n_seqs, n_heads, head_dim, head_dim, device="mps")
13
14out, state = gdn.gated_delta_net(q, k, v, g, beta, state) # (1, 1, 32, 128), (1, 32, 128, 128)