☰
LTX-2 量化感知蒸馏(QAD)实战:在原生 LTX 训练循环中用 Model Optimizer 融合 NVFP4 校准与知识蒸馏
2026/9/27 5:39:41 网站建设 项目流程
  • 人工智能
  • 大模型
  • 模型优化
  • 模型量化
  • 模型压缩

【免费下载链接】Model-Optimizer

A unified library of SOTA model optimization techniques like quantization, distillation, pruning, neural architecture search, speculative decoding, etc. It compresses deep learning models for downstream deployment frameworks like TensorRT-LLM, TensorRT, vLLM, etc. to optimize inference speed.

项目地址:https://gitcode.com/GitHub_Trending/te/Model-Optimizer
点击查看免费下载

本文基于 NVIDIA Model Optimizer 仓库中的 QAD 示例(examples/windows/diffusers/qad_example),完整讲解如何为 LTX-2 视频生成 DiT 模型搭建Quantization-Aware Distillation(量化感知蒸馏)训练管线:以 LTX 官方训练器(ltx-trainer)为主干,叠加 ModelOpt 的mtq.quantize(PTQ 校准)、mtd.convert(知识蒸馏封装)与 NVFP4 量化,最终产出可直接加载进 ComfyUI 的推理检查点。读完本文,你将掌握该示例的目录结构、配置参数、训练命令、检查点合并全流程,以及每个环节在源码中的具体实现。

示例概览:为什么在蒸馏训练中同时做量化

量化感知蒸馏的核心思想是:让学生模型以量化(fake quantization)形态训练,同时用全精度的教师模型通过蒸馏损失引导其输出,从而在权重被压缩到 NVFP4 等低位格式后仍保留原始模型的生成质量。该示例把这一思路落地到 LTX-2 上,组合了两类组件:

  • LTX packages(ltx-trainer等):提供训练循环、数据集、训练策略(masked loss、音视频分支拆分)、flow matching 时间步采样;
  • NVIDIA ModelOpt:提供 PTQ 校准(mtq.quantize)、蒸馏转换(mtd.convert)与 NVFP4 量化配置。

组合损失沿用了完整蒸馏训练器的思想,由任务损失与蒸馏损失加权而成:

L_total = α × L_task + (1−α) × L_distill

其中α即示例中的kd_loss_weight:0表示纯任务损失,1表示纯蒸馏损失。

重要提醒:第三方许可边界。LTX-2 是 Lightricks 提供的第三方模型与软件包(ltx-core、ltx-pipelines、ltx-trainer),不受NVIDIA Model Optimizer 的 Apache 2.0 许可约束。安装和使用这些包必须遵守 LTX Community License Agreement;任何基于 LTX-2 使用 Model Optimizer 产生的衍生模型或微调权重(包括量化、蒸馏后的检查点)仍受该协议约束,而不属于 Apache 2.0。示例脚本sample_example_qad_diffusers.py在启动时也会发出相同的UserWarning(见 sample_example_qad_diffusers.py)。

注意事项。本示例定位为演示 QAD 管线的示例脚本,已在 Linux RTX 5090 上验证可运行,但该配置下会遇到OOM(显存不足),实际使用时需按自己的 GPU 显存调整批次、梯度累积与 FSDP 分片策略。若需要完整阶段的 QAD 实现(LTX-2 DiT + ModelOpt 量化、完整校准选项、检查点续训、多节点训练),请参考仓库中的全量蒸馏训练器 distillation_trainer.py 及其文档 examples/diffusers/distillation/README.md。

环境要求与安装

运行本示例需要:

  • Python 3.10+
  • 支持 CUDA 的 GPU
  • Accelerate(用于 FSDP 多 GPU 训练,此链接为第三方文档,仓库内以accelerate依赖形式提供)

创建虚拟环境并安装依赖:

python -m venv .venv .venv\Scripts\activate # Windows # source .venv/bin/activate # Linux/macOS pip install -r requirements.txt

