博客

Blackwell 上的加速扩散模型:使用 Diffusers 和 TorchAO 实现 MXFP8 和 NVFP4

用于图像和视频生成的扩散模型正日益普及,能够提供超逼真的视觉媒体内容。然而,它们的采用往往受到内存和计算需求的极大限制。量化对于高效部署这些模型至关重要。

在本文中,我们将展示如何通过 diffuserstorchao 在 NVIDIA B200 上对 Flux.1-DevQwenImageLTX-2 模型实现可复现的端到端推理加速,其中 MXFP8 可达 1.26 倍,NVFP4 可达 1.68 倍。我们还概述了如何利用选择性量化、CUDA 图 (CUDA Graphs) 和 LPIPS 作为衡量指标,来迭代优化这些模型的精度与性能。复现本文实验的代码库位于此处

目录:

  • MXFP8 和 NVFP4 背景介绍
  • Diffusers 和 TorchAO 基本用法
  • 基准测试结果
  • 技术考量

MXFP8 和 NVFP4 背景介绍

MXFP8 和 NVFP4 是 NVIDIA Blackwell 架构(如 B200 GPU)原生支持的微缩放 (microscaling) 格式。与缩放整个张量的标准量化不同,微缩放将元素分组为小块(例如 16 或 32 个值),并共享一个高精度的缩放因子。这允许在保留动态范围和精度的同时,显著降低位深度。

  • MXFP8 (OCP Microscaling FP8):一种来自开放计算项目 (OCP) 的 8 位行业标准格式 (E4M3/E5M2)。它使用 32 的块大小和 8 位缩放。它提供了“最佳平衡点”,在几乎不损失视觉质量(更低的 LPIPS)的情况下比 BF16 推理更快,并且通常在较小的批量大小下实现最低延迟。
  • NVFP4 (NVIDIA FP4):一种由 Blackwell Tensor Cores 独特加速的 4 位浮点格式 (E2M1)。它使用 16 的块大小和 FP8 缩放因子。它提供了最高的理论吞吐量和最低的内存占用(比 BF16 小约 3.5 倍),使其成为高批量、计算密集型工作负载的理想选择。

更多详情请参考此文章

Diffusers 和 TorchAO 基本用法

先决条件

NVFP4 需要至少 10.0 的 CUDA 计算能力。因此,请确保您拥有符合要求的 GPU。本文档中的基准测试是在 B200 机器 (B200 DGX) 上进行的。

对于虚拟环境,您可以使用 conda

conda create -n nvfp4 python=3.11 -y

conda activate nvfp4

pip install --pre torch --index-url
https://download.pytorch.org/whl/nightly/cu130

pip install --pre torchao --index-url
https://download.pytorch.org/whl/nightly/cu130

pip install --pre mslk --index-url
https://download.pytorch.org/whl/nightly/cu130

pip install diffusers transformers accelerate sentencepiece protobuf av imageio-ffmpeg

撰写本文时,PyTorch、TorchAO 和 MSLK 的每日构建版本分别为 2.12.0.dev20260315+cu1300.17.0.dev20260316+cu1302026.3.15+cu130

某些模型要求用户在 Hugging Face Hub 平台上进行身份验证。因此,如果尚未完成,请务必在运行示例之前执行 hf auth login

基本用法

通过在 Diffusers 中的原生集成,使用 TorchAO 的 NVFP4 量化配置非常直接。

from diffusers import DiffusionPipeline, TorchAoConfig, PipelineQuantizationConfig

import torch

from torchao.prototype.mx_formats.inference_workflow import (
    NVFP4DynamicActivationNVFP4WeightConfig,
)

config = NVFP4DynamicActivationNVFP4WeightConfig(
    use_dynamic_per_tensor_scale=True, use_triton_kernel=True,
)
pipe_quant_config = PipelineQuantizationConfig(
    quant_mapping={"transformer": TorchAoConfig(config)}
)

pipe = DiffusionPipeline.from_pretrained(
    "black-forest-labs/FLUX.1-dev", 
    torch_dtype=torch.bfloat16,
    quantization_config=pipe_quant_config
).to("cuda")
pipe.transformer.compile_repeated_blocks(fullgraph=True)

