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

使用tflite_model_maker训练如何在Tensorboard展示检测结果及img_summary_steps报错

问题原因

你触发报错的核心原因是tflite_model_maker 0.3.2版本中,sample_image配置项预期接收预处理完成的图像张量,你直接传入本地图片路径字符串,底层推理逻辑无法识别字符串输入,因此触发类型不匹配报错。另外默认配置里num_classes为COCO数据集的90类,如果你用自定义数据集没有修改该参数,也会触发维度不匹配问题。

解决方法

方案1:使用内置img_summary_steps逻辑

步骤1:预处理示例图片为张量

import tensorflow as tf
from PIL import Image
import numpy as np

# 按模型要求预处理示例图片
img_path = "/content/data2/IMG_0187.JPG"
img = Image.open(img_path).convert("RGB")
# 和spec.config.image_size保持一致,efficientdet_lite0默认是320*320
img = img.resize((320, 320))
img_arr = np.array(img) / 255.0
# 增加batch维度,转为float32张量
sample_img_tensor = tf.convert_to_tensor(img_arr[np.newaxis, ...], dtype=tf.float32)

步骤2:修正配置项

spec = model_spec.get('efficientdet_lite0')

# 传入预处理后的张量,而非文件路径
setattr(spec.config, "sample_image", sample_img_tensor)
setattr(spec.config, "profile", True)
# 不要设置为1,避免过密输出拖慢训练,这里设为每1个epoch输出一次,可自行调整
setattr(spec.config, "img_summary_steps", len(train_data)//8)
# 自定义数据集必须修改为你的实际类别数
setattr(spec.config, "num_classes", 你的自定义类别数量)

步骤3:启动训练并查看TensorBoard

训练代码保持不变,训练启动后在Colab中运行以下命令加载TensorBoard:

%load_ext tensorboard
%tensorboard --logdir {spec.config.model_dir}

可视化结果会出现在TensorBoard的Images标签页下。


方案2:自定义可视化回调(兼容性更强)

如果内置逻辑仍有兼容问题,可以自己实现TensorBoard图像写入逻辑,不受版本限制:

import tensorflow as tf
import datetime
from tensorflow.keras.callbacks import TensorBoard

# 日志存储路径
log_dir = "logs/fit/" + datetime.datetime.now().strftime("%Y%m%d-%H%M%S")
file_writer = tf.summary.create_file_writer(log_dir + "/detection_results")
tb_callback = TensorBoard(log_dir=log_dir, histogram_freq=1)

# 自定义每个epoch结束后写入检测结果的回调
def log_detection_result(epoch, logs):
    # 推理示例图片
    pred = model.predict(sample_img_tensor)[0]
    # 归一化检测框坐标(tf.image.draw_bounding_boxes要求坐标为[y_min, x_min, y_max, x_max],范围0-1)
    boxes = pred['detection_boxes'][np.newaxis, ...]
    # 用红色绘制检测框
    img_with_boxes = tf.image.draw_bounding_boxes(
        sample_img_tensor, 
        boxes, 
        colors=[[1.0, 0.0, 0.0]]
    )
    # 写入TensorBoard日志
    with file_writer.as_default():
        tf.summary.image("检测结果示例", img_with_boxes, step=epoch)

image_callback = tf.keras.callbacks.LambdaCallback(on_epoch_end=log_detection_result)

# 训练时传入回调
model = object_detector.create(
    train_data,                               
    model_spec=spec,
    batch_size=8,
    train_whole_model=False,
    epochs = 100,
    validation_data = valid,
    do_train = True,
    callbacks=[tb_callback, image_callback]
)

内容的提问来源于stack exchange,提问作者Roy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 15:45:02