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

gated_group_rmsnorm

gated_group_rmsnorm()​

max.nn.state_space.gated_group_rmsnorm(y, gate, norm_weight, eps, group_size)

source

Fuses the gated group-RMSNorm (HF Zamba2RMSNormGated with norm_before_gate=False) into a single dispatch.

Collapses cast(y->f32) -> silu(gate)*y -> group rms_norm -> *norm_weight -> cast into one op. y and gate are [N, intermediate]; norm_weight is fp32 [intermediate]. Each contiguous group_size slice of the intermediate axis is normalized independently. Returns the model dtype (y.dtype), so the downstream out_proj cast is a no-op.

Parameters:

  • y (TensorValue) – The [N, intermediate] SSD scan output (model dtype).
  • gate (TensorValue) – The [N, intermediate] gate projection (any float dtype; may be a strided split view of the fused in-proj).
  • norm_weight (TensorValue) – The [intermediate] fp32 RMSNorm weight.
  • eps (float) – The epsilon inside rsqrt(mean_sq + eps).
  • group_size (int) – The width of each normalized group (intermediate // n_groups).

Returns:

The [N, intermediate] result in y.dtype.

Return type:

TensorValue