如何分块加载大型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
相关产品推荐
相关产品推荐

