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_e4m3fntorch.float8_e5m2cum_m 是分段长度的累积和。

  • b (torch.Tensor) – 列主输入张量形状 (batch_size, n, k),数据类型为 torch.float8_e4m3fntorch.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.int32m_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.bfloat16torch.float16

返回值:

out – 输出张量,形状为 (cum_m, n)

返回值类型:

torch.Tensor

注意

在调用此函数之前,应将 m_indptr 中的每个值填充到 4 的倍数,以适应内核的要求。