You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

无法将PyTorch MaskRCNN模型转为Scripted Module并在LibTorch加载

问题:MaskRCNN模型转LibTorch格式后无法加载

问题描述

使用Python将torchvision的maskrcnn_resnet50_fpn模型转为Scripted Module后,在C中调用torch::jit::load时崩溃,报错torch::jit::ErrorReport。Python转换代码可正常运行并输出结果,但C加载失败。

Python转换代码:

loaded_model = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=False)

# loaded_model.cpu()
loaded_model.eval()
example = torch.rand(1, 3, 256, 256)
scripted_model = torch.jit.script(loaded_model)
out = loaded_model(example)
scripted_model.save('../models/vanila_rcnn.pt')
out[0]["boxes"]

C++加载代码:

int main(int argc, const char* argv[]) {
    std::string _path = "C:\\Projects\\AnatomySegmTorch\\models\\vanila_rcnn.pt";
    torch::jit::script::Module module;
    //torch::NoGradGuard no_grad; //stops grad calculate
    try {
        module = torch::jit::load(_path);
    }
    catch (const c10::Error& ) {
        std::cerr << "error loading the model\n";
        return -1;
    }

   // Create a vector of inputs.
    std::vector<torch::jit::IValue> inputs;
    inputs.push_back(torch::ones({ 1, 3, 256, 256 }));

    // Execute the model and turn its output into a tensor.
    at::Tensor output = module.forward(inputs).toTensor(); 
    return 0;
}

解决方法

1. 严格匹配PyTorch与LibTorch版本

加载崩溃最常见的原因是版本不兼容,Python端使用的PyTorch版本必须和C++端的LibTorch版本完全一致,包括CUDA版本(若使用GPU)。例如PyTorch 2.0.1对应LibTorch 2.0.1,不能混用不同大版本或小版本。

2. 改用torch.jit.trace转换模型

TorchVision检测模型(如MaskRCNN)包含大量动态分支逻辑,torch.jit.script无法完全兼容,改用torch.jit.trace更适合这类模型。修改Python转换代码如下:

import torch
import torchvision

loaded_model = torchvision.models.detection.maskrcnn_resnet50_fpn(pretrained=False)
loaded_model.eval()

# 构造与推理时维度一致的示例输入
example = torch.rand(1, 3, 256, 256)
# 使用trace转换模型,传入示例输入捕捉计算图
traced_model = torch.jit.trace(loaded_model, example)
# 验证转换后的模型输出
out = traced_model(example)
print(out[0]["boxes"])
# 保存trace后的模型
traced_model.save('../models/vanila_rcnn.pt')

注意:若需要支持可变输入尺寸,可结合torch.jit.trace_module或调整模型动态逻辑,确保推理时输入尺寸与trace时一致。

3. 修正C++代码的输出处理逻辑

MaskRCNN的输出不是单一Tensor,而是包含boxes、labels等字段的字典列表,不能直接转为at::Tensor。修改C++代码如下:

#include <torch/script.h>
#include <iostream>
#include <vector>

int main(int argc, const char* argv[]) {
    std::string _path = "C:\\Projects\\AnatomySegmTorch\\models\\vanila_rcnn.pt";
    torch::jit::script::Module module;
    try {
        module = torch::jit::load(_path);
        module.eval(); // 加载后需设置为评估模式
    }
    catch (const c10::Error& e) {
        std::cerr << "error loading the model: " << e.what() << "\n";
        return -1;
    }

    // 创建与trace时维度一致的输入
    std::vector<torch::jit::IValue> inputs;
    inputs.push_back(torch::ones({1, 3, 256, 256}).to(torch::kCPU));

    // 执行推理并解析输出
    auto output_list = module.forward(inputs).toList();
    auto result_dict = output_list.get(0).toGenericDict();

    // 提取boxes张量
    at::Tensor boxes = result_dict.at("boxes").toTensor();
    std::cout << "Detected boxes:\n" << boxes << std::endl;

    return 0;
}

4. 额外注意事项

  • 若使用GPU版本LibTorch,需将模型和输入都移至CUDA设备:module.to(torch::kCUDA);、inputs.push_back(torch::ones(...).to(torch::kCUDA));。
  • 转换模型时必须保持eval模式,避免BatchNorm、Dropout等层的训练态行为干扰。
  • 若trace时出现动态逻辑警告,可尝试用torch.jit.script辅助转换,或修改模型中的动态分支(如将if-else替换为torch.where)。

内容的提问来源于stack exchange,提问作者Andrey Taranov

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.13 20:15:37