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)
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).
- y (TensorValue) – The
-
Returns:
-
The
[N, intermediate]result iny.dtype. -
Return type: