flashinfer.gemm.ragged_block_scaled_bmm¶
- flashinfer.gemm.ragged_block_scaled_bmm(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, m_indptr: Tensor, max_m: int, max_m_device: Tensor | None = None, transpose_a: bool = False, transpose_b: bool = True, out_dtype: dtype | None = None, segment_alignment: int = 128, backend: Literal['cutile'] = 'cutile') Tensor¶
Ragged block-scaled batched matrix multiplication (FP8 block-scaled).
Like
ragged_bmm()but with per-block scale tensors applied toaandb(block-scaled FP8 inputs dequantized toout_dtype).- Parameters:
a (torch.Tensor) – Block-scaled (FP8) batched inputs;
a’s M dimension is segmented bym_indptr.b (torch.Tensor) – Block-scaled (FP8) batched inputs;
a’s M dimension is segmented bym_indptr.a_scale (torch.Tensor) – Per-block scale tensors for
aandb.b_scale (torch.Tensor) – Per-block scale tensors for
aandb.m_indptr (torch.Tensor) – Segment offsets, shape
(Q + 1,), int32.max_m (int) – Maximum segment length (host int, used for grid sizing).
max_m_device (Optional[torch.Tensor]) – Optional device-side copy of
max_m.transpose_a (bool) – Whether
a/bare stored transposed.transpose_b (bool) – Whether
a/bare stored transposed.out_dtype (Optional[torch.dtype]) – Output dtype (e.g.
torch.bfloat16).segment_alignment (int) – Row alignment the caller guarantees for every
m_indptrsegment offset (default 128). Bounds the largest internal tile (BLOCK_Mmust divide it); pass 256, with 256-aligned segments, to enable the large-M fast path. It is a caller contract — it cannot be checked at runtime without a host sync that would break CUDA-graph capture.backend (str) – Implementation backend. Currently only
"cutile"is supported.
- Returns:
The ragged block-scaled output tensor.
- Return type:
torch.Tensor