Jetson Nano部署PyTorch模型遇版本冲突问题求助
PyTorch模型部署Jetson Nano时加载权重失败问题
环境信息
- 训练环境:PyTorch 2.0.0、Python 3.8
- Jetson Nano设备配置:
NVIDIA Jetson Nano Developer Kit JetPack: 4.6.2 OS: Ubuntu 18.04.6 LTS Kernel Version: 4.9.253-tegra CUDA: 10.2.300 CUDNN: 8.2.1.32 TensorRT: 8.2.1.8 Vision Works: 1.6.0.501 VPI: 1.2.3 Vulcan: 1.2.70
- Jetson预装PyTorch版本:1.10.2
加载权重代码
import torch import torchvision from torchvision.models.detection.faster_rcnn import FastRCNNPredictor NUM_CLASSES = 2 CLASSES = ['__background__', 'license-plate'] DEVICE = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu') OUT_DIR = '' WEIGHTS_PATH = 'best_model.pth' def create_model(num_classes, pretrained=True): # Load Faster RCNN pre-trained model model = torchvision.models.detection.fasterrcnn_resnet50_fpn() # Get the number of input features in_features = model.roi_heads.box_predictor.cls_score.in_features # define a new head for the detector with required number of classes model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes) return model def main(): checkpoint = torch.load(WEIGHTS_PATH, map_location=DEVICE) model = create_model(num_classes=NUM_CLASSES, coco_model=False) model.load_state_dict(checkpoint['model_state_dict']) model.to(DEVICE).eval() model_scripted = torch.jit.script(model) model_scripted.save('model_scripted.pt') main()
报错信息
Traceback (most recent call last): File "jetson_inference.py", line 34, in <module> main() File "jetson_inference.py", line 29, in main model.load_state_dict(checkpoint['model_state_dict']) File "/home/vha/.local/lib/python3.6/site-packages/torch/nn/modules/module.py", line 1483, in load_state_dict self.__class__.__name__, "\n\t".join(error_msgs))) RuntimeError: Error(s) in loading state_dict for FasterRCNN: Missing key(s) in state_dict: "backbone.fpn.inner_blocks.0.weight", "backbone.fpn.inner_blocks.0.bias", "backbone.fpn.inner_blocks.1.weight", "backbone.fpn.inner_blocks.1.bias", "backbone.fpn.inner_blocks.2.weight", "backbone.fpn.inner_blocks.2.bias", "backbone.fpn.inner_blocks.3.weight", "backbone.fpn.inner_blocks.3.bias", "backbone.fpn.layer_blocks.0.weight", "backbone.fpn.layer_blocks.0.bias", "backbone.fpn.layer_blocks.1.weight", "backbone.fpn.layer_blocks.1.bias", "backbone.fpn.layer_blocks.2.weight", "backbone.fpn.layer_blocks.2.bias", "backbone.fpn.layer_blocks.3.weight", "backbone.fpn.layer_blocks.3.bias", "rpn.head.conv.weight", "rpn.head.conv.bias". Unexpected key(s) in state_dict: "backbone.fpn.inner_blocks.0.0.weight", "backbone.fpn.inner_blocks.0.0.bias", "backbone.fpn.inner_blocks.1.0.weight", "backbone.fpn.inner_blocks.1.0.bias", "backbone.fpn.inner_blocks.2.0.weight", "backbone.fpn.inner_blocks.2.0.bias", "backbone.fpn.inner_blocks.3.0.weight", "backbone.fpn.inner_blocks.3.0.bias", "backbone.fpn.layer_blocks.0.0.weight", "backbone.fpn.layer_blocks.0.0.bias", "backbone.fpn.layer_blocks.1.0.weight", "backbone.fpn.layer_blocks.1.0.bias", "backbone.fpn.layer_blocks.2.0.weight", "backbone.fpn.layer_blocks.2.0.bias", "backbone.fpn.layer_blocks.3.0.weight", "backbone.fpn.layer_blocks.3.0.bias", "rpn.head.conv.0.0.weight", "rpn.head.conv.0.0.bias".
解决方案
错误根源是PyTorch/torchvision版本差异导致模型结构参数命名不一致:
- 训练用的高版本torchvision中,FPN模块的inner_blocks、layer_blocks及RPN的conv层是直接的Conv2d层,参数名无
.0.0后缀 - Jetson上的低版本torchvision中,这些模块被包装在Sequential容器内,参数名多了
.0.0层级
方法1:修改权重字典键名适配低版本
加载权重前手动修改参数键,移除多余的.0.0后缀:
def main(): checkpoint = torch.load(WEIGHTS_PATH, map_location=DEVICE) # 调整权重键名 new_state_dict = {} for k, v in checkpoint['model_state_dict'].items(): new_k = k.replace('.0.0', '') new_state_dict[new_k] = v checkpoint['model_state_dict'] = new_state_dict model = create_model(num_classes=NUM_CLASSES, pretrained=True) # strict=False跳过非核心参数差异 model.load_state_dict(checkpoint['model_state_dict'], strict=False) model.to(DEVICE).eval() model_scripted = torch.jit.script(model) model_scripted.save('model_scripted.pt')
方法2:在训练环境导出兼容模型
在训练服务器上先导出脚本模型,再复制到Jetson使用,规避版本结构差异:
# 训练服务器执行 model = create_model(num_classes=NUM_CLASSES) model.load_state_dict(torch.load('best_model.pth')['model_state_dict']) model.eval() torch.jit.script(model).save('model_scripted.pt')
Jetson上直接加载导出的模型:
# Jetson执行 model = torch.jit.load('model_scripted.pt') model.to(DEVICE).eval()
方法3:升级Jetson的PyTorch版本
JetPack4.6.2(CUDA10.2)最高支持PyTorch1.13.1,可参考NVIDIA官方提供的Jetson PyTorch安装包升级,直接匹配训练环境的模型结构。
内容的提问来源于stack exchange,提问作者Hiran Hasanka
相关产品推荐
相关产品推荐