仓库中 requirements.txt 的实际内容对应文档中的依赖表:

包来源
ltx-coregit+https://github.com/Lightricks/LTX-2.git#subdirectory=packages/ltx-core
ltx-pipelinesgit+https://github.com/Lightricks/LTX-2.git#subdirectory=packages/ltx-pipelines
ltx-trainergit+https://github.com/Lightricks/LTX-2.git#subdirectory=packages/ltx-trainer
nvidia-modelopt[hf]PyPI(nvidia-modelopt[hf],含 HF 相关扩展)

若环境中尚未具备以下基础包,可一并安装:

pip install torch accelerate safetensors pyyaml

项目文件布局

示例目录共 5 个文件,职责划分如下:

文件说明
sample_example_qad_diffusers.py主脚本:QAD 训练与推理检查点生成(约 1007 行)
ltx2_qad.yamlLTX 训练配置(模型、数据、优化、QAD 选项)
fsdp_custom.yamlAccelerate FSDP 多 GPU 训练配置
requirements.txt第三方依赖清单
README.md本文所依据的使用说明

主脚本从ltx_trainer导入LtxvTrainer、PrecomputedDataset、load_transformer、SAMPLERS、get_training_strategy,从modelopt.torch导入mtd(distill)、mto(opt)、mtq(quantization)以及NVFP4_DEFAULT_CFG,构成"原生 LTX 训练 + ModelOpt 量化蒸馏"的组合(见 sample_example_qad_diffusers.py)。

使用步骤

1. 准备数据集

运行 LTX 预处理脚本,从视频中提取 latent 与文本 embedding(与 LTX 训练管线一致):

python scripts/process_dataset.py /path/to/dataset.json \ --resolution-buckets 384x256x97 \ --output-dir /path/to/preprocessed \ --model-path /path/to/ltx2/checkpoint.safetensors \ --text-encoder-path /path/to/gemma \ --batch-size 4 \ --with-audio \ --decode

参数说明:

  • 位置参数:数据集元数据文件路径(带 caption 与视频路径的 CSV/JSON/JSONL)。
  • 必填:--resolution-buckets(分辨率分桶)、--model-path、--text-encoder-path。
  • 可选:--output-dir(默认.precomputed,位于数据集目录下)、--batch-size(默认 1)、--with-audio(启用音频分支)、--decode(解码并保存视频以便人工核验)。

随后在配置文件(步骤 2)中把data.preprocessed_data_root设为与--output-dir相同的路径。

Slurm 集群场景:用srun配合torchrun运行同一脚本,需从 Slurm 读取MASTER_ADDR、MASTER_PORT、WORLD_SIZE,并传--nnodes=$SLURM_NNODES与--nproc_per_node=8。

2. 配置路径与 QAD 超参数

编辑 ltx2_qad.yaml,至少设置三个路径:

  • model.model_path—— 基础 LTX 检查点路径(如.safetensors);
  • model.text_encoder_path—— Gemma 文本编码器路径;
  • data.preprocessed_data_root—— 预处理后的 LTX 数据集路径。

qad段可按需调整:calib_size、kd_loss_weight、exclude_blocks、skip_inference_ckpt。

可通过 YAML 控制的超参数

以下参数均可写入ltx2_qad.yaml;QAD 专属选项还可从 CLI 覆盖(见步骤 3)。文档给出的默认值如下:

段键默认值(示例)说明
qadcalib_size512PTQ 校准批次数(越多 scale 估计越准,但启动越慢)
qadkd_loss_weight0.5组合损失中蒸馏损失的权重;0= 仅任务损失,1= 仅蒸馏损失
qadexclude_blocks[0, 1, 46, 47]排除量化的 Transformer block 索引(如首尾若干层)
qadskip_inference_ckptfalse为true时训练结束后不构建推理检查点
optimizationlearning_rate1e-6学习率(QAD/蒸馏通常取低值)
optimizationsteps300总训练步数
optimizationbatch_size1每设备批次大小
optimizationgradient_accumulation_steps4梯度累积步数(有效批 = batch_size × 累积 × GPU 数)
optimizationoptimizer_type"adamw"优化器类型
checkpointsinterval100每 N 步保存检查点;null表示禁用
(根)output_dir"outputs/ltx2_qad"检查点与日志输出目录

