You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何高效将自定义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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.21 21:27:05