You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将预训练PyTorch模型封装为torch.nn.Module并转换为TorchScript

PyTorch模型转TorchScript:自定义Module类加载预训练权重问题

问题背景

手里只有预训练模型文件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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 15:59:51