使用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
相关产品推荐
相关产品推荐

