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.float32dtype。这些通常由原始 float32 张量的逐令牌量化生成。b_scale (torch.Tensor) – 张量 b 的缩放因子,形状为
(batch_size, n // block_size, k // block_size),具有torch.float32dtype。这些通常由每个组的原始 float32 张量的逐块量化生成。masked_m (torch.Tensor) – 掩码张量,形状为
(batch_size,),具有torch.int32dtype。每个元素指定每个组中要相乘的有效行数。例如,如果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用于 a 和per_block_cast_to_fp8用于 b)生成缩放因子该函数内部使用 DeepGEMM 后端进行优化的 FP8 计算
所有输入张量必须位于同一 CUDA 设备上
块大小由
scale_granularity_mnk参数确定