flashinfer.gemm.group_gemm_fp8_nt_groupwise¶
- flashinfer.gemm.group_gemm_fp8_nt_groupwise(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, m_indptr: Tensor, scale_granularity_mnk: Tuple[int, int, int] = (1, 128, 128), scale_major_mode: Literal['MN', 'K'] = 'MN', mma_sm: int = 1, out: Tensor | None = None, out_dtype: dtype | None = None) Tensor¶
使用组间缩放执行 FP8 数据类型的组 GEMM。目前仅支持 NVIDIA Blackwell 架构。
- 参数:
a (torch.Tensor) – 行主输入张量形状
(cum_m, k),数据类型为torch.float8_e4m3fn或torch.float8_e5m2。cum_m是分段长度的累积和。b (torch.Tensor) – 列主输入张量形状
(batch_size, n, k),数据类型为torch.float8_e4m3fn或torch.float8_e5m2。a_scale (torch.Tensor) – a 的列主缩放张量,形状为
(cum_m, k // block_size)如果 scale_major_mode 为K或形状为(k // block_size, cum_m)如果 scale_major_mode 为MN,数据类型为torch.float32。b_scale (torch.Tensor) – b 的行主缩放张量,形状为
(batch_size, n // block_size, k // block_size)如果 scale_major_mode 为K形状为(batch_size, k // block_size, n // block_size)如果 scale_major_mode 为MN,数据类型为torch.float32。m_indptr (torch.Tensor) – 分段长度的 indptr,形状为
(batch_size + 1,),数据类型为torch.int32。m_indptr中的每个元素必须是 4 的倍数。scale_granularity_mnk (Tuple[int, int, int]) – 缩放张量的粒度,(m_granularity, n_granularity, k_granularity)。
scale_major_mode (Literal["MN", "K"]) – 缩放张量的布局模式,MN 表示 MN 主缩放,形状为
(k // block_size, *),K 表示 K 主缩放,形状为(*, k // block_size)mma_sm (int) – 用于 MMA 操作的 SM 数量,必须为 1 或 2。当每组的行数 (M) 很大(>= 256)时,2 更快。
out (Optional[torch.Tensor]) – 输出张量,形状为
(cum_m, n)。如果未指定,我们将显式创建一个输出张量。out_dtype (Optional[torch.dtype]) – 输出张量的数据类型,必须为
torch.bfloat16或torch.float16。
- 返回值:
out – 输出张量,形状为
(cum_m, n)。- 返回值类型:
torch.Tensor
注意
在调用此函数之前,应将
m_indptr中的每个值填充到 4 的倍数,以适应内核的要求。