pipe_call_kwargs = {
    "prompt": "A cat holding a sign that says hello world",
    "height": 1024,
    "width": 1024,
    "guidance_scale": 3.5,
    "num_inference_steps": 28,
    "max_sequence_length": 512,
    "num_images_per_prompt": 1,
    "generator": torch.manual_seed(0),
}
result = pipe(**pipe_call_kwargs)
image = result.images[0]
image.save("my_image.png")

上述代码片段对模型中的每个 torch.nn.Linear 层进行了量化。

在本文中,我们始终使用 fullgraph=True 的区域编译,因为它显著缩短了编译时间并产生了几乎等同于全模型编译的结果。有关区域编译的更多信息,请访问此处

方案选择

下面的代码片段展示了如何使用 TorchAO 配置 MXFP8 和 NVFP4 推理。

# MXFP8

quant_config = MXDynamicActivationMXWeightConfig(
    activation_dtype=torch.float8_e4m3fn,
    weight_dtype=torch.float8_e4m3fn,
    kernel_preference=KernelPreference.AUTO,
)

# NVFP4

quant_config = NVFP4DynamicActivationNVFP4WeightConfig(
    use_dynamic_per_tensor_scale=True,
    use_triton_kernel=True,
)

基准测试结果

Flux.1-Dev

FLUX.1-dev 基准测试期间使用了以下推理参数。

{
    "prompt": "A cat holding a sign that says hello world",
    "height": 1024,
    "width": 1024,
    "guidance_scale": 3.5,
    "num_inference_steps": 28,
    "max_sequence_length": 512,
}

性能和峰值内存

首先,我们展示了不同设置和基准测试下的延迟和峰值内存消耗,MXFP8 可实现最高 1.26 倍加速,NVFP4 可实现最高 1.59 倍加速。请注意,这些结果使用了选择性量化,即我们排除了某些层不进行量化。我们将在本文后续部分讨论更多关于选择性量化的内容。

Flux-1.dev 使用 MXFP8 和 NVFP4 量化的性能和峰值内存

量化模式 Batch Size 延迟 (s) 内存 (GB) 相比 BF16 加速比
无 (None) 1 2.10 38.34 1.00
MXFP8 1 1.75 26.90 1.21
NVFP4 1 1.41 21.33 1.50
无 (None) 4 7.87 44.39 1.00
MXFP8 4 6.36 32.95 1.24
NVFP4 4 5.09 27.39 1.55
无 (None) 8 15.57 53.00 1.00
MXFP8 8 12.40 41.56 1.26
NVFP4 8 9.81 36.00 1.59

NVIDIA B200,选择性量化,使用区域编译的 torch.compile;batch_size=1 使用 torch.compile(..., mode='reduce-overhead')。量化模式“无”表示不进行量化。

准确度

为测试提示词生成的 MXFP8 和 NVFP4 图像与 bfloat16 基准非常接近。

为了进行更彻底的精度评估,我们计算了 bfloat16 图像(基准)与 MXFP8|NVFP4 图像(实验组)之间的平均 LPIPS 分数,该分数取自 Drawbench 数据集中的提示词平均值。

Flux-1.dev 使用 MXFP8 和 NVFP4 量化的平均 LPIPS 分数

量化模式 Drawbench 上的平均 LPIPS
无 (None) 0
MXFP8 0.11
NVFP4 0.44

NVIDIA B200,选择性量化,使用区域编译的 torch.compile。

LPIPS 分数为零表示“图像完全一致”,LPIPS 分数越低表示感知相似度越高。我们用于计算平均 LPIPS 分数的代码位于此处。请参阅本文后面的 LPIPS 部分,了解有关使用 LPIPS 进行准确度评估的更多详情。

LTX-2

