PyTorch LSTM模型GPU训练后CPU推理:Checkpoint加载失败求助
搞定这个模型加载错误的小妙招
嘿,这个问题我太熟了!你遇到的RuntimeError根源很明确:训练时定义的SentimentLSTM模型,和你推理时写的模型,内部嵌入层的变量名不一致。训练时模型里的嵌入层叫encoder(所以state_dict里的键是encoder.weight),但推理时你把这个层命名成了embedding(模型期望找embedding.weight),两边对不上自然就报错了。
两种快速解决方法
方法1:统一模型的变量名(推荐)
最稳妥的方式是让推理时的模型定义和训练时完全一模一样——毕竟训练时的模型结构才是和checkpoint匹配的。
比如如果你训练时的SentimentLSTM类里嵌入层是这么写的:
# 训练时的模型定义片段 self.encoder = nn.Embedding(vocab_size, embedding_dim)
那推理时的模型里也得用self.encoder,而不是self.embedding。把推理代码里的模型定义改成和训练时完全一致的结构,再加载checkpoint就没问题了。
方法2:手动修改state_dict的键名
要是不想改模型定义,也可以在加载时手动把checkpoint里的键名改过来,让它和推理模型匹配:
# 加载checkpoint并调整键名 device = torch.device('cpu') checkpoint = torch.load('lstmmodelgpu.tar', map_location=device) # 重命名state_dict里的键 adjusted_state_dict = {} for key, value in checkpoint['model_state_dict'].items(): if key == 'encoder.weight': adjusted_state_dict['embedding.weight'] = value else: adjusted_state_dict[key] = value # 加载调整后的state_dict model.load_state_dict(adjusted_state_dict) model.eval()
额外提醒
PyTorch的state_dict是靠模型层的变量名来绑定参数的,所以不仅是嵌入层,模型的所有层的变量名、层数、维度、甚至dropout这类超参数,训练和推理时都得完全一致,不然还会出现类似的加载错误。你之前保存checkpoint的代码torch.save({'model_state_dict': model.state_dict()},'lstmmodelgpu.tar')是没问题的,不用改~
内容的提问来源于stack exchange,提问作者Sijan Bhandari
相关产品推荐
相关产品推荐

