如何将IterableDataset转换为Dataset?实现大数据集处理后采样保存
解决IterableDataset转常规Dataset并保存采样的问题
方法一:直接采样N个元素转成Dataset(高效推荐)
因为你的需求是采样一小部分保存,没必要遍历全部流式数据,直接取shuffle后的前N个元素转成列表,再生成常规Dataset即可:
from datasets import Dataset # 加载流式数据集并做字段转换 ds = datasets.load_dataset("XYZ", name="ABC", split="train", streaming=True) ds = ds.map(_transform_record) # 打乱数据集(设置seed保证结果可复现) ds_shuffled = ds.shuffle(seed=42) # 提取前N个样本转成列表 sampled_samples = list(ds_shuffled.take(N)) # 转换为常规Dataset并保存 sampled_ds = Dataset.from_list(sampled_samples) sampled_ds.save_to_disk("path/to/save/sampled_dataset")
方法二:遍历全部流式数据转成全量Dataset
如果需要先把整个流式数据集转成常规Dataset再操作,可以用Dataset.from_generator(),但要注意必须传入生成器函数(不能直接传生成器对象,否则序列化失败):
from datasets import Dataset # 加载并转换流式数据集 ds = datasets.load_dataset("XYZ", name="ABC", split="train", streaming=True) ds = ds.map(_transform_record) # 定义生成器函数遍历流式数据集 def dataset_generator(): for sample in ds: yield sample # 转换为常规Dataset full_ds = Dataset.from_generator(dataset_generator) # 采样并保存 sampled_ds = full_ds.shuffle(seed=42).select(range(N)) sampled_ds.save_to_disk("path/to/save/sampled_dataset")
为什么你之前的尝试失败?
Dataset.from_generator()要求传入的是可序列化的生成器函数,而直接传iter(ds)得到的是一个生成器对象,无法被序列化,因此会报错。用函数包裹遍历逻辑就能解决这个问题。
内容的提问来源于stack exchange,提问作者Zach Moshe
相关产品推荐
相关产品推荐

