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 satisfy0 <= index < G;-1padding 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.bfloat16is supported. Ignored when out is supplied.validate_indices (bool, optional) – Validate expert-index values, sortedness and boundary alignment. Defaults to
Falseto avoid a GPU-to-CPU synchronization per call. Enable for new routing data outside CUDA graph capture. Ignored withskip_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 ... )