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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 10:32:31