PyTorch中CrossEntropyLoss输入维度问题:二分类batch=1时报错求助
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

