带批量归一化的VGG-11图像角点排序模型无学习效果技术问询
我仔细看了你的代码和问题描述,发现几个关键问题导致模型完全无法学习,下面逐个分析并给出修复方案:
核心问题分析
梯度断裂(最致命):你在
forward函数中使用了torch.argsort,这是一个不可微分的操作——它的输出无法反向传播梯度到前面的卷积层和全连接层,导致模型参数根本不会更新。这就是为什么准确率一直停留在随机水平(全对概率1/24≈4%,单角点正确概率1/4=25%),模型相当于在瞎猜。全连接层重复定义:代码里连续两次定义了
self.layer12,第二次定义直接覆盖了第一次,导致原本的Linear(4096, 4096)层被丢弃,模型的特征表达能力大幅削弱。损失函数与任务不匹配:用
MSELoss(均方误差损失)处理离散类别预测任务(预测0-3的角点编号)是不合适的。MSE是回归损失,适合连续值预测,而这里应该用分类专用的交叉熵损失。输出逻辑错误:
torch.argsort的输出是模型最后一层输出的排序索引,这和你的标签含义完全不匹配。标签是每个位置对应的正确角点编号,而不是模型输出值的排序结果,这导致模型输出和标签的对齐逻辑完全错误。
具体解决方案
1. 修复梯度断裂与输出逻辑
移除forward中的torch.argsort操作,让模型直接输出最后一层的原始logits(类别得分),在评估阶段再用argmax得到预测的类别编号:
def forward(self, x): l1 = self.layer1(x) l2 = self.layer2(l1) l3 = self.layer3(l2) l4 = self.layer4(l3) l5 = self.layer5(l4) l6 = self.layer6(l5) l7 = self.layer7(l6) l8 = self.layer8(l7) l9 = self.layer9(l8) l10 = self.layer10(l9) l11 = self.layer11(l10) l12 = self.layer12(l11) l13 = self.layer13(l12) # 修复后的全连接层 return l13
2. 修复全连接层重复定义
把重复的layer12改成layer13,保证模型的全连接层完整,同时调整最后一层的输出维度(4个位置×4个类别=16):
self.layer12 = nn.Sequential( nn.Linear(4096, 4096), ActFunc ) self.layer13 = nn.Sequential( nn.Linear(4096, 16) # 输出4个位置的4个类别得分 )
3. 调整损失函数与训练逻辑
将损失函数改为CrossEntropyLoss,并调整训练时的输入输出形状以匹配分类任务:
# 初始化损失函数和优化器 optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) # 建议调整学习率 loss_objective = nn.CrossEntropyLoss() # 训练循环内的修改 for x, y in training_generator: x = x.to(device) y = y.to(device) optimizer.zero_grad() predictions = model(x) # shape: (batch_size, 16) # 重塑为(batch_size, 4, 4),对应4个位置,每个位置4个类别得分 predictions = predictions.view(-1, 4, 4) # 交叉熵损失要求标签为Long类型 loss = loss_objective(predictions, y.type(torch.LongTensor).to(device)) loss.backward() optimizer.step()
4. 修改评估函数逻辑
在评估阶段,对每个位置的logits取argmax得到预测类别,再和标签比较:
def evaluate(self, test_generator, device="cuda"): batchfullacc = 0 batchacc = torch.zeros(4, device=device) total_samples = 0 with torch.no_grad(): self.eval() for images, labels in test_generator: images = images.to(device) labels = labels.to(device) predictions = self.forward(images) predictions = predictions.view(-1, 4, 4) # 对每个位置取argmax得到预测的类别编号 pred_classes = torch.argmax(predictions, dim=2) # 计算单个角点的准确率 batchacc += (pred_classes == labels).sum(dim=0) # 计算全对的准确率 batchfullacc += (pred_classes == labels).all(dim=1).sum() total_samples += images.size(0) self.train() # 回到训练模式 tl, tr, bl, br = (batchacc / total_samples).tolist() acc = (batchfullacc / total_samples).item() return tl, tr, bl, br, acc
修改后的完整训练代码
import torch import torch.nn as nn from IPython.display import clear_output EPOCHS = 20 # 增加epoch数量,给模型足够学习时间 activation_function = nn.LeakyReLU() LossList, AccuracyList, tllist, trlist, bllist, brlist = [], [], [], [], [], [] device = "cuda" if torch.cuda.is_available() else "cpu" model = MyModel(activation_function).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) loss_objective = nn.CrossEntropyLoss() for epoch in range(EPOCHS): model.train() epoch_loss = 0.0 for x, y in training_generator: x = x.to(device) y = y.to(device) optimizer.zero_grad() predictions = model(x) predictions = predictions.view(-1, 4, 4) loss = loss_objective(predictions, y.type(torch.LongTensor).to(device)) loss.backward() optimizer.step() epoch_loss += loss.item() * x.size(0) # 记录平均epoch损失 avg_epoch_loss = epoch_loss / len(training_generator.dataset) LossList.append(avg_epoch_loss) # 评估阶段 with torch.no_grad(): tl, tr, bl, br, acc = model.evaluate(test_generator, device=device) AccuracyList.append(acc) tllist.append(tl) trlist.append(tr) bllist.append(bl) brlist.append(br) # 每个epoch打印结果,方便观察 clear_output(wait=True) print( f'Epoch {epoch+1}/{EPOCHS}, \n ' f'Average Loss: {avg_epoch_loss:.4f}, \n ' f'Full Accuracy: {acc:.4f}, \n ' f'Top left Accuracy: {tl:.4f}, \n ' f'Top right Accuracy: {tr:.4f}, \n ' f'Bottom left Accuracy: {bl:.4f}, \n ' f'Bottom right Accuracy: {br:.4f}' ) print("-"*50)
额外建议
- 如果需要保证四个角点的预测是唯一的排列(每个类别0-3只出现一次),可以在模型能正常学习后,加入排列约束损失(比如用匈牙利算法计算最优匹配损失)。
- 可以尝试调整学习率、使用学习率调度器(如
ReduceLROnPlateau),或者换成SGD带动量的优化器,进一步提升性能。 - 检查训练集和测试集的标签是否正确,确保每个位置的标签确实对应正确的角点编号,避免数据本身的问题。
内容的提问来源于stack exchange,提问作者leis97

