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

执行训练代码时遇stack张量尺寸不匹配错误,求解决方案

解决PyTorch训练中stack expects each tensor to be equal size错误
  • 排查数据加载的批次对齐问题
    这个错误最常见的原因是同一个batch内的样本张量维度不统一,默认DataLoader会尝试直接堆叠这些张量导致失败。解决办法是自定义collate_fn,对批次内的样本做padding处理,统一维度:
    以NLP任务为例,示例代码如下:

    from torch.nn.utils.rnn import pad_sequence
    
    def custom_collate_fn(batch):
        # 提取batch中每个样本的input_ids和attention_mask
        input_ids = [item['input_ids'] for item in batch]
        attention_masks = [item['attention_mask'] for item in batch]
        
        # 将序列padding到当前batch的最大长度,padding_value设为pad token的ID(通常是0)
        padded_input_ids = pad_sequence(input_ids, batch_first=True, padding_value=0)
        padded_attention_masks = pad_sequence(attention_masks, batch_first=True, padding_value=0)
        
        return {
            'input_ids': padded_input_ids,
            'attention_mask': padded_attention_masks
        }
    

    创建DataLoader时指定该函数:

    train_loader = DataLoader(train_dataset, batch_size=32, collate_fn=custom_collate_fn)
    
  • 检查模型内部的堆叠逻辑
    如果模型中手动调用了torch.stack()操作,要确认传入该函数的所有张量形状完全一致。比如处理可变长度输入时,需先通过全局平均池化、自适应池化等方式将可变维度转为固定维度,再进行堆叠或拼接。

  • 验证数据集预处理的一致性
    检查数据集__getitem__方法的输出,确保每个样本的张量维度符合规范。比如文本样本的input_ids必须是一维张量,图像样本必须是[通道数, 高度, 宽度]的三维张量,过滤掉预处理后维度异常的样本,避免破坏批次的形状统一性。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 10:52:22