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

PyTorch中CrossEntropyLoss输入维度问题:二分类batch=1时报错求助

Fixing the Dimension Mismatch Error with CrossEntropyLoss

Hey there, let's break down what's going wrong here and how to fix it quickly.

The Root Cause of Your Error

Your logit tensor has a shape of (2,), but PyTorch's nn.CrossEntropyLoss expects specific dimensions for both logits and labels—especially when working with a batch size of 1. Let's clarify the exact requirements:

CrossEntropyLoss Input Dimension Requirements

For standard single-label classification (like your binary task):

  • Logits: Must be a tensor with shape (N, C), where:
    • N = batch size (in your case, 1)
    • C = number of classes (2 for binary classification)
      Higher-dimensional inputs (e.g., for image tasks: (N, C, H, W)) are allowed too, but the core rule is that the class dimension sits at index 1.
  • Labels: Must be a tensor with shape (N,), where each element is the class index (0 or 1 for your binary task). PyTorch will also accept (N, 1) but automatically squeezes it to (N,).

In your code, your logit is (2,)—it's missing the batch dimension (N=1). This makes PyTorch interpret it as a batch of 2 samples, each with 1 class. Your label has shape (1,), which doesn't match the inferred batch size of 2, hence the dimension out-of-range error.

The Fix: Add the Batch Dimension to Your Logits

All you need to do is add a batch dimension to your logit tensor using unsqueeze(0):

import torch
import torch.nn as nn

# Your original inputs
b_logits = torch.tensor([0.1198, 0.1911], device='cuda:0')
b_labels = torch.tensor([1], device='cuda:0')

# Add batch dimension to logits (shape becomes (1, 2))
b_logits = b_logits.unsqueeze(0)

loss_criterion = nn.CrossEntropyLoss()
loss_criterion.cuda()
loss = loss_criterion(b_logits, b_labels)
print(loss)

Quick Recap

  • Always ensure your logits have a batch dimension first, followed by the number of classes.
  • For batch size 1, it's easy to overlook the batch dimension—unsqueeze(0) is your go-to fix here.
  • Labels should match the batch size in their first dimension, with each value being a valid class index.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 14:17:58