Nunchaku 4-bit扩散推理集成Diffusers:性能与内存的双重突破

1 阅读

引言

近年来,扩散模型在图像生成领域取得了显著进展,但庞大的模型参数和计算需求使得在消费级GPU上运行这些模型变得困难。以BF16精度加载现代文本到图像模型通常需要20-30 GB的显存,这超出了大多数用户硬件的承受范围。量化技术应运而生,Diffusers已集成多种量化后端,如bitsandbytes、GGUF、torchao和Quanto。然而,这些后端大多仅量化权重(weight-only),虽然能减少内存占用,但通常无法提升推理速度,甚至可能增加延迟。

SVDQuant方法及其参考实现Nunchaku引擎则采用了不同的策略:对Transformer主要层进行4-bit权重和激活(W4A4)量化,既降低内存占用,又加速去噪循环。此前,使用这些检查点需要单独的推理库。现在,随着Nunchaku Lite的推出,Diffusers用户可以直接通过from_pretrained()加载Nunchaku检查点,无需本地CUDA编译,这得益于kernels包的支持。此外,配套的diffuse-compressor工具包允许用户自行量化新架构并发布为标准的Diffusers仓库。

快速上手Nunchaku Lite

首先,安装必要的依赖包:

pip install -U diffusers transformers accelerate kernels bitsandbytes

然后,像加载其他Diffusers模型一样加载预量化管道:

import torch
from diffusers import ErnieImagePipeline

pipe = ErnieImagePipeline.from_pretrained(
    "lite-infer/ERNIE-Image-Turbo-nunchaku-lite-nvfp4_r32-bnb4-text-encoder",
    torch_dtype=torch.bfloat16,
).to("cuda")

image = pipe(
    prompt="A cinematic portrait of a red fox in a misty forest at sunrise, detailed fur, volumetric light",
    height=1024,
    width=1024,
    num_inference_steps=8,
    guidance_scale=1.0,
    generator=torch.Generator("cuda").manual_seed(42),
).images[0]
image.save("output.png")

无需自定义管道类或单独的推理引擎,也无需本地编译。NVFP4内核首次使用时从Hub下载。该检查点将Nunchaku NVFP4 Transformer与bitsandbytes NF4文本编码器配对,在RTX 5090上生成1024x1024图像约需1.7秒,峰值内存约12 GB,而BF16管道约需24 GB。更多细节可参考官方Diffusers文档。

注意:NVFP4检查点需要NVIDIA Blackwell GPU(RTX 50系列、RTX PRO 6000、B200)。对于早期GPU,请使用INT4变体。

背景:SVDQuant与Nunchaku

SVDQuant是Nunchaku引擎背后的量化方法。标准4-bit量化对扩散Transformer而言颇具挑战,因为权重和激活都包含大量离群值。SVDQuant通过将激活离群值转移到权重中,用一个小型16-bit低秩分支表示每个权重矩阵中最难的部分,并将剩余残差量化为4-bit。Nunchaku通过融合内核加速4-bit路径和低秩分支。

Nunchaku将低秩下投影与量化内核融合,低秩上投影与4-bit计算内核融合,消除了16-bit分支的内存访问开销。

引入Nunchaku Lite

原始Nunchaku引擎通过模型特定的融合执行路径(如融合QKV投影、融合GELU/MLP内核)获得高速度。这些优化与每个架构的模块布局和检查点格式紧密相关,因此支持新模型家族通常需要模型特定的集成工作。

Nunchaku Lite是Diffusers中的新集成路径。它允许Diffusers加载Nunchaku风格的检查点,无需自定义管道或单独推理引擎。在底层,Nunchaku Lite在加载检查点之前,用运行时SVDQ/AWQ线性层替换标准Diffusers模型中的相关nn.Linear模块。CUDA内核通过kernels包从Hub获取。使用两种内核家族:

  • svdq_w4a4:4-bit权重和激活,带SVDQuant低秩校正。用于Transformer的注意力与MLP投影,这些地方几乎消耗了所有计算,提供INT4和NVFP4变体。
  • awq_w4a16:4-bit权重,16-bit激活,用于自适应归一化和调制投影,如FLUX的adanorm_single/adanorm_zero或Qwen-Image调制层。这些层受内存限制且对精度敏感,AWQ在保持精度的同时节省内存和空间。

权衡之处在于,没有架构特定的融合内核和模块,Nunchaku Lite无法达到原始Nunchaku引擎的加速效果。然而,基础实现仍能提供约30%的加速,同时保持相同的显存减少水平。

Diffusers中的原生加载

如果您使用过bitsandbytes或torchao,会感觉机制相似。Nunchaku Lite模型仓库是普通的Diffusers仓库,唯一特殊之处是Transformer的config.json中的quantization_config块:

