【Bug已解决】Adding a model 解决方案
2026/8/8 8:12:13 网站建设 项目流程

【Bug已解决】Adding a model 解决方案

一、现象长什么样

当你按 Transformers 的"添加新模型"流程,把一份第三方权重接进PreTrainedModel子类时,常遇到这一类失败:

# 现象 A:初始化权重全零 / 没被 Xavier 初始化 UserWarning: You are using the default `init_weights` which is not recommended. # 或者更糟:forward 输出全是同一个常数,因为权重没初始化 # 现象 B:load_weight 时 key 对不上 RuntimeError: Error(s) in loading state_dict for MyModel: Missing key(s) in state_dict: "model.layers.0.self_attn.q_proj.weight". Unexpected key(s) in state_dict: "transformer.h.0.attn.q_proj.weight". # 现象 C:save_pretrained 后再 from_pretrained 失败 KeyError: 'base_model_prefix' # 或保存出的 config.json 缺关键字段,导致二次加载崩溃 # 现象 D:CI 的 slow 测试直接报错 ValueError: Could not find dummy objects for model 'my_model'. # 官方 CI 要求提供 _dummy_xxx 输入,否则集成测试跑不起来

这些都不是"模型数学写错",而是集成骨架没搭全:权重初始化、key 命名、base_model_prefix、dummy 测试对象,缺一个就卡一个。

二、背景

把一个新模型(Adding a model)接进 transformers,需要的不只是modeling_xxx.py里的网络代码,还有一整套"契约":

  • _init_weights/init_weights:保证加载前权重被合理初始化。
  • base_model_prefix:告诉PreTrainedModel顶层容器叫什么(如"model""transformer"),所有get_input_embeddingstie_weights、state_dict key 前缀都依赖它。
  • 权重 key 命名:必须与 checkpoint 里的 key 完全一致(前缀、层级)。
  • _tied_weights_keys/tie_weights:词嵌入共享时声明。
  • dummy 测试对象:官方 CI 用_init_dummy_inputs等做无权重快速测试。

漏掉其中任何一项,都会在上文的现象里以不同形式炸出来。这些问题与 596(特定模型 D_Nikud 的 Auto 注册)是不同层面:596 是"Auto 体系认不认得你",601 是"模型自身骨架对不对"。

三、根因

把常见失败归到四类根因:

  1. 没实现_init_weights,或没调init_weights()PreTrainedModel.__init__默认不会自动初始化子模块权重(除非你覆盖了_init_weights并在__init__末尾调self.init_weights())。如果漏了,权重保持nn.Linear的默认(也可能被加载流程跳过)→ 全零或常数,forward 输出退化。

  2. base_model_prefix与权重 key 前缀不一致。 若你写base_model_prefix = "model",但 checkpoint 的 key 是transformer.h.0...,加载时所有 key 都"Missing/Unexpected"。反之保存时也会写出错误前缀,二次加载即KeyError

  3. tie_weights相关 key 未声明。 当lm_head.weightembed_tokens.weight共享,却没在_tied_weights_keys里声明,保存 checkpoint 时可能重复保存或漏保存,导致加载维度错乱。

  4. 缺 dummy 测试对象。 官方 CI 的models/__init__.py测试会尝试无权重构造模型并跑 dummy 输入。若没提供_init_dummy_inputs或对应的ModelTester,CI 直接ValueError: Could not find dummy objects

四、最小可运行复现

下面用纯 Python 模拟"base_model_prefix 与 key 前缀不一致导致 load 失败"的判定:

from typing import Dict, List class _PretendModel: def __init__(self, base_model_prefix: str): self.base_model_prefix = base_model_prefix def expected_keys(self, layer_keys: List[str]) -> List[str]: # 真实 transformers 会把 base_model_prefix 作为 state_dict 顶层前缀 return [f"{self.base_model_prefix}.{k}" for k in layer_keys] def load_state_dict(model, checkpoint_keys: List[str], model_keys: List[str]): missing = [k for k in model_keys if k not in checkpoint_keys] unexpected = [k for k in checkpoint_keys if k not in model_keys] return missing, unexpected # 情景 1:prefix 不一致 model = _PretendModel(base_model_prefix="model") layer_keys = ["layers.0.self_attn.q_proj.weight"] model_keys = model.expected_keys(layer_keys) # ["model.layers.0...q_proj.weight"] checkpoint_keys = ["transformer.h.0.attn.q_proj.weight"] # 错误前缀 missing, unexpected = load_state_dict(model, checkpoint_keys, model_keys) print("missing:", missing) print("unexpected:", unexpected) assert missing and unexpected, "复现失败:应当出现 key 不匹配" # 情景 2:prefix 一致则正常 model2 = _PretendModel(base_model_prefix="transformer") model_keys2 = model2.expected_keys(["h.0.attn.q_proj.weight"]) ckpt2 = ["transformer.h.0.attn.q_proj.weight"] m2, u2 = load_state_dict(model2, ckpt2, model_keys2) print("prefix 一致时 missing/unexpected:", m2, u2) # [] [] assert not m2 and not u2

运行后,情景 1 报missing/unexpected(key 前缀不符),情景 2 正常,正好对应现象 B 的Missing/Unexpected key(s)

五、解决方案(第一层:最小直接修复)

最直接:补齐骨架三个关键点——_init_weightsbase_model_prefix、key 命名对齐:

