For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo function
reduce_rms_norm_shape
def reduce_rms_norm_shape[dtype: DType](input: T, gamma: T, epsilon: Float32, weight_offset: Scalar[dtype]) -> IndexList[T.rank]
Computes the output shape for the mo.reduce.rms_norm graph op.
Args:
- βinput (
T): Input tensor normalized across the last dimension. - βgamma (
T): Per-column scale weights applied after normalization. - βepsilon (
Float32): Small constant added inside the RMS normalization square root for numerical stability. - βweight_offset (
Scalar[dtype]): Scalar offset added togammabefore scaling.
Returns:
IndexList[T.rank]: The output shape, which matches the input shape.