Mike0021/zonos2
3
1from typing import Tuple2 3import torch4 5from .base import BaseOP6 7 8class RMSNorm(BaseOP):9 def __init__(self, size: int, eps: float) -> None:10 from flashinfer import rmsnorm11 12 self.eps = eps13 self.weight = torch.empty(size)14 self.rmsnorm = rmsnorm15 16 def forward(self, x: torch.Tensor) -> torch.Tensor:17 return self.rmsnorm(x, self.weight, self.eps)18 19 def forward_inplace(self, x: torch.Tensor) -> None:20 self.rmsnorm(x, self.weight, self.eps, out=x)21 22 23class RMSNormFused(BaseOP):24 def __init__(self, size: int, eps: float, elementwise_affine: bool = True) -> None:25 from flashinfer import fused_add_rmsnorm, rmsnorm26 27 self.eps = eps28 self.elementwise_affine = elementwise_affine29 self._size = size30 31 if elementwise_affine:32 self.weight = torch.empty(size)33 # When elementwise_affine=False, we use a ones buffer created lazily34 # to ensure correct device/dtype35 36 self.rmsnorm = rmsnorm37 self.fused_add_rmsnorm = fused_add_rmsnorm38 self._ones_buffer: torch.Tensor | None = None39 40 def _get_weight(self, x: torch.Tensor) -> torch.Tensor:41 if self.elementwise_affine:42 return self.weight43 # Use cached ones buffer, recreate if needed for device/dtype match44 if self._ones_buffer is None or self._ones_buffer.device != x.device:45 self._ones_buffer = torch.ones(self._size, device=x.device, dtype=x.dtype)46 return self._ones_buffer47 48 def forward(49 self, x: torch.Tensor, residual: torch.Tensor | None = None50 ) -> Tuple[torch.Tensor, torch.Tensor]:51 weight = self._get_weight(x)52 if residual is None:53 return self.rmsnorm(x, weight, self.eps), x54 self.fused_add_rmsnorm(x, residual, weight, self.eps)55 return x, residual56 