news 2026/9/15 18:23:01

DiffSynth-Studio 怎么启用 Split Training 两阶段训练降低显存并加速训练

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
DiffSynth-Studio 怎么启用 Split Training 两阶段训练降低显存并加速训练

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):

  1. --dataset_repeat改为 1,避免冗余计算——第一阶段缓存的中间结果与模型参数无关,可被后续任意 epoch 复用,多跑几遍只会重复写盘;
  2. --output_path改为第一阶段计算结果的保存路径(本例为./models/train/Qwen-Image-LoRA-splited-cache);
  3. 追加参数--task "sft:data_process"
  4. --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_tasktorch.no_grad()遍历数据集,对每个样本执行拆分后剩下的单元(本例为 VAE 编码、文本编码),并把结果torch.save.pth文件写入--output_path下按process_index划分的子目录。因此这一步完成后,./models/train/Qwen-Image-LoRA-splited-cache/下会生成按进程分目录的.pth缓存文件,占用磁盘空间与数据量正相关,这是该方案「用磁盘换显存」的代价。

第二阶段:--task "sft:train"读取缓存训练 LoRA

相对普通训练命令,第二阶段的修改点也是 4 处:

  1. --dataset_base_path改为第一阶段的--output_path(即缓存目录),数据加载器直接读第一阶段产物;
  2. 删除--dataset_metadata_path——缓存中已包含预处理结果,不再需要 metadata 文件;
  3. 追加参数--task "sft:train"
  4. --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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/15 18:21:51

用MySQL和ODBC构建Cadence CIS统一元器件库管理系统

搞硬件设计的兄弟应该都有这种体会:原理图库和PCB封装库要是乱起来,那真是灾难。一个项目里同一个电阻,有人用R0603,有人用RESC1608,还有人直接画个矩形框当电阻用,到了做BOM的时候,采购拿着Exc…

作者头像 李华
网站建设 2026/9/15 18:20:14

Flutter与HarmonyOS跨端日期格式化解决方案

1. 跨端开发中的日期格式化痛点在Flutter与HarmonyOS 6.0的混合开发场景下,日期格式化这个看似简单的功能却暗藏玄机。我最近在开发一个便签类应用时,就遇到了这样的典型问题:当同一条数据需要在Android、iOS和HarmonyOS三端显示时&#xff0…

作者头像 李华
网站建设 2026/9/15 18:19:38

Windows虚拟内存设置指南:页面文件原理与16G/32G配置实操

干这行十几年,被同事喊去救急的场景里,出现频率最高的不是服务器宕机,而是 Windows 突然弹一句“系统虚拟内存太低”。机型五花八门,处理流程倒是出奇一致:先怀疑物理内存不够用,加一条内存条,然…

作者头像 李华
网站建设 2026/9/15 18:17:50

RTMPy:面向地震逆时偏移的GPU加速Python实现

简介:本资源是一个基于Python实现的GPU加速地震波建模与逆时偏移(RTM)开源工具包,面向地球物理勘探、油气成像及计算地球科学领域的研究人员与高校师生,解决传统CPU计算下RTM算法耗时长、建模效率低的核心痛点。压缩包…

作者头像 李华