IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/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. /max/get-started.md).

Mojo function

batch_matmul_shape

def batch_matmul_shape[rank: Int, a_type: DType, b_type: DType](a: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=a.static_spec], b: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=b.static_spec]) -> IndexList[rank]

Computes the output shape for the mo.batch_matmul graph op.

Parameters:

  • ​rank (Int): Tensor rank of the batched matmul operands and output.
  • ​a_type (DType): Element type of the a input tensor.
  • ​b_type (DType): Element type of the b input tensor.

Args:

Returns:

IndexList[rank]: The output shape of the batched matmul.