You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何从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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 07:07:47