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

PyTorch中自定义state_dict()与load_state_dict()实现嵌套模块权重预处理

自定义PyTorch Module的state_dict与load_state_dict实现权重预处理

我有一组嵌套的torch.nn.Module类,需要在保存其中一个嵌套类的权重前进行预处理。能否重写state_dict()函数,将预处理逻辑嵌入自定义实现中?

示例代码(原代码存在语法错误,已标注):

# 注意:原代码存在大小写错误(Class → class,module → Module)及方法参数缺失问题
class A(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.b1 = B1()
        self.b2 = B2()

class B1(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.var = torch.nn.Parameter(torch.Tensor((3, 5), dtype=float))

class B2(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.var = torch.nn.Parameter(torch.Tensor((3, 5), dtype=float))

    # 原代码缺失self参数,且父类方法调用格式错误
    def state_dict():
        # 我想要像这样重写默认state_dict,但无法生效。是否有可行方案?
        bool_var = self.var.bool().cpu().numpy()
        state_dict1 = super.state_dict()
        state_dict1.update({'var': bool_var})
        return state_dict1

    # 原代码缺失self参数,父类方法调用格式错误
    def load_state_dict(state_dict):
        state_dict['var'] = state_dict['var'].float()
        super.load_state_dict(state_dict)
        return

具体需求:针对B2类,保存权重时将var转为bool类型,加载时转回float类型。由于训练时该变量必须以float类型参与计算,无法默认设为bool类型。此前我通过硬编码方式处理变量保存,现在希望通过重写state_dict()和load_state_dict()自动处理转换逻辑。之前看到相关讨论称无法自定义,但希望PyTorch新版本已支持或存在其他实现方式。


可行实现方案

PyTorch支持重写state_dict()和load_state_dict()方法,只需遵循正确的方法签名并处理好父类调用逻辑即可。以下是修正后的完整代码:

import torch

class A(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.b1 = B1()
        self.b2 = B2()

class B1(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.var = torch.nn.Parameter(torch.randn(3, 5, dtype=torch.float32))

class B2(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.var = torch.nn.Parameter(torch.randn(3, 5, dtype=torch.float32))

    def state_dict(self, destination=None, prefix='', keep_vars=False):
        # 调用父类方法获取默认state_dict
        base_state_dict = super().state_dict(destination, prefix, keep_vars)
        # 对var进行预处理:转bool→CPU→numpy
        processed_var = self.var.bool().cpu().numpy()
        # 注意前缀:嵌套Module中PyTorch会自动添加前缀,需用prefix拼接键名
        base_state_dict[prefix + 'var'] = processed_var
        return base_state_dict

    def load_state_dict(self, state_dict, strict=True):
        # 取出预处理后的var,转回float类型的Tensor
        if 'var' in state_dict:
            state_dict['var'] = torch.tensor(state_dict['var'], dtype=torch.float32)
        # 调用父类方法加载处理后的state_dict
        super().load_state_dict(state_dict, strict)

关键注意点

  1. 方法签名匹配:重写state_dict()时必须保留父类的全部参数(destination, prefix, keep_vars),尤其是prefix,PyTorch递归处理嵌套Module时会自动添加前缀,确保键名匹配。
  2. 父类方法调用:必须使用super().state_dict(...)和super().load_state_dict(...)的格式调用父类方法,不能省略参数。
  3. 序列化兼容性:转换后的numpy数组可被torch.save正常序列化,无需额外处理。
  4. 测试验证:可通过以下代码验证逻辑是否生效:
# 创建模型实例
model = A()
# 保存state_dict
saved_dict = model.state_dict()
# 检查B2的var是否为bool类型的numpy数组
print(type(saved_dict['b2.var']))  # <class 'numpy.ndarray'>
print(saved_dict['b2.var'].dtype)  # bool

# 创建新实例加载权重
new_model = A()
new_model.load_state_dict(saved_dict)
# 检查B2的var是否恢复为float类型的Parameter
print(new_model.b2.var.dtype)  # torch.float32

内容的提问来源于stack exchange,提问作者Nagabhushan S N

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 16:33:40