如何从本地磁盘加载torchvision模型并解决权重加载报错问题
报错原因梳理
两个报错对应的根源分别是:
NotADirectoryError:你设置的TORCH_HOME路径不符合PyTorch的目录结构要求,PyTorch默认会在TORCH_HOME下找hub/checkpoints子目录存放权重文件,你直接把TORCH_HOME指向了权重文件所在的根目录,且该目录下直接存了后缀为.pth的文件,导致PyTorch误把pth文件当成目录尝试访问其下的hub文件夹,触发报错。URLError:你调用模型时传了pretrained=True参数,此时PyTorch会默认尝试联网下载预训练权重,而你当前运行环境无公网访问权限,域名解析失败触发报错。
正确加载步骤
按照以下方案修改代码即可正常加载本地权重:
- 优先推荐不需要调整目录结构的方案:初始化模型时关闭预训练开关,跳过联网下载逻辑,直接加载本地权重赋值
对应代码示例:import torch from torchvision.models.detection import fasterrcnn_resnet50_fpn DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # 初始化无预训练权重的空模型 model = fasterrcnn_resnet50_fpn(pretrained=False, pretrained_backbone=False).to(DEVICE) # 加载本地权重,根据你权重的实际存储结构调整是否需要取['state_dict']层 checkpoint = torch.load('../input/torchvision-fasterrcnn-resnet-50/model.pth.tar', map_location=DEVICE) model.load_state_dict(checkpoint['state_dict']) - 如果要使用
TORCH_HOME的方式加载,先调整本地目录结构:手动在../input/torchvision-fasterrcnn-resnet-50/下创建hub/checkpoints两级子目录,把你的fasterrcnn_resnet50_fpn_coco-258fb6c6.pth文件放到该checkpoints目录下,再运行你最开始的设置TORCH_HOME的代码即可。
内容的提问来源于stack exchange,提问作者Olli
相关产品推荐
相关产品推荐

