flashinfer.gemm.mm_fp4

flashinfer.gemm.mm_fp4(a: Tensor, b: Tensor, a_descale: Tensor, b_descale: Tensor, alpha: Tensor | None = None, out_dtype: dtype = torch.bfloat16, out: Tensor | None = None, block_size: int = 16, use_8x4_sf_layout: bool = False, backend: Literal['cudnn', 'trtllm', 'cutlass', 'auto'] = 'auto', use_nvfp4: bool = True) Tensor

MM FP4

参数:
  • a (torch.Tensor) – 输入张量,形状为 (m, k),fp4 e2m1fn_x2 或 uint8。

  • b (torch.Tensor) – Mat2 张量,形状为 (k, n),应为列主序,fp4 e2m1fn_x2 或 uint8。

  • a_descale (torch.Tensor) – A 的块缩放张量,形状为 (m, k // block_size),float8_e4m3fn 或 uint8。

  • b_descale (torch.Tensor) – B 的块缩放张量,形状为 (k, n // block_size),float8_e4m3fn 或 uint8。

  • alpha (Optional[torch.Tensor]) – 全局缩放张量,float 标量。

  • out_dtype (torch.dtype) – 输出 dtype,bf16 或 fp16。当 backend="trtllm" 时,仅支持 bf16

  • out (Optional[torch.Tensor]) – 输出张量,形状为 (m, n),bf16 或 fp16,默认为 None

  • block_size (int) – FP4 量化块大小,仅支持 16 和 32。nvfp4 量化时为 16。mxfp4 量化时为 32。

  • use_8x4_sf_layout (bool) – 是否使用 8x4 比例因子布局或 128x4 比例因子布局,默认为 False。

  • backend (Literal["cudnn", "trtllm", "cutlass", "auto"]) – 要使用的后端,默认为 "auto",它会根据当前的 CUDA 和 cuDNN 版本自动选择 "cudnn""cutlass" 之间的最佳后端。当 backend="auto" 时,永远不会选择 "trtllm" 后端,因为它需要不同的权重准备。

  • use_nvfp4 (bool) – 是否使用 nvfp4 量化或 mxfp4 量化,默认为 True。有关相关约束,请参阅 block_size 参数。

注意

当使用 cudnn/cutlass 后端时,a 和 b 都应使用 128x4 比例因子布局和 do_shuffle=False 进行量化。当使用 trtllm 后端时,b 必须使用 128x4 布局和 do_shuffle=True 进行量化。a 可以使用 128x4 或 8x4 布局(由 use_8x4_sf_layout 控制)进行量化,并且 do_shuffle=False

返回值:

out – 输出张量,形状为 (m, n),bf16 或 fp16。

返回值类型:

torch.Tensor

示例

>>> import torch
>>> from flashinfer import nvfp4_quantize, mm_fp4, SfLayout
>>> a = torch.randn([48, 128], device="cuda", dtype=torch.bfloat16)
>>> b = torch.randn([256, 128], device="cuda", dtype=torch.bfloat16)
>>> a_global_sf = (448 * 6) / a.float().abs().nan_to_num().max()
>>> b_global_sf = (448 * 6) / b.float().abs().nan_to_num().max()
>>> a_fp4, a_sf = nvfp4_quantize(a, a_global_sf, sfLayout=SfLayout.layout_128x4, do_shuffle=False)
>>> b_fp4, b_sf = nvfp4_quantize(b, b_global_sf, sfLayout=SfLayout.layout_128x4, do_shuffle=True)
>>> out = mm_fp4(a_fp4, b_fp4.T, a_sf, b_sf.T, 1.0/(a_global_sf * b_global_sf), torch.bfloat16, None, backend="trtllm")
>>> out.shape
torch.Size([48, 256])