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]