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

如何分块加载大型TensorFlow数据集(UCF101)且无重复?

无重复分块加载TensorFlow数据集的解决方案

要实现类似栈pop()的无重复分块加载,核心是通过维护偏移量,结合skip()和take()方法实现顺序读取,避免重复加载已处理的数据。以下是具体实现方式:

方法1:手动维护偏移量

直接通过变量记录已读取的数据量,每次读取时跳过已处理部分,再取指定数量的样本:

import tensorflow as tf
import tensorflow_datasets as tfds

# 加载数据集
(train, test) = tfds.load('ucf101', split=['train', 'test'], shuffle_files=False)
train_total = len(train)
offset = 0

# 第一次取2个样本
chunk_size = 2
chunk = train.skip(offset).take(chunk_size)
# 将chunk加载到内存(按需转换为numpy数组或其他格式)
chunk_data = list(chunk.as_numpy_iterator())
offset += chunk_size
print(f"已读取{offset}/{train_total}个样本")

# 第二次取3个样本
chunk_size = 3
chunk = train.skip(offset).take(chunk_size)
chunk_data = list(chunk.as_numpy_iterator())
offset += chunk_size
print(f"已读取{offset}/{train_total}个样本")

方法2:封装为工具类(更贴近pop()用法)

把偏移量维护、边界判断逻辑封装成类,调用时更简洁:

class DatasetChunker:
    def __init__(self, dataset):
        self.dataset = dataset
        self.offset = 0
        self.total = len(dataset)
    
    def pop(self, n):
        if self.offset >= self.total:
            return None  # 所有数据已读取完毕
        # 计算实际可取的样本数(避免超出剩余数据量)
        take_num = min(n, self.total - self.offset)
        # 跳过已读取部分,取指定数量样本并加载到内存
        chunk = self.dataset.skip(self.offset).take(take_num)
        chunk_data = list(chunk.as_numpy_iterator())
        # 更新偏移量
        self.offset += take_num
        return chunk_data

# 使用示例
train_chunker = DatasetChunker(train)
first_chunk = train_chunker.pop(2)  # 获取前2个样本
second_chunk = train_chunker.pop(3) # 获取接下来3个样本
remaining_chunk = train_chunker.pop(10000) # 获取剩余所有样本

关键说明

  • tf.data.Dataset.skip(k)会跳过前k个样本,保证每次读取的是未处理过的部分;
  • take(n)只取后续n个样本,结合skip()就能实现无重复的分块读取;
  • 样本真正加载到内存是在调用as_numpy_iterator()并转换为列表时,之前的skip()和take()只是构建数据流水线,不会占用大量内存,完美适配UCF101这类大型数据集。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 00:40:59