flashinfer.gemm.group_gemm_mxfp4_nt_groupwise

flashinfer.gemm.group_gemm_mxfp4_nt_groupwise(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, m_indptr: Tensor, mma_sm: int = 1, tile_m: int = 128, tile_n: int = 128, tile_k: int = 128, swap_ab: bool = True, out: Tensor | None = None, out_dtype: dtype | None = None) Tensor

使用组间缩放执行 MXFP4 数据类型的组 GEMM。目前仅支持 NVIDIA Blackwell 架构。

参数:
  • a (torch.Tensor) – 行主输入张量,形状 (cum_m, k),数据类型为 torch.float8_e4m3fntorch.float8_e5m2cum_m 是分段长度的累积和。

  • b (torch.Tensor) – 列主输入张量,形状 (batch_size, n, k // 2),数据类型为 torch.uint8

  • a_scale (torch.Tensor) – a 的列主缩放张量,形状 (cum_m_padded, k // 32),数据类型为 torch.uint8

  • b_scale (torch.Tensor) – b 的行主缩放张量,形状 (batch_size, n_padded, k // 32),数据类型为 torch.uint8

  • m_indptr (torch.Tensor) – 分段长度的 indptr,形状 (batch_size + 1,),数据类型为 torch.int32m_indptr 中的每个元素必须是 4 的倍数。

  • mma_sm (int) – 用于 MMA 操作的 SM 数量,必须为 1 或 2。当每组的行数 (M) 很大(>= 256)时,2 更快。

  • tile_m (int) – M 维度的切片大小,必须为 128。

  • tile_n (int) – N 维度的切片大小,必须为 64、128、192 或 256。

  • tile_k (int) – K 维度的切片大小,必须为 128 或 256。

  • swap_ab (bool) – 是否交换 A 和 B 张量。

  • out (Optional[torch.Tensor]) – 输出张量,形状 (cum_m, n)。 如果未指定,我们将显式创建一个输出张量。

  • out_dtype (Optional[torch.dtype]) – 输出张量的数据类型,必须为 torch.bfloat16torch.float16

返回值:

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

返回值类型:

torch.Tensor

注意

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