from __future__ import annotations """Small fused CUDA/Triton operators for GDN24 MAX TURBO. The runtime deliberately keeps a pure-Torch fallback so checkpoints remain portable. Triton is used only when it is already available through the CUDA PyTorch stack; it is not a checkpoint dependency. """ import torch from torch import nn import torch.nn.functional as F try: # Triton ships with CUDA PyTorch builds used by Colab. import triton import triton.language as tl _HAS_TRITON = True except Exception: # pragma: no cover - CPU/source validation path triton = None tl = None _HAS_TRITON = False if _HAS_TRITON: @triton.jit def _rmsnorm_kernel(x_ptr, w_ptr, y_ptr, n_cols: tl.constexpr, eps: tl.constexpr, BLOCK: tl.constexpr): row = tl.program_id(0) offs = tl.arange(0, BLOCK) mask = offs < n_cols x = tl.load(x_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32) w = tl.load(w_ptr + offs, mask=mask, other=0.0).to(tl.float32) var = tl.sum(x * x, axis=0) / n_cols rstd = tl.rsqrt(var + eps) y = x * rstd * (1.0 + w) tl.store(y_ptr + row * n_cols + offs, y, mask=mask) @triton.jit def _add_rmsnorm_kernel( x_ptr, update_ptr, w_ptr, sum_ptr, norm_ptr, n_cols: tl.constexpr, eps: tl.constexpr, BLOCK: tl.constexpr, ): row = tl.program_id(0) offs = tl.arange(0, BLOCK) mask = offs < n_cols x = tl.load(x_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32) u = tl.load(update_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32) w = tl.load(w_ptr + offs, mask=mask, other=0.0).to(tl.float32) s = x + u var = tl.sum(s * s, axis=0) / n_cols rstd = tl.rsqrt(var + eps) n = s * rstd * (1.0 + w) tl.store(sum_ptr + row * n_cols + offs, s, mask=mask) tl.store(norm_ptr + row * n_cols + offs, n, mask=mask) @triton.jit def _add_final_rmsnorm_kernel( x_ptr, update_ptr, w_ptr, norm_ptr, n_cols: tl.constexpr, eps: tl.constexpr, BLOCK: tl.constexpr, ): row = tl.program_id(0) offs = tl.arange(0, BLOCK) mask = offs < n_cols x = tl.load(x_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32) u = tl.load(update_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32) w = tl.load(w_ptr + offs, mask=mask, other=0.0).to(tl.float32) s = x + u var = tl.sum(s * s, axis=0) / n_cols rstd = tl.rsqrt(var + eps) n = s * rstd * (1.0 + w) tl.store(norm_ptr + row * n_cols + offs, n, mask=mask) @triton.jit def _silu_mul_kernel(a_ptr, b_ptr, out_ptr, n_elements, BLOCK: tl.constexpr): offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) mask = offs < n_elements a = tl.load(a_ptr + offs, mask=mask, other=0.0).to(tl.float32) b = tl.load(b_ptr + offs, mask=mask, other=0.0).to(tl.float32) # SiLU(a) = a * sigmoid(a) sig = 1.0 / (1.0 + tl.exp(-a)) out = a * sig * b tl.store(out_ptr + offs, out, mask=mask) def _can_triton(x: torch.Tensor) -> bool: return bool(_HAS_TRITON and x.is_cuda and x.is_contiguous()) def qwen_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: """Exact Qwen3.5 offset RMSNorm: norm(x) * (1 + weight).""" d = x.shape[-1] if _can_triton(x) and weight.is_cuda and weight.is_contiguous(): y = torch.empty_like(x) x2 = x.view(-1, d) y2 = y.view(-1, d) block = triton.next_power_of_2(d) _rmsnorm_kernel[(x2.shape[0],)](x2, weight, y2, n_cols=d, eps=float(eps), BLOCK=block) return y xf = x.float() y = xf * torch.rsqrt(xf.square().mean(dim=-1, keepdim=True) + eps) y = y * (1.0 + weight.float()) return y.to(dtype=x.dtype) def add_rmsnorm( x: torch.Tensor, update: torch.Tensor, weight: torch.Tensor, eps: float, ) -> tuple[torch.Tensor, torch.Tensor]: """Fuse residual addition with Qwen3.5 offset RMSNorm. Returns `(x + update, rmsnorm(x + update))`. """ d = x.shape[-1] if _can_triton(x) and update.is_contiguous() and weight.is_cuda and weight.is_contiguous(): summed = torch.empty_like(x) normed = torch.empty_like(x) x2 = x.view(-1, d) u2 = update.view(-1, d) s2 = summed.view(-1, d) n2 = normed.view(-1, d) block = triton.next_power_of_2(d) _add_rmsnorm_kernel[(x2.shape[0],)]( x2, u2, weight, s2, n2, n_cols=d, eps=float(eps), BLOCK=block ) return summed, normed summed = x + update xf = summed.float() normed = xf * torch.rsqrt(xf.square().mean(dim=-1, keepdim=True) + eps) normed = normed * (1.0 + weight.float()) return summed, normed.to(dtype=x.dtype) def add_final_rmsnorm( x: torch.Tensor, update: torch.Tensor, weight: torch.Tensor, eps: float, ) -> torch.Tensor: """Fuse final residual addition and RMSNorm when the unnormalized sum is not needed.""" d = x.shape[-1] if _can_triton(x) and update.is_contiguous() and weight.is_cuda and weight.is_contiguous(): normed = torch.empty_like(x) x2 = x.view(-1, d) u2 = update.view(-1, d) n2 = normed.view(-1, d) block = triton.next_power_of_2(d) _add_final_rmsnorm_kernel[(x2.shape[0],)]( x2, u2, weight, n2, n_cols=d, eps=float(eps), BLOCK=block ) return normed summed = x + update xf = summed.float() normed = xf * torch.rsqrt(xf.square().mean(dim=-1, keepdim=True) + eps) normed = normed * (1.0 + weight.float()) return normed.to(dtype=x.dtype) def silu_mul(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: """Fused SwiGLU pointwise core: SiLU(a) * b.""" if _can_triton(a) and b.is_contiguous() and a.shape == b.shape: out = torch.empty_like(a) n = a.numel() block = 256 _silu_mul_kernel[(triton.cdiv(n, block),)](a, b, out, n_elements=n, BLOCK=block) return out return F.silu(a) * b class SuperRMSNorm(nn.Module): """Checkpoint-compatible replacement for Qwen3_5RMSNorm.""" def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = float(eps) # Qwen3.5 uses zero-centered weights and multiplies by (1 + weight). self.weight = nn.Parameter(torch.zeros(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: return qwen_rmsnorm(x, self.weight, self.eps) def extra_repr(self) -> str: return f"{tuple(self.weight.shape)}, eps={self.eps}" class SuperSwiGLUMLP(nn.Module): """Checkpoint-compatible Qwen3.5 dense MLP with a fused SwiGLU pointwise kernel.""" def __init__(self, config): super().__init__() self.hidden_size = config.hidden_size self.intermediate_size = config.intermediate_size self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) self.hidden_act = str(config.hidden_act) def forward(self, x: torch.Tensor) -> torch.Tensor: gate = self.gate_proj(x) up = self.up_proj(x) if self.hidden_act == "silu": hidden = silu_mul(gate, up) else: from transformers.activations import ACT2FN hidden = ACT2FN[self.hidden_act](gate) * up return self.down_proj(hidden)