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_matmul_fused_partial_rms_norm_shape
def composite_matmul_fused_partial_rms_norm_shape[dtype: DType](input: T, weight: T, gamma: T, epsilon: Float32, weight_offset: Scalar[dtype]) -> IndexList[T.rank]
Computes the output shape for the mo.composite.matmul_fused_partial_rms_norm graph op.
Args:
- input (
T): Input activation tensorxof the GEMV. - weight (
T): Weight matrixWof shape(N, K)(rank 2). - gamma (
T): RMS normalization scale vector (rank 1). - epsilon (
Float32): Small constant added to the squared mean before the reciprocal square root for numerical stability. - weight_offset (
Scalar[dtype]): Reserved for API consistency with other RMS norm ops; not consumed by the shape function.
Returns:
IndexList[T.rank]: The output shape, which matches the input shape.