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

如何使用TensorFlow Dataset API实现批量滑动窗口?

实现滑动窗口式的重叠批量(步长为1)

嘿,我完全理解你的需求——你想要把原来无重叠的固定批量改成每次滑动1个样本的重叠批量,也就是滑动窗口批量。在TensorFlow Dataset API中,直接用batch()方法做不到这点,不过我们可以用window()方法来实现,下面是修改后的代码和详细说明:

核心修改思路

原来的dataset.batch(batch_size)是将连续的batch_size个样本分组,没有重叠。要实现滑动步长为1的窗口,我们需要:

  1. 使用window()生成滑动窗口的子数据集
  2. 将每个子数据集合并成一个完整的批量

修改后的完整代码

def tfrecords_train_input(input_dir, examples, epochs, nsensors, past, future, features, batch_size, threads, shuffle, record_type):
    filenames = sorted(
        [os.path.join(input_dir, f) for f in os.listdir(input_dir)]
    )
    num_records = 0
    for fn in filenames:
        for _ in tf.python_io.tf_record_iterator(fn):
            num_records += 1
    print("Number of files to use:", len(filenames), "/ Total records to use:", num_records)
    
    dataset = tf.data.TFRecordDataset(filenames)
    
    # Parse records
    read_proto = partial(record_type().read_proto, nsensors=nsensors, past=past, future=future, features=features)
    # Parallelize Data Transformation on available GPU
    dataset = dataset.map(map_func=read_proto, num_parallel_calls=threads)
    
    # Cache data
    dataset = dataset.cache()
    
    # Repeat for epochs
    dataset = dataset.repeat(epochs)
    
    # 核心修改:替换batch为滑动窗口实现
    # size=batch_size:每个窗口包含的样本数(即批量大小)
    # shift=1:每次滑动1个样本
    # drop_remainder=True:丢弃最后一个不足batch_size的窗口(和原batch行为一致)
    dataset = dataset.window(size=batch_size, shift=1, drop_remainder=True)
    # 将每个窗口的子数据集合并成一个批量张量
    dataset = dataset.flat_map(lambda window: window.batch(batch_size))
    
    # Efficient Pipelining
    dataset = dataset.prefetch(2)
    
    iterator = dataset.make_one_shot_iterator()
    return iterator

关键修改点解释

  • window(size=batch_size, shift=1, drop_remainder=True):

    • size:指定每个滑动窗口包含的样本数量,也就是你需要的批量大小(这里是4)
    • shift:指定每次窗口滑动的步长,设为1就实现了每次只滑动1个样本的需求
    • drop_remainder:如果设为True,会丢弃最后一个样本数量不足batch_size的窗口;如果你的场景允许最后一个小批量,可以改成False
  • flat_map(lambda window: window.batch(batch_size)):
    window()方法返回的是一个个子数据集(每个子数据集对应一个滑动窗口),我们需要用flat_map把这些子数据集转换成批量张量,这样最终得到的就是你想要的重叠批量形式。

验证效果

假设你的原始数据集是[Img0, Img1, Img2, Img3, Img4, Img5, Img6, Img7],修改后生成的批量会是:

  • Batch1: [Img0, Img1, Img2, Img3]
  • Batch2: [Img1, Img2, Img3, Img4]
  • Batch3: [Img2, Img3, Img4, Img5]
  • Batch4: [Img3, Img4, Img5, Img6]
  • Batch5: [Img4, Img5, Img6, Img7]
    完全符合你的需求!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 08:22:45