如何通过代码调用TensorFlow Object Detection API的model_main_tf2且避免程序退出
解决方案1:直接调用main函数,绕过tf.app.run()
tf.compat.v1.app.run() 内部执行逻辑是先解析命令行参数,再调用传入的main函数,最后执行sys.exit()退出进程,而model_main_tf2.main()本身不会主动退出进程,直接调用即可:
import sys from object_detection import model_main_tf2 # 原有参数配置逻辑保留 model_main_tf2.FLAGS.pipeline_config_path = pipeline_config_path model_main_tf2.FLAGS.model_dir = model_path # 可按需配置其他训练参数,例如训练步数 # model_main_tf2.FLAGS.num_train_steps = 10000 # 直接调用main函数,传入argv参数即可,无需通过app.run封装 model_main_tf2.main(sys.argv)
该方案改动最小,原有训练逻辑完全保留,训练结束后程序会继续执行后续代码。
解决方案2:直接调用核心训练接口,绕开model_main_tf2
如果需要更灵活的训练逻辑控制,可以直接调用TFOD API内部的训练核心接口model_lib_v2.train_loop,无需依赖model_main_tf2的FLAGS和入口逻辑:
from object_detection.utils import config_util from object_detection import model_lib_v2 # 可选:先加载配置自定义修改 # configs = config_util.get_configs_from_pipeline_file(pipeline_config_path) # 按需修改配置,例如修改训练步数、优化器参数等 # configs['train_config'].num_steps = 20000 # config_util.save_pipeline_config(configs, model_path) # 直接调用训练核心逻辑 model_lib_v2.train_loop( pipeline_config_path=pipeline_config_path, model_dir=model_path, # 可按需传入自定义参数,覆盖pipeline.config中的配置 # train_steps=20000, # eval_training_data=False, use_tpu=False )
该方案灵活性更高,支持自定义训练过程中的各类回调、多阶段训练、自定义评估逻辑等需求,同样不会触发进程退出。
内容的提问来源于stack exchange,提问作者chko
相关产品推荐
相关产品推荐

