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

如何使用C_API导入训练好的tf.contrib.learn.DNNClassifier?

使用TensorFlow C API导入训练好的DNNClassifier的解决方案

我来帮你排查并解决加载SavedModel失败的问题,核心问题大概率出在标签使用错误以及代码细节的缺失上,下面一步步来梳理:

一、先搞清楚DNNClassifier保存的SavedModel关键信息

DNNClassifier作为Estimator的子类,保存的SavedModel默认使用的标签是"serve",而不是你代码里的"import"。你可以先通过TensorFlow自带的工具确认模型的元数据,在终端执行以下命令:

saved_model_cli show --dir path_to_trained_model/folder_with_trained_model --all

这个命令会输出模型的标签集合、输入输出签名等信息,比如你会看到类似这样的内容:

MetaGraphDef with tag-set: 'serve' contains the following SignatureDefs:
signature_def['serving_default']:
inputs['inputs'] tensor_info:
dtype: DT_FLOAT
shape: (-1, 你的特征数)
name: serving_default_inputs:0
outputs['classes'] tensor_info:
dtype: DT_STRING
shape: (-1, 类别数)
name: StatefulPartitionedCall:0
outputs['scores'] tensor_info:
dtype: DT_FLOAT
shape: (-1, 类别数)
name: StatefulPartitionedCall:1

这些信息对你后续获取输入输出张量非常重要。

二、修正后的C API加载代码

下面是补全错误处理、修正标签后的完整代码片段:

#include <tensorflow/c/c_api.h>
#include <stdio.h>

int main() {
    const char* export_dir = "path_to_trained_model/folder_with_trained_model";
    // 替换为DNNClassifier默认的"serve"标签
    const char* tags[] = {"serve"};
    int num_tags = 1;

    // 初始化TensorFlow核心对象
    TF_Graph* graph = TF_NewGraph();
    TF_SessionOptions* session_options = TF_NewSessionOptions();
    TF_Buffer* run_options = TF_NewBuffer(); // 如果不需要自定义运行选项,也可以传NULL
    TF_Status* status = TF_NewStatus();

    // 加载SavedModel到会话和图中
    TF_Session* session = TF_LoadSessionFromSavedModel(
        session_options,
        run_options,
        export_dir,
        tags,
        num_tags,
        graph,
        NULL, // 不需要获取meta_graph_def的话传NULL即可
        status
    );

    // 检查加载是否成功
    if (TF_GetCode(status) != TF_OK) {
        fprintf(stderr, "加载SavedModel失败:%s\n", TF_Message(status));
        // 清理资源
        TF_DeleteStatus(status);
        TF_DeleteSessionOptions(session_options);
        TF_DeleteBuffer(run_options);
        TF_DeleteGraph(graph);
        return 1;
    }

    printf("SavedModel加载成功!\n");

    // 示例:获取输入操作(名称来自saved_model_cli的输出)
    TF_Operation* input_op = TF_GraphOperationByName(graph, "serving_default_inputs");
    if (input_op == NULL) {
        fprintf(stderr, "找不到输入操作,请检查saved_model_cli输出的输入名称\n");
        // 清理资源
        TF_DeleteSession(session, status);
        TF_DeleteStatus(status);
        TF_DeleteSessionOptions(session_options);
        TF_DeleteBuffer(run_options);
        TF_DeleteGraph(graph);
        return 1;
    }

    // 后续可以继续获取输出操作、构造输入张量、运行推理等

    // 最后记得清理所有资源
    TF_DeleteSession(session, status);
    TF_DeleteStatus(status);
    TF_DeleteSessionOptions(session_options);
    TF_DeleteBuffer(run_options);
    TF_DeleteGraph(graph);

    return 0;
}

三、常见的加载失败原因排查

如果还是失败,可以从这几个方向检查:

  • 模型路径错误:确保export_dir指向包含saved_model.pb文件和variables文件夹的目录,而不是父目录或子目录
  • 版本不兼容:C API的TensorFlow版本要和训练模型时的Python版本保持一致(比如都是2.x系列),跨大版本可能出现兼容性问题
  • 资源权限问题:确保程序有读取模型目录的权限

内容的提问来源于stack exchange,提问作者M.cat

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 10:09:14