如何使用tf.data API创建图像序列样本并实现图像拼接?
解决tf.data API拼接图像序列样本的问题
我完全懂你的痛点——想用tf.data构建灵活的图像序列数据管道,却卡在了window分组后的拼接步骤上,而且不想用TFRecords那种不够灵活还占空间的方案对吧?别担心,只需要在window之后加几步简单的处理,就能得到你想要的N x W x H x T x C格式的批次。
先给你梳理核心思路:dataset.window()返回的是Dataset of Datasets,每个子Dataset对应一个图像序列窗口,我们需要把这些子Dataset里的图像元素打包成单个张量,再调整维度顺序就能符合要求。
第一步:调整图像加载函数
你原来的load_and_process_image里多了一个不必要的维度(shape=(IMG_WIDTH, IMG_HEIGHT, 1, 3)),这会增加后续拼接的复杂度,先把它去掉,让单张图像的形状保持为(W, H, C):
import tensorflow as tf from glob import glob IMG_WIDTH = 256 IMG_HEIGHT = 256 def load_and_process_image(path): img = tf.io.read_file(path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, [IMG_WIDTH, IMG_HEIGHT]) # 保持单张图像形状为(256, 256, 3),去掉多余的维度 return img
第二步:完善create_dataset函数
接下来修改create_dataset,重点处理window后的拼接和维度调整:
def create_dataset(files, time_distance=8, frame_step=1, batch_size=32): dataset = tf.data.Dataset.from_tensor_slices(files) # 并行加载图像,提升效率 dataset = dataset.map(load_and_process_image, num_parallel_calls=tf.data.AUTOTUNE) # 生成包含time_distance张图像的窗口,shift=frame_step控制序列步长,drop_remainder=True避免最后一个不完整的窗口 dataset = dataset.window(time_distance, shift=frame_step, drop_remainder=True) # 将每个窗口的子Dataset转为扁平结构,并打包成序列张量(形状为(T, W, H, C)) dataset = dataset.flat_map(lambda window: window.batch(time_distance)) # 调整维度顺序:从(T, W, H, C)转为(W, H, T, C),满足你需要的格式 dataset = dataset.map(lambda seq: tf.transpose(seq, perm=[1, 2, 0, 3])) # 打包成批次,并开启预取优化 dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE) return dataset
关键步骤解释
flat_map(lambda window: window.batch(time_distance)):每个窗口是一个子Dataset,用batch(time_distance)把窗口里的T张图像打包成一个形状为(T, 256, 256, 3)的张量,flat_map则把这些张量平铺到主Dataset中。tf.transpose(seq, perm=[1, 2, 0, 3]):把时间维度(原来的第0维)移到第2位,让张量形状变成(256, 256, T, 3),后续batch之后就会得到(N, 256, 256, T, 3)的批次,完全符合你的需求。- 加入
num_parallel_calls和prefetch是为了让数据加载和模型训练并行,避免训练过程中等待数据。
测试验证
你可以用下面的代码测试一下输出形状是否符合预期:
files = sorted(glob('some/path/*.jpg')) dataset = create_dataset(files) # 取第一个批次查看形状 for batch in dataset.take(1): print(batch.shape) # 输出示例:(32, 256, 256, 8, 3)(对应batch_size=32,time_distance=8)
这样整个数据管道既灵活又高效,完全不需要依赖TFRecords就能生成符合要求的图像序列样本。
内容的提问来源于stack exchange,提问作者Marc Seibert
相关产品推荐
相关产品推荐

