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

如何可视化TensorFlow Model Garden中pipeline.config配置的数据增强结果?

可视化TensorFlow Model Garden目标检测训练中的数据增强结果

核心说明

训练过程默认不会保存增强后的数据,想要查看数据增强效果,需要通过代码实现可视化,有两种实用方案可选:

方案一:修改model_main_tf2.py实时查看训练中的增强样本

直接在训练入口文件中添加可视化代码,训练时就能生成增强后的样本图:

  1. 打开model_main_tf2.py,找到加载训练数据集的代码段(通常是调用dataset_builder.build或定义input_fn的部分)。
  2. 在数据集迭代器生成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()
  1. 运行训练脚本,当前目录下会生成前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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 00:50:23