资讯详情

资讯详情

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.mdmodelscope 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 的模型第二阶段只训练 DiTtext_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_dtypetorch.bfloat16, devicecuda, model_configs[ ModelConfig(model_idQwen/Qwen-Image, origin_file_patterntransformer/diffusion_pytorch_model*.safetensors), ModelConfig(model_idQwen/Qwen-Image, origin_file_patterntext_encoder/model*.safetensors), ModelConfig(model_idQwen/Qwen-Image, origin_file_patternvae/diffusion_pytorch_model.safetensors), ], tokenizer_configModelConfig(model_idQwen/Qwen-Image, origin_file_patterntokenizer/), ) pipe.load_lora(pipe.dit, ./models/train/Qwen-Image-LoRA-splited/epoch-4.safetensors) prompt a dog image pipe(prompt, seed0) image.save(split_training_Qwen-Image.jpg)脚本会把第二阶段输出的 LoRA文档以epoch-4.safetensors为例对应--num_epochs 5的最后一个 epoch请按实际保存的 checkpoint 文件名替换加载到 DiT 上用 prompta 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),仅供参考
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →