如何在C++中调用TFLite模型并获取推理输出数组?
如何在C++中获取TFLite模型的推理结果?
核心实现步骤
对应Python中get_tensor的操作,在C++的TFLite Interpreter里可以通过以下步骤获取并输出推理结果:
- 获取输出张量信息:通过
interpreter->outputs()[0]拿到第一个输出张量的索引(多输出模型可遍历outputs()数组),再用interpreter->tensor(output_idx)获取张量的维度、数据类型等元信息。 - 获取输出数据指针:使用
interpreter->typed_output_tensor<T>(output_idx)模板函数(T为输出张量的数据类型,比如你的场景中是float),直接获取指向输出数组的指针。 - 遍历输出数据:根据张量维度计算总元素数,遍历指针指向的数组完成打印或后续格式转换(比如转成LibTorch张量)。
修改后的完整C++代码
#include <iostream> #include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/model.h" #include "tensorflow/lite/kernels/register.h" #include <fstream> #include <vector> #include <numeric> #include <opencv2/opencv.hpp> int main() { // 加载模型 std::unique_ptr<tflite::FlatBufferModel> model = tflite::FlatBufferModel::BuildFromFile("model.tflite"); // 初始化解释器 std::unique_ptr<tflite::Interpreter> interpreter; tflite::ops::builtin::BuiltinOpResolver resolver; tflite::InterpreterBuilder(*model.get(), resolver)(&interpreter); // 分配张量内存 if (interpreter->AllocateTensors() != kTfLiteOk) { fprintf(stderr, "Failed to allocate tensor\n"); exit(-1); } // 配置解释器 interpreter->SetAllowFp16PrecisionForFp32(true); interpreter->SetNumThreads(1); // 获取输入张量维度 int input_idx = interpreter->inputs()[0]; auto height = interpreter->tensor(input_idx)->dims->data[1]; auto width = interpreter->tensor(input_idx)->dims->data[2]; auto channels = interpreter->tensor(input_idx)->dims->data[3]; // 加载输入图像 cv::Mat image = cv::imread("image.jpg"); if (image.empty()) { fprintf(stderr, "Failed to load image\n"); exit(-1); } // 调整图像尺寸匹配模型输入 cv::resize(image, image, cv::Size(width, height)); // 若模型要求RGB格式,需添加格式转换:cv::cvtColor(image, image, cv::COLOR_BGR2RGB) // 将图像数据拷贝到输入张量 image.convertTo(image, CV_32F, 1.0 / 255.0); // 归一化到[0,1] std::memcpy(interpreter->typed_input_tensor<float>(0), image.data, image.total() * image.elemSize()); // 执行推理 if (interpreter->Invoke() != kTfLiteOk) { fprintf(stderr, "Inference failed\n"); exit(-1); } // 获取推理结果 int output_idx = interpreter->outputs()[0]; TfLiteTensor* output_tensor = interpreter->tensor(output_idx); float* output_data = interpreter->typed_output_tensor<float>(0); // 计算输出张量总元素数 int output_num_elements = 1; for (int i = 0; i < output_tensor->dims->size; ++i) { output_num_elements *= output_tensor->dims->data[i]; } // 打印输出维度和数据 std::cout << "Output tensor dimensions: "; for (int i = 0; i < output_tensor->dims->size; ++i) { std::cout << output_tensor->dims->data[i] << " "; } std::cout << "\nOutput data: "; for (int i = 0; i < output_num_elements; ++i) { std::cout << output_data[i] << " "; // 限制打印数量,避免数据过多 if (i == 9) { std::cout << "..."; break; } } std::cout << std::endl; // 若需转换为Torch Tensor(需链接LibTorch) // torch::Tensor result = torch::from_blob(output_data, {output_tensor->dims->data[0], output_tensor->dims->data[1]}, torch::kFloat32); return 0; }
关键注意事项
- 数据类型匹配:
typed_output_tensor<T>中的T必须与模型输出张量类型一致,可通过output_tensor->type查看(比如kTfLiteFloat32对应float,kTfLiteInt32对应int)。 - 图像格式兼容:OpenCV默认读取BGR格式,若模型要求RGB输入,需添加
cv::cvtColor转换。 - 多输出处理:多输出模型可遍历
interpreter->outputs()数组,逐个处理每个输出索引。
内容的提问来源于stack exchange,提问作者Wang
相关产品推荐
相关产品推荐

