如何使用TensorFlow Dataset API实现批量滑动窗口?
实现滑动窗口式的重叠批量(步长为1)
嘿,我完全理解你的需求——你想要把原来无重叠的固定批量改成每次滑动1个样本的重叠批量,也就是滑动窗口批量。在TensorFlow Dataset API中,直接用batch()方法做不到这点,不过我们可以用window()方法来实现,下面是修改后的代码和详细说明:
核心修改思路
原来的dataset.batch(batch_size)是将连续的batch_size个样本分组,没有重叠。要实现滑动步长为1的窗口,我们需要:
- 使用
window()生成滑动窗口的子数据集 - 将每个子数据集合并成一个完整的批量
修改后的完整代码
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
相关产品推荐
相关产品推荐

