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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 06:45:12