训练AlexNet遇Target与Input尺寸不匹配错误的求助
解决方案
一、数据加载方式的正确性
数据加载的整体逻辑合理:
- 使用
ImageFolder按目录结构加载分类数据集的方式正确,但需确保./images目录下存在两个子目录(对应正常/黄疸两类),例如./images/normal/和./images/jaundice/,ImageFolder会自动将子目录名映射为标签(标签顺序按子目录名称排序)。 random_split拆分数据集、DataLoader封装的逻辑无问题。
二、尺寸不匹配问题的核心修复
1. 修复模型定义函数
当前alexnet()函数未返回模型实例,且无需在修改层时启用eval()模式和torch.no_grad(),修改后:
def alexnet(): model = models.alexnet(pretrained=True) # 新版torchvision建议替换为weights=AlexNet_Weights.DEFAULT nn_filters = model.classifier[6].in_features model.classifier[6] = nn.Linear(nn_filters, 1) model = model.to(device) return model # 必须返回模型实例 # 初始化模型 model = alexnet()
2. 修正损失函数的使用
BCEWithLogitsLoss要求输入顺序为模型输出的logits在前,标签在后,且无需提前对logits做sigmoid或round处理(损失函数内部会自动计算sigmoid),这是尺寸不匹配的核心原因。修改训练循环中的损失计算部分:
for i,(images, labels) in tqdm(enumerate(train_loader),total = len(train_loader)): model.train() images = images.to(device) labels = labels.unsqueeze(1).type(torch.float32).to(device) output_logits = model(images) # 直接用logits计算损失,顺序为(模型输出, 标签) loss = loss_fn(output_logits, labels) # 预测结果仅用于评估,无需参与损失计算 output_pred = torch.round(torch.sigmoid(output_logits)) optimizer_fn.zero_grad() loss.backward() optimizer_fn.step()
3. 修复测试循环的错误
测试循环存在3个关键错误:未处理测试集标签、损失计算顺序错误、使用训练集标签计算准确率,修改后:
model.eval() with torch.no_grad(): cum_loss = 0.0 cum_acc = 0.0 for images_test, labels_test in test_loader: images_test = images_test.to(device) # 统一处理测试集标签格式 labels_test = labels_test.unsqueeze(1).type(torch.float32).to(device) test_logits = model(images_test) test_pred = torch.round(torch.sigmoid(test_logits)) # 修正损失计算顺序 test_loss = loss_fn(test_logits, labels_test) cum_loss += test_loss.item() # 转换为数值避免累积张量 # 使用测试集标签计算准确率 cum_acc += accuracy_fn(y_true=labels_test, y_pred=test_pred) # 计算平均损失和准确率 avg_test_loss = cum_loss / len(test_loader) avg_test_acc = cum_acc / len(test_loader) print(f'Epoch {epoch+1}, Test Loss {avg_test_loss:.4f}, accuracy {avg_test_acc:.4f}')
4. 其他优化建议
- 将
model.train()移至epoch循环的开头,无需每个batch重复调用 - 训练循环中可累积训练损失,最后打印平均训练损失而非最后一个batch的损失
内容的提问来源于stack exchange,提问作者Divij Jasuja
相关产品推荐
相关产品推荐

