如何从TensorFlow Object Detection API模型获取模型摘要?
问题描述
使用TensorFlow Object Detection API通过pipeline配置加载SSD模型时,调用summary()方法报错,代码及报错情况如下:
import tensorflow as tf from object_detection.utils import config_util from object_detection.builders import model_builder tf.keras.backend.clear_session() pipeline_config = 'models/research/object_detection/configs/tf2/ssd_resnet50_v1_fpn_640x640_coco17_tpu-8.config' configs = config_util.get_configs_from_pipeline_file(pipeline_config) model_config = configs['model'] detection_model = model_builder.build(model_config=model_config, is_training=False) detection_model.summary() # AttributeError: 'SSDMetaArch' object has no attribute 'summary'
尝试预处理输入并遍历模型变量调用summary(),仍无法解决:
image, shapes = detection_model.preprocess(tf.zeros([1, 640, 640, 3])) prediction_dict = detection_model.predict(image, shapes) detection_model.postprocess(prediction_dict, shapes) for variable in detection_model.variables: variable.summary() # Throws AttributeError, since none of these have .summary()
使用TensorFlow版本:2.10.1
解决方案
方法1:将模型包装为标准Keras Functional模型
SSDMetaArch并非原生Keras Model实例,无法直接调用summary()。可以通过构建输入输出链路,将其包装为Keras模型:
import tensorflow as tf from object_detection.utils import config_util from object_detection.builders import model_builder tf.keras.backend.clear_session() pipeline_config = 'models/research/object_detection/configs/tf2/ssd_resnet50_v1_fpn_640x640_coco17_tpu-8.config' configs = config_util.get_configs_from_pipeline_file(pipeline_config) model_config = configs['model'] detection_model = model_builder.build(model_config=model_config, is_training=False) # 定义输入张量 input_tensor = tf.keras.Input(shape=(640, 640, 3), dtype=tf.float32) # 模拟模型完整前向流程 preprocessed_img, shapes = detection_model.preprocess(input_tensor) prediction_dict = detection_model.predict(preprocessed_img, shapes) output_detections = detection_model.postprocess(prediction_dict, shapes) # 包装为Keras模型 keras_detection_model = tf.keras.Model(inputs=input_tensor, outputs=output_detections) keras_detection_model.summary()
方法2:直接打印模型变量信息
如果仅需查看模型权重、偏置等变量的基本信息,可直接遍历变量并打印名称与形状:
for var in detection_model.variables: print(f"变量名:{var.name} | 形状:{var.shape}")
方法3:用TensorBoard可视化模型结构
通过TensorBoard追踪模型计算图,实现可视化查看:
# 生成示例输入 sample_input = tf.zeros([1, 640, 640, 3]) # 开启图追踪 tf.summary.trace_on(graph=True, profiler=True) # 运行一次模型前向流程 preprocessed, shapes = detection_model.preprocess(sample_input) preds = detection_model.predict(preprocessed, shapes) detection_model.postprocess(preds, shapes) # 将追踪结果写入日志 with tf.summary.create_file_writer('./detection_model_logs').as_default(): tf.summary.trace_export( name='ssd_model_trace', step=0, profiler_outdir='./detection_model_logs' )
执行完代码后,在终端运行tensorboard --logdir=./detection_model_logs,打开浏览器访问对应地址即可查看完整模型图。
内容的提问来源于stack exchange,提问作者wordhydrogen
相关产品推荐
相关产品推荐

