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
Ais flattened along its M dimension withm_indptrdefining 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 pertranspose_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/bare stored transposed.transpose_b (bool) – Whether
a/bare 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