如何可视化TensorFlow Model Garden中pipeline.config配置的数据增强结果?
可视化TensorFlow Model Garden目标检测训练中的数据增强结果
核心说明
训练过程默认不会保存增强后的数据,想要查看数据增强效果,需要通过代码实现可视化,有两种实用方案可选:
方案一:修改model_main_tf2.py实时查看训练中的增强样本
直接在训练入口文件中添加可视化代码,训练时就能生成增强后的样本图:
- 打开
model_main_tf2.py,找到加载训练数据集的代码段(通常是调用dataset_builder.build或定义input_fn的部分)。 - 在数据集迭代器生成batch数据后,插入以下可视化代码(仅处理前几个batch,避免拖慢训练):
import matplotlib.pyplot as plt import matplotlib.patches as patches # 假设dataset是训练数据集的迭代器 for idx, batch in enumerate(dataset): if idx > 5: # 只可视化前5个batch break # 取出单张图像和标注框 img_tensor = batch['image'][0] boxes_tensor = batch['groundtruth_boxes'][0] # 还原归一化的图像(根据你pipeline.config中的normalize参数调整均值/标准差) mean = tf.constant([0.485, 0.456, 0.406]) std = tf.constant([0.229, 0.224, 0.225]) img = img_tensor.numpy() * std + mean img = tf.clip_by_value(img, 0, 1).numpy() # 确保像素值在0-1区间 # 绘制图像和标注框 fig, ax = plt.subplots(1, dpi=100) ax.imshow(img) img_h, img_w = img.shape[:2] # 将归一化的框坐标转为像素坐标 for box in boxes_tensor.numpy(): ymin, xmin, ymax, xmax = box rect = patches.Rectangle( (xmin * img_w, ymin * img_h), (xmax - xmin) * img_w, (ymax - ymin) * img_h, linewidth=2, edgecolor='red', facecolor='none' ) ax.add_patch(rect) plt.axis('off') plt.savefig(f'train_aug_sample_{idx}.png', bbox_inches='tight') plt.close()
- 运行训练脚本,当前目录下会生成前5个batch的增强样本图。如果是远程训练,直接下载图片查看即可。
方案二:单独写脚本离线验证增强效果(无需修改训练文件)
这种方法更灵活,不用改动训练代码,直接加载你的配置和数据集验证:
创建visualize_augmentation.py,写入以下代码:
import tensorflow as tf from object_detection.utils import config_util from object_detection.builders import dataset_builder import matplotlib.pyplot as plt import matplotlib.patches as patches # 替换为你的pipeline.config路径 CONFIG_PATH = './pipeline.config' # 加载配置并构建训练数据集 configs = config_util.get_configs_from_pipeline_file(CONFIG_PATH) train_dataset = dataset_builder.build( input_config=configs['train_input_config'], config=configs['train_config'], is_training=True ) # 可视化增强样本 for idx, batch in enumerate(train_dataset): if idx > 5: break img_tensor = batch['image'][0] boxes_tensor = batch['groundtruth_boxes'][0] # 还原归一化图像 mean = tf.constant([0.485, 0.456, 0.406]) std = tf.constant([0.229, 0.224, 0.225]) img = img_tensor.numpy() * std + mean img = tf.clip_by_value(img, 0, 1).numpy() # 绘制图像和标注框 fig, ax = plt.subplots(1, dpi=100) ax.imshow(img) img_h, img_w = img.shape[:2] for box in boxes_tensor.numpy(): ymin, xmin, ymax, xmax = box rect = patches.Rectangle( (xmin * img_w, ymin * img_h), (xmax - xmin) * img_w, (ymax - ymin) * img_h, linewidth=2, edgecolor='red', facecolor='none' ) ax.add_patch(rect) plt.axis('off') plt.savefig(f'offline_aug_sample_{idx}.png', bbox_inches='tight') plt.close()
运行脚本后,就能在当前目录看到增强后的样本图,验证你的random_horizontal_flip、random_image_scale等策略是否生效。
注意事项
- 如果你在
pipeline.config中自定义了normalize_image的均值和标准差,一定要替换代码中的对应值,否则图像还原会失真。 - 数据增强是动态实时生成的,训练过程中不会自动保存增强数据,必须主动通过代码提取样本。
内容的提问来源于stack exchange,提问作者Katzoos
相关产品推荐
相关产品推荐

