flashinfer.gemm.ragged_scaled_bmm¶
- flashinfer.gemm.ragged_scaled_bmm(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, m_indptr: Tensor, max_m: int, block_scale_type: str, transpose_a: bool = False, transpose_b: bool = True, static_persistent: bool = True, swizzled_layout_a: bool = True, a_global_scale: Tensor | None = None, b_global_scale: Tensor | None = None, backend: Literal['cutile'] = 'cutile') Tensor¶
Ragged block-scaled batched matrix multiplication (FP8/FP4 block-scaled).
Like
ragged_block_scaled_bmm()but block-scaled withmma_scaled(Blackwell-only) and supporting NVFP4/MXFP4/MXFP8/mixed plus optional global scales. MatrixAis a ragged stack(total_m, K_A)partitioned bym_indptr;Bis batched(Q, N, K_B). Output is(total_m, N)(float32).- Parameters:
a (torch.Tensor) – Ragged block-scaled input (FP8/FP4), shape
(total_m, K_A), segmented bym_indptr.b (torch.Tensor) – Batched block-scaled 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.m_indptr (torch.Tensor) – Segment offsets, shape
(Q + 1,); each entry must be a multiple of 128.max_m (int) – Upper bound on any single segment length (host int, used for grid sizing).
block_scale_type (str) – One of
"nvfp4","mxfp4","mxfp8","mixed".transpose_a (bool) – Whether
a/bare stored transposed (only NT is supported).transpose_b (bool) – Whether
a/bare stored transposed (only NT is supported).static_persistent (bool) – Kept for API compatibility with the ocean signature.
swizzled_layout_a (bool) – Whether
a_scaleuses the swizzled layout (onlyTrueis supported).a_global_scale (Optional[torch.Tensor]) – Optional scalar global scale for
a.b_global_scale (Optional[torch.Tensor]) – Optional per-batch
(Q,)global scale forb.backend (str) – Implementation backend. Currently only
"cutile"is supported.
- Returns:
The ragged block-scaled output tensor
Cof shape(total_m, N).- Return type:
torch.Tensor