tolivert's picture
deploy: financial_rag streamlit app
4e316d6
Raw
History Blame Contribute Delete
2.03 kB
"""
RMSNorm (Root Mean Square Layer Normalization).
Unlike LayerNorm, RMSNorm skips mean-centering and only rescales by the
root-mean-square of activations: norm(x) = x / sqrt(mean(x^2) + eps).
This is cheaper (no mean subtraction, no variance) and performs
comparably for transformer pre-norm architectures.
Computation is done in float32 regardless of input dtype to avoid
numerical instability in half-precision, then cast back.
Two weight conventions exist across model families:
- Standard (LLaMA, Qwen3): output = norm(x) * weight
- Unit-offset (Gemma3): output = norm(x) * (1 + weight)
Gemma initialises weights to zero so the initial scale is 1.0;
the `add_unit_offset` flag selects this mode.
"""
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(self,
emb_dim: int,
eps: float = 1e-6,
dtype=None,
add_unit_offset: bool = False,
with_scale: bool = True):
super().__init__()
self.eps = eps
self.add_unit_offset = add_unit_offset
self.with_scale = with_scale
if self.add_unit_offset:
# Gemma3-style: weights stored as offset from 1.0
self.scale = nn.Parameter(torch.zeros(emb_dim, dtype=dtype))
else:
self.scale = nn.Parameter(torch.ones(emb_dim, dtype=dtype))
def _norm(self, x):
norm_x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return norm_x
def forward(self, x):
input_dtype = x.dtype
# Force float32 for variance calculation (Crucial for stability)
x = x.to(torch.float32)
out = self._norm(x)
if self.add_unit_offset:
# Gemma3-style: weights stored as offset from 1.0
# Gemma4 removed the unitoffset
return (out * (1.0 + self.scale.float())).to(input_dtype)
if not self.with_scale:
return out.to(input_dtype)
return out.to(input_dtype) * self.scale