DiffSynth-Studio Template 模型训练实战:从 model.py 组件设计到 FLUX.2 全流程训练与上传
DiffSynth-Studio Template 模型训练实战从 model.py 组件设计到 FLUX.2 全流程训练与上传【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-StudioTemplate模板模型是 DiffSynth-Studio Diffusion Templates 框架中用于实现可控生成的核心载体它以轻量模型的形式接管 Diffusion Pipeline 的部分输入参数如 KV-Cache、LoRA从而为文生图、图像编辑等任务注入额外的控制能力。本文以官方文档 Template 模型训练 为骨架结合框架源码diffsynth/diffusion/template.py、diffsynth/diffusion/training_module.py与真实训练脚本系统讲解如何构建一个新的 Template 模型、如何基于 FLUX.2 训练它、如何在低显存设备上通过两阶段拆分训练压缩开销以及最终如何将模型打包上传到魔搭社区。读完本文你将能够独立完成一个可发布的 Template 模型的从 0 到 1 全流程。一、当前训练支持概况DiffSynth-Studio 目前已为 black-forest-labs/FLUX.2-klein-base-4B 提供了全面的 Templates 训练支持更多模型的适配正在持续推进中。围绕这一基础模型官方在魔搭社区发布了多款预训练 Template 模型包括DiffSynth-Studio/Template-KleinBase4B-Aesthetic美学DiffSynth-Studio/Template-KleinBase4B-Brightness亮度DiffSynth-Studio/Template-KleinBase4B-Age年龄DiffSynth-Studio/Template-KleinBase4B-ControlNet结构控制DiffSynth-Studio/Template-KleinBase4B-Edit图像编辑DiffSynth-Studio/Template-KleinBase4B-Inpaint重绘DiffSynth-Studio/Template-KleinBase4B-PandaMeme、Sharpness、SoftRGB、Upscaler、ContentRef等上述模型在 FLUX.2 文档 的“模型总览”表格中均配有对应的推理示例examples/flux2/model_inference/与model_inference_low_vram/与全量训练脚本examples/flux2/model_training/full/及训练后验证脚本examples/flux2/model_training/validate_full/是学习与复现 Template 训练的最佳参考。二、基于预训练 Template 模型继续训练如果你希望基于官方已经预训练好的 Template 模型做继续训练例如在自己的数据集上微调亮度控制能力无需从头构建模型文件打开 FLUX.2 文档 中的“模型总览”表格在表格中找到目标 Template 模型例如DiffSynth-Studio/Template-KleinBase4B-Brightness点击对应的“全量训练”代码链接即可得到官方编写的训练脚本修改--dataset_base_path、--dataset_metadata_path、--output_path等参数后即可开始训练。在继续训练场景下--template_model_id_or_path直接填写官方模型 ID末尾带:框架会自动从魔搭下载权重并继续训练这一点在“低显存训练”一节的示例脚本中有完整体现。三、构建新的 Template 模型model.py 组件格式一个 Template 模型与一个模型库或一个本地文件夹绑定模型库中包含代码文件model.py作为唯一入口。一个完整的model.py模板如下import torch class CustomizedTemplateModel(torch.nn.Module): def __init__(self): super().__init__() torch.no_grad() def process_inputs(self, xxx, **kwargs): yyy xxx return {yyy: yyy} def forward(self, yyy, **kwargs): zzz yyy return {zzz: zzz} class DataProcessor: def __call__(self, www, **kwargs): xxx www return {xxx: xxx} TEMPLATE_MODEL CustomizedTemplateModel TEMPLATE_MODEL_PATH model.safetensors TEMPLATE_DATA_PROCESSOR DataProcessor文件末尾的三个全局变量是框架约定的三个关键入口全局变量作用是否必填TEMPLATE_MODELTemplate 模型的代码实现类继承torch.nn.Module必填TEMPLATE_MODEL_PATH预训练权重的相对路径字符串或字符串列表可选TEMPLATE_DATA_PROCESSOR训练阶段的数据预处理算子将数据集样本计算为process_inputs的输入可选从源码看load_template_model 会通过importlib动态加载指定目录下的model.py若TEMPLATE_MODEL_PATH不为None则通过load_model加载预训练权重否则直接实例化TEMPLATE_MODEL()并进行随机初始化同时迁移到指定 dtype 与 device。此外加载完成后会调用check_template_model_formattemplate.py做格式校验——process_inputs与forward必须存在且必须包含**kwargs否则直接抛出NotImplementedError。3.1 TEMPLATE_MODELprocess_inputs 与 forward 的职责拆分TEMPLATE_MODEL需实现两个函数二者共同构成完整的 Template 模型推理过程process_inputs必须带有torch.no_grad()装饰器进行不包含梯度的计算。典型用途是解析标量参数如亮度数值scale、调用预训练编码器提取特征、调用基础模型 Pipeline 的文本编码器等forward包含训练模型所需的全部梯度计算过程其输入与process_inputs的输出相同即process_inputs返回的字典会被解包后作为forward的输入。将推理过程拆分为两个阶段的设计是为了在训练中更容易适配两阶段拆分训练process_inputs的计算结果与模型参数无关或仅依赖冻结模型可以被缓存到磁盘并在多个 epoch 间复用从而大幅降低显存占用、提升训练速度。除**kwargs外框架预留了以下特殊参数pipe如需在process_inputs与forward中和基础模型 Pipeline 交互例如调用 Pipeline 中的文本编码器在输入参数中增加字段pipe即可use_gradient_checkpointing/use_gradient_checkpointing_offload如需在训练中启用 Gradient Checkpointing在forward的输入参数中增加这两个字段model_id当多个 Template 模型同时存在时框架通过model_id区分不同模型产生的 Template Inputs因此不要在process_inputs与forward的输入参数中使用该字段名。从框架实现看call_single_side 会依次执行model.process_inputs(pipepipe, **inputs)与model.forward(pipepipe, **cache)前者返回的字典会作为后者的关键字参数这意味着process_inputs返回的字段名必须与forward的形参名严格对应。3.2 TEMPLATE_MODEL_PATH权重加载的三种写法TEMPLATE_MODEL_PATH是模型预训练权重文件相对model.py所在目录的相对路径支持三种写法写法一单个权重文件TEMPLATE_MODEL_PATH model.safetensors写法二多个分片文件列表TEMPLATE_MODEL_PATH [ model-00001-of-00003.safetensors, model-00002-of-00003.safetensors, model-00003-of-00003.safetensors, ]写法三随机初始化TEMPLATE_MODEL_PATH None当模型尚未训练、需要随机初始化参数或该组件本身不包含可训练参数时将其设置为None或不设置即可。仓库自带的亮度控制示例 examples/flux2/model_training/scripts/brightness/model.py 正是采用这一写法并在注释中提示“训练完成后请修改此参数”TEMPLATE_MODEL_PATH None # You should modify this parameter after training。3.3 TEMPLATE_DATA_PROCESSOR训练数据集的输入计算训练 Template 模型需要构建包含template_inputs字段的数据集。需要注意metadata.json中的template_inputs并不是直接输入给 Template 模型process_inputs的参数而是提供给TEMPLATE_DATA_PROCESSOR的输入参数由TEMPLATE_DATA_PROCESSOR计算出真正输入给process_inputs的参数。示例一元数据中直接携带标量参数以亮度控制模型DiffSynth-Studio/Template-KleinBase4B-Brightness为例其输入参数是scale图像亮度数值可以直接写在metadata.json中此时TEMPLATE_DATA_PROCESSOR只需透传[ { image: images/image_1.jpg, prompt: a cat, template_inputs: {scale: 0.2} }, { image: images/image_2.jpg, prompt: a dog, template_inputs: {scale: 0.6} } ]class DataProcessor: def __call__(self, scale, **kwargs): return {scale: scale} TEMPLATE_DATA_PROCESSOR DataProcessor示例二元数据中填写图像路径训练时动态计算也可以在metadata.json中填写图像路径在训练过程中直接计算scale[ { image: images/image_1.jpg, prompt: a cat, template_inputs: {image: /path/to/your/dataset/images/image_1.jpg} }, { image: images/image_2.jpg, prompt: a dog, template_inputs: {image: /path/to/your/dataset/images/image_2.jpg} } ]class DataProcessor: def __call__(self, image, **kwargs): image Image.open(image) image np.array(image) return {scale: image.astype(np.float32).mean() / 255} TEMPLATE_DATA_PROCESSOR DataProcessor仓库中的真实实现与此完全对应DataAnnotator.__call__打开图像并计算scale mean / 255随后ValueFormatModel.process_inputs将标量转换为张量forward再为 DiT 的每个 block 生成 KV-Cache见 brightness/model.py。3.4 推理与训练的数据流差异推理时Template Input 先后经过TEMPLATE_MODEL的process_inputs和forward得到 Template Cache训练时Template Input 不再来自用户输入而是从数据集中获取先经TEMPLATE_DATA_PROCESSOR计算再进入TEMPLATE_MODEL关于 Template Input、Template Model、Template Cache、Template Pipeline 的完整架构与“模型能力媒介”KV-Cache、LoRA、Residual 等的设计动机可参考 Diffusion Templates 架构详解。四、训练 Template 模型4.1 “可训练”的充分条件Template 模型“可训练”的充分条件是Template Cache 中的变量计算与基础模型 Pipeline 完全解耦。这些变量在推理过程中输入给基础模型 Pipeline 后不会参与任何 Pipeline Unit 的计算而是直达model_fn。只有满足这一条件训练时梯度才能安全地流经 Template 模型而不会污染基础模型Cache 也可以在训练阶段被高效复用。从框架源码看训练时 load_training_template_model 会向 Pipeline 头部追加两个单元GeneralUnit_TemplateProcessInputs调用TEMPLATE_DATA_PROCESSOR与GeneralUnit_TemplateForward执行process_inputsforward随后才执行基础模型的 Pipeline Units——这与“Template 计算先于基础模型、其结果作为 Pipeline 输入参数”的设计是一致的。4.2 关键训练参数以基础模型black-forest-labs/FLUX.2-klein-base-4B为例训练脚本examples/flux2/model_training/train.py中与 Template 训练直接相关的参数如下参数含义与取值--extra_inputs额外输入。训练文生图模型的 Template 时填template_inputs训练图像编辑模型的 Template 时需填edit_image,template_inputs--template_model_id_or_pathTemplate 模型的魔搭模型 ID 或本地路径。框架优先匹配本地路径本地不存在则从魔搭下载填写模型 ID 时以:结尾例如DiffSynth-Studio/Template-KleinBase4B-Brightness:--remove_prefix_in_ckpt保存模型文件时移除的 state dict 变量名前缀填pipe.template_model.即可--trainable_models可训练模型填template_model表示训练整个 Template 模型若只需训练其中某个组件则填template_model.xxx,template_model.yyy逗号分隔--extra_inputs的解析实现在 parse_extra_inputs除 ControlNet 系列字段会被特殊聚合外其余字段如template_inputs、edit_image会直接写入inputs_shared作为 Pipeline 的共享输入训练模块的get_pipeline_inputs会将其与prompt、input_image、height、width等常规输入合并后送入模型见 train.py。4.3 完整样例训练脚本以下脚本会自动下载一个样例数据集随机初始化模型权重后开始训练亮度控制模型完整脚本亦可参考 examples/flux2/model_training/full/Template-KleinBase4B-Brightness.shmodelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include flux2/Template-KleinBase4B-Brightness/* --local_dir ./data/diffsynth_example_dataset accelerate launch examples/flux2/model_training/train.py \ --dataset_base_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness \ --dataset_metadata_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness/metadata.jsonl \ --extra_inputs template_inputs \ --max_pixels 1048576 \ --dataset_repeat 50 \ --model_id_with_origin_paths black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors \ --template_model_id_or_path examples/flux2/model_training/scripts/brightness \ --tokenizer_path black-forest-labs/FLUX.2-klein-4B:tokenizer/ \ --learning_rate 1e-4 \ --num_epochs 2 \ --remove_prefix_in_ckpt pipe.template_model. \ --output_path ./models/train/Template-KleinBase4B-Brightness_example \ --trainable_models template_model \ --use_gradient_checkpointing \ --find_unused_parameters脚本要点说明--template_model_id_or_path指向本地目录examples/flux2/model_training/scripts/brightness内含model.py因为该目录下TEMPLATE_MODEL_PATH None所以框架会随机初始化模型权重开始训练--model_id_with_origin_paths使用“模型 ID : 文件通配符”的格式分别加载 FLUX.2 的文本编码器text encoder、TransformerDiT与 VAE--remove_prefix_in_ckpt pipe.template_model.保证保存的权重文件中不包含pipe.template_model.前缀便于后续直接以 Template 模型权重形式加载--find_unused_parameters用于规避 DDP 训练中未使用参数导致的报错。五、与基础模型 Pipeline 组件交互Diffusion Template 框架允许 Template 模型与基础模型 Pipeline 进行交互。例如你可能需要使用基础模型 Pipeline 中的 text encoder 对文本进行编码此时在process_inputs和forward中使用预留字段pipe即可import torch class CustomizedTemplateModel(torch.nn.Module): def __init__(self): super().__init__() self.xxx xxx() torch.no_grad() def process_inputs(self, text, pipe, **kwargs): input_ids pipe.tokenizer(text) text_emb pipe.text_encoder(text_emb) return {text_emb: text_emb} def forward(self, text_emb, pipe, **kwargs): kv_cache self.xxx(text_emb) return {kv_cache: kv_cache} TEMPLATE_MODEL CustomizedTemplateModel从 template.py 的调用链可以看到pipe由框架自动注入call_single_side调用model.process_inputs(pipepipe, **inputs)与model.forward(pipepipe, **cache)因此只要在方法形参中声明pipe即可拿到基础 Pipeline 实例访问其tokenizer、text_encoder、dit、vae等组件。注意pipe.text_encoder这类基础模型组件默认是冻结的训练时通过switch_pipe_to_training_mode中的freeze_except冻结所有非trainable_models参数见 training_module.py因此放在torch.no_grad()的process_inputs中调用是安全且高效的做法。六、使用非训练的模型组件部分参数冻结在设计 Template 模型时如果希望使用预训练模型作为特征提取器、且不希望在训练过程中更新这部分参数例如一个来自外部库的 image encoder可以这样组织import torch class CustomizedTemplateModel(torch.nn.Module): def __init__(self): super().__init__() self.image_encoder XXXEncoder.from_pretrained(xxx) self.mlp MLP() torch.no_grad() def process_inputs(self, image, **kwargs): emb self.image_encoder(image) return {emb: emb} def forward(self, emb, **kwargs): kv_cache self.mlp(emb) return {kv_cache: kv_cache} TEMPLATE_MODEL CustomizedTemplateModel此时需在训练命令中通过参数--trainable_models template_model.mlp设置为仅训练mlp部分image_encoder的参数将保持冻结不参与梯度更新。框架在保存检查点时只会写入可训练参数结合--remove_prefix_in_ckpt处理命名前缀因此包含冻结组件的模型在推理前需要将非训练参数重新打包进权重文件具体做法见第八节“上传 Template 模型”的打包代码。七、在低显存的设备上训练7.1 两阶段拆分训练框架支持将 Template 模型的训练拆分为两个阶段第一阶段进行无梯度计算process_inputs及 VAE 编码、文本编码等与去噪模型无关的前处理第二阶段进行梯度更新forward及 DiT 相关计算。这一机制的原理与算法细节参见两阶段拆分训练文档。以下是官方提供的两阶段样例脚本第一阶段数据预处理生成 Cachemodelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include flux2/Template-KleinBase4B-Brightness/* --local_dir ./data/diffsynth_example_dataset accelerate launch examples/flux2/model_training/train.py \ --dataset_base_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness \ --dataset_metadata_path data/diffsynth_example_dataset/flux2/Template-KleinBase4B-Brightness/metadata.jsonl \ --extra_inputs template_inputs \ --max_pixels 1048576 \ --dataset_repeat 1 \ --model_id_with_origin_paths black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors \ --template_model_id_or_path DiffSynth-Studio/Template-KleinBase4B-Brightness: \ --tokenizer_path black-forest-labs/FLUX.2-klein-4B:tokenizer/ \ --learning_rate 1e-4 \ --num_epochs 2 \ --remove_prefix_in_ckpt pipe.template_model. \ --output_path ./models/train/Template-KleinBase4B-Brightness_full_cache \ --trainable_models template_model \ --use_gradient_checkpointing \ --find_unused_parameters \ --task sft:data_process第二阶段训练读取 Cache 并更新梯度accelerate launch examples/flux2/model_training/train.py \ --dataset_base_path ./models/train/Template-KleinBase4B-Brightness_full_cache \ --extra_inputs template_inputs \ --max_pixels 1048576 \ --dataset_repeat 50 \ --model_id_with_origin_paths black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors \ --template_model_id_or_path DiffSynth-Studio/Template-KleinBase4B-Brightness: \ --tokenizer_path black-forest-labs/FLUX.2-klein-4B:tokenizer/ \ --learning_rate 1e-4 \ --num_epochs 2 \ --remove_prefix_in_ckpt pipe.template_model. \ --output_path ./models/train/Template-KleinBase4B-Brightness_full \ --trainable_models template_model \ --use_gradient_checkpointing \ --find_unused_parameters \ --task sft:train两阶段的关键差异第一阶段--dataset_repeat改为1避免重复计算相同的前处理--output_path指向 Cache 存储路径并追加--task sft:data_process此时--model_id_with_origin_paths只需包含前处理所需模型text encoder、VAE第二阶段--dataset_base_path改为第一阶段的输出目录删除--dataset_metadata_path--model_id_with_origin_paths只需包含训练所需模型transformer/DiT并追加--task sft:train。两阶段拆分训练可以降低显存需求、提高训练速度训练过程无损精度但需要较大的硬盘空间用于存储 Cache 文件。从源码看这一机制由DiffusionTrainingModule.split_pipeline_units实现training_module.py根据--task后缀分别保留“与模型无关的单元”第一阶段或“与模型相关的单元”第二阶段并由launch_data_process_task与launch_training_task分别驱动见 train.py。7.2 进一步降低显存FP8 精度如需进一步减少显存需求可开启 FP8 精度在两阶段训练中分别添加参数--fp8_models black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors--fp8_models black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors注意事项FP8 精度只能在非训练模型组件上启用即不需要梯度回传的组件FP8 量化存在少量误差需在精度与显存之间权衡取舍。八、上传 Template 模型完成训练后按照以下步骤可将 Template 模型上传到魔搭社区供更多人下载使用。Step 1在model.py中填入训练好的模型文件名TEMPLATE_MODEL_PATH model.safetensorsStep 2上传model.pymodelscope upload user_name/your_model_id /path/to/your/model.py model.py --token ms-xxx其中--token ms-xxx在 https://modelscope.cn/my/access/token 获取。Step 3确认并打包模型文件确认要上传的模型文件例如epoch-1.safetensors、step-2000.safetensors。注意DiffSynth-Studio 保存的模型文件中只包含可训练的参数。如果模型中包含非训练参数如第六节中的 image encoder则需要重新将非训练的模型参数打包后才能用于推理可通过以下代码完成from diffsynth.diffusion.template import load_template_model, load_state_dict from safetensors.torch import save_file import torch model load_template_model(path/to/your/template/model, torch_dtypetorch.bfloat16, devicecpu) state_dict load_state_dict(path/to/your/ckpt/epoch-1.safetensors, torch_dtypetorch.bfloat16, devicecpu) state_dict.update(model.state_dict()) save_file(state_dict, model.safetensors)这段代码的逻辑是先用load_template_model加载model.py定义的完整模型含非训练组件再加载训练产生的检查点仅含可训练参数将两者 state dict 合并后写入model.safetensors。Step 4上传模型文件modelscope upload user_name/your_model_id /path/to/your/model/epoch-1.safetensors model.safetensors --token ms-xxxStep 5验证模型推理效果上传完成后使用TemplatePipeline结合基础模型 Pipeline 进行验证from diffsynth.diffusion.template import TemplatePipeline from diffsynth.pipelines.flux2_image import Flux2ImagePipeline, ModelConfig import torch # Load base model pipe Flux2ImagePipeline.from_pretrained( torch_dtypetorch.bfloat16, devicecuda, model_configs[ ModelConfig(model_idblack-forest-labs/FLUX.2-klein-4B, origin_file_patterntext_encoder/*.safetensors), ModelConfig(model_idblack-forest-labs/FLUX.2-klein-base-4B, origin_file_patterntransformer/*.safetensors), ModelConfig(model_idblack-forest-labs/FLUX.2-klein-4B, origin_file_patternvae/diffusion_pytorch_model.safetensors), ], tokenizer_configModelConfig(model_idblack-forest-labs/FLUX.2-klein-4B, origin_file_patterntokenizer/), ) # Load Template model template_pipeline TemplatePipeline.from_pretrained( torch_dtypetorch.bfloat16, devicecuda, model_configs[ ModelConfig(model_iduser_name/your_model_id) ], ) # Generate an image image template_pipeline( pipe, prompta cat, seed0, cfg_scale4, height1024, width1024, template_inputs[{xxx}], ) image.save(image.png)验证时需注意TemplatePipeline通过ModelConfig(model_iduser_name/your_model_id)从魔搭加载 Template 模型包含model.py与权重文件框架内部会调用load_template_model完成动态加载template_inputs为列表列表元素为 Template 模型的输入字典例如亮度控制模型的{scale: 0.2}其字段与process_inputs的形参一一对应TemplatePipeline.__call__template.py会把 Template Cache 中与pipe.__call__输入参数同名的字段注入基础 Pipeline并支持通过negative_template_inputs传递负向 Template 输入多个 Template 模型的输出会由merge_template_cache统一合并KV-Cache 按序列维度拼接、LoRA 合并、text_embedding沿序列拼接重复字段以第一个模型为准。九、总结Template 模型训练的整体流程可以概括为一条主线编写model.py定义TEMPLATE_MODEL、TEMPLATE_MODEL_PATH、TEMPLATE_DATA_PROCESSOR→ 构建含template_inputs的数据集 → 通过examples/flux2/model_training/train.py启动训练单阶段或两阶段拆分→ 打包并上传魔搭 → 用TemplatePipeline验证。整个过程建立在process_inputs无梯度前处理与forward梯度计算职责分离的设计之上这也是它天然适配两阶段拆分训练、能够在低显存设备上高效运行的根本原因。对于想要深入了解框架内部实现的读者建议进一步阅读diffsynth/diffusion/template.pyTemplate 模型动态加载、格式校验、TemplatePipeline 调度与 Cache 合并逻辑diffsynth/diffusion/training_module.pyload_training_template_model、parse_extra_inputs、switch_pipe_to_training_mode等训练集成逻辑examples/flux2/model_training/train.pyFLUX.2 训练脚本入口与任务分发examples/flux2/model_training/scripts/brightness/model.py一个可直接参考的完整亮度控制 Template 模型实现两阶段拆分训练拆分训练的计算图算法与原理Diffusion Templates 架构详解Template 框架的整体架构与模型能力媒介设计。【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →