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

训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 20:45:14