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 @ Bwith a per-batch M mask.Computes a batched GEMM where each batch
qonly produces the firstmasked_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)iftranspose_a), float16 or bfloat16, contiguous.b (torch.Tensor) – Batched input, shape
(Q, K, N)(or(Q, N, K)iftranspose_b), same dtype asa, contiguous.masked_m (torch.Tensor) – Per-batch row count, shape
(Q,), int32.transpose_a (bool) – Whether
a/bare stored transposed (see shapes above).transpose_b (bool) – Whether
a/bare 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
Cof shape(Q, M, N).- Return type:
torch.Tensor