flashinfer.gemm.ragged_bmm

flashinfer.gemm.ragged_bmm(a: Tensor, b: 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, backend: Literal['cutile'] = 'cutile') → Tensor

Ragged batched matrix multiplication with non-even M segments.

Matrix A is flattened along its M dimension with m_indptr defining the per-group segment boundaries (grouped/variable-length GEMM).

Parameters:
  • a (torch.Tensor) – Flattened batched input; the M dimension is segmented by m_indptr.

  • b (torch.Tensor) – Batched weights, shape (Q, K, N) (or transposed per transpose_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; defaults to a.dtype.

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

Returns:

The ragged output tensor.

Return type:

torch.Tensor