flashinfer.gemm.group_gemm_fp8_nt_groupwise_contiguous

flashinfer.gemm.group_gemm_fp8_nt_groupwise_contiguous(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, m_indices: Tensor, scale_granularity_mnk: Tuple[int, int, int] = (1, 128, 128), out: Tensor | None = None, out_dtype: dtype | None = None, validate_indices: bool = False) → Tensor

Compute contiguous grouped FP8 GEMM using CuTe DSL on SM100/SM103.

Each row of a is multiplied by the transposed expert matrix selected by m_indices. Input A uses per-row, 128-element K-block scales; B uses 128x128 block scales. The output uses bfloat16.

Parameters:
  • a (torch.Tensor) – Contiguous FP8 E4M3 input of shape (M, K).

  • b (torch.Tensor) – Contiguous FP8 E4M3 expert weights of shape (G, N, K).

  • a_scale (torch.Tensor) – Contiguous float32 scales of shape (M, K // 128).

  • b_scale (torch.Tensor) – Contiguous float32 scales of shape (G, N // 128, K // 128).

  • m_indices (torch.Tensor) – Contiguous int32 expert indices of shape (M,), sorted in nondecreasing order. Internal expert boundaries must align to 128 rows; the final expert may end in a partial tile. All values must satisfy 0 <= index < G; -1 padding is unsupported.

  • scale_granularity_mnk (Tuple[int, int, int], optional) – Scale granularity. Only (1, 128, 128) is supported.

  • out (Optional[torch.Tensor], optional) – Contiguous bfloat16 output of shape (M, N). Allocated if omitted.

  • out_dtype (Optional[torch.dtype], optional) – Output dtype when allocating; only torch.bfloat16 is supported. Ignored when out is supplied.

  • validate_indices (bool, optional) – Validate expert-index values, sortedness and boundary alignment. Defaults to False to avoid a GPU-to-CPU synchronization per call. Enable for new routing data outside CUDA graph capture. Ignored with skip_check=True.

Returns:

The supplied or allocated bfloat16 output of shape (M, N).

Return type:

torch.Tensor

Notes

Requires nvidia-cutlass-dsl. N and K must be positive multiples of 128, and G must be positive. M may be zero, in which case no kernel is launched. All tensors must be on the same CUDA device and at least 16-byte aligned.

Index values are unchecked unless validate_indices=True. Violating their preconditions results in undefined behavior, including incorrect results or invalid memory accesses.

Execution uses PyTorch’s current stream for a.device. Compilation is cached by device, weight shape, and M’s 128-row alignment class; warm both classes used by a workload before CUDA graph capture.

Examples

>>> from flashinfer.gemm import group_gemm_fp8_nt_groupwise_contiguous
>>> # a/b are FP8; a_scale/b_scale are float32 block scales.
>>> out = group_gemm_fp8_nt_groupwise_contiguous(
...     a, b, a_scale, b_scale, m_indices, validate_indices=True
... )