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

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内所有样本维度一致:

  1. 编写自定义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
  1. 修改训练和验证集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
    )
)
  1. 变长数据处理可选方案:如果显存充足,也可以在数据预处理阶段就把所有样本统一padding到模型支持的最大长度(BERT基础版为512),不需要动态padding,但会增加无效计算量,优先选动态padding方案。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 09:51:45