Qwen3.5-9B-SpeedX9-GDN32 / fused_ops.py
summerMC's picture
Upload Qwen3.5 UNI MAX 9B checkpoint with benchmark results
efc38b5 verified
Raw History Blame Contribute Delete
7.76 kB
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)