Expand description
RMSNorm Triton kernels.
RMSNorm normalises each row by its root-mean-square (no mean subtraction): rms[m] = sqrt( (1/N) * Σ_n x[m,n]² + eps ) y[m,n] = x[m,n] / rms[m] * γ[n]
Grid: [M] — one CTA per row. Layout identical to LayerNorm.
Structs§
- RmsNorm
Backward - RMSNorm backward pass.
- RmsNorm
Forward - RMSNorm forward pass.
Functions§
- rms_
norm_ backward - RMSNorm backward pass.
- rms_
norm_ forward - RMSNorm forward pass.