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 的倍数,以适应内核的要求。