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)
关键注意点
- 方法签名匹配:重写
state_dict()时必须保留父类的全部参数(destination,prefix,keep_vars),尤其是prefix,PyTorch递归处理嵌套Module时会自动添加前缀,确保键名匹配。 - 父类方法调用:必须使用
super().state_dict(...)和super().load_state_dict(...)的格式调用父类方法,不能省略参数。 - 序列化兼容性:转换后的numpy数组可被
torch.save正常序列化,无需额外处理。 - 测试验证:可通过以下代码验证逻辑是否生效:
# 创建模型实例 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
相关产品推荐
相关产品推荐

