flashinfer.fp4_quantization.nvfp4_batched_quantize

flashinfer.fp4_quantization.nvfp4_batched_quantize(a, a_global_sf, sf_vec_size=16)

将批处理输入张量量化为 NVFP4 格式。

参数:
  • a (torch.Tensor) – 输入张量,形状为 [B, M, K],数据类型为 fp16/bf16。

  • a_global_sf (torch.Tensor) – 全局缩放因子,形状为 [1],数据类型为 float32。

  • sf_vec_size (int, 可选) – 缩放因子向量大小。默认为 16。

返回值:

包含一个元组
  • 量化后的张量,形状为 [B, M, K/2],dtype 为 FLOAT4_E2M1X2

  • 缩放因子张量,其形状由布局和 sf_vec_size 决定

返回值类型:

Tuple[torch.Tensor, torch.Tensor]