1.首先官网上下载libtorch,放到当前项目下

2.将pytorch训练好的模型使用torch.jit.trace导出为.pt格式

 import torch
from skimage import io, transform, color
import numpy as np
import os
import torch.nn.functional as F
import warnings
warnings.filterwarnings("ignore")
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") labels = ['cock', 'drawing', 'neutral', 'porn', 'sexy']
path = "test/n_1.jpg"
im = io.imread(path)
if im.shape[2] == 4:
im = color.rgba2rgb(im) im = transform.resize(im, (224, 224))
im = np.transpose(im, (2, 0, 1))
dummy_input = np.expand_dims(im, 0)
inp = torch.from_numpy(dummy_input)
inp = inp.float()
model = torch.load(
"models/resnet50-epoch-0-accu-0.9213857428381079.pth", map_location='cpu')
traced_script_module = torch.jit.trace(model, inp)
output = model(inp)
probs = F.softmax(output).detach().numpy()[0]
pred = np.argmax(probs) traced_script_module.save("models/traced_resnet_model.pt")

torchscript加载.pt模型

 // One-stop header.
#include <torch/script.h> // headers for opencv
#include <opencv2/highgui/highgui.hpp>
#include <opencv2/imgproc/imgproc.hpp>
#include <opencv2/opencv.hpp> #include <cmath>
#include <iostream>
#include <memory>
#include <string>
#include <vector> #define kIMAGE_SIZE 224
#define kCHANNELS 3
#define kTOP_K 1 //print top k predicted results bool LoadImage(std::string file_name, cv::Mat &image)
{
image = cv::imread(file_name); // CV_8UC3
if (image.empty() || !image.data)
{
return false;
}
cv::cvtColor(image, image, CV_BGR2RGB);
// scale image to fit
cv::Size scale(kIMAGE_SIZE, kIMAGE_SIZE);
cv::resize(image, image, scale); // convert [unsigned int] to [float]
image.convertTo(image, CV_32FC3,1.0/255); return true;
} bool LoadImageNetLabel(std::string file_name,
std::vector<std::string> &labels)
{
std::ifstream ifs(file_name);
if (!ifs)
{
return false;
}
std::string line;
while (std::getline(ifs, line))
{
labels.push_back(line);
}
return true;
} int main(int argc, const char *argv[])
{
if (argc != 3)
{
std::cerr << "Usage:classifier <path-to-exported-script-module> <path-to-lable-file> " << std::endl;
return -1;
} //load model
torch::jit::script::Module module = torch::jit::load(argv[1]);
// to GPU
// module->to(at::kCUDA);
std::cout << "== ResNet50 loaded!\n"; //load labels(classes names)
std::vector<std::string> labels;
if (LoadImageNetLabel(argv[2], labels))
{
std::cout << "== Label loaded! Let's try it\n";
}
else
{
std::cerr << "Please check your label file path." << std::endl;
return -1;
} std::string file_name = "";
cv::Mat image;
while (true)
{
std::cout << "== Input image path: [enter q to exit]" << std::endl;
std::cin >> file_name;
if (file_name == "Q" || file_name == "q")
{
break;
}
if (LoadImage(file_name, image))
{
//read image tensor
auto input_tensor = torch::from_blob(
image.data, {1, kIMAGE_SIZE, kIMAGE_SIZE, kCHANNELS});
input_tensor = input_tensor.permute({0, 3, 1, 2});
input_tensor[0][0] = input_tensor[0][0].sub_(0.485).div_(0.229);
input_tensor[0][1] = input_tensor[0][1].sub_(0.456).div_(0.224);
input_tensor[0][2] = input_tensor[0][2].sub_(0.406).div_(0.225);
// to GPU
// input_tensor = input_tensor.to(at::kCUDA); torch::Tensor out_tensor = module.forward({input_tensor}).toTensor(); auto results = out_tensor.sort(-1, true);
auto softmaxs = std::get<0>(results)[0].softmax(0);
auto indexs = std::get<1>(results)[0]; for (int i = 0; i < kTOP_K; ++i)
{
auto idx = indexs[i].item<int>();
std::cout << " ============= Top-" << i + 1 << " =============" << std::endl;
std::cout << " Label: " << labels[idx] << std::endl;
std::cout << " With Probability: "
<< softmaxs[i].item<float>() * 100.0f << "%" << std::endl;
}
}
else
{
std::cout << "Can't load the image, please check your path." << std::endl;
}
}
}

CMakeLists.txt编译

 cmake_minimum_required(VERSION 2.8)
project(predict_demo)
SET(CMAKE_CXX_FLAGS ${CMAKE_CXX_FLAGS} "-std=c++11 -O3") set(OpenCV_DIR /home/buyizhiyou/opencv-3.4./build)
find_package(OpenCV REQUIRED)
find_package(Torch REQUIRED) # 添加头文件
include_directories( ${OpenCV_INCLUDE_DIRS} ) add_executable(resnet_demo resnet_demo.cpp)
target_link_libraries(resnet_demo ${TORCH_LIBRARIES} ${OpenCV_LIBS})
set_property(TARGET resnet_demo PROPERTY CXX_STANDARD )

运行

./resnet_demo   models/traced_resnet_model.pt  labels.txt

