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 to a and b (block-scaled FP8/FP4 inputs). Each batch q only produces the first masked_m[q] rows of the output. Blackwell-only (uses mma_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 a and b.

  • b_scale (torch.Tensor) – MX-swizzled per-block scale tensors for a and b.

  • 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 / b are stored transposed (only NT is supported).

  • transpose_b (bool) – Whether a / b are 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 C of shape (Q, M, N).

Return type:

torch.Tensor