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

鸢尾花数据集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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 00:41:59