Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
18 lines
504 B
Python
18 lines
504 B
Python
"""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
|