BERT预训练DataLoader报batch元素尺寸不一致RuntimeError排查
问题背景
- 训练带MLM和NSP任务的BERT预训练模型时触发RuntimeError,未定位到根因
- 运行环境配置:Python 3.8.10、PyTorch 1.8.0,训练使用IMDB数据集
- 已尝试排查方向:更换依赖版本、检查数据集内容,确认数据集为变长序列,不清楚该类变长数据的标准处理方式,需要对应处理技巧指导
预训练实现代码
def pretraining( model: MLMandNSPmodel, model_name: str, train_dataset: PretrainDataset, val_dataset: PretrainDataset, ): # Below options are just our recommendation. You can choose different options if you want. batch_size = 8 learning_rate = 1e-4 optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate) epochs = 200 # 200 if you want to feel the effect of pretraining steps_per_a_epoch: int=2000 steps_for_val: int=200 ### YOUR CODE HERE # pretraining(model, model_name, train_dataset, val_dataset) MLM_train_losses: List[float] = None MLM_val_losses: List[float] = None NSP_train_losses: List[float] = None NSP_val_losses: List[float] = None MLM_train_losses = [] MLM_val_losses = [] NSP_train_losses = [] NSP_val_losses = [] print('') print(train_dataset) print(val_dataset) print('') train_data_iterator = iter( torch.utils.data.dataloader.DataLoader(train_dataset, batch_size=batch_size, num_workers=2, shuffle=False)) # train_data_iterator = torch.utils.data.dataloader.DataLoader(train_dataset, batch_size=batch_size, num_workers=2, shuffle=False) eval_data_iterator = iter( torch.utils.data.dataloader.DataLoader(val_dataset, batch_size=batch_size, num_workers=2, shuffle=False)) # eval_data_iterator = torch.utils.data.dataloader.DataLoader(val_dataset, batch_size=batch_size, num_workers=2, shuffle=False) loss_log = tqdm(total=0, bar_format='{desc}') i = 0 for epoch in trange(epochs, desc="Epoch", position=0): i += 1 # Run batches for 'steps_per_a_epoch' times MLM_loss = 0 NSP_loss = 0 model.train() for step in trange(steps_per_a_epoch, desc="Training steps"): optimizer.zero_grad() src, mlm, mask, nsp = next(train_data_iterator) mlm_loss, nsp_loss = calculate_losses(model, src, mlm, mask, nsp) MLM_loss += mlm_loss NSP_loss += nsp_loss loss = mlm_loss + nsp_loss loss.backward() optimizer.step() des = 'Loss: {:06.4f}'.format(loss.cpu()) loss_log.set_description_str(des) # Calculate training loss MLM_loss = MLM_loss / steps_per_a_epoch NSP_loss = NSP_loss / steps_per_a_epoch MLM_train_losses.append(float(MLM_loss.data)) NSP_train_losses.append(float(NSP_loss.data)) # Calculate valid loss model.eval() valid_mlm_loss = 0. valid_nsp_loss = 0. for step in trange(steps_for_val, desc="Evaluation steps"): src, mlm, mask, nsp = next(eval_data_iterator) mlm_loss, nsp_loss = calculate_losses(model, src, mlm, mask, nsp) valid_mlm_loss += mlm_loss valid_nsp_loss += nsp_loss valid_mlm_loss = valid_mlm_loss / steps_for_val valid_nsp_loss = valid_nsp_loss / steps_for_val MLM_val_losses.append(float(valid_mlm_loss.data)) NSP_val_losses.append(float(valid_nsp_loss.data)) torch.save(model.state_dict(), os.path.join('/home/ml/Desktop/song/HW3/hw3/',model_name + str(i)+'.pth')) ### END YOUR CODE assert len(MLM_train_losses) == len(MLM_val_losses) == epochs and \ len(NSP_train_losses) == len(NSP_val_losses) == epochs assert all(isinstance(loss, float) for loss in MLM_train_losses) and \ all(isinstance(loss, float) for loss in MLM_val_losses) and \ all(isinstance(loss, float) for loss in NSP_train_losses) and \ all(isinstance(loss, float) for loss in NSP_val_losses) return MLM_train_losses, MLM_val_losses, NSP_train_losses, NSP_val_losses
完整报错信息
(hw3) ml@automl03:~/Desktop/song/HW3/hw3$ python pretrain.py ======MLM & NSP Pretraining====== <__main__.PretrainDataset object at 0x7fa72117eb80> <__main__.PretrainDataset object at 0x7fa70d0a2a30> <torch.utils.data.dataloader._MultiProcessingDataLoaderIter object at 0x7fa70c8b3640> <torch.utils.data.dataloader._MultiProcessingDataLoaderIter object at 0x7fa70c8b3670> Training steps: 0%| | 0/2000 [00:00<?, ?it/s] Epoch: 0%| | 0/200 [00:00<?, ?it/s] Traceback (most recent call last): File "pretrain.py", line 527, in <module> pretrain_model() File "pretrain.py", line 505, in pretrain_model = pretraining(model, model_name, train_dataset, val_dataset) File "pretrain.py", line 316, in pretraining src, mlm, mask, nsp = next(train_data_iterator) File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 517, in __next__ data = self._next_data() File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 1199, in _next_data return self._process_data(data) File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 1225, in _process_data data.reraise() File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/_utils.py", line 429, in reraise raise self.exc_type(msg) RuntimeError: Caught RuntimeError in DataLoader worker process 0. Original Traceback (most recent call last): File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/_utils/worker.py", line 202, in _worker_loop data = fetcher.fetch(index) File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/_utils/fetch.py", line 35, in fetch return self.collate_fn(data) File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 83, in default_collate return [default_collate(samples) for samples in transposed] File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 83, in <listcomp> return [default_collate(samples) for samples in transposed] File "/home/ml/anaconda3/envs/hw3/lib/python3.8/site-packages/torch/utils/data/_utils/collate.py", line 81, in default_collate raise RuntimeError('each element in list of batch should be of equal size') RuntimeError: each element in list of batch should be of equal size
问题根因
报错触发点在PyTorch DataLoader默认的批次拼接函数default_collate,该函数默认要求同一个batch内所有样本的张量维度完全一致,才能拼接成规则的批次张量。IMDB是变长文本数据集,同一个batch内不同样本的序列长度不统一,默认拼接逻辑无法处理维度不一致的序列,直接抛出长度不匹配的RuntimeError。
解决方法
核心是给DataLoader传入自定义的批次拼接函数(collate_fn),实现动态padding逻辑,保证同一个batch内所有样本维度一致:
- 编写自定义collate_fn,以当前batch内最长序列长度为基准,将短序列补PAD token到等长,同步处理MLM标签、注意力掩码、NSP标签:padding位置的MLM标签设为-1(计算交叉熵损失时会自动忽略该位置,不参与损失计算),注意力掩码的padding位置设为0(标识该位置是填充,不需要做注意力计算)。参考实现:
import torch from torch.nn.utils.rnn import pad_sequence # 替换成自己数据集里[PAD]对应的token id PAD_ID = 0 def bert_collate_fn(batch): src_samples, mlm_samples, mask_samples, nsp_samples = zip(*batch) # 按batch内最大长度做动态padding src_padded = pad_sequence(src_samples, batch_first=True, padding_value=PAD_ID) mlm_padded = pad_sequence(mlm_samples, batch_first=True, padding_value=-1) mask_padded = pad_sequence(mask_samples, batch_first=True, padding_value=0) nsp_tensor = torch.stack(nsp_samples) return src_padded, mlm_padded, mask_padded, nsp_tensor
- 修改训练和验证集DataLoader初始化代码,传入自定义collate_fn:
train_data_iterator = iter( torch.utils.data.dataloader.DataLoader( train_dataset, batch_size=batch_size, num_workers=2, shuffle=False, collate_fn=bert_collate_fn ) ) eval_data_iterator = iter( torch.utils.data.dataloader.DataLoader( val_dataset, batch_size=batch_size, num_workers=2, shuffle=False, collate_fn=bert_collate_fn ) )
- 变长数据处理可选方案:如果显存充足,也可以在数据预处理阶段就把所有样本统一padding到模型支持的最大长度(BERT基础版为512),不需要动态padding,但会增加无效计算量,优先选动态padding方案。
内容的提问来源于stack exchange,提问作者pasongsong
相关产品推荐
相关产品推荐

