DiffSynth-Studio 怎么启用 Split Training 两阶段训练降低显存并加速训练
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
在 DiffSynth-Studio 中训练扩散模型时,VAE 编码、文本编码这类预处理计算与去噪模型参数无关,同一个数据样本在多个 epoch 里会重复执行完全相同的计算。Split Training(两阶段拆分训练)就是针对这一点:框架自动分析 Pipeline 计算图,把与可训练模型无关的计算拆到第一阶段存盘,第二阶段直接读缓存继续训练。官方文档说明其效果是「降低显存占用、提高计算速度,额外占用磁盘空间」(Model Training 低显存训练表)。
需要注意的是:官方文档 明确标注 Split Training 是实验性功能(experimental feature),尚未经历大规模验证;低显存训练表格中也提示「部分模型的两阶段训练未经验证,谨慎使用」。本文以文档中给出的 Qwen-Image LoRA 训练为例,给出完整的两阶段操作路径。
启用前提:确认当前训练脚本支持--task拆分
目前 Split Training 支持两类任务(Split_Training.md):
- 标准监督训练(Standard Supervised Training,即
--task "sft:data_process"/--task "sft:train"); - 直接蒸馏训练(Direct Distillation Training,对应
direct_distill:data_process/direct_distill:train)。
--task参数是训练命令的控制开关,默认值为sft。以 Qwen-Image 为例,仓库中对应的任务映射在 train.py 中可以看到:"sft:data_process"与"direct_distill:data_process"映射到数据预处理入口,"sft:train"、"sft"、"direct_distill:train"、"direct_distill"映射到训练入口。其他模型是否支持,可查看该模型的训练文档,或运行python xxx/train.py -h查看支持的--task取值(见 Model Training 的脚本参数说明)。
准备工作:数据集与模型
以 Qwen-Image 为例,先下载示例数据集(命令来自 Split_Training.md):
modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "qwen_image/Qwen-Image/*" --local_dir ./data/diffsynth_example_dataset模型通过--model_id_with_origin_paths从 ModelScope 远程加载(transformer、text_encoder、vae 三个组件),默认下载到./models路径。训练框架基于accelerate,训练命令统一用accelerate launch启动;多 GPU 等配置可用accelerate config或--config_file设置,本文沿用示例脚本的默认配置。
第一阶段:--task "sft:data_process"预处理并缓存中间结果
相对普通 LoRA 训练命令,第一阶段的修改点共 4 处(Split_Training.md):
--dataset_repeat改为 1,避免冗余计算——第一阶段缓存的中间结果与模型参数无关,可被后续任意 epoch 复用,多跑几遍只会重复写盘;--output_path改为第一阶段计算结果的保存路径(本例为./models/train/Qwen-Image-LoRA-splited-cache);- 追加参数
--task "sft:data_process"; - 用
--offload_models列出当前阶段不需要 forward 计算的模型,格式与--model_id_with_origin_paths一致。本例第一阶段不跑 DiT,所以 offload 的是 transformer。
也可以直接从--model_id_with_origin_paths中删掉不需要 forward 的模型来省显存,但必须确保这些模型不会在 Pipeline 中被间接调用,这需要你了解 Pipeline 内部细节;不确定时保留加载并走--offload_models。
完整命令(与文档及仓库示例脚本 Qwen-Image-LoRA.sh 一致):
accelerate launch examples/qwen_image/model_training/train.py \ --dataset_base_path data/diffsynth_example_dataset/qwen_image/Qwen-Image \ --dataset_metadata_path data/diffsynth_example_dataset/qwen_image/Qwen-Image/metadata.csv \ --max_pixels 1048576 \ --dataset_repeat 1 \ --model_id_with_origin_paths "Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors,Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors" \ --offload_models "Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors" \ --learning_rate 1e-4 \ --num_epochs 5 \ --remove_prefix_in_ckpt "pipe.dit." \ --output_path "./models/train/Qwen-Image-LoRA-splited-cache" \ --lora_base_model "dit" \ --lora_target_modules "to_q,to_k,to_v,add_q_proj,add_k_proj,add_v_proj,to_out.0,to_add_out,img_mlp.net.2,img_mod.1,txt_mlp.net.2,txt_mod.1" \ --lora_rank 32 \ --use_gradient_checkpointing \ --dataset_num_workers 8 \ --find_unused_parameters \ --task "sft:data_process"执行逻辑可在 runner.py 中核对:launch_data_process_task以torch.no_grad()遍历数据集,对每个样本执行拆分后剩下的单元(本例为 VAE 编码、文本编码),并把结果torch.save为.pth文件写入--output_path下按process_index划分的子目录。因此这一步完成后,./models/train/Qwen-Image-LoRA-splited-cache/下会生成按进程分目录的.pth缓存文件,占用磁盘空间与数据量正相关,这是该方案「用磁盘换显存」的代价。
第二阶段:--task "sft:train"读取缓存训练 LoRA
相对普通训练命令,第二阶段的修改点也是 4 处:
--dataset_base_path改为第一阶段的--output_path(即缓存目录),数据加载器直接读第一阶段产物;- 删除
--dataset_metadata_path——缓存中已包含预处理结果,不再需要 metadata 文件; - 追加参数
--task "sft:train"; --offload_models同样填本阶段不需要 forward 的模型:第二阶段只训练 DiT,text_encoder 与 vae 都不再 forward,因此 offload 这两者。
--dataset_repeat恢复为原值(本例 50),缓存对每个样本只需计算一次、可被所有 epoch 复用:
accelerate launch examples/qwen_image/model_training/train.py \ --dataset_base_path "./models/train/Qwen-Image-LoRA-splited-cache" \ --max_pixels 1048576 \ --dataset_repeat 50 \ --model_id_with_origin_paths "Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors,Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors" \ --offload_models "Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors" \ --learning_rate 1e-4 \ --num_epochs 5 \ --remove_prefix_in_ckpt "pipe.dit." \ --output_path "./models/train/Qwen-Image-LoRA-splited" \ --lora_base_model "dit" \ --lora_target_modules "to_q,to_k,to_v,add_q_proj,add_k_proj,add_v_proj,to_out.0,to_add_out,img_mlp.net.2,img_mod.1,txt_mlp.net.2,txt_mod.1" \ --lora_rank 32 \ --use_gradient_checkpointing \ --dataset_num_workers 8 \ --find_unused_parameters \ --task "sft:train"训练结束后,LoRA 权重按 epoch(或--save_steps间隔)保存在./models/train/Qwen-Image-LoRA-splited/。
验证结果:加载 LoRA 做推理出图
框架不记录 loss 值(loss 与扩散模型实际效果关系不大,见 Model Training 的训练注意事项),文档给出的验证方式是加载训练产物做一次推理。仓库提供了校验脚本 validate.py,内容如下:
from diffsynth.pipelines.qwen_image import QwenImagePipeline, ModelConfig import torch pipe = QwenImagePipeline.from_pretrained( torch_dtype=torch.bfloat16, device="cuda", model_configs=[ ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="transformer/diffusion_pytorch_model*.safetensors"), ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="text_encoder/model*.safetensors"), ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), ], tokenizer_config=ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="tokenizer/"), ) pipe.load_lora(pipe.dit, './models/train/Qwen-Image-LoRA-splited/epoch-4.safetensors') prompt = "a dog" image = pipe(prompt, seed=0) image.save('split_training_Qwen-Image.jpg')脚本会把第二阶段输出的 LoRA(文档以epoch-4.safetensors为例,对应--num_epochs 5的最后一个 epoch,请按实际保存的 checkpoint 文件名替换)加载到 DiT 上,用 prompt"a dog"出图并保存为split_training_Qwen-Image.jpg。判断标准就是该图片是否正常生成并体现出训练数据(示例数据集为qwen_image/Qwen-Image)的风格。
限制与注意事项
- 实验性功能:Split_Training.md 提示该功能尚未经历大规模验证,使用中遇到问题建议提交 issue。
- 适用任务:目前文档只覆盖
sft(标准监督训练)与direct_distill(直接蒸馏)两类任务的拆分;其他--task值没有给出拆分支持。 - 磁盘开销:第一阶段把每个样本的中间结果存为
.pth文件,数据量大时磁盘占用明显,规划--output_path时预留空间。 - offload 判断:
--offload_models填的是「当前阶段不需要 forward 计算的模型」。如果不确定某模型是否会在 Pipeline 中被间接调用,就不要从--model_id_with_origin_paths中删除它,改用 offload。 - 框架通用限制:训练框架不支持 batch size > 1(见 QA),两阶段脚本与此一致。
如果当前任务不在sft/direct_distill范围内,或目标模型的拆分逻辑无法确认,建议先用普通--task "sft"训练跑通,再考虑 Split Training。
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考