flashinfer.gemm.batch_deepgemm_fp8_nt_groupwise

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

使用 DeepGEMM 后端执行 FP8 数据类型的批量矩阵乘法。

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

该操作可以概念化为

>>> for i in range(num_groups):
>>>     output[i] = a[i][:masked_m[i]] @ b[i][:masked_m[i]].T

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

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

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

  • a_scale (torch.Tensor) – 张量 a 的缩放因子,形状为 (batch_size, 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 张量的逐块量化生成。

  • masked_m (torch.Tensor) – 掩码张量,形状为 (batch_size,),具有 torch.int32 dtype。每个元素指定每个组中要相乘的有效行数。例如,如果 masked_m[i] = j,则 a[i] 中的前 j 行将与 b 中的组 i 相乘。

  • expected_m (int) – 每个批次的 M 期望值的提示值(CPU 上的值),正确设置此值可以提高性能。

  • 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) – 预分配的输出张量,形状为 (batch_size, m, n)。如果未提供,将创建一个新的张量。

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

返回值:

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

返回值类型:

torch.Tensor

示例

>>> import torch
>>> from flashinfer.gemm import batch_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, n, k = 128, 2048, 4096
>>> group_size = 2
>>>
>>> # Create float32 inputs
>>> a = torch.rand((group_size, m, k), device="cuda", dtype=torch.float32)
>>> b = torch.rand((group_size, n, k), device="cuda", dtype=torch.float32)
>>> masked_m = torch.randint(0, m, (group_size,), device="cuda", dtype=torch.int32)
>>> a_fp8 = torch.empty_like(a, device="cuda", dtype=torch.float8_e4m3fn)
>>> a_scale = torch.empty((group_size, m, k // 128), device="cuda", dtype=torch.float32)
>>> b_fp8 = torch.empty_like(b, device="cuda", 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):
...    a_fp8[i], a_scale[i] = per_token_cast_to_fp8(a[i])
...    b_fp8[i], b_scale[i] = per_block_cast_to_fp8(b[i])
>>>
>>> expected_m = min(int(masked_m.float().mean()) + 1, m)
>>>
>>> # Perform batch GEMM
>>> result = batch_deepgemm_fp8_nt_groupwise(
...     a_fp8, b_fp8, a_scale, b_scale, masked_m, expected_m, out_dtype=torch.bfloat16
... )
>>> print(result.shape)  # torch.Size([2, 128, 2048])

注意

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

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

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

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

  • 块大小由 scale_granularity_mnk 参数确定