For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/get-started.md).
Mojo function
composite_rms_norm_fused_residual_add_shape
def composite_rms_norm_fused_residual_add_shape[dtype: DType, rank: Int](input: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input.static_spec], residual_input: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=residual_input.static_spec], gamma1: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gamma1.static_spec], gamma2: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gamma2.static_spec], epsilon1: Float32, epsilon2: Float32, weight_offset1: Scalar[dtype], weight_offset2: Scalar[dtype]) -> IndexList[rank]
Computes the output shape for the mo.composite.rms_norm_fused_residual_add graph op.
Parameters:
- βdtype (
DType): Element type of theinput,residual_input, and weight tensors. - βrank (
Int): Number of dimensions in theinputand output tensors.
Args:
- βinput (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input.static_spec]): Primary input tensor whose shape the output mirrors. - βresidual_input (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=residual_input.static_spec]): Residual tensor added to the normalizedinput. - βgamma1 (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gamma1.static_spec]): Per-column scale weights applied to the first RMS normalization. - βgamma2 (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=gamma2.static_spec]): 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 togamma1before scaling the first normalization. - βweight_offset2 (
Scalar[dtype]): Scalar offset added togamma2before scaling the second normalization.
Returns:
IndexList[rank]: The output shape, which matches the input shape.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!