import torch
import torchvision.transforms as T
from models.crnn import CRNN
model = CRNN(num_classes=5000)
model.load_state_dict(torch.load("crnn_best.pth", map_location="cpu"))
model.eval()
dummy_input = torch.randn(1, 1, 32, 280)
torch.onnx.export(
model,
dummy_input,
"crnn.onnx",
export_params=True,
opset_version=11,
do_constant_folding=True,
input_names=['input'],
output_names=['output'],
dynamic_axes={
'input': {0: 'batch', 3: 'width'},
'output': {0: 'batch', 1: 'seq_len'}
}
)
#include <onnxruntime/core/session/onnxruntime_cxx_api.h>
#include <opencv2/opencv.hpp>
#include <iostream>
#include <vector>
#include <string>
class CRNNOCR {
private:
Ort::Env env{ORT_LOGGING_LEVEL_WARNING, "CRNN_OCR"};
Ort::Session *session;
Ort::MemoryInfo memory_info = Ort::MemoryInfo::CreateCpu(
OrtAllocatorType::OrtArenaAllocator,
OrtMemType::OrtMemTypeDefault);
std::vector<std::string> char_dict = {"<blank>", "a", "b", ..., "一", "丁", ...};
public:
CRNNOCR(const std::string& model_path) {
Ort::SessionOptions session_options;
session_options.SetIntraOpNumThreads(1);
session_options.SetGraphOptimizationLevel(
GraphOptimizationLevel::ORT_ENABLE_ALL);
session = new Ort::Session(env, model_path.c_str(), session_options);
}
~CRNNOCR() {
delete session;
}
cv::Mat preprocess(cv::Mat& image) {
cv::Mat gray, resized;
if (image.channels() == 3)
cv::cvtColor(image, gray, cv::COLOR_BGR2GRAY);
else
gray = image;
int height = 32;
double ratio = static_cast<double>(height) / image.rows;
int width = static_cast<int>(image.cols * ratio);
cv::resize(gray, resized, cv::Size(width, height), 0, 0, cv::INTER_AREA);
return resized;
}
std::string decode_output(float* output, int seq_len) {
std::string text;
int prev_idx = -1;
for (int i = 0; i < seq_len; ++i) {
int idx = std::distance(output + i * 5000,
std::max_element(output + i * 5000, output + (i + 1) * 5000));
if (idx != 0 && idx != prev_idx)
text += char_dict[idx];
prev_idx = idx;
}
return text;
}
std::string predict(cv::Mat& img) {
auto input_tensor = preprocess(img);
input_tensor.convertTo(input_tensor, CV_32F, 1.0 / 255.0);
const int input_width = input_tensor.cols;
const int input_height = input_tensor.rows;
const int batch_size = 1;
const int channels = 1;
const int sequence_length = input_width / 4;
std::vector<int64_t> input_shape = {batch_size, channels, input_height, input_width};
auto allocator = Ort::AllocatorWithDefaultOptions();
size_t input_tensor_size = batch_size * channels * input_height * input_width;
Ort::Value input_tensor_value = Ort::Value::CreateTensor<float>(
memory_info,
input_tensor.ptr<float>(),
input_tensor_size,
input_shape.data(),
input_shape.size());
const char* input_names[] = {"input"};
const char* output_names[] = {"output"};
auto output_tensors = session->Run(
Ort::RunOptions{nullptr},
input_names, &input_tensor_value, 1,
output_names, 1);
auto* float_data = output_tensors[0].GetTensorMutableData<float>();
int output_seq_len = output_tensors[0].GetTensorTypeAndShapeInfo().GetShape()[1];
return decode_output(float_data, output_seq_len);
}
};
int main(int argc, char** argv) {
if (argc < 2) {
std::cerr << "Usage: " << argv[0] << " <image_path>\n";
return -1;
}
CRNNOCR ocr("crnn.onnx");
cv::Mat img = cv::imread(argv[1], cv::IMREAD_GRAYSCALE);
if (img.empty()) {
std::cerr << "Failed to load image.\n";
return -1;
}
auto start = std::chrono::steady_clock::now();
std::string result = ocr.predict(img);
auto end = std::chrono::steady_clock::now();
auto duration = std::chrono::duration_cast<std::chrono::milliseconds>(end - start);
std::cout << "Text: " << result << "\n";
std::cout << "Inference Time: " << duration.count() << " ms\n";
return 0;
}