如何基于自定义Hugging Face数据集创建无错PyTorch DataLoader
问题:自定义Hugging Face数据集封装进PyTorch DataLoader后触发NoneType错误
自定义Hugging Face数据集本身无None值,但封装到PyTorch DataLoader后运行失败,报错显示default_collate遇到<class 'NoneType'>类型数据,已确认数据集返回为字典,仍未定位问题。
依赖安装代码
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) # -- 获取数据集批次 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))=}') # - 定义批处理tokenize函数 def preprocess(examples): # 获取数据集原始文本并tokenize return tokenizer(examples["link"], padding="max_length", max_length=128, truncation=True, return_tensors="pt") def map(batch): # 对数据集批次应用preprocess 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')
报错信息
dataset=<datasets.iterable_dataset.IterableDataset object at 0x7e42c2f21d20> type(dataset)=<class 'datasets.iterable_dataset.IterableDataset'> dataset.__class__=<class 'datasets.iterable_dataset.IterableDataset'> isinstance(dataset, Dataset)=True iter(data_loader)=<torch.utils.data.dataloader._SingleProcessDataLoaderIter object at 0x7e42c2f21660> /usr/local/lib/python3.10/dist-packages/datasets/formatting/torch_formatter.py:68: UserWarning: To copy construct from a tensor, it is recommended to use sourceTensor.clone().detach() or sourceTensor.clone().detach().requires_grad_(True), rather than torch.tensor(sourceTensor). return torch.tensor(value, **{**default_dtype, **self.torch_tensor_kwargs}) --------------------------------------------------------------------------- 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 /usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/collate.py in <dictcomp>(.0) 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: /usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/collate.py in collate(batch, collate_fn_map) 149 --> 150 raise TypeError(default_collate_err_msg_format.format(elem_type)) 151 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) <ipython-input-6-1153c5915bd8> in <cell line: 49>() 47 num_workers=loader_opts.get('num_workers', 0), drop_last=False, sampler=None) 48 print(f'{iter(data_loader)=}') ---> 49 print(f'{next(iter(data_loader))=}') 50 print('Done\a') /usr/local/lib/python3.10/dist-packages/torch/utils/data/dataloader.py in __next__(self) 631 # TODO(https://github.com/pytorch/pytorch/issues/76750) 632 self._reset() # type: ignore[call-arg] --> 633 data = self._next_data() 634 self._num_yielded += 1 635 if self._dataset_kind == _DatasetKind.Iterable and \ /usr/local/lib/python3.10/dist-packages/torch/utils/data/dataloader.py in _next_data(self) 675 def _next_data(self): 676 index = self._next_index() # may raise StopIteration --> 677 data = self._dataset_fetcher.fetch(index) # may raise StopIteration 678 if self._pin_memory: 679 data = _utils.pin_memory.pin_memory(data, self._pin_memory_device) /usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/fetch.py in fetch(self, possibly_batched_index) 40 else: 41 data = next(self.dataset_iter) ---> 42 return self.collate_fn(data) 43 44 /usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/collate.py in default_collate(batch) 263 >>> default_collate(batch) # Handle `CustomType` automatically 264 """ --> 265 return collate(batch, collate_fn_map=default_collate_fn_map) /usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/collate.py in collate(batch, collate_fn_map) 128 except TypeError: 129 # The mapping type may not support `__init__(iterable)`. --> 130 return {key: collate([d[key] for d in batch], collate_fn_map=collate_fn_map) for key in elem} 131 elif isinstance(elem, tuple) and hasattr(elem, '_fields'): # namedtuple 132 return elem_type(*(collate(samples, collate_fn_map=collate_fn_map) for samples in zip(*batch))) /usr/local/lib/python3.10/dist-packages/torch/utils/data/_utils/collate.py in <dictcomp>(.0) 128 except TypeError: 129 # The mapping type may not support `__init__(iterable)`. --> 130 return {key: collate([d[key] for d in batch], collate_fn_map=collate_fn_map) for key in elem} 131 elif isinstance(elem, tuple) and hasattr(elem, '_fields'): # namedtuple 132 return elem_type(*(collate(samples, collate_fn_map=collate_fn_map) for samples in zip(*batch))) /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'>
问题分析与解决方案
核心原因
- 重复赋值覆盖处理结果:代码中先通过
batch.map得到处理后的数据集,随后又执行tokenized_batch = map(batch)重复处理,且remove_columns为空,导致原始字段与tokenize字段共存,引入无效值。 - 流式数据集处理不彻底:使用
streaming=True加载的IterableDataset,若未移除原始文本字段,且原始数据存在隐性空值,会导致collate时出现None。 - collate逻辑不兼容:PyTorch默认collate函数无法处理None类型,需确保所有样本字段均为张量、数组等可处理类型。
修复步骤
- 删除重复赋值:移除
tokenized_batch = map(batch),保留第一次batch.map的结果。 - 正确设置移除字段:将
remove_columns设为["link"],清理原始文本字段,避免干扰。 - 添加空值过滤:在
preprocess函数中过滤空链接,确保输入有效:def preprocess(examples): valid_links = [link for link in examples["link"] if link is not None and link.strip()] if not valid_links: return {"input_ids": torch.tensor([], dtype=torch.long), "attention_mask": torch.tensor([], dtype=torch.long)} return tokenizer(valid_links, padding="max_length", max_length=128, truncation=True, return_tensors="pt") - 适配IterableDataset:保持DataLoader的
sampler=None,无需额外设置采样器。
修复后核心代码片段
# -- 获取数据集批次 path, name = 'brando/debug0_af', 'debug0_af' remove_columns = ["link"] # 指定移除原始字段 dataset = load_dataset(path, name, streaming=True, split="train", token=token).with_format("torch") batch = dataset.take(batch_size) # - 定义批处理tokenize函数 def preprocess(examples): valid_links = [link for link in examples["link"] if link is not None and link.strip()] if not valid_links: return {"input_ids": torch.tensor([], dtype=torch.long), "attention_mask": torch.tensor([], dtype=torch.long)} return tokenizer(valid_links, padding="max_length", max_length=128, truncation=True, return_tensors="pt") # 仅执行一次map操作 tokenized_batch = batch.map(preprocess, batched=True, remove_columns=remove_columns) # 后续DataLoader代码不变
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

