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

PyTorch加载state_dict时权重转置导致形状不匹配的原因

问题:PyTorch state_dict加载时权重形状不匹配(转置问题)

背景信息

模型结构

>>> torch_model
TorchModel(
  (dropout): Dropout(p=0, inplace=False)
  (base): Sequential(
    (0): Flatten(start_dim=1, end_dim=-1)
    (1): Linear(in_features=18018, out_features=1024, bias=True)
    (2): ReLU()
    (3): Linear(in_features=1024, out_features=512, bias=True)
    (4): ReLU()
  )

权重保存代码

def save_blocks(self, path):
        inters_dir = self.__pthjoin(path, 'inters')
        integs_dir = self.__pthjoin(path, 'integs')
        outputs_dir = self.__pthjoin(path, 'outputs')
        for dir in [path, inters_dir, integs_dir, outputs_dir]:
            if not os.path.isdir(dir):
                os.mkdir(dir)
        torch.save(self.base.state_dict(), self.__pthjoin(path, 'base'))

模型加载代码

def __restore_from(self, path):
        base_dir = self.__pthjoin(path, 'base')
        self.base = self.init_base_block(num_features=len(self.statistics))
        self.base.load_state_dict(torch.load(base_dir))

加载时的错误信息

RuntimeError: Error(s) in loading state_dict for Sequential:
        size mismatch for 1.weight: copying a param with shape torch.Size([18018, 1024]) from checkpoint, the shape in current model is torch.Size([1024, 18018]).
        size mismatch for 3.weight: copying a param with shape torch.Size([1024, 512]) from checkpoint, the shape in current model is torch.Size([512, 1024]).

注:使用torch.save(torch_model, path)保存整个模型时加载正常,但使用state_dict加载出现权重转置的形状不匹配问题。


原因分析与解决方案

核心原因

问题出在加载时初始化的base模块,与保存权重时的base模块中Linear层的参数定义顺序不一致:

  • PyTorch的Linear层权重张量默认形状是(out_features, in_features),对应计算逻辑为y = x @ weight.T + bias。
  • 保存权重时,原模型的Linear层参数是正确的:比如第一个Linear层in_features=18018, out_features=1024,对应权重形状为(1024, 18018)。
  • 但调用init_base_block初始化新base时,误将in_features和out_features的顺序搞反了,导致新模型的Linear层权重形状变成(18018, 1024),与保存的权重形状完全转置,因此加载时触发形状不匹配错误。

为什么保存整个模型正常?

使用torch.save(torch_model, path)保存整个模型时,会同时保存模型的完整结构定义。加载时PyTorch会直接根据保存的结构重建模型,参数形状自然与训练时一致,因此不会出现形状不匹配问题。而state_dict仅保存参数张量,不包含模型结构信息,必须确保加载时的模型结构与保存时完全一致才能正常加载。

修复方法

检查init_base_block函数的实现,确保其中Linear层的in_features和out_features参数顺序,与训练时base模块的定义完全一致:
比如训练时第一个Linear层是in_features=18018, out_features=1024,初始化时必须严格保持这个参数顺序,不能颠倒。


内容的提问来源于stack exchange,提问作者Ludens Zhang

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 15:37:56