如何使用C_API导入训练好的tf.contrib.learn.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

