Premium problem96. RMSNorm From Scratch

Hard Locked

Implement RMSNorm over the last dimension:

rms = sqrt(mean(x**2) + eps)
y = x / rms * weight

Note what is missing compared with LayerNorm: RMSNorm does not subtract the mean and has no shift term. Dropping the re-centring is the whole saving, and modern large models use it for exactly that reason.

Input

x =
tensor([[ 1.,  2.,  3.],
        [-1.,  0.,  1.]])
weight = tensor([1., 1., 1.])
eps = 1e-06

Output

tensor([[ 0.4629,  0.9258,  1.3887],
        [-1.2247,  0.0000,  1.2247]])

Premium problem

This one's part of Premium. Unlock the full PyTorch track plus every other premium problem on the site.

Implement solve(...)