如何仅加载TensorFlow Datasets中lm1b数据集的小样本?
解决TensorFlow Datasets加载lm1b时全量加载及文件过多问题
问题根源
lm1b数据集按分片存储,但tfds.load默认会遍历所有分片生成全局元数据,同时文件级别的shuffle也会触发全量分片扫描,导致打开过多文件触发Errno 24,且即使只取少量样本也会先处理整个数据集。
解决方案
通过tfds.builder手动控制加载逻辑,限制仅读取所需分片,避免全量扫描:
方案1:指定单个分片加载小样本
直接选取某一个分片(如train_000)的前N个样本,避免遍历所有分片:
import tensorflow as tf import tensorflow_datasets as tfds # 构建数据集builder,替代tfds.load获得更精细控制 builder = tfds.builder("lm1b") # 配置读取参数,限制同时打开的文件数 read_config = tfds.ReadConfig( interleave_cycle_length=1, # 一次仅处理1个分片,避免同时打开大量文件 interleave_block_length=512, # 每个分片取512个样本 ) # 加载指定分片的前512个样本(lm1b的train分片格式为train_000到train_099) dataset = builder.as_dataset( split="train_000[:512]", read_config=read_config, shuffle_files=False # 关闭文件级shuffle,避免扫描所有分片 ) dataset = dataset.batch(256).prefetch(tf.data.AUTOTUNE) for example in dataset.take(10): print(example) break
方案2:随机选取一个分片加载小样本
如果需要随机样本但不想扫描全部分片,可先获取分片列表再随机选择:
import tensorflow as tf import tensorflow_datasets as tfds import random builder = tfds.builder("lm1b") # 获取所有训练分片的名称 train_splits = list(builder.info.splits["train"].split_infos.keys()) # 随机挑选一个分片 selected_split = random.choice(train_splits) # 加载该分片的前512个样本 dataset = builder.as_dataset( split=f"{selected_split}[:512]", shuffle_files=False, read_config=tfds.ReadConfig(interleave_cycle_length=1) ) dataset = dataset.batch(256).prefetch(tf.data.AUTOTUNE) for example in dataset.take(10): print(example) break
临时缓解文件过多错误
如果仍触发Errno 24,可临时提高系统文件描述符限制(仅Linux/macOS有效):
import resource # 将最大打开文件数临时提高到4096 resource.setrlimit(resource.RLIMIT_NOFILE, (4096, 4096))
关键注意点
- 调试阶段关闭
shuffle_files=False,避免tfds为了随机取样而扫描所有分片 - 若需要多个分片的小样本,可直接指定多个分片,如
split="train_000[:256]+train_001[:256]",仅会加载指定分片 - 确保数据集已完成下载,首次下载时tfds会生成全局索引,但后续指定分片加载不会重复扫描全量数据
内容的提问来源于stack exchange,提问作者jpjandrade
相关产品推荐
相关产品推荐

