Skip to main content

Module rmsnorm

Module rmsnorm 

Source
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§

RmsNormBackward
RMSNorm backward pass.
RmsNormForward
RMSNorm forward pass.

Functions§

rms_norm_backward
RMSNorm backward pass.
rms_norm_forward
RMSNorm forward pass.