flashinfer.gemm.group_deepgemm_fp8_nt_groupwise

flashinfer.gemm.group_deepgemm_fp8_nt_groupwise(a: Tensor, b: Tensor, a_scale: Tensor, b_scale: Tensor, m_indices: Tensor, scale_granularity_mnk: Tuple[int, int, int] = (1, 128, 128), out: Tensor | None = None, out_dtype: dtype | None = None)

使用 DeepGEMM 后端执行具有 FP8 数据类型的分组矩阵乘法。

此函数执行分组 GEMM 操作,其中张量 b 中的每个组与张量 a 中的相应行相乘。分组由 m_indices 张量确定,该张量指定每个行属于哪个组。这对于诸如专家混合 (MoE) 的场景特别有用,其中不同的令牌路由到不同的专家。

该操作可以概念化为

>>> for i in range(num_groups):
>>>    row_slice = slice(i * m_per_group, (i + 1) * m_per_group)
>>>    output[row_slice] = a[row_slice] @ b[i].T

目前仅支持 NVIDIA Blackwell (SM100) 架构。

参数:
  • a (torch.Tensor) – 输入张量 A 的形状为 (m, k),具有 FP8 数据类型 (torch.float8_e4m3fn)。此张量包含将与 b 中的不同组相乘的所有行。

  • b (torch.Tensor) – 输入张量 B 的形状为 (batch_size, n, k),具有 FP8 数据类型 (torch.float8_e4m3fn)。每个切片 b[i] 代表一个不同的组/专家,它将与 a 中的相应行相乘。

  • a_scale (torch.Tensor) – 张量 a 的缩放因子,形状为 (m, k // block_size),具有 torch.float32 dtype。这些通常由原始 float32 张量的逐令牌量化生成。

  • b_scale (torch.Tensor) – 张量 b 的缩放因子,形状为 (batch_size, n // block_size, k // block_size),具有 torch.float32 dtype。这些通常由为每个组的原始 float32 张量的逐块量化生成。

  • m_indices (torch.Tensor) – 组分配张量,形状为 (m,),具有 torch.int32 dtype。每个元素指定 a 中的相应行属于哪个组(在 b 中的索引)。例如,如果 m_indices[i] = j,则 a 中的行 i 将与 b 中的组 j 相乘。

  • scale_granularity_mnk (Tuple[int, int, int], optional) – 缩放因子的粒度,格式为 (m_granularity, n_granularity, k_granularity)。默认值为 (1, 128, 128),这意味着 a 的逐令牌缩放和 b 的 128x128 块缩放。

  • out (Optional[torch.Tensor], optional) – 预分配的输出张量,形状为 (m, n)。如果未提供,将创建一个新的张量。

  • out_dtype (Optional[torch.dtype], optional) – 输出张量的数据类型。如果提供了 out,则忽略此参数。默认值为 torch.bfloat16

返回值:

输出张量,形状为 (m, n),包含分组矩阵乘法的结果。

返回值类型:

torch.Tensor

示例

>>> import torch
>>> from flashinfer.gemm import group_deepgemm_fp8_nt_groupwise
>>> from flashinfer.utils import per_token_cast_to_fp8, per_block_cast_to_fp8
>>>
>>> # Setup: 2 groups, 128 tokens per group, 4096 hidden size, 2048 expert size
>>> m_per_group, n, k = 128, 2048, 4096
>>> group_size = 2
>>> m = m_per_group * group_size
>>>
>>> # Create float32 inputs
>>> a_f32 = torch.randn(m, k, device="cuda", dtype=torch.float32)
>>> b_f32 = torch.randn(group_size, n, k, device="cuda", dtype=torch.float32)
>>>
>>> # Quantize to FP8 with appropriate scaling
>>> a_fp8, a_scale = per_token_cast_to_fp8(a_f32)
>>> b_fp8 = torch.empty_like(b_f32, dtype=torch.float8_e4m3fn)
>>> b_scale = torch.empty((group_size, n // 128, k // 128), device="cuda", dtype=torch.float32)
>>> for i in range(group_size):
...     b_fp8[i], b_scale[i] = per_block_cast_to_fp8(b_f32[i])
>>>
>>> # Create group assignment
>>> m_indices = torch.empty(m, device="cuda", dtype=torch.int32)
>>> for i in range(group_size):
...     row_slice = slice(i * m_per_group, (i + 1) * m_per_group)
...     m_indices[row_slice] = i
>>>
>>> # Perform grouped GEMM
>>> result = group_deepgemm_fp8_nt_groupwise(
...     a_fp8, b_fp8, a_scale, b_scale, m_indices, out_dtype=torch.bfloat16
... )
>>> print(result.shape)  # torch.Size([256, 2048])

注意

  • 此函数需要 NVIDIA Blackwell (SM100) 架构

  • 应使用适当的量化函数(如 per_token_cast_to_fp8 用于 aper_block_cast_to_fp8 用于 b)生成缩放因子

  • 该函数内部使用 DeepGEMM 后端进行优化的 FP8 计算

  • 所有输入张量必须位于同一 CUDA 设备上

  • 块大小由 scale_granularity_mnk 参数确定