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])