PyTorch中model.state_dict()存储内容:哪些会被保存?哪些不会?
PyTorch模型state_dict存储规则解析
针对你给出的模型,调用torch.save({'model_state': model.state_dict()}, 'model.pt')时,组件的保存情况可以按以下规则区分:
会被保存的组件
所有作为模型属性的nn.Module子类实例,不管是否被使用,它们的可训练参数(权重、偏置等)都会被存入state_dict:
self.fc_direct:Linear层,直接在forward中使用,参数会被保存self.fc_indirect:Linear层,通过foo方法间接在forward中使用,参数会被保存self.fc_not_used:Linear层,完全未被调用,参数依然会被保存
不会被保存的组件
普通Python原生类型的属性(比如int、float),无论是否被方法调用,默认都不会被存入state_dict:
self.value_a:int类型,即使在forward中修改,也不会被保存self.value_b:int类型,在bar方法中使用,同样不会被保存self.value_c:int类型,完全未被使用,不会被保存
补充说明
如果需要让普通Python数值类型的属性被state_dict保存,可以通过以下方式注册:
- 若需要作为可训练参数:使用
self.register_parameter('param_name', torch.nn.Parameter(torch.tensor(100))) - 若不需要训练(比如固定统计量):使用
self.register_buffer('buffer_name', torch.tensor(100))
内容的提问来源于stack exchange,提问作者landings
相关产品推荐
相关产品推荐