from transformers import PreTrainedModel, PretrainedConfig import torch.nn as nn import torch.nn.functional as F class MyConfig(PretrainedConfig): model_type = "my_model" def __init__(self, hidden_size=768, vocab_size=32000, **kwargs): super().__init__(**kwargs) self.hidden_size = hidden_size self.vocab_size = vocab_size class MyModel(PreTrainedModel): config_class = MyConfig base_model_prefix = "model" # 关键 1:与 checkpoint key 前缀一致 _tied_weights_keys = ["lm_head.weight", "model.embed_tokens.weight"] def __init__(self, config: MyConfig): super().__init__(config) self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) self.layers = nn.ModuleList([nn.Linear(config.hidden_size, config.hidden_size) for _ in range(2)]) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) # 关键 2:初始化权重 self.init_weights() # 会调用下面的 _init_weights def _init_weights(self, module): # 关键 3:明确初始化,避免全零 if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=0.02) if module.bias is not None: module.bias.data.zero_() def forward(self, input_ids): x = self.embed_tokens(input_ids) for layer in self.layers: x = F.relu(layer(x)) return self.lm_head(x) # 关键 4:保存/加载 key 前缀一致 model = MyModel(MyConfig()) model.save_pretrained("./my_ckpt") # 写出 model.* / lm_head.* m2 = MyModel.from_pretrained("./my_ckpt") # 前缀对齐,正常加载

第一层让用户加载/保存/二次加载都正常,且权重被正确初始化。

六、解决方案(第二层:结构性改进)

把"新模型骨架检查"做成ModelScaffoldValidator,在 CI 或加载前自动校验骨架完整性:

from dataclasses import dataclass from typing import List, Type @dataclass class ModelScaffoldValidator: """校验一个新模型类是否满足 transformers 集成骨架契约。""" required_attrs: List[str] = None def __post_init__(self): self.required_attrs = [ "base_model_prefix", "config_class", "_init_weights", "init_weights", ] def check(self, model_cls: Type) -> List[str]: problems: List[str] = [] for attr in self.required_attrs: if not hasattr(model_cls, attr): problems.append(f"缺少 {attr}") # base_model_prefix 必须是非空字符串 prefix = getattr(model_cls, "base_model_prefix", None) if not isinstance(prefix, str) or not prefix: problems.append("base_model_prefix 必须是非空字符串") # 必须有 dummy 测试入口(官方 CI 需要) if not hasattr(model_cls, "_init_dummy_inputs") and \ not hasattr(model_cls, "dummy_inputs"): problems.append("缺少 dummy 测试对象(CI slow 测试会失败)") return problems def assert_ready(self, model_cls: Type): probs = self.check(model_cls) if probs: raise RuntimeError("模型骨架不完整:\n" + "\n".join(probs)) # 使用 from my_modeling import MyModel ModelScaffoldValidator().assert_ready(MyModel) # 不抛异常即骨架完整

ModelScaffoldValidator把"集成骨架"从"靠经验记忆"变成"可自动检查",作者每次加模型先跑一遍,缺什么一目了然。

七、解决方案(第三层:断言 / CI 守护)

用 pytest 固化"骨架契约",任何一项缺失都红灯:

import pytest from transformers import PreTrainedModel from my_modeling import MyModel, MyConfig def test_has_base_model_prefix(): assert isinstance(MyModel.base_model_prefix, str) and MyModel.base_model_prefix, \ "base_model_prefix 缺失或为空,会导致 state_dict key 前缀错误" def test_init_weights_initializes(): m = MyModel(MyConfig(hidden_size=64, vocab_size=100)) w = m.layers[0].weight.data # 不应是全零(初始化生效) assert w.abs().sum() > 0, "_init_weights 未生效,权重可能全零" def test_save_load_roundtrip(): import tempfile, os m = MyModel(MyConfig(hidden_size=64, vocab_size=100)) d = tempfile.mkdtemp() m.save_pretrained(d) m2 = MyModel.from_pretrained(d) # key 前缀应一致,能正常加载 assert m2.base_model_prefix == m.base_model_prefix def test_has_dummy_inputs_for_ci(): assert hasattr(MyModel, "_init_dummy_inputs") or hasattr(MyModel, "dummy_inputs"), \ "缺少 dummy 测试对象,官方 CI 的 slow 测试会 ValueError"

CI 跑pytest tests/test_model_scaffold.py,以后只要有人加模型漏了base_model_prefix_init_weights,测试立刻拦截。

八、排查清单

当你"Adding a model"遇到加载/保存/CIT 失败,按顺序查:

  1. 权重全零或输出常数 → 检查是否实现_init_weights并在__init__self.init_weights()
  2. Missing/Unexpected key→ 比对base_model_prefix与 checkpoint 真实 key 前缀,必须逐字符一致。
  3. 二次加载KeyError: base_model_prefix→ 保存前确认base_model_prefix已设且config_class正确。
  4. 共享 embedding 维度错乱 → 在_tied_weights_keys声明lm_head.weightembed_tokens.weight
  5. CI slow 测试Could not find dummy objects→ 补_init_dummy_inputsdummy_inputs

九、小结

"Adding a model" 卡住的往往不是网络数学,而是集成骨架契约_init_weights初始化、base_model_prefix与 key 前缀一致、_tied_weights_keys共享声明、dummy 测试对象。这四样缺一个就以一种具体现象炸出来。

  • 第一层:补齐_init_weights+self.init_weights()base_model_prefix对齐、key 命名一致,立即能保存/加载/二次加载。
  • 第二层:用ModelScaffoldValidator自动校验骨架完整性,作者不再靠记忆。
  • 第三层:pytest 断言"prefix 非空、权重已初始化、save/load 往返、有 dummy 对象",防止回归。

记住:加模型先搭骨架,再填数学;骨架四件套(_init_weights/base_model_prefix/_tied_weights_keys/ dummy)齐了,集成基本不会翻车。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询