鸢尾花数据集mini-batch梯度下降:准确率低、损失停滞问题排查
鸢尾花数据集Mini-Batch梯度下降模型性能瓶颈问题解析
问题描述
在经典鸢尾花数据集上实现mini-batch梯度下降时,模型准确率始终卡在75-80%,损失值停滞在0.45左右,即使将迭代次数拉到10000也无改善。同时对损失计算方式存疑,使用torch.max()时出现错误:
Expected floating point type for target with class probabilities, got Long
模型定义代码
class NeuralNetwork(nn.Module): def __init__(self): super().__init__() self.linear_stack = nn.Sequential( nn.Linear(4,128), nn.ReLU(), nn.Linear(128,64), nn.ReLU(), nn.Linear(64,3), ) def forward(self, x): logits = self.linear_stack(x) return logits
训练循环代码(batch size=10)
lr = 0.01 model = NeuralNetwork() optim = torch.optim.Adam(model.parameters(), lr=lr) loss = torch.nn.CrossEntropyLoss() n_iters = 1000 steps = n_iters/10 LOSS = [] for epochs in range(n_iters): for i,(inputs, labels) in enumerate(train_loader): out = model(inputs) train_labels = transform_label(labels) l = loss(out, train_labels) l.backward() #update weights optim.step() optim.zero_grad() LOSS.append(l.item()) if epochs%steps == 0: print(f"\n epoch: {int(epochs+steps)}/{n_iters}, loss: {sum(LOSS)/len(LOSS)}") #if i % 1 == 0: #print(f" steps: {i+1}, loss : {l.item()}")
训练输出
epoch: 100/1000, loss: 1.0636296272277832
epoch: 400/1000, loss: 0.5142968013338076
epoch: 500/1000, loss: 0.49906910391073867
epoch: 900/1000, loss: 0.4586030915751588
epoch: 1000/1000, loss: 0.4543738731996598
性能瓶颈的关键原因及修复方案
- 数据未标准化:鸢尾花的4个特征数值范围差异明显,直接输入会导致梯度更新失衡,模型训练不稳定。必须对特征做标准化处理,比如计算训练集特征的均值和标准差,将每个特征转换为均值为0、方差为1的分布,可使用
torchvision.transforms.Normalize或手动实现。 - 训练循环逻辑错误:当前
n_iters=1000意味着训练1000个完整epoch,但仅记录每个epoch最后一个batch的损失,且打印的epoch数计算错误(epochs+steps无意义)。另外缺少验证环节,无法判断模型是过拟合还是欠拟合,建议在每个epoch后加入验证集的准确率和损失计算。 - 学习率过高:Adam优化器用0.01的学习率对鸢尾花这种小数据集来说过大,容易导致模型在最优解附近震荡,无法收敛到更低损失。建议将lr调整为0.001或0.0005。
- 模型结构冗余:鸢尾花是简单三分类任务,128+64的隐藏层神经元数量过多,易引发过拟合。可简化模型为
nn.Linear(4,16)+ ReLU +nn.Linear(16,3),或在隐藏层后加入nn.Dropout(0.2)抑制过拟合。 - 标签处理不当:
transform_label若将原始的[0,1,2]类别索引转成one-hot向量,会与CrossEntropyLoss的要求冲突——该损失函数直接接受Long类型的类别索引作为目标,无需额外映射。若映射后的标签是one-hot形式,会导致损失计算异常,甚至影响模型收敛。
损失计算相关问题解答
- 当前损失计算是否可行? 只要
train_labels是Long类型的类别索引(0、1、2),使用CrossEntropyLoss完全可行,该损失函数会自动将模型输出的logits与类别索引计算交叉熵损失。 torch.max()报错原因? CrossEntropyLoss的目标参数有两种合法形式:一是Long类型的类别索引,二是Float类型的类别概率分布。若你用torch.max()处理原始Long类型标签,错误地将其转换为不符合要求的形式,就会触发该报错。正确做法是直接使用原始类别索引作为损失函数的目标,无需额外处理。
内容的提问来源于stack exchange,提问作者rndnewbiee22
相关产品推荐
相关产品推荐

