PyTorch使用deepcopy备份模型触发RuntimeError报错求助
PyTorch使用copy.deepcopy备份模型报错的解决方法
问题重现
尝试用copy.deepcopy备份PyTorch模型:
backup_model = copy.deepcopy(model)
出现错误:
RuntimeError: Only Tensors created explicitly by the user (graph leaves) support the deepcopy protocol at the moment
使用PyTorch版本为1.11.0。
错误原因
PyTorch中,模型训练过程里的部分参数张量会成为计算图的中间节点(非叶子张量),而copy.deepcopy目前仅支持复制用户显式创建的叶子张量,因此触发该错误。你参考的论坛方案仅适用于参数全为叶子张量的静态模型场景,实际训练中的模型往往存在非叶子张量,导致报错。
解决方案
方法1:使用官方推荐的state_dict(最可靠)
PyTorch官方建议通过state_dict保存和加载模型参数,这种方式不涉及计算图,稳定性和效率更高:
# 保存模型参数状态 backup_state = model.state_dict() # 初始化新模型并加载参数 # 假设你的模型类为YourModel backup_model = YourModel() backup_model.load_state_dict(backup_state) # 按需切换模型模式 backup_model.train() # 或 backup_model.eval()
方法2:转换参数为叶子张量后再deepcopy
如果必须使用copy.deepcopy,可以先将模型参数从计算图中分离,转为叶子张量:
import copy # 切换模型到评估模式,关闭动态计算模块(如Dropout) model.eval() # 将所有参数从计算图中分离,转为叶子张量 for param in model.parameters(): param.detach_() # 执行deepcopy backup_model = copy.deepcopy(model)
注意:这种方法仅适用于简单模型,复杂模型可能仍存在隐藏的非叶子张量,可靠性不如state_dict。
内容的提问来源于stack exchange,提问作者postnubilaphoebus
相关产品推荐
相关产品推荐

