PyTorch中Target与Input尺寸不匹配报错原因及解决方法
PyTorch多分类训练报错:Target size与input size不匹配问题解决
问题现象
运行以下多分类模型训练代码时,在loss = loss_fn(y_logits, y_blob_train)行触发ValueError:
Target size (torch.Size([800])) must be the same as input size (torch.Size([800, 4]))
尝试用squeeze()处理y_logits后问题仍存在。
代码片段
# Fit the multi-class model to the data torch.manual_seed(42) torch.cuda.manual_seed(42) # Set number of epochs epochs = 100 # Put the data to the target device X_blob_train, y_blob_train = X_blob_train.to(device), y_blob_train.to(device) X_blob_test, y_blob_test = X_blob_test.to(device), y_blob_test.to(device) # Loop through data for epoch in range(epochs): ### Training model_4.train() y_logits = model_4(X_blob_train) y_pred = torch.softmax(y_logits, dim = 1).argmax(dim = 1) print(y_logits.shape, y_blob_train.shape) loss = loss_fn(y_logits, y_blob_train) # <<- ValueError acc = accuracy_fn(y_true = y_blob_train, y_pred = y_pred) optimizer.zero_grad() loss.backward() optimizer.step() ### Testing model_4.eval() with torch.inference_mode(): test_logits = model_4(X_blob_test) test_preds = torch.softmax(test_logits, dim = 1).argmax(dim = 1) test_loss = loss_fn(test_logits, y_blob_test) test_acc = accuracy_fn(y_true=y_blob_test, y_pred = test_preds) # Print out whats happening if epoch % 10 == 0: print(f'Epoch: {epochs} | Loss: {loss:.4f}, Acc: {acc:.2f}% | Test Loss: {test_loss:.4f}, Test acc: {test_acc:.2f}%')
报错原因
核心问题是损失函数的输入格式要求与当前数据不匹配:
- 模型输出
y_logits形状为[800, 4](800是批量大小,4是类别数),是多分类任务的标准logits输出。 - 目标标签
y_blob_train形状为[800],是类别索引格式(每个样本对应一个类别编号)。 - 如果此时使用的是
BCELoss、MSELoss这类损失函数,它们要求目标标签的形状必须和输入logits完全一致(即[800,4]),因此触发形状不匹配报错。 squeeze()仅能压缩维度为1的轴,y_logits的第二个维度是4,无法被压缩,所以该操作无效。
解决方法
根据多分类任务的标准做法,有两种可行方案:
方案1:使用正确的多分类损失函数(推荐)
将损失函数替换为CrossEntropyLoss,它专门为多分类任务设计,支持类别索引格式的目标标签:
# 替换原来的loss_fn定义 loss_fn = torch.nn.CrossEntropyLoss()
注意:CrossEntropyLoss内部已整合softmax计算,无需手动对y_logits调用torch.softmax()(代码中计算y_pred时调用softmax不影响损失计算)。
方案2:将目标标签转换为one-hot编码(仅特殊场景使用)
如果因需求必须使用BCELoss等需要同形状输入的损失函数,可将类别索引格式的标签转换为one-hot编码:
# 在将标签送入损失函数前,转换为one-hot编码,并转为float类型(匹配logits的数据类型) y_blob_train_onehot = torch.nn.functional.one_hot(y_blob_train, num_classes=4).float() loss = loss_fn(y_logits, y_blob_train_onehot)
测试集的标签也需要做同样处理。
内容的提问来源于stack exchange,提问作者Mubassir Habib 1812927042
相关产品推荐
相关产品推荐

