Diffusers 中的 TorchAO 量化实战:从 INT4/FP8 权重压缩到序列化加载
发布时间:2026/9/12 16:33:42来源:尧图网络
Diffusers 中的 TorchAO 量化实战从 INT4/FP8 权重压缩到序列化加载【免费下载链接】diffusers Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers导读本文基于 Diffusers 官方文档 torchao.md 与仓库量化实现系统讲解如何利用 PyTorch 官方的 TorchAO 高性能量化库在 Diffusers 中对任意模态图像、视频、音频的扩散模型做权重压缩与加速推理。读完本文你将掌握TorchAoConfig与PipelineQuantizationConfig的配置方式、INT4/INT8/FP8 等量化类型的适用场景、torch.compile的叠加加速、以及量化模型序列化与反序列化的完整流程与踩坑规避。一、TorchAO 量化在 Diffusers 中的定位TorchAOPyTorch 官方量化与稀疏化仓库提供基于量化与稀疏性的高性能数据类型与优化方案可用于 PyTorch 模型的训练与推理。其量化能力在 Diffusers 中有一个关键优势只要模型支持通过 Accelerate 加载并且内部包含torch.nn.Linear层就可以被 TorchAO 量化——与模型所属模态图像、视频、音频无关。这一点在 Diffusers 的量化器实现中得到了印证。查看 torchao_quantizer.py 中的check_if_quantized_param方法可以看到量化器只对两类情况放行参数属于nn.Linear模块、且参数名是weight。这意味着量化粒度是层级的所有nn.Linear的权重会被替换为低比特张量子类tensor subclass而其他层保持原精度。# src/diffusers/quantizers/torchao/torchao_quantizer.py节选 module, tensor_name get_module_from_name(model, param_name) return isinstance(module, torch.nn.Linear) and (tensor_name weight)环境要求PyTorch 2.5torchao 的整数量化如 int8依赖此版本在 PyTorch 2.6 中进一步引入 int1–int7 支持见源码中SUPPORTED_TORCH_DTYPES_FOR_QUANTIZATION的定义分支。torchao 0.15.0TorchAoConfig.post_init与量化器的validate_environment都会强制检查该最低版本低于此版本会直接抛出ValueError/RuntimeError。安装命令uv pip install -U torch torchao二、配置量化AOBaseConfig 与 TorchAoConfigTorchAO 的每种量化 dtype 都以一个独立的AOBaseConfig子类实例呈现如Int4WeightOnlyConfig、Int8WeightOnlyConfig、Float8WeightOnlyConfig。这种设计暴露了更多可调参数如group_size、version配置更灵活。Diffusers 用TorchAoConfig包裹这些AOBaseConfig。它的定义位于 quantization_config.py核心参数参数类型说明quant_typeAOBaseConfig指定量化类型的配置实例例如Int4WeightOnlyConfig(group_size128)modules_to_not_convertlist[str]默认None保持原精度、不参与量化的模块名列表适用于模型要求某些模块如 embedding 层、输出头保留原始精度的场景TorchAoConfig的两个实现细节值得注意类型校验post_init会检查quant_type必须是AOBaseConfig实例传入字符串或数字会直接TypeError。这一点在 test_torchao.py 的test_post_init_check中有明确覆盖TorchAoConfig(int4_weight_only)与TorchAoConfig(42)均报错。序列化格式to_dict会将quant_type序列化为{default: config_to_dict(...)}的字典结构from_dict再通过config_from_dict反序列化从而支持将量化配置写入模型的config.json并在加载时自动还原。流水线级量化PipelineQuantizationConfig对于完整流水线pipelineDiffusers 提供PipelineQuantizationConfig见 pipe_quant_config.py以按组件component粒度配置量化。其核心参数quant_mapping是一个从组件名到量化配置的映射例如{transformer: TorchAoConfig(...)}。以一个 FLUX.1-dev 的完整加载示例源自官方文档为例import torch from diffusers import DiffusionPipeline, PipelineQuantizationConfig, TorchAoConfig from torchao.quantization import Int8WeightOnlyConfig pipeline_quant_config PipelineQuantizationConfig( quant_mapping{transformer: TorchAoConfig(Int8WeightOnlyConfig(group_size128, version2))} ) pipeline DiffusionPipeline.from_pretrained( black-forest-labs/FLUX.1-dev, quantization_configpipeline_quant_config, dtypetorch.bfloat16, device_mapcuda # 或 mps、xpu、cpu )要点说明device_mapcuda时量化在 GPU 上逐层进行速度快但加载期间需要额外的 GPU 显存同时容纳原始权重与量化后的权重若显存本就紧张可能出现 OOM。pipeline_quant_config中映射到的组件这里仅transformer会被量化文本编码器等未映射组件保持原精度如需量化文本编码器在quant_mapping中增加对应键即可。PipelineQuantizationConfig的校验逻辑_validate_init_args要求quant_backend与quant_mapping二选一若使用quant_backendtorchaoquant_kwargs的方式Diffusers 会为流水线中的每个组件统一实例化TorchAoConfig。模块级量化AutoModel.from_pretrained如果不走完整流水线也可以直接量化单个模型组件如 transformerTorchAoConfig的 docstring 给出了示范quantization_config.pyfrom diffusers import FluxTransformer2DModel, TorchAoConfig from torchao.quantization import Int8WeightOnlyConfig quantization_config TorchAoConfig(Int8WeightOnlyConfig()) transformer FluxTransformer2DModel.from_pretrained( black-forest-labs/Flux.1-Dev, subfoldertransformer, quantization_configquantization_config, torch_dtypetorch.bfloat16, )整数量化的 dtype 约束必须 bfloat16源码中update_torch_dtype方法揭示了一个重要限制对于Int*/Uint*开头的整数量化配置当前仅支持torch_dtypetorch.bfloat16若传入其他 dtype 会发出告警并强制覆盖为 bfloat16若未指定 dtype也会自动兜底为 bfloat16。原因是量化后的 Linear 运算需要统一的计算精度类型。三、显存紧张时的替代加载策略device_mapcuda虽然最快但会临时占用额外显存。若模型已占用大部分显存加载可能直接 OOM。此时官方推荐去掉device_map加载完成后再把流水线整体搬到 GPUpipeline DiffusionPipeline.from_pretrained( black-forest-labs/FLUX.1-dev, quantization_configpipeline_quant_config, torch_dtypetorch.bfloat16, ) pipeline.to(cuda) # 或 mps、xpu、cpu不带device_map时Diffusers 在CPU 上执行量化速度较慢但避免了量化期间的 GPU 显存尖峰。若希望进一步压低 GPU 显存可以改用DiffusionPipeline.enable_model_cpu_offload()将未使用的组件卸载到 CPU。需要留意的是源码中的额外约束如果使用了包含cpu/disk目标的device_map即 offload 场景预量化pre-quantized模型目前不支持validate_environment会直接抛错此外被 offload 到 CPU/disk 的模块会自动加入modules_to_not_convert即 offload 出去的权重不参与量化。四、叠加 torch.compile 加速推理TorchAO 量化与torch.compile兼容二者叠加可用一行代码获得推理加速也可参考 fp16.md 中关于torch.compile的通用说明。import torch from diffusers import DiffusionPipeline, PipelineQuantizationConfig, TorchAoConfig from torchao.quantization import Int4WeightOnlyConfig pipeline_quant_config PipelineQuantizationConfig( quant_mapping{transformer: TorchAoConfig(Int4WeightOnlyConfig(group_size128))} ) pipeline DiffusionPipeline.from_pretrained( black-forest-labs/FLUX.1-dev, quantization_configpipeline_quant_config, dtypetorch.bfloat16, device_mapcuda # 或 mps、xpu、cpu ) pipeline.transformer.compile(transformer, modemax-autotune, fullgraphTrue)编译要点modemax-autotune会对多种编译配置做网格搜索以寻找最优 kernelfullgraphTrue要求整个前向图无 graph break避免性能回退。编译是形状敏感的一旦输入分辨率与编译时不一致会触发重新编译见 fp16.md 的相关提示。官方文档提供的 Flux 与 CogVideoX 推理速度/显存基准见其对应 PR 中的附表不同硬件上的更多基准可在 torchao 仓库的torchao/quantization目录下查阅。FP8 的使用建议[!TIP] torchao 的 FP8 后训练量化方案适合计算能力compute capability不低于 8.9的 GPU如 RTX-4090、Hopper 等。FP8 在图像与视频生成中通常能取得速度、显存与生成质量的最佳平衡。如果你的 GPU 兼容建议将 FP8 与torch.compile组合使用。从源码看TorchAO 量化器对编译是友好的is_compileable属性返回True且量化后的Linear权重是 TorchAO 的张量子类可被torch.compile的图捕获。五、支持的量化类型与适用场景TorchAO 支持两类量化范式1. 仅权重量化weight-only权重以低比特 dtype 存储但计算仍以更高精度如bfloat16执行。效果是降低权重占用的内存但激活计算的内存峰值保持不变。2. 动态激活量化weight dynamic activation权重存为低比特同时在前向过程中对激活即时量化以进一步省内存。效果是同时压低权重与激活的内存开销但可能带来一定的质量折损官方建议对不同的模型做充分测试后再选用。常用量化配置汇总完整列表以官方 torchao 文档为准类别配置类整数量化Int4WeightOnlyConfig、Int8WeightOnlyConfig、Int8DynamicActivationInt8WeightConfig8 位浮点量化Float8WeightOnlyConfig、Float8DynamicActivationFloat8WeightConfig无符号整数量化IntxWeightOnlyConfig配合dtypetorch.uint4等使用从源码中的_TRAINABLE_QUANTIZATION_CONFIGS可以看到Diffusers 还区分了可训练与不可训练的量化类型Int8WeightOnlyConfig、Int8DynamicActivationInt8WeightConfig、Float8WeightOnlyConfig等 5 类配置的模型可以继续训练/微调is_trainableTrue而 4-bit 无符号等类型通常只用于推理。另外量化器还会根据位宽调整推理内存估算get_cuda_warm_up_factorint4 权重按 1/8 折算、int8 权重按 1/4 折算因为量化张量内部用子张量元数据表示element_size()反映的仍是原始 dtype 而非真实位宽。六、量化模型的序列化与反序列化6.1 常规流程save_pretrained / from_pretrained要保存某个量化 dtype 的模型先按目标 dtype 加载再调用ModelMixin.save_pretrainedimport torch from diffusers import AutoModel, TorchAoConfig from torchao.quantization import Int8WeightOnlyConfig quantization_config TorchAoConfig(Int8WeightOnlyConfig()) transformer AutoModel.from_pretrained( black-forest-labs/Flux.1-Dev, subfoldertransformer, quantization_configquantization_config, dtypetorch.bfloat16, ) transformer.save_pretrained(/path/to/flux_int8wo, safe_serializationFalse)加载已序列化的量化模型用ModelMixin.from_pretrainedimport torch from diffusers import FluxPipeline, AutoModel transformer AutoModel.from_pretrained(/path/to/flux_int8wo, dtypetorch.bfloat16, use_safetensorsFalse) pipe FluxPipeline.from_pretrained(black-forest-labs/Flux.1-Dev, transformertransformer, dtypetorch.bfloat16) pipe.to(cuda) # 或 mps、xpu、cpu prompt A cat holding a sign that says hello world image pipe(prompt, num_inference_steps30, guidance_scale7.0).images[0] image.save(output.png)序列化相关源码要点is_serializable要求huggingface_hub 0.25.0否则保存会被阻止。若模型含 offload 模块且未通过modules_to_not_convert显式声明保存也会被阻止因为 offload 的模块未量化无法可靠重载。torchao 0.16.0 后支持 safetensors 序列化supports_safetensors_serializationTrue量化张量子类会先被flatten_tensor_state_dict展平为普通张量再写入加载时再通过unflatten_tensor_state_dict依据 safetensors 元数据还原见get_state_dict_and_metadata/maybe_update_state_dict。注意模型根部的无前缀张量如 Wan 的scale_shift_table会绕过还原逻辑直接加载。6.2 torch 2.6.0 的兼容问题uint4 无法直接加载在torch2.6.0下部分量化方法如uint4weight-only保存正常但加载会报UnpicklingError。官方给出的绕行方案是手动用torch.load读入 state dict 再装载进空权重模型。注意该方案要求torch.load(..., weights_onlyFalse)只应在权重来自可信来源时使用import torch from accelerate import init_empty_weights from diffusers import FluxPipeline, AutoModel, TorchAoConfig from torchao.quantization import IntxWeightOnlyConfig # 保存模型 transformer AutoModel.from_pretrained( black-forest-labs/Flux.1-Dev, subfoldertransformer, quantization_configTorchAoConfig(IntxWeightOnlyConfig(dtypetorch.uint4)), dtypetorch.bfloat16, ) transformer.save_pretrained(/path/to/flux_uint4wo, safe_serializationFalse, max_shard_size50GB) # ... # 手动加载 state_dict torch.load(/path/to/flux_uint4wo/diffusion_pytorch_model.bin, weights_onlyFalse, map_locationcpu) with init_empty_weights(): transformer AutoModel.from_config(/path/to/flux_uint4wo/config.json) transformer.load_state_dict(state_dict, strictTrue, assignTrue)[!TIP] 上述示例中的AutoModelAPI 需要PyTorch 2.6才能获得完整支持。新版本 PyTorch 下Diffusers 还会通过torch.serialization.add_safe_globals注册UintxTensor、NF4Tensor、Float8AQTTensorImpl等量化张量类型见 torchao_quantizer.py 的_update_torch_safe_globals使 torchao 量化张量能被 PyTorch 2.6 的默认weights_onlyTrue安全加载。七、测试与验证仓库中的量化覆盖Diffusers 仓库为 TorchAO 提供了分层测试可作为自行验证的参照配置层测试tests/quantization/torchao/test_torchao.py 覆盖TorchAoConfig的to_dict格式含quant_type与quant_method字段、AOBaseConfig类型校验、repr输出并要求 torchao 0.15.0。模型层 / 流水线层测试tests/models/testing_utils/quantization.py 与 tests/pipelines/testing_utils/quantization.py 覆盖了TorchAoConfig(Int8WeightOnlyConfig())在真实小模型如hf-internal-testing/tiny-flux-pipe上的端到端生成结果比对以及PipelineQuantizationConfig的quant_mapping/quant_backend两种用法。这些测试同时验证了文档中的核心行为量化配置可被正确序列化进 config、整数量化强制 bfloat16、量化后流水线输出可复现等。八、小结与进一步阅读在 Diffusers 中使用 TorchAO 的标准流程可归纳为四步选定量化类型AOBaseConfig→ 包进TorchAoConfig组件级或PipelineQuantizationConfig流水线级→ 在from_pretrained中传入并选好device_map/ dtype → 按需叠加torch.compile与enable_model_cpu_offload。对于显存敏感或低算力 GPUFP8 torch.compile是文档推荐的组合对于需要长期复用的量化产物建议遵循保存前确认 torchao/huggingface_hub 版本、加载时注意 uint4 与 PyTorch 版本兼容性两条纪律。相关参考量化总览与更多后端bitsandbytes、GGUF 等见 docs/source/en/quantization 目录。TorchAoConfig定义见 quantization_config.py。TorchAO 量化器实现见 torchao_quantizer.py。官方 TorchAO 文档与 Diffusers-TorchAO 示例仓库提供了更完整的量化方法说明与端到端示例。【免费下载链接】diffusers Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网