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

文本分类多分类任务报错:输入batch_size(2)与目标batch_size(4)不匹配

文本分类任务报错:输入与目标batch_size不匹配

这是一个3标签多分类文本任务,已经试过把标签转成整数、检查损失函数,但还是卡在这里,核心评估代码和报错信息如下:

def evaluate(model, dataloader_val):
    model.eval()
    model.train(False)
    
    loss_val_total = 0
    predictions, true_vals = [], []
    
    for batch in dataloader_val:
        
        batch = tuple(b.to(device) for b in batch)
        
        inputs = {'input_ids':      batch[0],
                  'attention_mask': batch[1],
                  'labels':         batch[2],
                 }

        with torch.no_grad():        
            outputs = model(**inputs)
            
        loss = outputs[0]
        
        logits = outputs[1]
        loss_val_total += loss.item()

        probs = torch.argmax(logits, dim = 1).detach().cpu().numpy()
        label_ids = inputs['labels'].cpu().numpy()
        predictions.append(probs)
        true_vals.append(label_ids)
        
    loss_val_avg = loss_val_total/len(dataloader_val) 
    
    predictions = np.concatenate(predictions, axis=0)
    true_vals = np.concatenate(true_vals, axis=0)
    
    ### after evaluating we resume model training
    model.train(True)
    
    return loss_val_avg, predictions, true_vals

报错信息:

ValueError                                Traceback (most recent call last)
<ipython-input-55-a095c6ad8f10> in <module>
     44                      }       
     45 
---> 46             outputs = model(**inputs)
ValueError: Expected input batch_size (2) to match target batch_size (4).

排查方向

  • 检查数据加载器:确认dataloader_val的batch_size设置是否和训练集一致,有没有构建时参数配置错误
  • 验证batch维度:在循环内添加print(batch[0].shape, batch[2].shape),查看input_ids的batch维度(第一个数值)和labels的维度是否匹配,比如input_ids是(2, 512)但labels是(4,)就会触发该错误
  • 排查标签预处理:确认Dataset类中有没有错误重复标签、拆分单个样本标签的情况,导致labels的batch_size异常翻倍
  • 核对模型配置:如果用HuggingFace预训练模型,确认初始化时num_labels=3已正确设置;如果是自定义模型,检查forward函数中计算损失时logits与labels的维度是否匹配

内容的提问来源于stack exchange,提问作者Andreea-Codrina Moldovan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 10:45:55