IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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).

Python function

rms_norm

rms_norm()​

max.experimental.nn.norm.rms_norm(input, weight, epsilon, weight_offset=0.0, multiply_before_cast=False)

source

Computes root mean square normalization over the last dimension of input.

The output is input / rms(input) * (weight + weight_offset) where rms(x) = sqrt(mean(x ** 2) + epsilon). Reduction runs over the last axis of input and is broadcast back across the leading axes. See Root Mean Square Layer Normalization for the original formulation.

Two variants are supported through weight_offset and multiply_before_cast:

  • Llama-style (default): weight_offset=0 and multiply_before_cast=False. The normalized input is cast to the output dtype before multiplication by the weight.
  • Gemma-style: weight_offset=1 and multiply_before_cast=True. The weight is treated as 1 + weight and multiplication runs in the reduction dtype before casting back.
import numpy as np

from max.dtype import DType
from max.engine import InferenceSession
from max.graph import DeviceRef, Graph, ops

device = DeviceRef.CPU()
with Graph("rms_norm_example") as graph:
    x = ops.constant([[3.0, 4.0]], DType.float32, device=device)
    weight = ops.constant([1.0, 1.0], DType.float32, device=device)
    y_llama = ops.rms_norm(x, weight, epsilon=1e-6)
    y_gemma = ops.rms_norm(
        x, weight, epsilon=1e-6,
        weight_offset=1.0, multiply_before_cast=True,
    )
    graph.output(y_llama, y_gemma)

model = InferenceSession().load(graph)
llama, gemma = model.execute()
assert np.allclose(llama.to_numpy(), [[0.848528, 1.131371]], atol=1e-4)

Parameters:

  • input (Tensor) – The tensor to normalize. Reduction runs over the last axis.
  • weight (Tensor) – The scale applied after normalization. A 1-D tensor whose shape matches the last dimension of input.
  • epsilon (float) – A small positive constant added to the mean of squares for numerical stability.
  • weight_offset (float) – A value added to weight before scaling. Use 1.0 for Gemma-style normalization and 0.0 otherwise. Defaults to 0.0.
  • multiply_before_cast (bool) – Whether to multiply by the (offset) weight before casting the normalized input back to the output dtype. Llama-style sets this to False. Defaults to False.

Returns:

A TensorValue with the same shape and dtype as input.

Raises:

ValueError – If weight does not match the last dimension of input.

Return type:

Tensor