TensorFlow C_API加载Faster R-CNN模型报错‘Bitcast未注册’求解
解决TensorFlow C_API加载Faster R-CNN模型时的Bitcast算子未注册问题
问题描述
使用Model Zoo预训练模型训练自定义Faster R-CNN模型后,通过TensorFlow C_API加载SavedModel进行预测,加载代码如下:
TF_SessionOptions* session_opts = TF_NewSessionOptions(); TF_Buffer* run_opts = NULL; TF_Status* status = TF_NewStatus(); TF_Graph* graph = TF_NewGraph(); TF_Buffer* meta_graph = TF_NewBuffer(); const char* saved_model_dir = "/saved_model"; const char* tags[] = {"serve"}; int ntags = 1; TF_Session* sess = TF_LoadSessionFromSavedModel(session_opts, run_opts, saved_model_dir, tags, ntags, graph, NULL, status);
加载时出现以下错误:
F tensorflow/compiler/tf2xla/xla_op_registry.cc:230] Non-OK-status: lookup_status status: NOT_FOUND: Op type not registered 'Bitcast' in binary running on lin-VirtualBox. Make sure the Op and Kernel are registered in the binary running in this process. Note that if you are loading a saved graph which used ops from tf.contrib, accessing (e.g.) `tf.contrib.resampler` should be done before importing the graph, as contrib ops are lazily registered when the module is first accessed. Aborted (core dumped)
环境为TensorFlow 2.7.0、Ubuntu 20.04,已尝试重新冻结模型但问题仍存在。
解决方案
1. 确保TensorFlow C_API库与训练版本严格一致
训练使用的是TensorFlow 2.7.0,必须下载对应版本的TensorFlow C_API库(包含libtensorflow.so和头文件)来编译你的C++代码。算子未注册的常见原因是C_API库版本与训练环境TF版本不匹配,Bitcast算子在不同TF版本的C_API支持上存在差异,版本一致是基础前提。
2. 导出模型时禁用XLA编译
错误信息指向tf2xla/xla_op_registry.cc,说明模型训练或导出时可能启用了XLA优化,但TensorFlow C_API默认未加载XLA相关算子。导出SavedModel时添加禁用XLA的配置:
import tensorflow as tf from object_detection.utils import config_util from object_detection.builders import model_builder # 加载模型配置与训练权重 configs = config_util.get_configs_from_pipeline_file('path/to/your/pipeline.config') model_config = configs['model'] detection_model = model_builder.build(model_config=model_config, is_training=False) detection_model.load_weights('path/to/your/checkpoint') # 禁用XLA并导出SavedModel tf.config.optimizer.set_jit(False) tf.saved_model.save(detection_model, '/saved_model')
3. 转换为冻结图(Frozen Graph)加载
尝试将训练好的模型转换为冻结图(.pb格式),冻结图会把算子和权重打包为单个文件,降低版本兼容性问题。转换代码示例:
import tensorflow as tf from object_detection.exporter import export_inference_graph export_inference_graph( input_type='image_tensor', pipeline_config_path='path/to/pipeline.config', trained_checkpoint_prefix='path/to/model.ckpt', output_directory='path/to/frozen_graph' )
然后使用C_API加载冻结图,示例代码:
TF_Graph* graph = TF_NewGraph(); TF_Status* status = TF_NewStatus(); TF_Buffer* graph_def = TF_NewBuffer(); // 读取冻结图文件 FILE* f = fopen("path/to/frozen_inference_graph.pb", "rb"); fseek(f, 0, SEEK_END); long fsize = ftell(f); fseek(f, 0, SEEK_SET); char* data = (char*)malloc(fsize); fread(data, fsize, 1, f); fclose(f); graph_def->data = data; graph_def->length = fsize; graph_def->data_deallocator = [](void* data, size_t length) { free(data); }; // 导入图结构 TF_ImportGraphDefOptions* opts = TF_NewImportGraphDefOptions(); TF_GraphImportGraphDef(graph, graph_def, opts, status); TF_DeleteImportGraphDefOptions(opts); TF_DeleteBuffer(graph_def); if (TF_GetCode(status) != TF_OK) { printf("%s", TF_Message(status)); return 1; } // 创建会话 TF_SessionOptions* sess_opts = TF_NewSessionOptions(); TF_Session* sess = TF_NewSession(graph, sess_opts, status);
4. 手动注册XLA算子(备选方案)
如果必须保留XLA优化,需要在C++代码中手动注册Bitcast等XLA相关算子,同时编译时链接TensorFlow的XLA库(如libtensorflow_compiler.so),示例代码片段:
#include "tensorflow/compiler/tf2xla/xla_op_registry.h" #include "tensorflow/core/framework/op_kernel.h" // 注册Bitcast算子的XLA实现 REGISTER_XLA_OP(Name("Bitcast"), tensorflow::XlaBitcastOp);
此方法会增加编译复杂度,优先推荐前三种方案。
内容的提问来源于stack exchange,提问作者Vin
相关产品推荐
相关产品推荐

