| 20 | |
| 21 | |
| 22 | class RMSNorm(nn.Module): |
| 23 | def __init__(self, dim: int, eps: float = 1e-6): |
| 24 | super().__init__() |
| 25 | self.eps = eps |
| 26 | self.weight = nn.Parameter(torch.ones(dim)) |
| 27 | |
| 28 | def forward(self, x: torch.Tensor) -> torch.Tensor: |
| 29 | norm = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) |
| 30 | return x * norm * self.weight |
| 31 | |
| 32 | |
| 33 | def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0) -> torch.Tensor: |