c++ 使用torchscript 加载训练好的pytorch模型的更多相关文章

  1. vue中加载three.js的gltf模型

    vue中加载three.js的gltf模型 一.开始引入three.js相关插件.首先利用淘宝镜像,操作命令为: cnpm install three //npm install three也行 二. ...

  2. pytorch 加载训练好的模型做inference

    前提: 模型参数和结构是分别保存的 1. 构建模型(# load model graph) model = MODEL() 2.加载模型参数(# load model state_dict) mode ...

  3. Tensorflow加载预训练模型和保存模型(ckpt文件)以及迁移学习finetuning

    转载自:https://blog.csdn.net/huachao1001/article/details/78501928 使用tensorflow过程中,训练结束后我们需要用到模型文件.有时候,我 ...

  4. Tensorflow加载预训练模型和保存模型

    转载自:https://blog.csdn.net/huachao1001/article/details/78501928 使用tensorflow过程中,训练结束后我们需要用到模型文件.有时候,我 ...

  5. 关于Tensorflow 加载和使用多个模型的方式

    在Tensorflow中,所有操作对象都包装到相应的Session中的,所以想要使用不同的模型就需要将这些模型加载到不同的Session中并在使用的时候申明是哪个Session,从而避免由于Sessi ...

  6. [原][osgearth]earth文件加载道路一初步看见模型道路

    时间是2017年2月5日17:16:32 由于OE2.9还没有发布,但是我又急于使用OE的道路. 所以,我先编译了正在github上调试中的OE2.9 github网址是:https://github ...

  7. Three.js中加载外部fbx格式的模型素材

    index.html部分: index.js部分: Scene.js部分:

  8. 学习笔记TF016:CNN实现、数据集、TFRecord、加载图像、模型、训练、调试

    AlexNet(Alex Krizhevsky,ILSVRC2012冠军)适合做图像分类.层自左向右.自上向下读取,关联层分为一组,高度.宽度减小,深度增加.深度增加减少网络计算量. 训练模型数据集 ...

  9. 深度学习原理与框架-猫狗图像识别-卷积神经网络(代码) 1.cv2.resize(图片压缩) 2..get_shape()[1:4].num_elements(获得最后三维度之和) 3.saver.save(训练参数的保存) 4.tf.train.import_meta_graph(加载模型结构) 5.saver.restore(训练参数载入)

    1.cv2.resize(image, (image_size, image_size), 0, 0, cv2.INTER_LINEAR) 参数说明:image表示输入图片,image_size表示变 ...

随机推荐

  1. 剑指offer:链表中环的入口结点

    题目描述: 给一个链表,若其中包含环,请找出该链表的环的入口结点,否则,输出null. 思路分析: 这道题首先需要判断链表是否存在环,很快就能想到用快慢指针来判断. 由于快慢指针的相遇位置并不一定为链 ...

  2. DELPHI7 ADO二层升三层新增LINUX服务器方案

    DELPHI7 ADO二层升三层新增LINUX服务器方案 引子:笔者曾经无数次在用户的LINUX服务器上创建一个WINDOWS虚拟机,用于运行自己DELPHI开发中间件. 现在再不需要如此麻烦了. 咏 ...

  3. Linux 常用操作和命令

    腾讯云部署 java web 环境:https://blog.csdn.net/niceLiuSir/article/details/78879844 Tomcat部署和配置:https://blog ...

  4. 安卓 android studio 报错 The specified Android SDK Build Tools version (27.0.3) is ignored, as it is below the minimum supported version (28.0.3) for Android Gradle

    今天将项目迁移到另一台笔记本,进行build出现以下问题,导致build失败 报错截图: 大致意思,目前使用的build工具版本27.0.3不合适.因为当前使用Gradle插件版本是3.2.1,这个版 ...

  5. 分类的性能评估:准确率、精确率、Recall召回率、F1、F2

    import numpy as np import pandas as pd from sklearn.feature_extraction.text import TfidfVectorizer f ...

  6. Java 解析XML数据

    实例一:获取指定两个标签之间的数据 XML数据格式: <?xml version="1.0" encoding="utf-8"?> <soap ...

  7. [LeetCode] 21. Merge Two Sorted Lists 合并有序链表

    Merge two sorted linked lists and return it as a new list. The new list should be made by splicing t ...

  8. 【Tools】HP/惠普v285w 量产工具

    前段时间朋友说自己u盘坏了,让帮忙看看.看下图是这个u盘. 坏的问题:往里面复制东西,提示:请去掉写保护或使用另一张磁盘.但是能正常从里面读取出来数据. 无论更换电脑,还是使用网上的修改注册表等方式皆 ...

  9. consul异地多数据中心以及集群部署方案

    consul异地多数据中心以及集群部署方案目的实现consul 异地多数据中心环境部署,使得一个数据中心的服务可以从另一个数据中心的consul获取已注册的服务地址 环境准备两台 linux服务器,外 ...

  10. 「模拟赛20191019」A 简单DP

    题目描述 给一个\(n\times m\)的网格,每个格子上有一个小写字母. 对于所有从左上角\((1,1)\)到右下角\((n,m)\)只向下或向右走的路径构成的集合,判断是否存在两条走法不同的路径 ...