从磁盘加载HuggingFace Dataset后用Generator迭代异常,如何解决?
HuggingFace Datasets加载磁盘后shuffle重复输出问题的解决
问题分析
你遇到的问题核心在于复用了同一个有状态的numpy随机生成器,以及HuggingFace Datasets中内存数据集与磁盘加载数据集(ArrowDataset)的shuffle实现差异:
- 每次运行脚本时,你初始化一个固定seed的生成器,第一个循环中多次调用
shuffle会持续推进生成器的内部状态。 - 当从磁盘加载数据集后,ArrowDataset的
shuffle方法对生成器的处理逻辑和内存中的Dataset不同,在生成器状态被推进到特定阶段后,每次shuffle生成的索引列表首个元素固定,导致输出重复。 - 虽然每次运行脚本都会重新保存数据集,但生成器的状态复用依然会触发这个问题。
解决方案
方案1:每次shuffle使用独立的生成器实例
避免复用同一个生成器,在每次需要shuffle时创建新的生成器:
import datasets import numpy as np X = np.arange(1000) ds = datasets.Dataset.from_dict(mapping={"X":X}) ds.save_to_disk("tmp") print("First loop") for _ in range(10): # 每次shuffle创建新生成器 generator = np.random.default_rng(0) print(next(ds.shuffle(generator=generator).iter(batch_size=1)), end=", ") print("") print("Second loop") ds = datasets.Dataset.load_from_disk("tmp") for _ in range(10): # 每次shuffle创建新生成器 generator = np.random.default_rng(0) print(next(ds.shuffle(generator=generator).iter(batch_size=1)), end=", ") print("")
方案2:使用seed参数代替传入生成器
让HuggingFace Datasets内部管理随机生成器,通过指定不同的seed保证每次shuffle结果不同:
import datasets import numpy as np X = np.arange(1000) ds = datasets.Dataset.from_dict(mapping={"X":X}) ds.save_to_disk("tmp") print("First loop") for i in range(10): print(next(ds.shuffle(seed=i).iter(batch_size=1)), end=", ") print("") print("Second loop") ds = datasets.Dataset.load_from_disk("tmp") for i in range(10): print(next(ds.shuffle(seed=i+10).iter(batch_size=1)), end=", ") print("")
方案3:重置生成器状态(不推荐)
如果你必须复用生成器,可以在第一个循环后重置其状态,但这种方式依赖numpy内部实现,兼容性较差:
import datasets import numpy as np generator = np.random.default_rng(0) # 保存初始状态 initial_state = generator.bit_generator.state X = np.arange(1000) ds = datasets.Dataset.from_dict(mapping={"X":X}) ds.save_to_disk("tmp") print("First loop") for _ in range(10): print(next(ds.shuffle(generator=generator).iter(batch_size=1)), end=", ") print("") # 重置生成器到初始状态 generator.bit_generator.state = initial_state print("Second loop") ds = datasets.Dataset.load_from_disk("tmp") for _ in range(10): print(next(ds.shuffle(generator=generator).iter(batch_size=1)), end=", ") print("")
说明
方案1和方案2是更可靠的选择,因为它们从根源上避免了生成器状态复用导致的意外行为,同时适配HuggingFace Datasets对不同存储类型数据集的shuffle实现。
内容的提问来源于stack exchange,提问作者LudvigH
相关产品推荐
相关产品推荐