对于 LTX-2,我们在 VAE 上启用了平铺 (tiling) 以使内存需求可控。获得结果时使用了以下推理时间参数。

 {
        "prompt": (
              "INT. HOME OFFICE - DAY. Soft natural daylight lights a desk with an open laptop. The camera holds a steady medium shot. A small real house cat sits naturally on all fours in front of the laptop, much smaller than the desk and computer. The cat looks at the screen curiously. Suddenly, with a soft magical sparkle effect, a pair of tiny reading glasses appears in midair and gently lands on the cat's face. A faint whimsical chime sound plays. The cat pauses for a split second, then begins pressing the keyboard clumsily with one paw, producing rapid typing sounds. The laptop screen glow reflects softly on the cat's fur while light playful music continues."
        ),
        "negative_prompt": "worst quality, inconsistent motion, blurry, jittery, distorted",
        "width": 768,
        "height": 512,
        "num_frames": 121,
        "frame_rate": 24.0,
        "num_inference_steps": 40,
        "guidance_scale": 4.0,
}

性能和峰值内存

LTX-2 使用 MXFP8 和 NVFP4 量化的性能和峰值内存

量化模式 Batch Size 延迟 (s) 内存 (GB) 加速比
无 (None) 1 16.230 72.77 1.00
MXFP8 1 13.724 54.54 1.18
NVFP4 1 10.374 45.72 1.56
无 (None) 4 61.591 87.61 1.00
MXFP8 4 50.956 69.38 1.21
NVFP4 4 36.963 60.56 1.67
无 (None) 8 122.427 107.40 1.00
MXFP8 8 102.546 89.18 1.19
NVFP4 8 72.689 80.36 1.68

NVIDIA B200,选择性量化,使用区域编译的 torch.compile。量化模式“无”表示不进行量化。

准确度

请查看此链接获取测试提示词的视频结果对比。在提示词数据集上计算评估分数(正如我们对 Flux-1.dev 所做的那样)留待未来研究。

QwenImage

获得结果时使用了以下推理时间参数。

 {
    "prompt": "A cat holding a sign that says hello world",
    "negative_prompt": " ",
    "height": 1024,
    "width": 1024,
    "true_cfg_scale": 4.0,
    "num_inference_steps": 50,
}

性能和峰值内存

QwenImage 使用 MXFP8 和 NVFP4 量化的性能和峰值内存

量化模式 Batch Size 延迟 (s) 内存 (GB) 加速比
无 (None) 1 7.454 62.21 1.00
MXFP8 1 6.430 55.65 1.16
NVFP4 1 5.369 52.45 1.39
无 (None) 4 26.779 75.52 1.00
MXFP8 4 21.835 68.97 1.23
NVFP4 4 18.279 65.76 1.47
无 (None) 8 52.095 92.47 1.00
MXFP8 8 41.569 85.91 1.25
NVFP4 8 34.969 82.7 1.49

NVIDIA B200,选择性量化,使用区域编译的 torch.compile,batch_size=1 使用 torch.compile(..., mode='reduce-overhead')。量化模式“无”表示不进行量化。

准确度

为测试提示词生成的 MXFP8 和 NVFP4 图像与 bfloat16 基准非常接近,NVFP4 显示出与 MXFP8 相比略大的差异。

在下表中,我们报告了与 Flux.1-Dev 类似的 LPIPS 分数。

QwenImage 使用 MXFP8 和 NVFP4 量化的平均 LPIPS 分数

量化模式 Drawbench 上的平均 LPIPS
无 (None) 0
MXFP8 0.34
NVFP4 0.41

注:在我们的实验中,我们发现 QwenImage 对量化比 Flux.1-Dev 更敏感,这一点从 QwenImage 的平均 MXFP8 LPIPS 分数(0.34)高于 Flux-1.Dev(0.11)可以得到证明。通过更激进的选择性量化或更先进的数值算法(如 GPTQ、QAT 等)进一步降低 QwenImage 的平均 LPIPS 分数留待未来研究。

技术考量

在本节中,我们将分享如何使用选择性量化、CUDA 图和 LPIPS 来迭代本文中提出的性能和精度指标。

通过选择性量化优化精度和性能

我们使用选择性量化来优化延迟(所有模型)和 LPIPS(Flux-1.dev),根据以下两个简单启发式规则跳过层:

  1. 如果 torch.nn.Linear 的权重或激活形状太小而无法从量化中获益 min(M, K, N) < 1024),则跳过它。这是为了确保矩阵乘法量化带来的加速大于量化激活带来的额外开销(更多背景信息:此处)。
    • 使用 torchao 工具查找模型中权重和激活形状的教程位于此处。请注意,即使权重很大,较小的激活形状也可能使量化不划算。
  2. 如果该层可能对模型精度产生重大贡献(如嵌入层、归一化层),则跳过它。
    • 要将其应用于您的模型,您可以打印模型 (print(model)) 并手动检查 FQN,然后根据您对模型架构的了解,跳过您认为可能影响精度的 FQN。

