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

如何用交错Hugging Face数据集创建PyTorch DataLoader及报错解决

问题:交错Hugging Face数据集接入PyTorch DataLoader报错解析

问题场景

使用交错Hugging Face流式数据集,对样本分词后送入PyTorch DataLoader时触发TypeError,但使用单一数据集(如c4、wiki-text)时无此问题,期望无需编写自定义collate_function解决。

复现代码

# -*- coding: utf-8 -*-
"""issues with dataloader and custom data sets

Automatically generated by Colaboratory.

Original file is located at
    https://colab.research.google.com/drive/1sbs95as_66mtK9VK_vbaE9gLE-Tjof1-
"""

!pip install datasets
!pip install pytorch
!pip install transformers

token = None
batch_size = 10
from datasets import load_dataset
import torch
from transformers import GPT2Tokenizer, GPT2LMHeadModel

tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
if tokenizer.pad_token_id is None:
  tokenizer.pad_token = tokenizer.eos_token
probe_network = GPT2LMHeadModel.from_pretrained("gpt2")
device = torch.device(f"cuda:{0}" if torch.cuda.is_available() else "cpu")
probe_network = probe_network.to(device)

# -- Get batch from dataset
from datasets import load_dataset
# path, name = 'brando/debug1_af', 'debug1_af'
path, name = 'brando/debug0_af', 'debug0_af'
remove_columns = []
dataset = load_dataset(path, name, streaming=True, split="train", token=token).with_format("torch")
print(f'{dataset=}')
batch = dataset.take(batch_size)
# print(f'{next(iter(batch))=}')

# - Prepare functions to tokenize batch
def preprocess(examples):  # gets the raw text batch according to the specific names in table in data set & tokenize
    return tokenizer(examples["link"], padding="max_length", max_length=128, truncation=True, return_tensors="pt")
def map(batch):  # apply preprocess to batch to all examples in batch represented as a dataset
    return batch.map(preprocess, batched=True, remove_columns=remove_columns)
tokenized_batch = batch.map(preprocess, batched=True, remove_columns=remove_columns)
tokenized_batch = map(batch)
# print(f'{next(iter(tokenized_batch))=}')

from torch.utils.data import Dataset, DataLoader, SequentialSampler
dataset = tokenized_batch
print(f'{type(dataset)=}')
print(f'{dataset.__class__=}')
print(f'{isinstance(dataset, Dataset)=}')
# for i, d in enumerate(dataset):
#     assert isinstance(d, dict)
#     # dd = dataset[i]
#     # assert isinstance(dd, dict)
loader_opts = {}
classifier_opts = {}
# data_loader = DataLoader(dataset, shuffle=False, batch_size=loader_opts.get('batch_size', 1),
#                         num_workers=loader_opts.get('num_workers', 0), drop_last=False, sampler=SequentialSampler(range(512))  )
data_loader = DataLoader(dataset, shuffle=False, batch_size=loader_opts.get('batch_size', 1),
                    num_workers=loader_opts.get('num_workers', 0), drop_last=False, sampler=None)
print(f'{iter(data_loader)=}')
print(f'{next(iter(data_loader))=}')
print('Done\a')

报错信息

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
/usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/collate.py in collate(batch, collate_fn_map)
    126         try:
--> 127             return elem_type({key: collate([d[key] for d in batch], collate_fn_map=collate_fn_map) for key in elem})
    128         except TypeError:

9 frames
TypeError: default_collate: batch must contain tensors, numpy arrays, numbers, dicts or lists; found <class 'NoneType'>

During handling of the above exception, another exception occurred:

TypeError                                 Traceback (most recent call last)
/usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/collate.py in collate(batch, collate_fn_map)
    148                 return [collate(samples, collate_fn_map=collate_fn_map) for samples in transposed]
    149 
--> 150     raise TypeError(default_collate_err_msg_format.format(elem_type))
    151 
    152 

TypeError: default_collate: batch must contain tensors, numpy arrays, numbers, dicts or lists; found <class 'NoneType'>

问题原因

  • 交错数据集存在无效样本:交错数据集由多数据源合并而成,部分样本的link字段为None,分词后会生成包含None的结果,而PyTorch默认collate_fn无法处理None类型。单一数据集的字段完整性有保障,不会出现此类无效样本,因此无报错。
  • 代码冗余操作导致数据异常:代码中先执行tokenized_batch = batch.map(...),随后又重复调用tokenized_batch = map(batch),重复处理可能引入数据混乱;同时remove_columns为空,未移除原始数据集的无效字段,导致残留数据干扰后续处理。
  • 流式数据集的特性限制:流式数据集的map操作在batched=True时,不会自动过滤None值,会直接将无效数据传递到下游。

解决方案(无需自定义collate_fn)

  • 过滤无效样本:在处理数据集前,先过滤掉link字段为None的样本
    dataset = dataset.filter(lambda x: x["link"] is not None)
    
  • 修正remove_columns参数:移除原始数据集的link字段,避免残留无效数据
    remove_columns = ["link"]  # 替换原空列表,指定要移除的原始字段
    tokenized_batch = batch.map(preprocess, batched=True, remove_columns=remove_columns)
    
  • 删除冗余的map调用:移除重复的tokenized_batch = map(batch)语句,避免重复处理
  • 增强分词函数的鲁棒性:在preprocess中先校验输入有效性,确保仅对非空文本分词
    def preprocess(examples):
        # 过滤掉None值的link
        valid_examples = {"link": [link for link in examples["link"] if link is not None]}
        return tokenizer(valid_examples["link"], padding="max_length", max_length=128, truncation=True, return_tensors="pt")
    

内容的提问来源于stack exchange,提问作者Charlie Parker

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 22:48:12