如何将自训练SSD目标检测模型集成至TensorFlow目标计数API
已训练SSD模型集成到目标计数API的实现步骤
一、导出训练模型为冻结图
你手头的模型检查点(.ckpt文件)无法直接被计数API加载,需要先导出为frozen_inference_graph.pb格式:
- 用TFODCourse提供的导出脚本(如
export_tflite_graph_tf2.py,根据你使用的TensorFlow版本选择),执行以下命令(替换路径为你的实际路径):
python export_tflite_graph_tf2.py \ --pipeline_config_path=./your_training_config/pipeline.config \ --trained_checkpoint_dir=./your_checkpoint_dir \ --output_directory=./saved_frozen_graph
导出完成后,output_directory下会生成frozen_inference_graph.pb文件,这是计数API需要的模型文件。
二、加载模型图到计数API
计数API的核心函数cumulative_object_counting_y_axis需要detection_graph参数,添加以下代码加载冻结图:
import tensorflow as tf from object_detection.utils import label_map_util # 加载冻结图 detection_graph = tf.Graph() with detection_graph.as_default(): od_graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile('./saved_frozen_graph/frozen_inference_graph.pb', 'rb') as fid: serialized_graph = fid.read() od_graph_def.ParseFromString(serialized_graph) tf.import_graph_def(od_graph_def, name='')
三、生成类别索引映射
计数API需要category_index来关联检测出的类别ID和名称,从你训练用的label_map.pbtxt生成:
category_index = label_map_util.create_category_index_from_labelmap( './your_training_config/label_map.pbtxt', use_display_name=True )
四、调用计数核心函数
准备好所有参数后,直接调用API的核心函数:
# 替换为你的实际参数 cumulative_object_counting_y_axis( input_video='./input_video.mp4', detection_graph=detection_graph, category_index=category_index, is_color_recognition_enabled=False, roi=400, # 计数线的Y轴坐标,根据视频场景调整 deviation=20, # 目标穿过计数线的允许偏差范围 custom_object_name='你的目标类别名', # 如"car"、"person" targeted_objects=[1] # 要计数的类别ID,对应label_map中的ID )
五、环境与依赖适配
- 确保安装API依赖:
opencv-python、tensorflow(建议用TF2.x,API兼容TF1模式)、numpy。 - 把计数API中的
vis_util.py放到项目路径下,或者确保代码能正确导入vis_util模块(核心函数依赖其中的可视化方法)。
关键注意事项
- 版本兼容:API基于TF1.x编写,用
tf.compat.v1.Session运行,TF2训练的模型导出冻结图时需确保兼容TF1图结构。 - 类别ID匹配:
targeted_objects必须和label_map.pbtxt中的类别ID一致,否则会出现计数错误。 - ROI线调整:
roi参数决定计数位置,需根据视频中目标的移动路径设置合适的Y坐标。
内容的提问来源于stack exchange,提问作者gestogloria
相关产品推荐
相关产品推荐

