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).

Mojo function

composite_rms_norm_fused_residual_add_shape

def composite_rms_norm_fused_residual_add_shape[dtype: DType](input: T, residual_input: T, gamma1: T, gamma2: T, epsilon1: Float32, epsilon2: Float32, weight_offset1: Scalar[dtype], weight_offset2: Scalar[dtype]) -> IndexList[T.rank]

Computes the output shape for the mo.composite.rms_norm_fused_residual_add graph op.

Args:

  • ​input (T): Primary input tensor whose shape the output mirrors.
  • ​residual_input (T): Residual tensor added to the normalized input.
  • ​gamma1 (T): Per-column scale weights applied to the first RMS normalization.
  • ​gamma2 (T): Per-column scale weights applied to the second RMS normalization.
  • ​epsilon1 (Float32): Small constant added inside the first RMS normalization square root for numerical stability.
  • ​epsilon2 (Float32): Small constant added inside the second RMS normalization square root for numerical stability.
  • ​weight_offset1 (Scalar[dtype]): Scalar offset added to gamma1 before scaling the first normalization.
  • ​weight_offset2 (Scalar[dtype]): Scalar offset added to gamma2 before scaling the second normalization.

Returns:

IndexList[T.rank]: The output shape, which matches the input shape.

Was this page helpful?