from __future__ import annotations import torch from torch import nn from torch.nn.attention import SDPBackend class RMSNorm(nn.Module): """Zero-centered RMSNorm as used by Qwen3-Next.""" def __init__(self, hidden_size: int, eps: float) -> None: super().__init__() self.weight = nn.Parameter(torch.zeros(hidden_size)) self.eps = eps def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: input_dtype = hidden_states.dtype normalized = hidden_states.float() normalized = normalized * torch.rsqrt( normalized.square().mean(dim=-1, keepdim=True) + self.eps ) normalized = normalized * (1.0 + self.weight.float()) return normalized.to(dtype=input_dtype) def attention_sdpa_backends(device: torch.device) -> list[SDPBackend]: if device.type != "cuda": return [SDPBackend.MATH] return [ SDPBackend.CUDNN_ATTENTION, SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, ]