from __future__ import annotations import torch import torch.nn.functional as F from torch import Tensor, nn __all__ = [ "RMSNorm", "SwiGLU", ] class RMSNorm(nn.RMSNorm): def __init__( self, hidden_size: int, eps: float = 1e-5, *, device: torch.device | str | None = None, dtype: torch.dtype | None = None, ) -> None: if hidden_size <= 0: raise ValueError(f"hidden_size must be positive, got {hidden_size}") if eps <= 0.0: raise ValueError(f"eps must be positive, got {eps}") super().__init__( normalized_shape=hidden_size, eps=eps, elementwise_affine=True, device=device, dtype=dtype, ) self.hidden_size = hidden_size class SwiGLU(nn.Module): def __init__( self, hidden_size: int, intermediate_size: int, *, device: torch.device | str | None = None, dtype: torch.dtype | None = None, ) -> None: super().__init__() if hidden_size <= 0: raise ValueError(f"hidden_size must be positive, got {hidden_size}") if intermediate_size <= 0: raise ValueError( f"intermediate_size must be positive, got {intermediate_size}" ) self.hidden_size = hidden_size self.intermediate_size = intermediate_size self.gate_up_proj = nn.Linear( in_features=hidden_size, out_features=2 * intermediate_size, bias=False, device=device, dtype=dtype, ) self.down_proj = nn.Linear( in_features=intermediate_size, out_features=hidden_size, bias=False, device=device, dtype=dtype, ) def forward(self, hidden_states: Tensor) -> Tensor: gate, up = self.gate_up_proj(hidden_states).chunk(2, dim=-1) hidden_states = F.silu(gate) * up return self.down_proj(hidden_states) def extra_repr(self) -> str: return ( f"hidden_size={self.hidden_size}, " f"intermediate_size={self.intermediate_size}, " "bias=False" )