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 to a and b (block-scaled FP8 inputs dequantized to out_dtype).

Parameters:
  • a (torch.Tensor) – Block-scaled (FP8) batched inputs; a’s M dimension is segmented by m_indptr.

  • b (torch.Tensor) – Block-scaled (FP8) batched inputs; a’s M dimension is segmented by m_indptr.

  • a_scale (torch.Tensor) – Per-block scale tensors for a and b.

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

  • 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 / b are stored transposed.

  • transpose_b (bool) – Whether a / b are stored transposed.

  • out_dtype (Optional[torch.dtype]) – Output dtype (e.g. torch.bfloat16).

  • segment_alignment (int) – Row alignment the caller guarantees for every m_indptr segment offset (default 128). Bounds the largest internal tile (BLOCK_M must 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