如何将预训练PyTorch模型封装为torch.nn.Module并转换为TorchScript
问题背景
手里只有预训练模型文件model.pth,想自定义torch.nn.Module类来封装模型并转换为TorchScript。之前用HuggingFace的Diffusers或Transformers类完成过转换,但自己写的类加载权重时报错,想确认这种场景下自定义封装类的可行性,以及需要注意的事项。
复现代码
import torch PATH = 'model.pth' pretrained_dict = torch.load(PATH) for key in list(pretrained_dict.keys()): print(key) class MyModel(torch.nn.Module): def __init__(self): super().__init__() def forward(self, x): return 0 model = MyModel() model.load_state_dict(pretrained_dict) model.eval() example_input = torch.rand(1, 3, 224, 224) torch_script = torch.jit.trace(model, example_input)
运行报错信息
tok_embeddings.weight norm.weight output.weight layers.0.attention.wq.weight layers.0.attention.wk.weight layers.0.attention.wv.weight layers.0.attention.wo.weight layers.0.feed_forward.w1.weight layers.0.feed_forward.w2.weight layers.0.feed_forward.w3.weight layers.0.attention_norm.weight layers.0.ffn_norm.weight layers.1.attention.wq.weight layers.1.attention.wk.weight layers.1.attention.wv.weight layers.1.attention.wo.weight layers.1.feed_forward.w1.weight layers.1.feed_forward.w2.weight layers.1.feed_forward.w3.weight layers.1.attention_norm.weight layers.1.ffn_norm.weight layers.2.attention.wq.weight layers.2.attention.wk.weight layers.2.attention.wv.weight layers.2.attention.wo.weight ... layers.31.feed_forward.w3.weight layers.31.attention_norm.weight layers.31.ffn_norm.weight rope.freqs --> 17 model.load_state_dict(pretrained_dict) 18 model.eval() 19 example_input = torch.rand(1, 3, 224, 224) File ~/text-generation-webui-main/installer_files/env/lib/python3.10/site-packages/torch/nn/modules/module.py:2041, in Module.load_state_dict(self, state_dict, strict) 2036 error_msgs.insert( 2037 0, 'Missing key(s) in state_dict: {}. '.format( 2038 ', '.join('"{}"'.format(k) for k in missing_keys))) 2040 if len(error_msgs) > 0: -> 2041 raise RuntimeError('Error(s) in loading state_dict for {}:\n\t{}'.format( 2042 self.__class__.__name__, "\n\t".join(error_msgs))) 2043 return _IncompatibleKeys(missing_keys, unexpected_keys) RuntimeError: Error(s) in loading state_dict for MyModel: Unexpected key(s) in state_dict: "tok_embeddings.weight", "norm.weight", "output.weight", "layers.0.attention.wq.weight", "layers.0.attention.wk.weight", "layers.0.attention.wv.weight", "layers.0.attention.wo.weight", "layers.0.feed_forward.w1.weight", "layers.0.feed_forward.w2.weight", "layers.....
解决方案与注意事项
可行性:完全可以自定义封装类
只要能1:1还原预训练模型的网络结构,就能用自定义Module类加载权重并转换为TorchScript。你当前的报错核心是自定义的MyModel和预训练权重的结构完全不匹配——你的类里没有定义任何对应层,权重自然无法加载。
关键注意事项及修复步骤
完全还原模型的网络结构
从权重的key能看出,这是类似LLaMA的Transformer架构,包含词嵌入层tok_embeddings、32个Transformer层layers、归一化层norm、输出层output,还有RoPE频率参数rope.freqs。必须在MyModel的__init__方法里定义所有子模块,结构和命名要和权重key完全对应。示例结构定义(需根据实际模型参数调整):
import torch import torch.nn as nn import torch.nn.functional as F class Attention(nn.Module): def __init__(self, hidden_size): super().__init__() self.wq = nn.Linear(hidden_size, hidden_size) self.wk = nn.Linear(hidden_size, hidden_size) self.wv = nn.Linear(hidden_size, hidden_size) self.wo = nn.Linear(hidden_size, hidden_size) def forward(self, x): q = self.wq(x) k = self.wk(x) v = self.wv(x) attn = F.softmax(q @ k.transpose(-2, -1) / (x.size(-1)**0.5), dim=-1) output = attn @ v return self.wo(output) class FeedForward(nn.Module): def __init__(self, hidden_size, intermediate_size): super().__init__() self.w1 = nn.Linear(hidden_size, intermediate_size) self.w2 = nn.Linear(intermediate_size, hidden_size) self.w3 = nn.Linear(hidden_size, intermediate_size) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x)) class TransformerLayer(nn.Module): def __init__(self, hidden_size, intermediate_size): super().__init__() self.attention_norm = nn.LayerNorm(hidden_size) self.attention = Attention(hidden_size) self.ffn_norm = nn.LayerNorm(hidden_size) self.feed_forward = FeedForward(hidden_size, intermediate_size) def forward(self, x): x = x + self.attention(self.attention_norm(x)) x = x + self.feed_forward(self.ffn_norm(x)) return x class MyModel(nn.Module): def __init__(self, vocab_size, hidden_size, num_layers, intermediate_size): super().__init__() self.tok_embeddings = nn.Embedding(vocab_size, hidden_size) self.layers = nn.ModuleList([TransformerLayer(hidden_size, intermediate_size) for _ in range(num_layers)]) self.norm = nn.LayerNorm(hidden_size) self.output = nn.Linear(hidden_size, vocab_size) self.rope = nn.Parameter(torch.load(PATH)['rope.freqs'], requires_grad=False) def forward(self, x): x = self.tok_embeddings(x) for layer in self.layers: x = layer(x) x = self.norm(x) return self.output(x)匹配模型超参数
需要获取原模型的核心超参数:词汇表大小vocab_size、隐藏层大小hidden_size、层数num_layers、中间层大小intermediate_size等。可以从模型原始文档获取,也可以通过权重形状反推(比如tok_embeddings.weight的形状是[vocab_size, hidden_size])。严格匹配子模块命名
权重key的层级(如layers.0.attention.wq.weight)对应模块的嵌套结构:model.layers[0].attention.wq,子模块名称必须和权重key完全一致,不能随意改名。处理非可训练参数
比如rope.freqs属于固定参数,无需训练,定义为nn.Parameter并设置requires_grad=False,直接从权重加载即可。验证前向逻辑正确性
加载权重后,用和原模型一致的输入测试输出是否匹配(若有原模型推理结果),确保前向计算逻辑完全对齐,避免转TorchScript后出现结果偏差。TorchScript转换细节
- 使用
torch.jit.trace时,输入要和模型实际推理的输入形状、类型一致(比如文本模型输入应为token整数张量,而非示例中的图像张量)。 - 若模型包含动态控制流(如可变循环次数、条件分支),建议用
torch.jit.script替代trace,避免遗漏逻辑。
- 使用
内容的提问来源于stack exchange,提问作者timyau