"quantization_config": {
    "quant_method": "nunchaku_lite",
    "compute_dtype": "bfloat16",
    "svdq_w4a4": {
        "precision": "nvfp4",
        "group_size": 16,
        "rank": 32,
        "targets": [
            "layers.0.self_attention.to_q",
            "layers.0.self_attention.to_k",
            "..."
        ]
    },
    "awq_w4a16": {
        "precision": "int4",
        "group_size": 64,
        "targets": [
            "adaLN_modulation.1",
            "..."
        ]
    }
}

此配置告诉Diffusers哪些模块被量化、使用哪种方案,以及实例化哪个Nunchaku Lite运行时层(SVDQW4A4LinearAWQW4A16Linear)。

由于量化模型保持与密集模型完全相同的模块结构,所有下游组件(调度器、LoRA加载钩子、卸载、torch.compile)都将Nunchaku Lite模型视为普通Diffusers模型。

硬件支持

Nunchaku Lite根据GPU代次和检查点精度使用不同的内核变体:

方案 精度 支持的GPU
svdq_w4a4 nvfp4 Blackwell(RTX 50系列、RTX PRO 6000、B200)
svdq_w4a4 int4 Turing / Ampere / Ada(RTX 30 & 40系列、A100、L40S)
awq_w4a16 int4 Turing / Ampere / Ada(RTX 30 & 40系列、A100、L40S)

Volta和Hopper GPU目前不支持4-bit内核。量化器在加载时会验证GPU的CUDA能力,并抛出明确错误,而不是产生错误输出。

进一步提升速度与降低内存

Nunchaku Lite可以与其他Diffusers内存和速度优化结合使用。

torch.compile:编译Transformer可将端到端加速从1.35倍提升到1.8倍:

pipe.transformer.compile(fullgraph=True)

# 或使用 compile_repeated_blocks() 加快编译
pipe.transformer.compile_repeated_blocks(fullgraph=True)

量化文本编码器:Transformer并非唯一占用大量内存的组件。T5或Qwen3等文本编码器可能占用数GB。使用bitsandbytes NF4进一步量化文本编码器,在我们的基准测试中峰值显存减少约22%。

卸载:Diffusers的卸载辅助函数如enable_model_cpu_offload()enable_sequential_cpu_offload()在需要将管道适配到较小GPU时照常工作。

基准测试

以下所有数字均在NVIDIA RTX PRO 6000(Blackwell)上以1024x1024分辨率测量,使用rootonchair/ERNIE-Image-Turbo-nunchaku-lite-int4-bnb4-text-encoder

端到端延迟与内存

配置 完整管道 去噪循环 峰值显存 加速比
BF16基线 3.00 s 2.86 s 31.1 GB 1.0x
Nunchaku Lite NVFP4 2.27 s 2.13 s 20.6 GB 1.35x
Nunchaku Lite NVFP4 + torch.compile 1.68 s 1.53 s 20.6 GB 1.8x
Nunchaku Lite NVFP4 + NF4文本编码器 2.29 s 2.13 s 16.0 GB 1.35x

如上所示,Nunchaku将峰值显存减少高达50%,同时延迟改善约30%。剩余开销主要来自额外的内核启动,torch.compile可以缓解,将完整管道降至1.68秒,比BF16基线快1.8倍。

图像质量

在相同种子和设置下,BF16与4-bit输出质量对比显示,Nunchaku Lite在保持图像质量接近BF16原始版本的同时,实现了显著的性能提升。

量化自定义模型

Nunchaku Lite在Diffusers中的支持是架构无关的,diffuse-compressor工具包提供了端到端的SVDQuant工作流:校准、量化、打包和发布。

以下以量化FLUX.2 Klein 4B为例,涵盖主要步骤:检查模型、校准和量化Transformer、将结果打包为Diffusers管道,然后验证并推送到Hub。完整教程涵盖每个标志的细节。

1. 检查将被量化的内容

通用扫描器遍历模型并决定目标:重复Transformer块栈中的兼容线性层成为SVDQ W4A4目标,识别的调制线性层成为AWQ W4A16目标,其余保持密集。

python examples/text_to_image/quantize_hf.py black-forest-labs/FLUX.2-klein-4B \
  --precision int4 --rank 32 --inspect-config

量化前务必阅读此报告。对于FLUX.2 Klein 4B,预期结果为100个SVDQ目标、3个AWQ目标和6个密集外部线性层,无缺失模式或重复名称。

2. 运行量化

以下命令对Transformer运行SVDQuant,并将量化检查点写入outputs/checkpoints/svdq-int4_r32-flux-2-klein-4b.safetensors

