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

TF pipeline动态提取图像patch并打平数据集的实现方案

TensorFlow动态提取图像patch的流水线实现方案

核心解决思路

你现有方案的问题是没有将单张图像提取出的patch集合拆分为独立的数据集元素,导致shuffle和batch的作用对象不符合预期。通过unbatch()操作拆分patch集合,即可实现单patch级别的混洗与采样,全程流式处理无需将所有patch存入内存。

完整实现代码

PATCH_SIZE = 64

def extract_patches(img, patch_size=PATCH_SIZE, stride=PATCH_SIZE//2):
    # 输入单张图像 shape: (256, 512, 1)
    n_channels = img.shape[-1]
    # 增加batch维度适配tf.image.extract_patches接口要求
    img = tf.expand_dims(img, axis=0)  
    patches = tf.image.extract_patches(
        img,
        sizes=[1, patch_size, patch_size, n_channels],
        strides=[1, stride, stride, n_channels],
        rates=[1, 1, 1, 1],
        padding='VALID'
    )
    # 转换为(单图patch总数, patch_size, patch_size, 通道数)格式
    return tf.reshape(patches, (-1, patch_size, patch_size, n_channels))

batch_size = 8
dataset = (tf.data.Dataset.from_tensor_slices(tf.cast(imgs, tf.float32))
            # 逐张图像提取patch,输出每个元素为对应图像的所有patch集合
            .map(extract_patches, num_parallel_calls=tf.data.AUTOTUNE, deterministic=False)
            # 拆分patch集合,每个数据集元素对应单个patch,shape为(64, 64, 1)
            .unbatch()
            # 单patch级别混洗,buffer_size可根据内存情况调整,数值越大混洗效果越好
            .shuffle(buffer_size=1000, reshuffle_each_iteration=True)
            # 按设定批次大小拼接
            .batch(batch_size)
            # 预取数据提升流水线运行效率
            .prefetch(tf.data.AUTOTUNE)
          )

方案说明

  • 第一版原有实现的问题:shuffle和batch操作的对象是「单图对应的所有patch组成的集合」,所以batch后会多一个维度,输出shape为(batch_size, 单图patch数, 64, 64, 1)
  • 第二版原有实现的问题:先batch多张图像再提取patch,会直接把批次内所有图像的patch拼在一起,导致单批次patch数远大于设定的batch_size
  • 新增的unbatch()是整个方案的核心,它会把map阶段输出的每个patch集合拆分为独立的样本,整个流水线全程流式处理,不需要一次性把所有patch加载到内存中,最终输出的批次shape完全符合预期的(batch_size, 64, 64, 1)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 06:15:03