flashinfer.gemm.masked_bmm

flashinfer.gemm.masked_bmm(a: Tensor, b: Tensor, masked_m: Tensor, transpose_a: bool = False, transpose_b: bool = False, out: Tensor | None = None, backend: Literal['cutile'] = 'cutile') → Tensor

Masked batched matrix multiplication C = A @ B with a per-batch M mask.

Computes a batched GEMM where each batch q only produces the first masked_m[q] rows of the output; rows beyond the mask are left unspecified (callers typically zero them). This is the grouped/masked GEMM used by MoE-style expert routing.

Parameters:
  • a (torch.Tensor) – Batched input, shape (Q, M, K) (or (Q, K, M) if transpose_a), float16 or bfloat16, contiguous.

  • b (torch.Tensor) – Batched input, shape (Q, K, N) (or (Q, N, K) if transpose_b), same dtype as a, contiguous.

  • masked_m (torch.Tensor) – Per-batch row count, shape (Q,), int32.

  • transpose_a (bool) – Whether a / b are stored transposed (see shapes above).

  • transpose_b (bool) – Whether a / b are stored transposed (see shapes above).

  • out (Optional[torch.Tensor]) – Optional output tensor, shape (Q, M, N). Allocated if omitted.

  • backend (str) – Implementation backend. Currently only "cutile" (the cuda.tile Python backend) is supported.

Returns:

The output tensor C of shape (Q, M, N).

Return type:

torch.Tensor