如何高效将自定义PyTorch Dataset转换为HuggingFace Dataset?
更优的PyTorch Dataset转HuggingFace Dataset方法
你当前的循环拼接方法在样本量较大时效率极低,每次concatenate_datasets都会复制现有数据,内存开销大且速度慢。这里提供两种更高效的方案:
方案1:批量收集拼接后转Dataset(内存充足时优先)
先遍历PyTorch Dataset收集所有样本的对应张量,统一拼接后再转成HuggingFace Dataset:
import torch from datasets import Dataset # 初始化字典存储各字段的张量列表 data_dict = {} for sample in myPytorchDataset: for key, tensor in sample.items(): if key not in data_dict: data_dict[key] = [] data_dict[key].append(tensor) # 对每个字段的张量列表进行纵向拼接 for key in data_dict: data_dict[key] = torch.cat(data_dict[key], dim=0) # 直接转成HuggingFace Dataset hf_dataset = Dataset.from_dict(data_dict)
方案2:用生成器逐样本拆分overflow条目(内存有限时使用)
利用Dataset.from_generator,编写生成器函数遍历PyTorch Dataset,将每个样本拆分成独立的overflow条目逐个输出:
from datasets import Dataset def overflow_data_generator(): for sample in myPytorchDataset: # 获取当前样本拆分出的overflow数量 num_splits = sample["input_ids"].shape[0] # 逐个yield拆分后的条目 for idx in range(num_splits): yield { "input_ids": sample["input_ids"][idx].numpy(), "attention_mask": sample["attention_mask"][idx].numpy(), "bbox": sample["bbox"][idx].numpy(), "labels": sample["labels"][idx].numpy(), "pixel_values": sample["pixel_values"][idx].numpy() } # 从生成器创建Dataset hf_dataset = Dataset.from_generator(overflow_data_generator)
补充说明
- 你之前调用
Dataset.from_dict(myPytorchDataset)报错,是因为该方法需要传入键对应数组/张量的字典,而非PyTorch Dataset实例; Dataset.from_generator(myPytorchDataset)报错则是因为该方法需要传入生成器函数,不能直接传Dataset对象。
内容的提问来源于stack exchange,提问作者Pavan K
相关产品推荐
相关产品推荐