仓库中 ltx2_qad.yaml 的完整内容还包含文档未列出的训练细节,可作为扩展参考:

  • model.training_mode: "full"(非 LoRA);
  • training_strategy.name: "text_to_video"、first_frame_conditioning_p: 0.1;
  • optimization.max_grad_norm: 1.0、scheduler_type: "linear"、enable_gradient_checkpointing: true;
  • acceleration.mixed_precision_mode: "bf16",且注释明确quantization留空——量化由 ModelOpt 接管,而不是 LTX 自带的 quant 方案;load_text_encoder_in_8bit: true用于省显存;
  • validation段定义了推理验证用的 prompts、negative_prompt、video_dims: [768, 448, 89]、guidance_scale: 3.5、inference_steps: 50等;
  • flow_matching.timestep_sampling_mode: "shifted_logit_normal";
  • hub、wandb段控制 HF Hub 推送与 W&B 日志。

主脚本解析 QAD 参数时遵循"CLI 覆盖 YAML,YAML 覆盖默认值"的优先级:只有当 CLI 参数仍等于默认值(512/0.5/[0, 1, 46, 47])时才回退读取qad段(见 sample_example_qad_diffusers.py)。

3. 运行 QAD 训练

使用 Accelerate 配合仓库提供的 FSDP 配置启动:

accelerate launch --config_file fsdp_custom.yaml sample_example_qad_diffusers.py train \ --config ltx2_qad.yaml \

QAD 专属参数可在命令行追加覆盖,例如:

accelerate launch --config_file fsdp_custom.yaml sample_example_qad_diffusers.py train \ --config ltx2_qad.yaml \ --calib-size 512 \ --kd-loss-weight 0.5 \ --exclude-blocks 0 1 46 47 \ --skip-inference-ckpt

训练结束后,检查点保存在output_dir下(如outputs/ltx2_qad/checkpoints/),格式为 safetensors,并附带可选的 amax 与 modelopt state 文件。主脚本还支持仅构建推理检查点的子命令:

python sample_example_qad_diffusers.py create-inference \ --trained path/to/model_weights_step_02200.safetensors \ --base path/to/ltx2/base.safetensors \ --output path/to/inference.safetensors

fsdp_custom.yaml(见 fsdp_custom.yaml)的关键设置包括:distributed_type: FSDP、fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP、fsdp_transformer_layer_cls_to_wrap: BasicAVTransformerBlock、fsdp_state_dict_type: SHARDED_STATE_DICT、fsdp_sync_module_states: true、fsdp_cpu_ram_efficient_loading: true、num_processes: 8。需要强调的是:主脚本在 FSDP 包装(accelerator.prepare)之前先完成量化和蒸馏封装,这与通常的"先 prepare 再转换"顺序不同,是让mtq/mtd的模块替换与 FSDP 分片正确共存的关键设计。

4. 创建推理检查点(ComfyUI 兼容)

ComfyUI 是基于节点的扩散模型运行界面(支持 Stable Diffusion、LTX 等),导出的检查点可在其中加载,通过 prompt 与工作流生成图像或视频。要将训练产物合并为单个 ComfyUI 兼容的检查点,可使用 PTQ 检查点合并器:

python -m ltx2.tools.ptq.checkpoint_merger \ --artefact /path/to/amax_artifact.json \ --checkpoint /path/to/ltx2_qad_bf16.safetensors \ --config /path/to/config.yaml \ --output /path/to/comfyui_checkpoints/nvfp4_qad_inference.safetensors

