DiffSynth-Studio 模型接入指南:从模型结构代码到显存管理的五步集成实战
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
本篇技术指南围绕DiffSynth-Studio官方开发者文档《接入模型结构》展开,系统讲解如何将任意扩散模型(DiT、Text Encoder、VAE、ControlNet 等)接入该框架,供Pipeline等模块统一调用。读者完成阅读后,将掌握五步完整的接入流程:编写模型结构代码、编写 state dict 格式转换器、注册模型 Config、用ModelPool验证加载、并为模型接入显存管理方案,从而让新模型能够直接复用框架的加载、量化、LoRA 与显存管理能力。
一、接入总览:一切模型都汇聚到 ModelPool
DiffSynth-Studio采用"模型结构代码统一管理 + 配置驱动加载"的架构:所有模型结构的实现统一放在diffsynth/models目录下,每个.py文件实现一个模型结构;所有模型文件则通过 diffsynth/models/model_loader.py 中的ModelPool类来加载。接入一个新模型时,核心工作就是为它补齐"结构代码 + 格式转换 + 配置注册"三件套。
从 model_loader.py 中ModelPool.auto_load_model的实现可以看到加载的完整链路:
def auto_load_model(self, path, vram_config=None, vram_limit=None, clear_parameters=False, state_dict=None, quantize=None): print(f"Loading models from: {json.dumps(path, indent=4)}") if vram_config is None: vram_config = self.default_vram_config() model_hash = hash_model_file(path) loaded = False for config in MODEL_CONFIGS: if config["model_hash"] == model_hash: model = self.load_model_file(config, path, vram_config, vram_limit=vram_limit, state_dict=state_dict, quantize=quantize) ...也就是说:框架先对模型文件计算哈希,再与MODEL_CONFIGS中逐条注册的model_hash比对,命中后按该配置指定的model_class、state_dict_converter、extra_kwargs完成实例化与权重加载。这条"哈希 → 匹配 → 加载"的链路就是下文所有接入步骤的落脚点。
二、Step 1:集成模型结构代码
2.1 代码放置位置与目录结构
所有模型结构实现统一放在diffsynth/models目录下,接入新模型时在此路径下新建.py文件。以 Qwen-Image 系列为例,目录结构如下:
diffsynth/models/ ├── general_modules.py ├── model_loader.py ├── qwen_image_controlnet.py ├── qwen_image_dit.py ├── qwen_image_text_encoder.py ├── qwen_image_vae.py └── ...2.2 方式一:原生 PyTorch 代码(推荐)
绝大多数情况下,建议以原生 PyTorch 代码形式集成模型,让模型结构类直接继承torch.nn.Module:
import torch class NewDiffSynthModel(torch.nn.Module): def __init__(self, dim=1024): super().__init__() self.linear = torch.nn.Linear(dim, dim) self.activation = torch.nn.Sigmoid() def forward(self, x): x = self.linear(x) x = self.activation(x) return x重要原则:删掉额外依赖。如果模型实现包含额外的第三方包依赖,强烈建议将其删除,否则会给整个项目带来沉重的包依赖问题。仓库中 Qwen-Image 的 Blockwise ControlNet 就是以这种方式集成的,代码非常轻量,可参考 diffsynth/models/qwen_image_controlnet.py。从该文件可以看到,整个 ControlNet 仅依赖torch与框架内部的RMSNorm(来自 general_modules.py),并通过BlockWiseControlBlock与QwenImageBlockWiseControlNet两个类组织,且提供了零初始化权重的init_weight方法与按块前向的blockwise_forward接口——这正是为了在推理时逐块注入 DiT 而设计,体现了"轻量、可被框架编排"的集成风格。
2.3 方式二:包装 Huggingface Library 风格模型
如果模型已被 Huggingface 生态(transformers、diffusers等)集成,可以更简单地包装。这类模型在 Huggingface Library 中的加载方式通常为:
from transformers import XXX_Model model = XXX_Model.from_pretrained("path_to_your_model")但DiffSynth-Studio不支持通过from_pretrained加载模型,因为这与显存管理等功能存在冲突(from_pretrained会在内部自行处理权重加载与设备放置,无法被框架的加载管线接管)。因此需要将模型结构改写成以下格式:
import torch class DiffSynth_XXX_Model(torch.nn.Module): def __init__(self): super().__init__() from transformers import XXX_Config, XXX_Model config = XXX_Config(**{ "architectures": ["XXX_Model"], "other_configs": "Please copy and paste the other configs here.", }) self.model = XXX_Model(config) def forward(self, x): outputs = self.model(x) return outputs其中XXX_Config为模型对应的 Config 类。例如Qwen2_5_VLModel对应的 Config 类是Qwen2_5_VLConfig,可通过查阅其源代码找到。Config 内部的参数通常可以在模型库的config.json文件中找到,DiffSynth-Studio不会读取config.json文件,因此需要把其中的内容手动复制粘贴到代码中。
仓库中的典型范例是 diffsynth/models/qwen_image_text_encoder.py:该类在__init__中构造Qwen2_5_VLConfig,将architectures、hidden_size、intermediate_size、num_hidden_layers、rope_scaling、vision_config等完整配置逐项写入,再以Qwen2_5_VLModel(config)实例化并包装为torch.nn.Module子类。
注意事项:在少数情况下,transformers与diffusers的版本更新会导致部分模型无法导入。因此,如果可能的话,仍建议优先采用 2.2 节的原生 PyTorch 集成方式,以彻底解耦对 Huggingface 生态的版本依赖。
三、Step 2:模型文件格式转换(state dict converter)
3.1 为什么需要转换
开源社区中开发者提供的模型文件格式多种多样,有时需要对模型文件格式进行转换,以形成格式正确的 state dict。常见于以下几种情况:
- 模型文件由不同代码库构建:例如 Wan2.1-T2V-1.3B 官方仓库与 Diffusers 重封装仓库的权重组织结构不一致;
- 模型在接入中做了修改:例如 Qwen-Image 的 Text Encoder 在 qwen_image_text_encoder.py 中增加了
model.前缀,导致键名需要重映射; - 模型文件包含多个模型:例如 Wan2.1-VACE-14B 的 VACE Adapter 与基础 DiT 模型混合存储在同一组模型文件中,加载时需要按需分离。
3.2 转换逻辑放在哪里
DiffSynth-Studio专门增加了diffsynth/utils/state_dict_converters模块,用于在模型加载过程中进行文件格式转换。之所以采用"加载时转换"而非"重新封装模型文件",是出于对模型原作者意愿的尊重:如果对模型文件进行重新封装(例如 Qwen-Image 的 ComfyUI 重封装仓库),虽然调用更方便,但流量(模型页面浏览量、下载量等)会被引向他处,模型原作者也会失去删除模型的权力。因此在框架内部完成转换,可以保证始终直接使用原作者发布的模型文件。
3.3 一个 10 行代码的转换器范例
转换逻辑本身非常简单,以 Qwen-Image 的 Text Encoder 为例,对应实现见 diffsynth/utils/state_dict_converters/qwen_image_text_encoder.py,只需 10 行代码:
def QwenImageTextEncoderStateDictConverter(state_dict): state_dict_ = {} for k in state_dict: v = state_dict[k] if k.startswith("visual."): k = "model." + k elif k.startswith("model."): k = k.replace("model.", "model.language_model.") state_dict_[k] = v return state_dict_这段逻辑把官方权重中visual.*前缀的键改写为model.visual.*,把model.*改写为model.language_model.*,从而与 qwen_image_text_encoder.py 中包装的Qwen2_5_VLModel的参数结构对齐。
转换器会在加载管线中被调用:在 diffsynth/core/loader/model.py 中,加载后的 state dict 会先经过state_dict_converter(state_dict)转换,再执行model.load_state_dict(state_dict, assign=True),随后统一调用model.to(dtype=..., device=...)。从源码注释可以看出,转换器的存在正是为了兼容各种复杂格式的权重文件。
四、Step 3:编写模型 Config
4.1 字段说明
模型 Config 位于 diffsynth/configs/model_configs.py,用于识别模型类型并完成加载。需要填写的字段如下:
| 字段 | 是否必填 | 说明 |
|---|---|---|
model_hash | 必填 | 模型文件哈希值,通过hash_model_file函数获取;该哈希仅与模型文件中 state dict 的 keys 和张量 shape 有关,与文件中其他信息(如元数据、实际数值)无关 |
model_name | 必填 | 模型名称,供Pipeline识别所需模型;若不同结构的模型在Pipeline中发挥相同作用,可使用相同model_name;接入新模型时只需保证model_name与现有功能模型不同即可 |
model_class | 必填 | 模型结构导入路径,指向 Step 1 中实现的模型结构类,例如diffsynth.models.qwen_image_text_encoder.QwenImageTextEncoder |
state_dict_converter | 可选 | 模型文件格式转换逻辑的导入路径,例如diffsynth.utils.state_dict_converters.qwen_image_text_encoder.QwenImageTextEncoderStateDictConverter |
extra_kwargs | 可选 | 模型初始化时需传入的额外参数,例如 Canny 与 Inpaint 两种 Blockwise ControlNet 共用QwenImageBlockWiseControlNet结构,但 Inpaint 版本还需additional_in_dim=4,这部分差异就通过extra_kwargs表达 |
以 Qwen-Image 系列在 model_configs.py 中的真实注册为例:
{ "model_hash": "8004730443f55db63092006dd9f7110e", "model_name": "qwen_image_text_encoder", "model_class": "diffsynth.models.qwen_image_text_encoder.QwenImageTextEncoder", "state_dict_converter": "diffsynth.utils.state_dict_converters.qwen_image_text_encoder.QwenImageTextEncoderStateDictConverter", }, { "model_hash": "a9e54e480a628f0b956a688a81c33bab", "model_name": "qwen_image_blockwise_controlnet", "model_class": "diffsynth.models.qwen_image_controlnet.QwenImageBlockWiseControlNet", "extra_kwargs": {"additional_in_dim": 4}, },所有系列的 Config 在文件末尾聚合成MODEL_CONFIGS元组(见 model_configs.py),并由 diffsynth/configs/init.py 导出,供ModelPool在加载时逐条匹配。仓库中已有 Wan、FLUX、FLUX2、LTX-2、MiniMax、SDXL 等数十个模型系列,新模型的 Config 注册到对应系列(或新建系列)后即可生效。
4.2 关于model_hash的计算
model_hash通过hash_model_file函数获得,其实现位于 diffsynth/core/loader/file.py:函数读取模型文件中每个张量的键名与 shape,按字典序拼接为字符串后进行 MD5 哈希。因此,同一份权重文件无论内容数值如何变化,只要键名与形状不变,哈希就保持不变;而键名或形状的任何变化都会导致哈希改变——这正是"格式转换器 + 哈希"配合工作的基础:哈希用于识别文件,转换器用于适配结构。
4.3 加载机制全流程演示
以下代码来自官方文档,可以快速理解模型是如何通过上述配置信息完成加载的(其中skip_model_initialization用于跳过随机初始化以加速加载并节省显存):
from diffsynth.core import hash_model_file, load_state_dict, skip_model_initialization from diffsynth.models.qwen_image_text_encoder import QwenImageTextEncoder from diffsynth.utils.state_dict_converters.qwen_image_text_encoder import QwenImageTextEncoderStateDictConverter import torch model_hash = "8004730443f55db63092006dd9f7110e" model_name = "qwen_image_text_encoder" model_class = QwenImageTextEncoder state_dict_converter = QwenImageTextEncoderStateDictConverter extra_kwargs = {} model_path = [ "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", ] if hash_model_file(model_path) == model_hash: with skip_model_initialization(): model = model_class(**extra_kwargs) state_dict = load_state_dict(model_path, torch_dtype=torch.bfloat16, device="cuda") state_dict = state_dict_converter(state_dict) model.load_state_dict(state_dict, assign=True) print("Done!")Q:上述代码的逻辑看起来很简单,为什么
DiffSynth-Studio中的这部分代码极为复杂?A:因为框架提供了激进的显存管理功能,它与模型加载逻辑深度耦合,导致框架结构复杂;但暴露给开发者的接口已经尽可能简化,即本文描述的这一套 Config 字段。
注意,model_configs.py中的model_hash并不是唯一的:同一模型文件中可能包含多个模型(例如 Wan2.1-VACE 中 Adapter 与 DiT 共存)。对于这种情况,请使用多个模型 Config 分别加载每个模型,并为每个 Config 编写相应的state_dict_converter来分离每个模型所需的参数。仓库中wan_series就是这一做法的典型:同一个model_hash(如7a513e1f257a861512b1afd387a8ecd9)同时注册了wan_video_dit与wan_video_vace两条 Config,分别由不同的state_dict_converter从同一组文件中抽取各自参数。
五、Step 4:检验模型能否被识别和加载
模型接入之后,可通过以下代码验证模型能否被正确识别和加载。以下代码会试图将模型加载到内存中:
from diffsynth.models.model_loader import ModelPool model_pool = ModelPool() model_pool.auto_load_model( [ "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", ], )如果模型能够被识别和加载,则终端会输出以下内容:
Loading models from: [ "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" ] Loaded model: { "model_name": "qwen_image_text_encoder", "model_class": "diffsynth.models.qwen_image_text_encoder.QwenImageTextEncoder", "extra_kwargs": null }从源码看,这段输出由 model_loader.py 中的打印语句直接产生。需要特别注意:如果哈希无法命中任何 Config,auto_load_model会抛出ValueError: Cannot detect the model type. File: ... Model hash: ...异常,并直接打印计算出的model_hash——此时只需将该哈希填入 Config 即可。
加载成功后,Pipeline通过ModelPool.fetch_model(model_name, index)按model_name取用模型(见 model_loader.py):同一model_name下可能加载了多个模型实例,index参数用于选择具体取哪个(例如多阶段流水线中的第一个或前几个模型)。
六、Step 5:编写模型显存管理方案
DiffSynth-Studio支持复杂的显存管理机制。接入模型后,如需让模型在低显存环境下运行,请参阅文档 docs/zh/Developer_Guide/Enabling_VRAM_management.md(英文版见 docs/en/Developer_Guide/Enabling_VRAM_management.md)。
从源码层面看,显存管理与模型加载是深度耦合的:在 model_loader.py 中,fetch_module_map会根据显存配置决定是否为模型包装AutoWrappedModule或使用 vram_management_module_maps.py 中注册的自定义模块映射;在 core/loader/model.py 中,load_model会根据module_map调用enable_vram_management完成层级的按需卸载与重载。因此,接入模型时通常无需改动加载逻辑,只需在显存管理模块映射表中注册结构对应的包装类,即可复用框架的 offload 能力。
七、接入流程自查清单
完成上述五个步骤后,可用以下清单快速自查:
- 结构代码:新模型的
.py文件已放入diffsynth/models/,类继承torch.nn.Module,不引入多余的外部依赖; - 格式转换:如权重键名与结构类参数名不一致,已在
diffsynth/utils/state_dict_converters/中编写转换函数; - Config 注册:已在 model_configs.py 中注册
model_hash、model_name、model_class(必要时含state_dict_converter与extra_kwargs),且model_hash与hash_model_file计算结果一致; - 加载验证:
ModelPool().auto_load_model(...)输出Loaded model且无ValueError; - 显存方案:如需低显存运行,已按 Enabling_VRAM_management.md 完成模块映射注册;
- Pipeline 集成:确认
Pipeline中能通过model_name正确获取到该模型,并完成端到端推理验证。
完成以上步骤后,新模型即可与 Qwen-Image、Wan、FLUX 等既有模型一样,在Pipeline中被from_pretrained统一调度,并完整享受框架提供的量化加载、LoRA 注入与显存管理等基础设施能力。
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考