Keras训练模型转TensorFlow后C++推理结果与Keras存在差异
Keras模型转TensorFlow后C++推理实现(含图像预处理)
我之前刚好落地过一套从Keras训练到TensorFlow C++推理的流程,结合你提到的图像预处理逻辑,给你整理一下具体的实现细节和关键点,应该能帮到你:
一、先回顾模型转换步骤
训练好Keras模型后,直接保存成TensorFlow的SavedModel格式就可以了,C++的TensorFlow API直接支持加载这个格式(因为你不需要量化,所以不用额外处理):
# 假设model是你训练好的Keras模型 model.save("path/to/your_saved_model")
执行完这个命令会生成一个目录,里面包含模型的结构和权重文件,C++端直接加载这个目录就行。
二、C++推理核心代码实现
你提到的ReadFile函数是图像预处理的核心,负责读取PNG、解码、resize和归一化,我把完整的函数实现和配套的推理流程整理好了:
1. 图像预处理函数(ReadFile)
这个函数会把PNG文件转换成符合模型输入要求的张量,包含了所有预处理步骤:
#include "tensorflow/core/framework/tensor.h" #include "tensorflow/core/lib/io/path.h" #include "tensorflow/core/platform/env.h" #include "tensorflow/core/public/session.h" #include "tensorflow/core/ops/image_ops.h" static tensorflow::Status ReadFile(tensorflow::Env* oEnv, const std::string& sFileName, tensorflow::Tensor* output) { uint64 nFileSize = 0; // 获取文件大小 TF_RETURN_IF_ERROR(oEnv->GetFileSize(sFileName, &nFileSize)); std::string oFile; oFile.resize(nFileSize); std::unique_ptr<tensorflow::RandomAccessFile> file; // 打开文件 TF_RETURN_IF_ERROR(oEnv->NewRandomAccessFile(sFileName, &file)); tensorflow::StringPiece data; // 读取文件内容 TF_RETURN_IF_ERROR(file->Read(0, nFileSize, &data, &oFile[0])); // 1. 解码PNG图像 tensorflow::Tensor png_tensor(tensorflow::DT_STRING, tensorflow::TensorShape({})); png_tensor.scalar<std::string>()() = data.ToString(); tensorflow::SessionOptions options; std::unique_ptr<tensorflow::Session> session(tensorflow::NewSession(options)); tensorflow::GraphDef graph; // 构建DecodePng节点 tensorflow::NodeDef decode_png_node; decode_png_node.set_name("decode_png"); decode_png_node.set_op("DecodePng"); decode_png_node.add_input(png_tensor.name()); // 设置输出通道数,根据你的图像类型调整(3是RGB,1是灰度) (*decode_png_node.mutable_attr())["channels"].set_i(3); graph.add_node()->CopyFrom(decode_png_node); // 2. Resize图像到模型输入尺寸 tensorflow::NodeDef resize_node; resize_node.set_name("resize"); resize_node.set_op("ResizeBilinear"); resize_node.add_input("decode_png:0"); // 设置目标尺寸,比如你的模型输入是224x224,这里就填对应数值 tensorflow::Tensor size_tensor(tensorflow::DT_INT32, tensorflow::TensorShape({2})); auto size_flat = size_tensor.flat<int32>(); size_flat(0) = 224; size_flat(1) = 224; resize_node.add_input(size_tensor.name()); (*resize_node.mutable_attr())["align_corners"].set_b(false); graph.add_node()->CopyFrom(resize_node); // 3. 归一化处理(和训练时的预处理逻辑对齐!) tensorflow::NodeDef normalize_node; normalize_node.set_name("normalize"); normalize_node.set_op("Div"); normalize_node.add_input("resize:0"); // 如果训练时是除以255,这里就用255.0f;如果是用均值方差归一化,要改成对应的计算逻辑 tensorflow::Tensor scale_tensor(tensorflow::DT_FLOAT, tensorflow::TensorShape({})); scale_tensor.scalar<float>()() = 255.0f; normalize_node.add_input(scale_tensor.name()); graph.add_node()->CopyFrom(normalize_node); // 构建会话并运行预处理流程 TF_RETURN_IF_ERROR(session->Create(graph)); std::vector<std::pair<std::string, tensorflow::Tensor>> inputs = { {png_tensor.name(), png_tensor}, {size_tensor.name(), size_tensor}, {scale_tensor.name(), scale_tensor} }; std::vector<tensorflow::Tensor> outputs; TF_RETURN_IF_ERROR(session->Run(inputs, {"normalize:0"}, {}, &outputs)); *output = outputs[0]; return tensorflow::Status::OK(); }
2. 完整的推理主流程
把预处理好的张量喂给模型,执行推理并输出结果:
int main(int argc, char* argv[]) { // 初始化TensorFlow会话 tensorflow::SessionOptions options; std::unique_ptr<tensorflow::Session> session(tensorflow::NewSession(options)); // 加载SavedModel模型 tensorflow::Status status = tensorflow::LoadSavedModel(options, tensorflow::RunOptions(), "path/to/your_saved_model", {"serve"}, &session); if (!status.ok()) { std::cerr << "模型加载失败: " << status.ToString() << std::endl; return 1; } // 读取并预处理图像 tensorflow::Tensor input_tensor; status = ReadFile(tensorflow::Env::Default(), "your_test_image.png", &input_tensor); if (!status.ok()) { std::cerr << "图像预处理失败: " << status.ToString() << std::endl; return 1; } // 调整张量形状为模型要求的batch格式(比如[1, 224, 224, 3],1是batch size) tensorflow::Tensor input_batch(tensorflow::DT_FLOAT, tensorflow::TensorShape({1, 224, 224, 3})); auto input_batch_flat = input_batch.flat<float>(); auto input_flat = input_tensor.flat<float>(); std::copy_n(input_flat.data(), input_flat.size(), input_batch_flat.data()); // 执行推理 // 注意:这里的输入输出节点名称要和你的SavedModel实际节点名一致 std::vector<std::pair<std::string, tensorflow::Tensor>> inputs = { {"serving_default_input_1:0", input_batch} // 替换成你的模型输入节点名 }; std::vector<tensorflow::Tensor> outputs; status = session->Run(inputs, {"StatefulPartitionedCall:0"}, {}, &outputs); // 替换成你的模型输出节点名 if (!status.ok()) { std::cerr << "推理失败: " << status.ToString() << std::endl; return 1; } // 处理并打印推理结果 auto output_flat = outputs[0].flat<float>(); std::cout << "推理结果: "; for (int i = 0; i < output_flat.size(); ++i) { std::cout << output_flat(i) << " "; } std::cout << std::endl; session->Close(); return 0; }
三、几个关键注意事项(我踩过的坑)
- 节点名称要对应:你可以用
tensorflow saved_model_cli show --dir path/to/your_saved_model --all命令查看SavedModel的输入输出节点名称,代码里的名称必须和实际一致,不然会报错。 - 预处理逻辑必须对齐:这个是最容易出错的地方!比如你训练Keras模型时用了
tf.keras.applications.resnet50.preprocess_input,那C++里也要做相同的均值减法和缩放,不能只简单除以255,否则推理结果会完全不对。 - 编译环境配置:编译C代码时要链接TensorFlow的C库,用CMake的话要设置好TensorFlow的include路径和lib路径,确保依赖正确。
内容的提问来源于stack exchange,提问作者Illuminati0x5B
相关产品推荐
相关产品推荐

