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