参数含义:

  • --artefact—— 校准/QAD 训练产生的 amax artifact JSON 路径;
  • --checkpoint—— 训练后的 QAD 权重(如本次运行得到的ltx2_qad_bf16.safetensors);
  • --config—— 合并器配置 YAML 路径;
  • --output—— 输出 ComfyUI 就绪的.safetensors文件路径。

该命令产出单个.safetensors文件,可直接在 ComfyUI 中加载。仓库主脚本内置的create_inference_checkpoint实现了等价逻辑,其 7 步流程(见 sample_example_qad_diffusers.py)可概括为:加载训练检查点(兼容 torch pickle 与 safetensors,通过文件头魔数自动探测格式)→ 提取 amax 并单独存为 JSON → 剔除 teacher/loss/quantizer 键 → 加载基础检查点并匹配 dtype → 为 transformer 键补model.diffusion_model.前缀 → 合并(基础模型非 transformer 权重 + 基础 embeddings_connector + 训练过的 transformer)→ 以原子写方式(写.tmp再 rename)保存 safetensors 并携带基础模型元数据。

工作原理:源码级拆解

文档给出的 5 步流程与 sample_example_qad_diffusers.py 的实现一一对应:

1. 模型加载—— 基础 Transformer 通过ltx_trainer.model_loader.load_transformer加载。

2. PTQ 校准——mtq.quantize使用 LTX 数据集与训练策略跑校准循环,NVFP4 配置排除敏感层与指定 block。量化配置由build_quant_config生成(sample_example_qad_diffusers.py):

  • NVFP4 核心理化格式:num_bits: (2, 1)(E2M1 主格式 + 1 位缩放)、block_sizes: {-1: 16, "type": "dynamic", "scale_bits": (4, 3)}、axis: None;
  • *weight_quantizer与*input_quantizer全部启用;
  • 一组敏感层模式被禁用量化,包括*patchify_proj*、*adaln_single*、*caption_projection*、*proj_out*,以及音频分支与音视频交叉注意力(av_ca_*)相关模式(见SENSITIVE_LAYER_PATTERNS,sample_example_qad_diffusers.py);
  • exclude_blocks指定的 block(默认首尾各两个[0, 1, 46, 47])整块排除量化。

_run_calibration(sample_example_qad_diffusers.py)用 LTX 的PrecomputedDataset+ 训练策略构造校准 DataLoader,按calib_size取批量样本执行前向(calibration_forward_loop),中途失败批次被记录跳过,失败比例过半才中止;校准完成后 rank 0 打印mtq.print_quant_summary。

3. 蒸馏设置—— 以同一检查点加载全精度教师模型,随后用 ModelOpt 的mtd.convert把量化后的学生包装为蒸馏模型。_setup_distillation(sample_example_qad_diffusers.py)的关键配置:

  • teacher_model以惰性 lambda 方式提供(教师加载在 CPU、bf16);
  • 蒸馏损失criterion为自定义DiffusionMSELoss(sample_example_qad_diffusers.py),针对 LTX-2 前向返回(video_pred, audio_pred)元组的新输出格式,按video_weight=0.95、audio_weight=0.05加权两个分支的 MSE;
  • loss_balancer使用 ModelOpt 的mtd.StaticLossBalancer(kd_loss_weight=...)(对应 loss_balancers.py 中的StaticLossBalancer),实现L_total = α·L_task + (1−α)·L_distill的组合;
  • expose_minimal_state_dict: False。

4. 训练—— 标准 LTX 训练循环,但覆写了_training_step(sample_example_qad_diffusers.py):先由训练策略算出hard_loss(任务损失),再unwrap_model后通过DistillationModel.compute_kd_loss(student_loss=hard_loss)注入蒸馏损失。

