"""RMSNorm used by attention, FFN, and the final LM stem.""" from __future__ import annotations import torch from torch import nn class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, x: torch.Tensor): dtype = x.dtype x = x.float() return (x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps)).to(dtype) * self.weight