如何使用Python读写TFOD2的pipeline.config配置文件
TensorFlow Object Detection pipeline.config 动态读写修改方案
1. 核心依赖说明
TensorFlow官方提供的pipeline.config是基于Protobuf格式定义的,我们可以直接用官方预生成的配置类读写,不需要自己写解析逻辑,操作前请确保已经正确安装TensorFlow Object Detection API。
2. 完整操作步骤
- 第一步:导入所需依赖
import tensorflow as tf from object_detection.utils import config_util from object_detection.protos import pipeline_pb2 from google.protobuf import text_format
- 第二步:读取配置文件
两种常用读取方式,可按需选择:
# 方式1:用官方封装的config_util工具读取,返回配置字典+配置proto对象 configs = config_util.get_configs_from_pipeline_file('path/to/your/pipeline.config') pipeline_proto = config_util.create_pipeline_proto_from_configs(configs) # 方式2:直接用protobuf解析原始配置文本 pipeline_proto = pipeline_pb2.TrainEvalPipelineConfig() with tf.io.gfile.GFile('path/to/your/pipeline.config', 'r') as f: text_format.Merge(f.read(), pipeline_proto)
- 第三步:动态修改参数
直接对proto对象的属性赋值即可,以下是常见修改场景示例:
# 修改训练批量大小 pipeline_proto.train_config.batch_size = 16 # 修改训练集标注文件、TFRecord路径 pipeline_proto.train_input_reader.label_map_path = 'path/to/label_map.pbtxt' pipeline_proto.train_input_reader.tf_record_input_reader.input_path[0] = 'path/to/train.tfrecord' # 修改评估任务的 checkpoint 路径 pipeline_proto.eval_config.checkpoint_path = 'path/to/model/ckpt-xxx' # 修改检测头的类别数(SSD模型示例,Faster RCNN对应修改faster_rcnn下的同名属性即可) pipeline_proto.model.ssd.num_classes = 5
- 第四步:落盘配置/直接传入运行接口
如果需要保存修改后的配置供后续使用:
# 导出为新的配置文件 with open('path/to/new_pipeline.config', 'w') as f: f.write(text_format.MessageToString(pipeline_proto)) # 也可以直接把修改后的configs传入官方训练/评估接口,无需落盘
3. 注意事项
- 修改参数前要对应你使用的模型结构找对应层级,例如Faster RCNN模型的参数存放在
pipeline_proto.model.faster_rcnn路径下,不要写错层级 - 列表类型的参数(如支持多输入的TFRecord路径)如果要新增元素,使用
append()方法操作,不要直接赋值 - 不确定参数层级时,可以先打印读取的proto对象查看完整结构:
print(pipeline_proto)
内容的提问来源于stack exchange,提问作者Bappy Ahmed
相关产品推荐
相关产品推荐