5. 检查点保存—— 覆写的_save_checkpoint(sample_example_qad_diffusers.py)做四件事:

  • 提取 amax:把含_amax的键值序列化为 JSON(amax_step_XXXXX.json),供后续 NVFP4 推理合并使用;
  • 过滤键:按QUANTIZER_KEYWORDS(_amax、_zero_point、input_quantizer、weight_quantizer、output_quantizer)、TEACHER_KEYWORDS(_teacher_model)、LOSS_KEYWORDS(_loss_modules)剔除教师/损失/量化器状态(is_removable_key,sample_example_qad_diffusers.py),只保留干净的学生权重;
  • dtype 对齐:与基础模型对照,把model.diffusion_model.前缀下的权重 dtype 对齐到基础模型(fp32 兜底转 bf16);
  • 原子保存:先写.tmp再 rename,并使用safetensors.save_file(不静默回退 pickle);随后另存modelopt_state_step_*.pth(含get_quantizer_state_dict得到的量化器状态),便于后续续训或进一步导出。

多节点安全方面,脚本通过is_global_rank0()(sample_example_qad_diffusers.py)保证只有全局 rank 0 写共享文件系统;FSDP 的get_state_dict集合通信在所有 rank 上执行。

进阶:从示例走向全阶段 QAD

若需要完整的 QAD 训练器能力(文档中明确说明本示例只是演示脚本),仓库的 examples/diffusers/distillation 提供了更完整的实现:

  • 全量训练器 distillation_trainer.py;
  • 完整配置示例 distillation_example.yaml,支持:
    • distillation.distillation_alpha(对应本文的kd_loss_weight)、distillation_loss_type(mse/cosine)、teacher_dtype、teacher_model_path;
    • 与 PTQ 工作流一致的全推理校准:calibration_prompts_file(默认使用提示词数据集)、calibration_size(默认 128,每个 prompt 跑完整去噪循环)、calibration_n_steps(默认 30)、calibration_guidance_scale(默认 4.0);
    • 检查点续训:resume_from_checkpoint("latest"或显式路径)、must_save_by(Slurm 时限前自动保存退出)、restore_quantized_checkpoint/save_quantized_checkpoint;
    • 自定义量化配置:在CUSTOM_QUANT_CONFIGS中注册后,YAML 里用quant_cfg: MY_FP8_CFG引用(如FP8_DEFAULT_CFG、INT8_DEFAULT_CFG、NVFP4_DEFAULT_CFG,后者定义于 config.py);
    • 多节点启动需在各节点设置NUM_NODES、GPUS_PER_NODE、NODE_RANK、MASTER_ADDR、MASTER_PORT,用accelerate launch --num_machines ... --num_processes ... --machine_rank ... --main_process_ip ... --main_process_port ...拉起,并可通过++distillation.distillation_alpha=0.6这类点号记法在 CLI 覆盖配置。

小结

本示例提供了一个"以 LTX 原生训练循环为主干、ModelOpt 只负责量化与蒸馏"的 QAD 最小可运行范式:PTQ 校准阶段用 NVFP4 配置配合敏感层/首尾 block 排除,蒸馏阶段用全精度教师 + 音视频加权 MSE +StaticLossBalancer组合损失,训练阶段通过覆写_training_step注入 KD 损失,保存阶段则通过键过滤、amax 提取、dtype 对齐与原子写产出干净且 ComfyUI 可加载的推理检查点。对于生产级需求,可直接迁移到 examples/diffusers/distillation 的全阶段蒸馏训练器上,复用其完整的校准选项、检查点续训与多节点支持。

  • 人工智能
  • 大模型
  • 模型优化
  • 模型量化
  • 模型压缩

【免费下载链接】Model-Optimizer

A unified library of SOTA model optimization techniques like quantization, distillation, pruning, neural architecture search, speculative decoding, etc. It compresses deep learning models for downstream deployment frameworks like TensorRT-LLM, TensorRT, vLLM, etc. to optimize inference speed.

项目地址:https://gitcode.com/GitHub_Trending/te/Model-Optimizer
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询