Skip to main content

//clika-runtime/io.clika.runtime/Ops/gatedRmsNorm

gatedRmsNorm

[common]
fun gatedRmsNorm(input: Tensor, gate: Tensor, normalizedShape: LongArray, weight: Tensor? = null, eps: Double? = null): Tensor

gatedRmsNorm(input: Tensor, gate: Tensor, normalizedShape: LongArray, weight: Tensor? = null, eps: Double? = null): the gated_rms_norm operator. Fused group RMS norm + post-norm silu(z) gate: out = rms_norm(x) · silu(z), the norm taken over the trailing normalized_shape group extent ({vd} for a per-head norm over [.., HV, vd]). weight is the optional [group] gamma applied after the normalization; an absent eps takes the runtime's rms-norm default.