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