flashinfer.gemm.masked_scaled_bmm¶
- flashinfer.gemm.masked_scaled_bmm(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, masked_m: Tensor, block_scale_type: str, max_m_device: Tensor | None = None, transpose_a: bool = False, transpose_b: bool = True, out_dtype: dtype | None = None, backend: Literal['cutile'] = 'cutile') Tensor¶
Masked block-scaled batched matrix multiplication (FP8/FP4 block-scaled).
Like
masked_bmm()but with per-block scale tensors applied toaandb(block-scaled FP8/FP4 inputs). Each batchqonly produces the firstmasked_m[q]rows of the output. Blackwell-only (usesmma_scaled).- Parameters:
a (torch.Tensor) – Block-scaled batched input (FP8/FP4), shape
(Q, M, K_A).b (torch.Tensor) – Block-scaled batched weights (FP8/FP4), shape
(Q, N, K_B).a_scale (torch.Tensor) – MX-swizzled per-block scale tensors for
aandb.b_scale (torch.Tensor) – MX-swizzled per-block scale tensors for
aandb.masked_m (torch.Tensor) – Per-batch row count, shape
(Q,), int32.block_scale_type (str) – One of
"nvfp4","mxfp4","mxfp8","mixed".max_m_device (Optional[torch.Tensor]) – Optional device-side scalar with
max(masked_m); computed on device when omitted (avoids a host sync).transpose_a (bool) – Whether
a/bare stored transposed (only NT is supported).transpose_b (bool) – Whether
a/bare stored transposed (only NT is supported).out_dtype (Optional[torch.dtype]) – Output dtype; defaults to
torch.bfloat16.backend (str) – Implementation backend. Currently only
"cutile"is supported.
- Returns:
The output tensor
Cof shape(Q, M, N).- Return type:
torch.Tensor