python examples/text_to_image/quantize_hf.py black-forest-labs/FLUX.2-klein-4B \
  --precision int4 \
  --output outputs/checkpoints/svdq-int4_r32-flux-2-klein-4b.safetensors

--precision int4替换为nvfp4以构建Blackwell原生权重。

3. 打包Diffusers管道

转换器将量化Transformer与基础管道的其他组件结合,将紧凑的nunchaku_lite配置写入transformer/config.json,并可选择将文本编码器转换为NF4:

python examples/convert_nunchaku_lite_diffusers.py \
  --checkpoint outputs/checkpoints/svdq-int4_r32-flux-2-klein-4b.safetensors \
  --model-id black-forest-labs/FLUX.2-klein-4B \
  --bnb4-text-encoder text_encoder \
  --compute-dtype bfloat16 \
  --output-dir outputs/diffusers/FLUX.2-klein-4B-nunchaku-lite-int4-bnb4-text-encoder

4. 加载、验证并推送到Hub

import torch
from diffusers import DiffusionPipeline

pipe = DiffusionPipeline.from_pretrained(
    "outputs/diffusers/FLUX.2-klein-4B-nunchaku-lite-int4-bnb4-text-encoder",
    device_map="cuda",
)
image = pipe(
    "A glass robot in a greenhouse, cinematic lighting",
    num_inference_steps=4, guidance_scale=1.0,
    generator=torch.Generator("cuda").manual_seed(12345),
).images[0]

输出满意后,运行pipe.push_to_hub("your-name/your-model-nunchaku-lite-int4")。其他用户即可使用上述from_pretrained()模式加载。

量化具有结构重写的模型

注意,通用路径假设架构可以在没有结构重写的情况下量化。为了额外加速,原始Nunchaku引擎将Diffusers层组重写为融合模块。通用路径无法自行推断这些更改,例如将独立的Q、K、V投影合并为一个模块,或将融合投影拆分为多个模块。

FLUX.1-dev的QKV投影是一个具体例子。Diffusers定义了三个独立模块:

self.to_q = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_k = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_v = torch.nn.Linear(query_dim, self.inner_dim, bias=bias)

Nunchaku FLUX模块将这些层组合成一个量化的to_qkv模块:

to_qkv = fuse_linears([other.to_q, other.to_k, other.to_v])
self.to_qkv = SVDQW4A4Linear.from_linear(to_qkv, **kwargs)

这个分组模块是必需的,因为Nunchaku的融合算子同时消费QKV投影、Q/K归一化和旋转嵌入。相比之下,默认Diffusers路径分别执行它们:

query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)

query = query.unflatten(-1, (attn.heads, -1))
key = key.unflatten(-1, (attn.heads, -1))
value = value.unflatten(-1, (attn.heads, -1))

query = attn.norm_q(query)
key = attn.norm_k(key)

if image_rotary_emb is not None:
    query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1)
    key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1)

Nunchaku路径将分组投影、归一化模块和旋转嵌入提供给一个融合算子:

qkv = fused_qkv_norm_rottary(
    hidden_states, attn.to_qkv, attn.norm_q, attn.norm_k, image_rotary_emb
)

这就是通用路径无法推断的结构重写。Diffusers有三个目标模块,参数前缀为to_qto_kto_v,而Nunchaku有一个分组模块to_qkv。模型特定的目标配置或适配器必须声明Q、K、V参数应按输出维度按顺序拼接,并加载到to_qkv中。

此类结构重写由量化期间的模型特定目标配置描述,并在加载检查点时由小型运行时适配器处理。FLUX.2 Klein 4B量化脚本提供了生成结构重写检查点的具体目标配置示例,而rootonchair/nunchaku-lite提供了加载分组QKV张量、拆分融合投影等融合操作所需的运行时适配器。完整工作流可参考“添加新模型”指南。

现成检查点

以下仓库可供立即使用:

结论

Nunchaku的SVDQuant内核是在消费级硬件上高效运行扩散Transformer的最有效方法之一,现已原生支持Diffusers。预量化检查点可通过from_pretrained()加载,diffuse-compressor工具包使新架构的量化无需等待引擎支持。通过量化权重和激活,W4A4路径降低了内存使用,同时改善了去噪延迟,图像质量接近BF16原始版本。

如果您量化并发布了新模型,我们期待您的分享。如有任何问题,欢迎加入我们的Discord。

更多资源:

致谢

感谢Diffusers维护者在整个集成过程中的审查和指导,以及MIT HAN Lab / Nunchaku团队对原始SVDQuant工作的贡献。感谢Marc Sun对博客文章的反馈,感谢Álvaro Somoza试用nunchaku-lite并提供反馈。rootonchair还感谢SilverAI支持这项工作并提供开发环境。