AzeezIsh commited on
Commit
dff7b2f
·
verified ·
1 Parent(s): 3f6ddf9

Uploaded using `kernel-builder`.

Browse files
build/torch214-cxx11-rocm72-x86_64-linux/__init__.py ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-License-Identifier: AGPL-3.0-only
2
+ """Atlas Gated DeltaNet kernels for AMD Strix Halo (gfx1151, RDNA3.5).
3
+
4
+ The kernels are the ones the Atlas engine ran for its MLPerf Inference v6.1
5
+ Strix Halo submission (Qwen3.6-27B, Atlas @ eabfa8f). The first four op
6
+ schemas match the GB10/SM121 CUDA build of ``Atlas-Inference/gdn``; ``layers``
7
+ runs the serve chain the submission dispatched:
8
+
9
+ prefill: causal_conv1d_update_prefill -> l2_norm (q, k) -> gdn_prefill
10
+ decode: causal_conv1d_update_l2norm_f32 -> gdn_decode_f32
11
+ """
12
+
13
+ from typing import Optional
14
+
15
+ import torch
16
+
17
+ from ._ops import ops
18
+
19
+ from . import layers
20
+
21
+ __all__ = [
22
+ "gdn_decode",
23
+ "gdn_prefill",
24
+ "gdn_prefill_fla",
25
+ "causal_conv1d_fwd",
26
+ "causal_conv1d_update",
27
+ "gdn_decode_f32",
28
+ "causal_conv1d_update_prefill",
29
+ "causal_conv1d_update_l2norm_f32",
30
+ "l2_norm",
31
+ "layers",
32
+ ]
33
+
34
+
35
+ def gdn_decode(
36
+ h_state: torch.Tensor,
37
+ query: torch.Tensor,
38
+ key: torch.Tensor,
39
+ value: torch.Tensor,
40
+ gate: torch.Tensor,
41
+ beta: torch.Tensor,
42
+ output: torch.Tensor,
43
+ ) -> None:
44
+ """Single-token GDN decode (in-place update of ``h_state`` and ``output``).
45
+
46
+ h_state : (B, num_v_heads, 128, 128) float32, in-place updated
47
+ query : (B, num_k_heads, 128) float32
48
+ key : (B, num_k_heads, 128) float32
49
+ value : (B, num_v_heads, 128) float32
50
+ gate : (B, num_v_heads) float32 (exp(g_t) decay)
51
+ beta : (B, num_v_heads) float32 (sigmoid(b_t))
52
+ output : (B, num_v_heads, 128) bfloat16, in-place written
53
+ """
54
+ ops.gdn_decode(h_state, query, key, value, gate, beta, output)
55
+
56
+
57
+ def gdn_prefill(
58
+ h_state: torch.Tensor,
59
+ query: torch.Tensor,
60
+ key: torch.Tensor,
61
+ value: torch.Tensor,
62
+ gate: torch.Tensor,
63
+ beta: torch.Tensor,
64
+ output: torch.Tensor,
65
+ ) -> None:
66
+ """Multi-token GDN prefill.
67
+
68
+ h_state : (B, num_v_heads, 128, 128) float32, in-place updated
69
+ query : (B, seq_len, num_k_heads, 128) bfloat16
70
+ key : (B, seq_len, num_k_heads, 128) bfloat16
71
+ value : (B, seq_len, num_v_heads, 128) bfloat16
72
+ gate : (B, seq_len, num_v_heads) float32
73
+ beta : (B, seq_len, num_v_heads) float32
74
+ output : (B, seq_len, num_v_heads, 128) bfloat16
75
+ """
76
+ ops.gdn_prefill(h_state, query, key, value, gate, beta, output)
77
+
78
+
79
+ def gdn_prefill_fla(
80
+ h_state: torch.Tensor,
81
+ query: torch.Tensor,
82
+ key: torch.Tensor,
83
+ value: torch.Tensor,
84
+ gate: torch.Tensor,
85
+ beta: torch.Tensor,
86
+ output: torch.Tensor,
87
+ ) -> None:
88
+ """Multi-token GDN prefill, FLA-chunked (64-token chunks, three kernels).
89
+
90
+ The prefill Atlas main serves on gfx1151. Same arguments and layouts as
91
+ ``gdn_prefill``; allocates its chunk scratch per call.
92
+ """
93
+ ops.gdn_prefill_fla(h_state, query, key, value, gate, beta, output)
94
+
95
+
96
+ def causal_conv1d_fwd(
97
+ x: torch.Tensor,
98
+ weight: torch.Tensor,
99
+ bias: Optional[torch.Tensor],
100
+ out: torch.Tensor,
101
+ ) -> None:
102
+ """Depthwise causal Conv1d + SiLU.
103
+
104
+ x : (B, D, L) bfloat16
105
+ weight : (D, d_conv) bfloat16, d_conv <= 8
106
+ bias : (D,) float32 or None
107
+ out : (B, D, L) bfloat16
108
+ """
109
+ ops.causal_conv1d_fwd(x, weight, bias, out)
110
+
111
+
112
+ def causal_conv1d_update(
113
+ conv_state: torch.Tensor,
114
+ x: torch.Tensor,
115
+ weight: torch.Tensor,
116
+ bias: Optional[torch.Tensor],
117
+ out: torch.Tensor,
118
+ ) -> None:
119
+ """Single-step causal Conv1d + SiLU (decode).
120
+
121
+ conv_state : (B, D, d_conv) float32, in-place updated (rolled left, last slot = x)
122
+ x : (B, D) bfloat16
123
+ weight : (D, d_conv) bfloat16
124
+ bias : (D,) float32 or None
125
+ out : (B, D) bfloat16
126
+ """
127
+ ops.causal_conv1d_update(conv_state, x, weight, bias, out)
128
+
129
+
130
+ def gdn_decode_f32(
131
+ h_state: torch.Tensor,
132
+ query: torch.Tensor,
133
+ key: torch.Tensor,
134
+ value: torch.Tensor,
135
+ gate: torch.Tensor,
136
+ beta: torch.Tensor,
137
+ output: torch.Tensor,
138
+ ) -> None:
139
+ """``gdn_decode`` with a float32 ``output`` (B, num_v_heads, 128): the
140
+ decode kernel Atlas serves (``gated_delta_rule_decode_f32``)."""
141
+ ops.gdn_decode_f32(h_state, query, key, value, gate, beta, output)
142
+
143
+
144
+ def causal_conv1d_update_prefill(
145
+ conv_state: torch.Tensor,
146
+ x: torch.Tensor,
147
+ weight: torch.Tensor,
148
+ bias: Optional[torch.Tensor],
149
+ out: torch.Tensor,
150
+ ) -> None:
151
+ """Multi-token causal Conv1d + SiLU scanned through a float32 window.
152
+
153
+ conv_state : (B, D, 4) float32, in-place updated (the last 4 inputs)
154
+ x : (B, S, D) bfloat16, token-major
155
+ weight : (D, 4) bfloat16
156
+ bias : (D,) float32 or None
157
+ out : (B, S, D) bfloat16
158
+ """
159
+ ops.causal_conv1d_update_prefill(conv_state, x, weight, bias, out)
160
+
161
+
162
+ def causal_conv1d_update_l2norm_f32(
163
+ conv_state: torch.Tensor,
164
+ x: torch.Tensor,
165
+ weight: torch.Tensor,
166
+ bias: Optional[torch.Tensor],
167
+ out: torch.Tensor,
168
+ qk_channels: int,
169
+ head_dim: int,
170
+ eps: float,
171
+ ) -> None:
172
+ """Single-step causal Conv1d + SiLU, then per-head L2 norm of the first
173
+ ``qk_channels`` channels (q and k); float32 ``out`` (B, D).
174
+
175
+ conv_state : (B, D, d_conv) float32, in-place updated
176
+ x : (B, D) bfloat16
177
+ head_dim : must be 128; ``qk_channels`` a multiple of 256
178
+ """
179
+ ops.causal_conv1d_update_l2norm_f32(conv_state, x, weight, bias, out, qk_channels, head_dim, eps)
180
+
181
+
182
+ def l2_norm(data: torch.Tensor, num_heads: int, head_dim: int, eps: float) -> None:
183
+ """In-place L2 norm of heads ``[0, num_heads)`` of every row of ``data``
184
+ (N, stride) bfloat16: ``x * rsqrt(sum(x^2) + eps)``."""
185
+ ops.l2_norm(data, num_heads, head_dim, eps)
build/torch214-cxx11-rocm72-x86_64-linux/_gdn_rocm_8d6df22.abi3.so ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:774be8a3f9c680fdab5f936a4494d9851b30798f6707b84b047899286fd04f66
3
+ size 530576
build/torch214-cxx11-rocm72-x86_64-linux/_ops.py ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from . import _gdn_rocm_8d6df22
3
+ ops = torch.ops._gdn_rocm_8d6df22
4
+
5
+ def add_op_namespace_prefix(op_name: str):
6
+ """
7
+ Prefix op by namespace.
8
+ """
9
+ return f"_gdn_rocm_8d6df22::{op_name}"
build/torch214-cxx11-rocm72-x86_64-linux/layers.py ADDED
@@ -0,0 +1,248 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-License-Identifier: AGPL-3.0-only
2
+ """Pure, stateless ``kernels`` layers that map the Atlas Gated DeltaNet
3
+ kernels onto Hugging Face ``transformers``. This build targets AMD Strix Halo
4
+ (gfx1151); the layer contract is identical to the GB10/SM121 CUDA build.
5
+
6
+ When a model is ``kernelize()``-d, ``kernels`` binds one of these classes'
7
+ ``forward`` onto the host module instance (``MethodType(layer.forward,
8
+ module)``), so ``self`` here is the *adopting* GatedDeltaNet module: it exposes
9
+ the host's submodules, parameters (``A_log``, ``dt_bias``), config scalars
10
+ (``num_v_heads``, ``key_dim`` ...), gated ``norm``, and ``out_proj``. We replace
11
+ only the two compute cores (causal conv1d + gated-delta-rule); everything else
12
+ is reused verbatim from the host.
13
+
14
+ Two host architectures share the exact same GDN core and differ only in their
15
+ input-projection layout, so the shared core lives in the module-level
16
+ ``_gdn_run`` and each ``forward`` only does its own projection preamble:
17
+
18
+ * ``GatedDeltaNet`` -> ``Qwen3NextGatedDeltaNet`` (Qwen3-Next-80B):
19
+ fused ``in_proj_qkvz`` + ``in_proj_ba``, split via the host's
20
+ ``fix_query_key_value_ordering``.
21
+ * ``Qwen3_5GatedDeltaNet`` -> ``Qwen3_5GatedDeltaNet`` (Qwen3.6-27B dense) and
22
+ ``Qwen3_5MoeGatedDeltaNet`` (Qwen3.6-35B-A3B): already-split
23
+ ``in_proj_qkv`` / ``in_proj_z`` / ``in_proj_b`` / ``in_proj_a``,
24
+ no ordering fixup.
25
+
26
+ ``kernels`` forbids extra class members and a custom ``__init__`` on a layer
27
+ (``_validate_layer``), which is why all helpers are module-level functions, not
28
+ methods. ``_validate_layer`` also requires the layer ``forward`` signature to
29
+ match the host's argument count exactly, so ``forward`` takes the same
30
+ ``**kwargs`` (``Unpack[TransformersKwargs]``) the host GDN layers carry in
31
+ transformers >= 5.10; the kernel path ignores those kwargs.
32
+
33
+ On Strix Halo (gfx1151), as on the DGX Spark, the upstream ``fla`` /
34
+ ``causal_conv1d`` fast paths have no build, so ``transformers`` silently falls
35
+ back to a slow pure-torch implementation. These kernels fill exactly that gap.
36
+
37
+ The conv / norm / recurrence chain is the one the Atlas Strix Halo serve
38
+ dispatches for a single sequence:
39
+ prefill: causal_conv1d_update_prefill -> l2_norm (q, k) -> gdn_prefill_fla
40
+ (FLA chunked, Atlas main's gfx1151 default; ATLAS_GDN_FLA_GFX=0
41
+ selects the split4 recurrence the MLPerf v6.1 submission ran)
42
+ decode: causal_conv1d_update_l2norm_f32 -> gdn_decode_f32 (fp32 into the norm)
43
+
44
+ Convention (pinned against transformers' torch reference):
45
+ * q,k are L2-normalized before the recurrence (the kernel applies 1/sqrt(d))
46
+ * gate = exp(g), g = -A_log.exp() * softplus(a + dt_bias) (per token)
47
+ * beta = sigmoid(b)
48
+ * the causal conv1d kernels apply SiLU internally
49
+ * recurrent state h: fp32 [B, num_v_heads, head_k_dim, head_v_dim]
50
+ """
51
+
52
+ import os
53
+ import torch
54
+ import torch.nn.functional as F
55
+ from torch import nn
56
+
57
+ from ._ops import ops
58
+
59
+
60
+ def _l2norm(x, eps: float = 1e-6):
61
+ # Matches transformers' fla-aligned l2norm (eps inside the rsqrt).
62
+ return x * torch.rsqrt((x * x).sum(dim=-1, keepdim=True) + eps)
63
+
64
+
65
+ def _apply_mask_to_padding_states(hidden_states, attention_mask):
66
+ if attention_mask is not None and attention_mask.shape[1] > 1 and attention_mask.shape[0] > 1:
67
+ dtype = hidden_states.dtype
68
+ hidden_states = (hidden_states * attention_mask[:, :, None]).to(dtype)
69
+ return hidden_states
70
+
71
+
72
+ def _state(states):
73
+ # transformers >= 5.17: dict keyed by state_idx; earlier: the tensor itself.
74
+ return states[0] if isinstance(states, dict) else states
75
+
76
+
77
+ def _gdn_run(host, hidden_states, cache_params, mixed_qkv, z, b, a):
78
+ """Shared GDN compute core (conv1d + gated delta rule + gated RMSNorm).
79
+
80
+ Architecture-agnostic: the caller supplies the already-projected tensors.
81
+
82
+ Parameters
83
+ ----------
84
+ host : the adopting GatedDeltaNet module (source of submodules/dims)
85
+ mixed_qkv : [B, conv_dim, S] bf16 (concatenated Q|K|V, pre-conv)
86
+ z : [B, S, num_v_heads, head_v_dim] gate for the output RMSNorm
87
+ b, a : [B, S, num_v_heads] raw beta / decay projections
88
+ """
89
+ batch_size, seq_len, _ = hidden_states.shape
90
+ layer = cache_params.layers[host.layer_idx] if cache_params is not None else None
91
+ has_prev = cache_params is not None and cache_params.has_previous_state(host.layer_idx)
92
+ # transformers >= 5.17 keys per-layer states by state_idx and can keep the
93
+ # full conv history for speculative rollback (record_past).
94
+ keyed = layer is not None and isinstance(layer.conv_states, dict)
95
+ record_past = bool(getattr(layer, "record_past", False))
96
+ use_precomputed_states = has_prev and seq_len == 1
97
+
98
+ conv_w = host.conv1d.weight.squeeze(1).contiguous() # [conv_dim, K], bf16
99
+ K = host.conv_kernel_size
100
+ qk_channels = 2 * host.key_dim
101
+
102
+ # --- causal conv1d (+SiLU), then L2 norm of q/k: the Atlas serve chain ---
103
+ if use_precomputed_states and not record_past:
104
+ # decode: roll the fp32 window, conv + SiLU + q/k L2 norm in one kernel,
105
+ # fp32 out (causal_conv1d_update_l2norm_f32).
106
+ conv_state = _state(layer.conv_states)
107
+ cs = conv_state.to(torch.float32).contiguous()
108
+ x_step = mixed_qkv[:, :, 0].contiguous() # [B, conv_dim]
109
+ conv_out = torch.empty(x_step.shape, device=x_step.device, dtype=torch.float32)
110
+ ops.causal_conv1d_update_l2norm_f32(
111
+ cs, x_step, conv_w, None, conv_out, qk_channels, host.head_k_dim, 1e-6
112
+ )
113
+ conv_state.copy_(cs.to(conv_state.dtype)) # persist rolled window
114
+ mixed_qkv = conv_out.unsqueeze(1) # [B, 1, conv_dim] fp32, q/k normalized
115
+ else:
116
+ # prefill (or a record_past step): the fp32 window is seeded from the
117
+ # cached left context, then causal_conv1d_update_prefill scans the new
118
+ # tokens with it in registers.
119
+ if keyed:
120
+ # Cached left context + new tokens (just the new tokens, zero-padded
121
+ # to K, on a fresh prefill); the cache itself is updated here.
122
+ full = cache_params.update_conv_state(mixed_qkv, host.layer_idx, conv_kernel_size=K)
123
+ left = full[..., : full.shape[-1] - seq_len]
124
+ else:
125
+ if cache_params is not None:
126
+ padded = F.pad(mixed_qkv, (K - mixed_qkv.shape[-1], 0))
127
+ cache_params.update_conv_state(padded, host.layer_idx)
128
+ left = mixed_qkv[..., :0]
129
+ left = left[..., -K:]
130
+ cs = F.pad(left, (K - left.shape[-1], 0)).to(torch.float32).contiguous() # [B, conv_dim, K]
131
+ x = mixed_qkv.transpose(1, 2).contiguous() # [B, S, conv_dim], token-major as Atlas lays it
132
+ conv_out = torch.empty_like(x)
133
+ ops.causal_conv1d_update_prefill(cs, x, conv_w, None, conv_out)
134
+ # q,k L2 norm in place (l2_norm_bf16), heads [0, 2*num_k_heads) of each token row
135
+ ops.l2_norm(conv_out.view(-1, conv_out.shape[-1]), 2 * host.num_k_heads, host.head_k_dim, 1e-6)
136
+ mixed_qkv = conv_out # [B, S, conv_dim] bf16, q/k normalized
137
+
138
+ query, key, value = torch.split(
139
+ mixed_qkv, [host.key_dim, host.key_dim, host.value_dim], dim=-1
140
+ )
141
+ # q/k stay at num_k_heads: the kernels map v-head -> k-head themselves.
142
+ query = query.reshape(batch_size, -1, host.num_k_heads, host.head_k_dim)
143
+ key = key.reshape(batch_size, -1, host.num_k_heads, host.head_k_dim)
144
+ value = value.reshape(batch_size, -1, host.num_v_heads, host.head_v_dim)
145
+
146
+ beta = b.sigmoid()
147
+ g = -host.A_log.float().exp() * F.softplus(a.float() + host.dt_bias)
148
+
149
+ if not use_precomputed_states:
150
+ # --- gated delta rule prefill (gated_delta_rule_prefill_split4) ---
151
+ qn, kn, vv = (t.contiguous() for t in (query, key, value))
152
+ gate = g.exp().float().contiguous() # [B, S, VH]
153
+ betaf = beta.float().contiguous()
154
+ if has_prev and keyed:
155
+ # continuation chunk: start from the cached state (h is in/out)
156
+ h = _state(layer.recurrent_states).to(torch.float32).clone().contiguous()
157
+ else:
158
+ h = torch.zeros(
159
+ batch_size, host.num_v_heads, host.head_k_dim, host.head_v_dim,
160
+ device=hidden_states.device, dtype=torch.float32,
161
+ )
162
+ core_attn_out = torch.empty(
163
+ batch_size, seq_len, host.num_v_heads, host.head_v_dim,
164
+ device=hidden_states.device, dtype=torch.bfloat16,
165
+ )
166
+ prefill = ops.gdn_prefill if os.environ.get("ATLAS_GDN_FLA_GFX") == "0" else ops.gdn_prefill_fla
167
+ prefill(h, qn, kn, vv, gate, betaf, core_attn_out)
168
+ if cache_params is not None:
169
+ cache_params.update_recurrent_state(h, host.layer_idx)
170
+ else:
171
+ # --- single-token recurrence, fp32 end to end (gated_delta_rule_decode_f32) ---
172
+ qn = query[:, 0].to(torch.float32).contiguous() # [B, QH, KD]
173
+ kn = key[:, 0].to(torch.float32).contiguous()
174
+ vv = value[:, 0].to(torch.float32).contiguous() # [B, VH, VD]
175
+ gate = g[:, 0].exp().float().contiguous() # [B, VH]
176
+ betaf = beta[:, 0].float().contiguous()
177
+ h = _state(layer.recurrent_states).to(torch.float32).contiguous()
178
+ out_t = torch.empty(
179
+ batch_size, host.num_v_heads, host.head_v_dim,
180
+ device=hidden_states.device, dtype=torch.float32,
181
+ )
182
+ ops.gdn_decode_f32(h, qn, kn, vv, gate, betaf, out_t)
183
+ cache_params.update_recurrent_state(h, host.layer_idx)
184
+ # fp32 into the gated norm, as Atlas feeds gated_rms_norm_f32
185
+ core_attn_out = out_t.unsqueeze(1) # [B, 1, VH, VD]
186
+
187
+ # --- gated RMSNorm + output projection (reused from host) ---
188
+ z_shape_og = z.shape
189
+ core_attn_out = core_attn_out.reshape(-1, core_attn_out.shape[-1])
190
+ z = z.reshape(-1, z.shape[-1])
191
+ core_attn_out = host.norm(core_attn_out, z)
192
+ core_attn_out = core_attn_out.reshape(z_shape_og)
193
+ core_attn_out = core_attn_out.reshape(core_attn_out.shape[0], core_attn_out.shape[1], -1)
194
+ return host.out_proj(core_attn_out.to(hidden_states.dtype))
195
+
196
+
197
+ class GatedDeltaNet(nn.Module):
198
+ """Drop-in for ``Qwen3NextGatedDeltaNet.forward`` (Qwen3-Next-80B).
199
+
200
+ Fused QKVZ / BA projections, split via the host's
201
+ ``fix_query_key_value_ordering``.
202
+ """
203
+
204
+ # Pure recurrent/conv kernels: no autograd, not torch.compile-traceable.
205
+ has_backward: bool = False
206
+ can_torch_compile: bool = False
207
+
208
+ def forward(self, hidden_states, cache_params=None, attention_mask=None, **kwargs):
209
+ hidden_states = _apply_mask_to_padding_states(hidden_states, attention_mask)
210
+
211
+ projected_states_qkvz = self.in_proj_qkvz(hidden_states)
212
+ projected_states_ba = self.in_proj_ba(hidden_states)
213
+ query, key, value, z, b, a = self.fix_query_key_value_ordering(
214
+ projected_states_qkvz, projected_states_ba
215
+ )
216
+ query, key, value = (x.reshape(x.shape[0], x.shape[1], -1) for x in (query, key, value))
217
+ mixed_qkv = torch.cat((query, key, value), dim=-1).transpose(1, 2) # [B, conv_dim, S]
218
+
219
+ return _gdn_run(self, hidden_states, cache_params, mixed_qkv, z, b, a)
220
+
221
+
222
+ class Qwen3_5GatedDeltaNet(nn.Module):
223
+ """Drop-in for the Qwen3.5/3.6 GDN layer.
224
+
225
+ Targets both ``Qwen3_5GatedDeltaNet`` (Qwen3.6-27B dense) and
226
+ ``Qwen3_5MoeGatedDeltaNet`` (Qwen3.6-35B-A3B); their GDN cores are identical.
227
+ Already-split projections: ``in_proj_qkv`` (Q|K|V), ``in_proj_z`` (gate),
228
+ ``in_proj_b`` (beta), ``in_proj_a`` (decay). No ordering fixup.
229
+ """
230
+
231
+ has_backward: bool = False
232
+ can_torch_compile: bool = False
233
+
234
+ def forward(self, hidden_states, cache_params=None, attention_mask=None, **kwargs):
235
+ hidden_states = _apply_mask_to_padding_states(hidden_states, attention_mask)
236
+ batch_size, seq_len, _ = hidden_states.shape
237
+
238
+ mixed_qkv = self.in_proj_qkv(hidden_states).transpose(1, 2) # [B, conv_dim, S]
239
+ z = self.in_proj_z(hidden_states).reshape(batch_size, seq_len, -1, self.head_v_dim)
240
+ b = self.in_proj_b(hidden_states) # [B, S, num_v_heads]
241
+ a = self.in_proj_a(hidden_states) # [B, S, num_v_heads]
242
+
243
+ return _gdn_run(self, hidden_states, cache_params, mixed_qkv, z, b, a)
244
+
245
+
246
+ # Qwen3.6-35B-A3B (MoE) shares the dense layer's GDN core verbatim. Expose the
247
+ # host class name so a single LayerRepository entry resolves for either model.
248
+ Qwen3_5MoeGatedDeltaNet = Qwen3_5GatedDeltaNet
build/torch214-cxx11-rocm72-x86_64-linux/metadata.json ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "gdn",
3
+ "id": "_gdn_rocm_8d6df22",
4
+ "version": 1,
5
+ "kernels-minver": "0.14.0",
6
+ "license": "AGPL-3.0-only",
7
+ "python-depends": [],
8
+ "kernel-depends": [],
9
+ "backend": {
10
+ "type": "rocm",
11
+ "archs": [
12
+ "gfx1151"
13
+ ]
14
+ },
15
+ "digest": {
16
+ "algorithm": "sha256",
17
+ "files": {
18
+ "__init__.py": "LiP481kYPHgVyWn/hX7QlwghehLrO6jOKeUkE8w0HNU=",
19
+ "_gdn_rocm_8d6df22.abi3.so": "d0voo/nGgP2rX5NqRJTZhRsweY9nB7hLBHiZKG/QT2Y=",
20
+ "_ops.py": "8RvxjaM9m2Nc3SoC6Nmkjrj2ERAl+o8Oi1pIpPO1bh0=",
21
+ "layers.py": "y5VWO3DlzuCO1r7wuJoYjbMkZImfPQXqON5MouNWnWQ="
22
+ }
23
+ },
24
+ "provenance": {
25
+ "kernel-builder": {
26
+ "version": "0.18.0-dev0",
27
+ "commit": "f13568d6c568d18b2da481edf66b83714300e441",
28
+ "dirty": false
29
+ },
30
+ "kernel": {
31
+ "commit": "8d6df2295686c12b0fa37fdd71eba3f65273fae0",
32
+ "dirty": false
33
+ }
34
+ }
35
+ }