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

求助:TFX管道Transform组件目标检测预处理代码编写

TFX Transform组件目标检测预处理实现示例

核心转换模块示例(transform_module.py)

这是你需要指定给TRANSFORM_MODULE的Python文件,实现图像缩放、带边界框同步调整的数据增强逻辑:

import tensorflow as tf
import tensorflow_transform as tft

# 定义目标图像尺寸
TARGET_HEIGHT = 640
TARGET_WIDTH = 640

def preprocessing_fn(inputs):
    """TFX Transform的核心预处理函数"""
    # 从输入中提取特征(需与你的TFRecord schema匹配)
    image = inputs['image/encoded']
    original_height = tf.cast(inputs['image/height'], tf.float32)
    original_width = tf.cast(inputs['image/width'], tf.float32)
    xmin = inputs['bbox/xmin']
    xmax = inputs['bbox/xmax']
    ymin = inputs['bbox/ymin']
    ymax = inputs['bbox/ymax']
    labels = inputs['bbox/label']

    # 解码图像
    image = tf.map_fn(
        lambda img: tf.image.decode_jpeg(img, channels=3),
        image,
        dtype=tf.uint8
    )

    # 计算缩放因子
    scale_height = TARGET_HEIGHT / original_height
    scale_width = TARGET_WIDTH / original_width

    # 缩放图像
    image = tf.image.resize(image, [TARGET_HEIGHT, TARGET_WIDTH])

    # 同步调整边界框坐标(基于原图尺寸缩放)
    xmin = xmin * scale_width
    xmax = xmax * scale_width
    ymin = ymin * scale_height
    ymax = ymax * scale_height

    # 仅在训练阶段应用数据增强
    split_handle = tft.get_split_handle()
    is_train = tf.equal(split_handle, 'train')

    def apply_augmentation():
        """训练阶段的增强操作:随机水平翻转"""
        # 随机翻转图像
        flipped_image = tf.image.random_flip_left_right(image)
        
        # 同步翻转边界框
        flipped_xmin = TARGET_WIDTH - xmax
        flipped_xmax = TARGET_WIDTH - xmin
        
        return flipped_image, flipped_xmin, flipped_xmax, ymin, ymax

    def no_augmentation():
        """评估阶段仅返回原始处理后的数据"""
        return image, xmin, xmax, ymin, ymax

    # 根据是否为训练分支选择操作
    image, xmin, xmax, ymin, ymax = tf.cond(
        is_train,
        apply_augmentation,
        no_augmentation
    )

    # 将图像归一化到[0,1]范围
    image = tf.cast(image, tf.float32) / 255.0

    # 整理输出特征(需与后续模型输入匹配)
    outputs = {
        'image': image,
        'bbox/xmin': xmin,
        'bbox/xmax': xmax,
        'bbox/ymin': ymin,
        'bbox/ymax': ymax,
        'bbox/label': labels
    }

    return outputs

关键实现要点

  • 边界框同步调整:所有针对图像的几何变换(缩放、翻转)都必须同步修改bbox坐标,否则会导致标注与图像不匹配。
  • 训练/评估分支区分:使用tft.get_split_handle()判断当前处理的是训练还是评估数据集,仅在训练阶段应用数据增强,保证评估数据的一致性。
  • TF图模式兼容:所有预处理操作必须使用TensorFlow的图操作(如tf.map_fn、tf.cond),不能使用eager模式代码,因为TFX Transform是基于TF图执行的。
  • 特征匹配schema:确保你的输入特征名称(如image/encoded、bbox/xmin)与TFRecord数据集及自动生成的schema完全一致。

管道集成调整

在你的管道代码中,指定TRANSFORM_MODULE为上述文件的绝对路径:

# 添加这行代码,替换为你的transform_module.py路径
TRANSFORM_MODULE = os.path.abspath('transform_module.py')

transform = Transform(
    examples=example_gen.outputs['examples'],
    schema=infer_schema.outputs['result'],
    module_file=TRANSFORM_MODULE
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 11:45:34