PyTorch中含可选成员(可设为None)的模型加载问题
定义PyTorch的torch.nn.Module子类A,初始化方法如下:
def __init__(self, additional_layer=False): ... if additional_layer: self.additional = nn.Sequential(nn.Linear(8,3)).to(self.device) else: self.additional = None ...
训练时设置additional_layer=True,训练完成后用torch.save保存模型的state_dict()。推理阶段加载模型时执行model.load_state_dict(best_model["my_model"]),出现如下错误:
RuntimeError: Error(s) in loading state_dict for A: Unexpected key(s) in state_dict: "additional.0.weight"
疑问:是否不允许使用可设为None的可选字段?该如何正确处理?
不是PyTorch不允许用可设为None的可选字段,报错核心原因是加载模型时的实例结构和训练时不一致:训练时模型包含additional子模块,而加载时如果初始化A用了additional_layer=False,模型里没有这个子模块,state_dict中的additional.0.weight等key找不到对应的参数,因此报错。
可以通过以下几种方式解决:
保持加载与训练的模型结构一致
加载模型时,必须用和训练时完全相同的初始化参数创建模型实例,确保结构匹配:# 创建和训练时结构一致的模型实例 model = A(additional_layer=True) # 正常加载state_dict model.load_state_dict(best_model["my_model"])忽略不匹配的参数(谨慎使用)
如果确实需要用additional_layer=False的结构加载,可以在load_state_dict中设置strict=False,跳过不匹配的key。但注意这会丢失额外层的参数,仅适用于确定不需要该层功能的场景:model = A(additional_layer=False) # strict=False忽略state_dict中不存在的key model.load_state_dict(best_model["my_model"], strict=False)统一模型结构(推荐)
避免将可选模块设为None,而是用**空操作层(如nn.Identity())**代替,确保无论参数如何,模型结构始终一致,从根源避免加载时的结构不匹配问题:def __init__(self, additional_layer=False): super().__init__() # 其他层定义... if additional_layer: self.additional = nn.Sequential(nn.Linear(8,3)) else: # 用恒等层代替None,不改变输入输出,同时保持结构一致 self.additional = nn.Identity() # 建议不要在初始化中直接调用.to(self.device),后续统一用model.to(device)移动设备这种方式下,无论
additional_layer取值如何,模型都包含additional子模块,state_dict的key始终一致,加载时无需额外处理。
内容的提问来源于stack exchange,提问作者affine_scheme