我们为每个模型使用的具体启发式方法如下:

  1. Flux-1.dev
  2. QwenImage
  3. LTX-2

为了量化选择性量化的影响,我们测量了纯 Bfloat16 生成的图像与使用 NVFP4 和 MXFP8 生成的图像之间的性能、内存和平均 LPIPS(使用 AlexNet)。

Flux-1.dev 全量化与选择性量化的影响

量化模式 LPIPS 延迟 (s) 内存 (GB)
MXFP8 + 全量化 0.138128 1.774 26.84
MXFP8 + 选择性量化 0.107562 1.746 26.90
NVFP4 + 全量化 0.479679 2.112 21.25
NVFP4 + 选择性量化 0.438337 2.076 21.33

(LPIPS 分数越低越好,LPIPS 约为 0.1 通常意味着图像几乎无法区分。LPIPS 计算代码可在此处获取)。

从上面的结果中我们可以看出,从量化中排除某些层(即“选择性量化”)在延迟、峰值内存消耗和 LPIPS 之间提供了最佳的权衡。因此,我们在本文报告的其余两个模型中遵循了选择性量化的配方。

我们使用了简单的启发式方法来寻找我们的选择性量化配方。对于选择性量化还有更先进的方法,例如此层敏感度研究

请注意,在迭代我们的选择性量化配方时,我们发现 TorchAO 的 NVFP4 张量量化内核存在性能差距。我们在此 PR 中改进了 NVFP4 性能,升级了 to_nvfp4 内核以使用 MSLK

使用 CUDA 图改善 CPU 开销

我们注意到,当使用 NVFP4 处理像 1 这样的小批量大小时,CPU 开销往往会对延迟改进产生不可忽视的影响。为了显著减少这种开销,我们使用了“reduce-overhead”编译模式,该模式启用了 CUDA 图。下面我们提供了应用 CUDA 图前后的配置跟踪。

为了将 torch.compile(..., mode='reduce-overhead')diffusers 库中的逐块编译干净地结合起来,我们必须将每个 Transformer 块包装在一个克隆其输入的函数中。执行此操作的 PR 位于此处,在 batch_size==1QwenImage + nvfp4 速度提升了 1.81 倍。

使用 LPIPS 评估图像生成准确度

我们使用 LPIPS (GitHub) 指标来比较由量化模型生成的图像与由基准 (bfloat16) 模型生成的图像的相似程度。伪代码如下:

lpips_scores = []

for text_prompt in dataset:
    generator = torch.Generator(device=device).manual_seed(seed)
    kwargs = {"prompt": prompt, "generator": generator, ...}
    image_baseline = pipe_bf16(**kwargs)
    image_quantized = pipe_quantized(**kwargs)
    lpips_score = calculate_lpips_score(image_baseline, image_quantized)
    lpips_scores.append(lpips_score)

lpips_mean = lpips_scores.sum() / len(lpips_scores)

我们实际使用的代码位于此处

图像对的示例 LPIPS 分数

本节提供图像对的示例 LPIPS 分数,以帮助将上述报告的 LPIPS 指标置于上下文中,并使读者能够理解“什么是好的 LPIPS 分数”。

下图使用 FLUX.1-dev 生成。左侧图像是基准 (bfloat16),右侧图像是对模型中每个 torch.nn.Linear 进行 MXFP8 量化后的结果。LPIPS 分数基于右侧图像(实验组)与左侧图像(基准组)的比较。

下面我们提供类似的对比,但右侧为 NVFP4 图像。

结论

在本文中,我们研究了 NVFP4 和 MXFP8 量化方案在主流图像和视频生成模型上的性能。我们展示了在速度、质量和内存之间提供合理权衡的方案。我们还发现了一些可能阻碍最佳性能的重要问题,并探讨了解决这些问题的方法。我们希望这些方案能帮助提升您的图像和视频生成工作负载的性能。

资源

所有输出结果均可在此找到