DiffSynth-Studio 模型训练完全指南:数据集准备、参数详解与低显存训练实战
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
本文是 DiffSynth-Studio 模型训练体系的实战指南,覆盖从训练目录结构、通用数据集格式、模型加载方式到训练参数、启动命令与低显存优化策略的完整链路。读完本文,你将能够基于examples下各模型的train.py与.sh脚本,独立完成全量训练(Full Training)与 LoRA 训练的配置、启动、断点续训与效果验证,并能针对显存瓶颈选择合适的技术组合。
DiffSynth-Studio 为 Diffusion 模型提供了一套统一的训练框架:每个模型架构的训练代码以独立的train.py编写在 examples 目录下,并配套.sh启动脚本与验证脚本。以 Z-Image 为例,其训练相关文件结构如下:
diffsynth/diffusion/ # 基础训练框架 examples/z_image/ ├── model_inference ├── model_inference_low_vram └── model_training ├── train.py # Z-Image 架构的模型训练代码在这里 ├── full │ └── Z-Image.sh # 启动全量训练 ├── validate_full │ └── Z-Image.py # 全量训练完成后,运行这个脚本加载模型,验证效果 ├── lora │ └── Z-Image.sh # 启动 LoRA 训练 └── validate_lora └── Z-Image.py # LoRA 训练完成后,运行这个脚本加载模型,验证效果从源码结构看,这套模式在仓库中高度统一:examples/qwen_image/model_training/train.py、examples/flux/model_training/train.py、examples/wanvideo/model_training/train.py等均为同一套框架的不同模型实现,因此掌握一个模型的训练流程即可举一反三。
训练脚本参数全解
训练脚本的参数由训练框架统一注册,定义于 diffsynth/diffusion/parsers.py,按功能分为以下分组。下表补充了源码中的默认值,帮助你理解参数缺省行为。
数据集基础配置
| 参数 | 说明 | 默认值 | | - | - | - | |--dataset_base_path| 数据集的根目录。 | 必填(空字符串) | |--dataset_metadata_path| 数据集的元数据文件路径,支持csv、json、jsonl三种格式。 |None| |--dataset_repeat| 每个 epoch 中数据集重复的次数,可用于放大单轮训练样本量。 |1| |--dataset_num_workers| 每个 DataLoader 的进程数量。 |0| |--data_file_keys| 元数据中需要加载的字段名称,通常是图像或视频文件的路径,以,分隔。 |"image,video"|
模型加载配置
| 参数 | 说明 | | - | - | |--model_paths| 要加载的模型路径,JSON 格式,可描述多文件分片模型。 | |--model_id_with_origin_paths| 带原始路径的模型 ID,例如"Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors",多个条目用逗号分隔。 | |--extra_inputs| 模型 Pipeline 所需的额外输入参数。例如训练图像编辑模型 Qwen-Image-Edit 时需要额外参数edit_image,以,分隔。 | |--fp8_models| 以 FP8 格式加载的模型,格式与--model_paths或--model_id_with_origin_paths一致。目前仅支持参数不被梯度更新的模型(不需要梯度回传,或梯度仅更新其 LoRA)。 | |--offload_models| 需要 offload 的模型,逗号分隔,仅用于拆分训练(split training)场景。 | |--quant_options| 对加载的模型进行动态量化。以;分隔多个条目,每条为<模型字符串>:<method>[/<exclude_modules>]。<模型字符串>需与--model_paths/--model_id_with_origin_paths中的条目一致,method为已注册的量化方法(如bitsandbytes_nf4),exclude_modules为可选的保持全精度的层。 | |--resume_from_checkpoint| 从 checkpoint 文件中加载模型权重并继续训练。目前仅支持非 LoRA 的单模型加载。 |
关于--quant_options有一个值得注意的实现细节:在 training_module.py 的parse_quant_options中,框架会检查所选量化后端的is_differentiable能力,若量化方法不可微,则会直接抛出ValueError,因为冻结的量化层必须能向 LoRA 分支传递梯度。因此训练场景下只能选择声明is_differentiable=True的量化方法。
训练基础配置
| 参数 | 说明 | 默认值 | | - | - | - | |--learning_rate| 学习率。LoRA 训练建议1e-4,全量训练建议1e-5。 |1e-4| |--num_epochs| 轮数(Epoch)。 |1| |--trainable_models| 可训练的模型,例如dit、vae、text_encoder,逗号分隔可同时训练多个。 |None| |--find_unused_parameters| DDP 训练中是否存在未使用的参数。少数模型包含不参与梯度计算的冗余参数,需开启这一设置避免多 GPU 训练报错。 |False(store_true) | |--weight_decay| 权重衰减大小。 |0.01| |--task| 训练任务,默认为sft。部分模型支持更多训练模式(如direct_distill、dmd2等),请参考各模型的文档。 |"sft"| |--customized_optimizer| 自定义优化器,例如bitsandbytes.optim.Adam8bit或torch.optim.Adam。 |None(默认torch.optim.AdamW) |
--customized_optimizer的实现位于 runner.py 的get_optimizer_class,通过importlib动态导入指定模块中的优化器类;在 runner.py 的launch_training_task中,优化器由model.trainable_modules()(即requires_grad=True的参数)构造,并配合ConstantLR调度器使用。
输出配置
| 参数 | 说明 | 默认值 | | - | - | - | |--output_path| 模型保存路径。 |./models| |--remove_prefix_in_ckpt| 在模型文件的 state dict 中移除前缀,如pipe.dit.。 |"pipe.dit."| |--save_steps| 保存模型的训练步数间隔。若留空,则每个 epoch 保存一次。 |None|
保存逻辑由 logger.py 中的on_step_end/on_epoch_end/save_model实现:设置--save_steps时按step-{num_steps}.safetensors命名保存;未设置时在 epoch 结束时保存为epoch-{epoch_id}.safetensors;训练结束若步数未被save_steps整除,还会补存最后一个 checkpoint。此外,训练启动时会把完整的命令行参数以training_args.json形式写入输出目录(见 runner.py),便于复现实验配置。
LoRA 配置
| 参数 | 说明 | 默认值 | | - | - | - | |--lora_base_model| LoRA 添加到哪个模型上,例如dit。 |None| |--lora_target_modules| LoRA 添加到哪些层上,逗号分隔。 |"q,k,v,o,ffn.0,ffn.2"| |--lora_rank| LoRA 的秩(Rank)。 |32| |--lora_checkpoint| LoRA 检查点的路径。若提供,LoRA 将从此检查点加载(用于断点续训)。 |None| |--preset_lora_path| 预置 LoRA 检查点路径。若提供,此 LoRA 将以融入基础模型的形式加载,用于 LoRA 差分训练。 |None| |--preset_lora_model| 预置 LoRA 融入的模型,例如dit。 |None|
LoRA 的注入基于peft的LoraConfig与inject_adapter_in_model(见 training_module.py),未指定lora_alpha时默认等于lora_rank。一个实用的细节:当--lora_target_modules留空时,框架会通过auto_detect_lora_target_modules自动搜索适合打 LoRA 的层(training_module.py),搜索策略以“ModuleList 且长度大于 1”识别为 block 边界,并在其内部寻找输入输出维度不小于 512 的 Linear 层。
梯度配置
| 参数 | 说明 | 默认值 | | - | - | - | |--use_gradient_checkpointing| 是否启用 gradient checkpointing。 |False(store_true) | |--use_gradient_checkpointing_offload| 是否将 gradient checkpointing 卸载到内存中。 |False(store_true) | |--gradient_accumulation_steps| 梯度累积步数,会被传入accelerate.Accelerator(gradient_accumulation_steps=...)。 |1|
CPU Offload 训练配置
| 参数 | 说明 | 默认值 | | - | - | - | |--enable_model_cpu_offload| 启用 CPU offload 训练:权重保留在 CPU,逐层加载到 GPU 进行计算。 |False(store_true) | |--enable_optimizer_cpu_offload| 当--enable_model_cpu_offload启用时,在 CPU 上执行 optimizer,所有参数都 offload 到 CPU。默认 False(可训练参数和 optimizer 留在 GPU)。 |False(store_true) | |--cpu_offload_split_threshold| (实验性)当--enable_model_cpu_offload启用时,参数总量超过此阈值(单位 MB)的模块会被递归拆分为子模块。None表示直接以叶子模块为单位 offload。 |None|
图像/视频宽高配置
| 参数 | 说明 | 默认值 | | - | - | - | |--height| 图像或视频的高度。将height和width留空以启用动态分辨率。 |None| |--width| 图像或视频的宽度。将height和width留空以启用动态分辨率。 |None| |--max_pixels| 图像或视频帧的最大像素面积。启用动态分辨率时,分辨率大于该值的图片会被缩小,小于该值的图片保持不变。 |1024*1024|
视频模型另有--num_frames(每段视频的帧数,默认 81,从视频前缀采样)等参数,见 parsers.py 的add_video_size_config。此外add_general_config(parsers.py)还提供了模板模型(--template_model_id_or_path、--enable_lora_hot_loading)与日志相关参数(--enable_tensorboard_log、--enable_swanlab_log、--enable_wandb_log、--enable_csv_log等)。
部分模型的训练脚本还包含额外的参数,详见各模型的文档,也可以通过python xxx/train.py -h查看当前脚本支持的全部参数。
准备数据集:通用数据集格式
DiffSynth-Studio 采用通用数据集格式(Universal Dataset):数据集由一系列数据文件(图像、视频等)与一份标注元数据文件组成,建议按如下结构组织:
data/example_image_dataset/ ├── metadata.csv ├── image_1.jpg └── image_2.jpg其中image_1.jpg、image_2.jpg为训练用图像数据,metadata.csv为元数据列表,例如:
image,prompt image_1.jpg,"a dog" image_2.jpg,"a cat"元数据中每一列是一个数据字段,image列指向图像文件路径,prompt列是训练文本提示词;字段名与--data_file_keys对应。通用数据集架构的实现位于 diffsynth/core/data/unified_dataset.py:UnifiedDataset.load_metadata会根据扩展名自动识别csv(pandas 读取)、json与jsonl(逐行 JSON)三种元数据格式;若--dataset_metadata_path为空,则会递归扫描base_path下的.pth缓存文件直接作为训练数据(即load_from_cache模式,常用于两阶段拆分训练的第二阶段)。
图像加载管线由 unified_dataset.py 的default_image_operator定义:路径先转为绝对路径,再经过LoadImage与ImageCropAndResize处理——后者正是--max_pixels、--height、--width参数生效的位置,且宽高会被对齐到 16 的倍数(height_division_factor=16)。视频数据则由default_video_operator支持 jpg/png 等静态图、gif 与 mp4/avi/mov 等常见视频格式的统一处理。
为方便测试,项目构建了样例数据集,可通过以下命令下载:
modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --local_dir ./data/diffsynth_example_dataset适用于 Qwen-Image、FLUX 等图像生成模型的训练。实际训练脚本还会用--include筛选子集,例如 Qwen-Image 的训练脚本 examples/qwen_image/model_training/full/Qwen-Image.sh 中为:
modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "qwen_image/Qwen-Image/*" --local_dir ./data/diffsynth_example_dataset加载模型:远程下载与本地路径两种方式
与推理时的模型加载类似,训练框架支持多种方式配置模型路径,且两种方式可以混用。
从远程下载模型并加载
如果在推理时通过以下设置加载模型:
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"), ]那么在训练时,填入以下参数即可加载对应的模型:
--model_id_with_origin_paths "Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors,Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors"模型文件默认下载到./models路径,该路径可通过环境变量 DIFFSYNTH_MODEL_BASE_PATH 修改。默认情况下,即使模型已经下载完毕,程序仍会向远程查询是否有遗漏文件;若要完全关闭远程请求,请将环境变量 DIFFSYNTH_SKIP_DOWNLOAD 设置为True。
从本地文件路径加载模型
如果从本地文件加载模型,例如推理时:
model_configs=[ ModelConfig([ "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00001-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00002-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00003-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00004-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00005-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00006-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00007-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00008-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00009-of-00009.safetensors" ]), ModelConfig([ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors", "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors", "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors", "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors" ]), ModelConfig("models/Qwen/Qwen-Image/vae/diffusion_pytorch_model.safetensors") ]那么训练时需设置为:
--model_paths '[ [ "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00001-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00002-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00003-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00004-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00005-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00006-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00007-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00008-of-00009.safetensors", "models/Qwen/Qwen-Image/transformer/diffusion_pytorch_model-00009-of-00009.safetensors" ], [ "models/Qwen/Qwen-Image/text_encoder/model-00001-of-00004.safetensors", "models/Qwen/Qwen-Image/text_encoder/model-00002-of-00004.safetensors", "models/Qwen/Qwen-Image/text_encoder/model-00003-of-00004.safetensors", "models/Qwen/Qwen-Image/text_encoder/model-00004-of-00004.safetensors" ], "models/Qwen/Qwen-Image/vae/diffusion_pytorch_model.safetensors" ]'请注意,--model_paths是 JSON 格式,其中不能出现多余的,,否则无法被正常解析。两种方式混用时,框架会先解析--model_paths再解析--model_id_with_origin_paths,最终统一为ModelConfig列表(见 training_module.py 的parse_model_configs)。
设置可训练模块:全量训练与 LoRA
训练框架支持任意模型的训练。以 Qwen-Image 为例,若全量训练其中的 DiT 模型,则需设置:
--trainable_models "dit"若训练 DiT 模型的 LoRA,则需设置:
--lora_base_model dit --lora_target_modules "to_q,to_k,to_v" --lora_rank 32(注意:不同模型的注意力层命名不同,如 Qwen-Image 使用to_q,to_k,to_v,而部分模型默认参数为q,k,v,o,ffn.0,ffn.2,请以目标模型实际层名为准。)
框架为技术探索保留了充分的发挥空间,支持同时训练任意多个模块,例如同时训练 text encoder、controlnet,以及 DiT 的 LoRA:
--trainable_models "text_encoder,controlnet" --lora_base_model dit --lora_target_modules "to_q,to_k,to_v" --lora_rank 32训练模式的切换由switch_pipe_to_training_mode完成(training_module.py):它将 scheduler 切换为训练时间表(set_timesteps(1000, training=True)),通过freeze_except冻结所有非--trainable_models指定的模块,再为lora_base_model注入 LoRA 层;若提供了--preset_lora_path,则先将预置 LoRA 融合进基础模型,用于差分训练场景。
此外,由于训练脚本中加载了多个模块(text encoder、dit、vae 等),保存模型文件时需要移除前缀。例如在全量训练 DiT 部分或训练 DiT 部分的 LoRA 时,请设置--remove_prefix_in_ckpt pipe.dit.。如果多个模块同时训练,则需开发者在训练完成后自行编写代码拆分模型文件中的 state dict。export_trainable_state_dict(training_module.py)会按“当前requires_grad=True的参数名”过滤 state dict 并移除指定前缀,这正是--remove_prefix_in_ckpt的底层实现。
启动训练程序
训练框架基于accelerate构建,训练命令按照如下格式编写:
accelerate launch xxx/train.py \ --xxx yyy \ --xxxx yyyy仓库为每个模型都编写了预置的训练脚本,详见各模型的文档。默认情况下,accelerate会按照~/.cache/huggingface/accelerate/default_config.yaml的配置进行训练,使用accelerate config可在终端交互式地配置,包括多 GPU 训练、DeepSpeed 等。
仓库为部分模型提供了推荐的accelerate配置文件,可通过--config_file设置。例如 Qwen-Image 模型的全量训练(完整示例见 examples/qwen_image/model_training/full/Qwen-Image.sh):
accelerate launch --config_file examples/qwen_image/model_training/full/accelerate_config_zero2offload.yaml examples/qwen_image/model_training/train.py \ --dataset_base_path data/example_image_dataset \ --dataset_metadata_path data/example_image_dataset/metadata.csv \ --max_pixels 1048576 \ --dataset_repeat 50 \ --model_id_with_origin_paths "Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors,Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors" \ --learning_rate 1e-5 \ --num_epochs 2 \ --remove_prefix_in_ckpt "pipe.dit." \ --output_path "./models/train/Qwen-Image_full" \ --trainable_models "dit" \ --use_gradient_checkpointing \ --find_unused_parameters对应的推荐配置 examples/qwen_image/model_training/full/accelerate_config_zero2offload.yaml 使用 DeepSpeed ZeRO-2(zero_stage: 2)并将 optimizer 与参数 offload 到 CPU(offload_optimizer_device: 'cpu'、offload_param_device: 'cpu'),配合mixed_precision: bf16与 8 进程(num_processes: 8),是典型的多卡全量训练基准配置。
训练入口代码的整体流程可以参考 examples/qwen_image/model_training/train.py:构建UnifiedDataset→ 构造继承自DiffusionTrainingModule的训练模块(解析模型配置、拆分 Pipeline units、切换训练模式)→ 创建ModelLogger→ 按--task选择launch_training_task或launch_data_process_task启动。训练主循环位于 runner.py:每个 batch 依次执行前向计算损失、accelerator.backward(loss)、优化器step()、调度器step(),并在每步/每 epoch 结束触发model_logger的保存回调。
训练注意事项
以下是训练实践中总结的关键经验,直接决定训练能否顺利进行与最终效果:
- 数据集的元数据除
csv格式外,还支持json、jsonl格式。关于如何选择最佳的元数据格式,请参考通用数据集文档。 - 通常训练效果与训练步数强相关,与 epoch 数量弱相关,因此更推荐使用
--save_steps按训练步数间隔保存模型文件,而非依赖默认的每 epoch 保存。 - 当数据量 ×
dataset_repeat超过 $10^9$ 时,已观测到数据集加载速度明显变慢,这似乎是 PyTorch 的一个 bug,尚不确定新版本是否已修复。 - 学习率
--learning_rate在 LoRA 训练中建议设置为1e-4,在全量训练中建议设置为1e-5。 - 训练框架不支持 batch size > 1,原因较为复杂。从实现看,runner.py 中
DataLoader使用collate_fn=lambda x: x[0],即每个 batch 只取单条样本,这是框架的既定设计。详见 Q&A 文档。 - 少数模型包含冗余参数,例如 Qwen-Image 的 DiT 部分最后一层的文本编码部分。训练这些模型时需设置
--find_unused_parameters以避免多 GPU 训练中报错。出于对开源社区模型兼容性的考虑,项目不打算删除这些冗余参数。 - Diffusion 模型的损失函数值与实际效果的关系不大,因此训练过程中不会记录损失函数值。建议把
--num_epochs设置为足够大的数值,边训边测,直至效果收敛后手动关闭训练程序。需要说明的是,虽然主循环不依赖 loss 值做早停,但 ModelLogger 仍支持通过--enable_tensorboard_log/--enable_swanlab_log/--enable_wandb_log/--enable_csv_log将 loss 写入日志平台,便于监控训练过程。 - 关于损失函数本身:SFT 训练使用 loss.py 中的
FlowMatchSFTLoss,它在min/max_timestep_boundary限定的范围内随机采样一个 timestep,为输入 latent 添加噪声得到训练样本,计算模型预测与training_target的 MSE 损失,并按training_weight(timestep)加权——这也解释了为何损失值对最终效果参考意义有限。 --use_gradient_checkpointing通常是开启的,除非 GPU 显存足够;--use_gradient_checkpointing_offload则按需开启,详见 diffsynth.core.gradient。- 如需加载前一次训练好的 checkpoint 并继续训练:使用
--lora_checkpoint加载 LoRA checkpoint(training_module.py 会经mapping_lora_state_dict兼容新旧 LoRA 键名),使用--resume_from_checkpoint加载基础模型 checkpoint,目前仅支持单模型的加载。
低显存训练:七种方案的技术原理与选型
框架支持多种方式减少训练所需的显存,下表汇总了各方案的开启动、技术原理、使用效果与适用时机:
| 名称 | 开启方式 | 技术原理 | 使用效果 | 何时启用 | 参考文档 | | - | - | - | - | - | - | | Gradient Checkpointing | 通过--use_gradient_checkpointing开启 | 在前向传播时不保留梯度相关参数,在反向传播时重新计算这些参数 | 显著减少显存占用,增加计算时间 | 在大部分情况下推荐开启 | 文档 | | Gradient Checkpointing Offload | 通过--use_gradient_checkpointing_offload开启 | 在 Gradient Checkpointing 的基础上,将 checkpoint 参数从显存移至内存中 | 进一步减少显存占用和增加计算时间,同时增加内存占用 | 仅推荐在视频生成模型的训练中考虑开启 | 文档 | | DeepSpeed | 通过accelerate config交互式配置 | DeepSpeed 支持将梯度、Optimizer 等参数分拆到多 GPU 上 | 减少显存占用,增加多 GPU 与多机之间的通信成本,增加计算时间 | 仅推荐在多 GPU 与多机集群训练中启用 | 文档 | | FP8 训练 | 通过--fp8_models设置切换 FP8 的模型组件 | 将模型参数以 FP8 精度存储在显存中,在推理时临时转换为更高精度;仅支持不需要梯度更新参数的模型 | 减少显存占用,少量增加计算时间,引入少量训练误差 | 仅推荐在text_encoder、vae等非训练模块上启用,也可在 LoRA 训练时对dit启用 | 文档 | | 自定义量化精度 | 通过--quant_options设置每个模型组件的量化配置 | FP8 训练的高阶版,将模型参数以任意量化精度存储在显存中 | 减少显存占用,少量增加计算时间,引入少量训练误差 | 仅推荐在text_encoder、vae等非训练模块上启用,也可在 LoRA 训练时对dit启用 | 文档 | | 两阶段拆分训练 | 较为复杂,请参考文档 | 将训练过程拆分为两个阶段:第一阶段进行无梯度计算并将中间结果保存至硬盘,第二阶段计算梯度并更新模型参数 | 减少显存占用,增加计算速度,占用额外硬盘空间 | 部分模型的两阶段训练功能未验证,请谨慎使用 | 文档 | | CPU Offload | 通过--enable_model_cpu_offload启用 | 在训练时将模型保存在内存中,逐层移至显存中进行前向和后向传播 | 减少显存占用,增加计算时间,增加内存占用 | 仅推荐在单 GPU 且显存极为有限的设备上启用 | 文档 |
补充几点源码级细节,帮助你做出更精准的选型:
- FP8 加载:
parse_vram_config(training_module.py)中,FP8 模式将参数的 offload/onload/preparing 精度统一设为torch.float8_e4m3fn,计算精度为bfloat16——即 FP8 只负责"存",计算时仍转换回更高精度,因此仅适合冻结模块。 - CPU Offload 与多卡兼容性:启用
--enable_model_cpu_offload时,runner.py 走OffloadTrainingManager分支(实现见 diffsynth/core/offload_training/manager.py),模型权重常驻内存、逐层进出 GPU,并通过--cpu_offload_split_threshold控制大模块的递归拆分粒度;该模式设计上面向单 GPU 场景。 - 量化与 DDP 的兼容:使用
--quant_options且配合多卡 DDP 训练时,runner.py 的exclude_quantized_params_from_ddp_sync会将由 tensor subclass 承载的冻结量化权重排除在 DDP 广播之外,避免广播桶无法展平导致的报错。 - DeepSpeed 与 gradient checkpointing:
initialize_deepspeed_gradient_checkpointing(runner.py)会读取 DeepSpeed 配置中的activation_checkpointing段并调用deepspeed.checkpointing.configure完成初始化,确保两种机制可以协同工作。
训练后的验证与进阶
全量训练完成后,运行对应validate_full目录下的脚本加载模型验证效果;LoRA 训练完成后则使用validate_lora目录下的脚本(如 examples/z_image/model_training/validate_full/Z-Image.py),验证脚本的加载方式与模型推理保持一致。若需更系统的训练方案,可进一步参考仓库中的 训练相关文档、LoRA 差分训练、直接蒸馏 Direct Distill 与 FP8 精度 等专题。
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考