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)
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=0andmultiply_before_cast=False. The normalized input is cast to the output dtype before multiplication by the weight. - Gemma-style:
weight_offset=1andmultiply_before_cast=True. The weight is treated as1 + weightand 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
weightbefore scaling. Use1.0for Gemma-style normalization and0.0otherwise. Defaults to0.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 toFalse.
-
Returns:
-
A
TensorValuewith the same shape and dtype asinput. -
Raises:
-
ValueError – If
weightdoes not match the last dimension ofinput. -
Return type: