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 with mma_scaled (Blackwell-only) and supporting NVFP4/MXFP4/MXFP8/mixed plus optional global scales. Matrix A is a ragged stack (total_m, K_A) partitioned by m_indptr; B is 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 by m_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 a and b.

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

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

  • transpose_b (bool) – Whether a / b are stored transposed (only NT is supported).

  • static_persistent (bool) – Kept for API compatibility with the ocean signature.

  • swizzled_layout_a (bool) – Whether a_scale uses the swizzled layout (only True is 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 for b.

  • backend (str) – Implementation backend. Currently only "cutile" is supported.

Returns:

The ragged block-scaled output tensor C of shape (total_m, N).

Return type:

torch.Tensor