如何用交错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
相关产品推荐
相关产品推荐

