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

📅 2026/8/8 8:12:14
【Bug已解决】Adding a model 解决方案
【Bug已解决】Adding a model 解决方案一、现象长什么样当你按 Transformers 的添加新模型流程把一份第三方权重接进PreTrainedModel子类时常遇到这一类失败# 现象 A初始化权重全零 / 没被 Xavier 初始化 UserWarning: You are using the default init_weights which is not recommended. # 或者更糟forward 输出全是同一个常数因为权重没初始化 # 现象 Bload_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. # 现象 Csave_pretrained 后再 from_pretrained 失败 KeyError: base_model_prefix # 或保存出的 config.json 缺关键字段导致二次加载崩溃 # 现象 DCI 的 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_embeddings、tie_weights、state_dict key 前缀都依赖它。权重 key 命名必须与 checkpoint 里的 key 完全一致前缀、层级。_tied_weights_keys/tie_weights词嵌入共享时声明。dummy 测试对象官方 CI 用_init_dummy_inputs等做无权重快速测试。漏掉其中任何一项都会在上文的现象里以不同形式炸出来。这些问题与 596特定模型 D_Nikud 的 Auto 注册是不同层面596 是Auto 体系认不认得你601 是模型自身骨架对不对。三、根因把常见失败归到四类根因没实现_init_weights或没调init_weights()。PreTrainedModel.__init__默认不会自动初始化子模块权重除非你覆盖了_init_weights并在__init__末尾调self.init_weights()。如果漏了权重保持nn.Linear的默认也可能被加载流程跳过→ 全零或常数forward 输出退化。base_model_prefix与权重 key 前缀不一致。 若你写base_model_prefix model但 checkpoint 的 key 是transformer.h.0...加载时所有 key 都Missing/Unexpected。反之保存时也会写出错误前缀二次加载即KeyError。tie_weights相关 key 未声明。 当lm_head.weight与embed_tokens.weight共享却没在_tied_weights_keys里声明保存 checkpoint 时可能重复保存或漏保存导致加载维度错乱。缺 dummy 测试对象。 官方 CI 的models/__init__.py测试会尝试无权重构造模型并跑 dummy 输入。若没提供_init_dummy_inputs或对应的ModelTesterCI 直接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 # 情景 1prefix 不一致 model _PretendModel(base_model_prefixmodel) 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 不匹配 # 情景 2prefix 一致则正常 model2 _PretendModel(base_model_prefixtransformer) 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/unexpectedkey 前缀不符情景 2 正常正好对应现象 B 的Missing/Unexpected key(s)。五、解决方案第一层最小直接修复最直接补齐骨架三个关键点——_init_weights、base_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_size768, vocab_size32000, **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, biasFalse) # 关键 2初始化权重 self.init_weights() # 会调用下面的 _init_weights def _init_weights(self, module): # 关键 3明确初始化避免全零 if isinstance(module, nn.Linear): module.weight.data.normal_(mean0.0, std0.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_size64, vocab_size100)) 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_size64, vocab_size100)) 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 测试会 ValueErrorCI 跑pytest tests/test_model_scaffold.py以后只要有人加模型漏了base_model_prefix或_init_weights测试立刻拦截。八、排查清单当你Adding a model遇到加载/保存/CIT 失败按顺序查权重全零或输出常数 → 检查是否实现_init_weights并在__init__调self.init_weights()。Missing/Unexpected key→ 比对base_model_prefix与 checkpoint 真实 key 前缀必须逐字符一致。二次加载KeyError: base_model_prefix→ 保存前确认base_model_prefix已设且config_class正确。共享 embedding 维度错乱 → 在_tied_weights_keys声明lm_head.weight与embed_tokens.weight。CI slow 测试Could not find dummy objects→ 补_init_dummy_inputs或dummy_inputs。九、小结Adding a model 卡住的往往不是网络数学而是集成骨架契约_init_weights初始化、base_model_prefix与 key 前缀一致、_tied_weights_keys共享声明、dummy 测试对象。这四样缺一个就以一种具体现象炸出来。第一层补齐_init_weightsself.init_weights()、base_model_prefix对齐、key 命名一致立即能保存/加载/二次加载。第二层用ModelScaffoldValidator自动校验骨架完整性作者不再靠记忆。第三层pytest 断言prefix 非空、权重已初始化、save/load 往返、有 dummy 对象防止回归。记住加模型先搭骨架再填数学骨架四件套_init_weights/base_model_prefix/_tied_weights_keys/ dummy齐了集成基本不会翻车。