AssertionError求助:capturable=False时state_steps不能为CUDA张量
解决PyTorch 1.12.0加载模型权重时的AssertionError问题
方法1:手动转换state_steps张量到CPU
当加载的权重中state_steps是CUDA张量时,遍历state_dict将其转为CPU张量后再加载:
import torch # 加载 checkpoint 文件 checkpoint = torch.load("your_model_checkpoint.pth") # 遍历并处理所有含"state_steps"的参数 for key in list(checkpoint["model_state_dict"].keys()): if "state_steps" in key: checkpoint["model_state_dict"][key] = checkpoint["model_state_dict"][key].cpu() # 加载处理后的权重到模型 model.load_state_dict(checkpoint["model_state_dict"])
方法2:修改循环层的capturable参数
如果模型包含LSTM/GRU等循环层,在定义时设置capturable=True,让PyTorch允许state_steps为CUDA张量:
# 以LSTM为例,修改前 # lstm_layer = torch.nn.LSTM(input_size=256, hidden_size=128) # 修改后 lstm_layer = torch.nn.LSTM(input_size=256, hidden_size=128, capturable=True)
注意:此修改仅对新初始化的模型生效,加载旧权重时仍需配合方法1处理,后续保存的模型则不会再触发该错误。
方法3:加载时统一映射到CPU再移回GPU
通过map_location将所有权重先加载到CPU,再将模型移至GPU,避免state_steps在GPU上的冲突:
# 加载权重到CPU checkpoint = torch.load("your_model_checkpoint.pth", map_location=torch.device('cpu')) # 加载权重到模型 model.load_state_dict(checkpoint["model_state_dict"]) # 再将模型移至GPU(如果需要) model = model.to(torch.device("cuda"))
内容的提问来源于stack exchange,提问作者StanWarne
相关产品推荐
相关产品推荐

