flashinfer.gemm.gemm_fp8_nt_groupwise¶
- flashinfer.gemm.gemm_fp8_nt_groupwise(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, scale_major_mode: Literal['MN', 'K'] = None, mma_sm: int = 1, scale_granularity_mnk: Tuple[int, int, int] = (1, 128, 128), out: Tensor | None = None, out_dtype: dtype | None = None, backend: Literal['cutlass', 'trtllm'] = 'cutlass') Tensor¶
使用组态缩放执行 FP8 数据类型的矩阵乘法。
此函数实现一个 GEMM 操作,允许对不同维度上的缩放粒度进行细粒度控制。目前仅支持 NVIDIA Blackwell 架构。
- 参数:
a (torch.Tensor) – 行主输入张量形状 (m, k),fp8 e4m3 或 fp8 e5m2。
b (torch.Tensor) – 列主输入张量形状 (n, k),fp8 e4m3 或 fp8 e5m2。
a_scale (torch.Tensor) –
- 如果后端是
cutlass a 的列主缩放张量,形状
(m, k // block_size)如果 scale_major_mode 是K或形状(k // block_size, m)如果 scale_major_mode 是MN- 如果后端是
trtllm scale_major_mode 应该为 None,缩放张量应该是 (m, k // block_size),在第一个维度上连续
- 如果后端是
b_scale (torch.Tensor) –
- 如果后端是
cutlass b 的行主缩放张量,形状
(n // block_size, k // block_size)如果 scale_major_k 是K或形状(k // block_size, n // block_size)如果 scale_major_mode 是MN- 如果后端是
trtllm scale_major_mode 应该为 None,缩放张量应该是 (k // block_size, n // block_size),在第一个维度上连续
- 如果后端是
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]) – 输出张量,形状 (m, n)。如果未指定,我们将显式创建一个输出张量。
out_dtype (Optional[torch.dtype]) – 如果未指定 out,我们将使用此 dtype 创建一个输出张量。默认为
torch.bfloat16。backend (Literal["cutlass", "trtllm"]) – 用于该操作的后端。默认为
"cutlass"。
- 返回值:
out – 输出张量,形状 (m, n)。
- 返回值类型:
torch.Tensor
注意
在调用此函数之前,
m应该填充为 4 的倍数,以适应内核的要求。