[论文解读] A simple squared-error reformulation for ordinal classification
本文提出了一种基于固定类别索引向量和Softmax激活隐藏层的简单平方误差重构方法,用于序数分类。通过最小化预测值与真实序数类别索引之间的平方误差,该方法生成校准良好的概率分布,并在糖尿病视网膜病变数据集上优于交叉熵损失和直接优化QWK的方法,实现了更高的加权kappa分数,且对标签噪声更具鲁棒性。
In this paper, we explore ordinal classification (in the context of deep neural networks) through a simple modification of the squared error loss which not only allows it to not only be sensitive to class ordering, but also allows the possibility of having a discrete probability distribution over the classes. Our formulation is based on the use of a softmax hidden layer, which has received relatively little attention in the literature. We empirically evaluate its performance on the Kaggle diabetic retinopathy dataset, an ordinal and high-resolution dataset and show that it outperforms all of the baselines employed.
研究动机与目标
- 解决标准交叉熵损失在序数分类中的局限性,即对所有误分类一视同仁,不考虑类别间的距离。
- 开发一种简单、单模型的深度学习方法,尊重类别顺序,并在序数类别上生成有意义的概率分布。
- 评估使用固定类别索引的平方误差损失是否能在真实世界的序数医学影像任务中优于传统的交叉熵损失和直接优化QWK的方法。
- 探索在不同目标下优化不同指标时,交叉熵与加权kappa性能之间的权衡。
提出的方法
- 该方法使用一个深层神经网络,其最后一层为Softmax激活,以生成类别概率,随后通过类别索引 a = [0, 1, ..., k-1] 的固定线性组合来预测期望的序数类别。
- 损失函数定义为平方误差 (c - a^T f(x))^2,其中 c 为真实类别索引,f(x) 为Softmax层的输出。
- 固定向量 a 不参与学习,保持可解释性并避免引入额外参数,同时网络学习预测出使距离真实序数索引最近的类别概率。
- 该方法在Kaggle糖尿病视网膜病变数据集上进行评估,该数据集为高分辨率的序数分类任务,包含五种疾病严重程度阶段。
- 将该方法与标准交叉熵训练以及使用Nesterov动量的SGD直接优化加权kappa(QWK)指标的方法进行比较。
- 通过直方图和箱线图分析概率分布,以评估预测的校准性和保守性。
实验结果
研究问题
- RQ1与标准交叉熵相比,使用固定类别索引向量的简单平方误差损失是否能提升深度神经网络在序数分类中的性能?
- RQ2所提出的方法是否比交叉熵或直接QWK优化产生更校准和更保守的概率分布?
- RQ3当优化目标不同时,交叉熵与加权kappa性能之间的权衡如何体现?
- RQ4考虑到该方法放松了对正确类别概率高度集中的要求,其在标签噪声下的鲁棒性如何?
主要发现
- 所提出的方法(命名为'fix a')在验证集上的加权kappa分数高于交叉熵和直接QWK优化,优于所有基线模型。
- 与QWK相比,'fix a'方法中正确类别概率接近零的估计数量更少(约125次),表明其校准性更好,且对错误预测的过度自信程度更低。
- 与交叉熵相比,'fix a'方法中正确类别的预测概率更常接近于1,表明其在正确预测上具有更高的置信度。
- 与交叉熵相比,'fix a'方法的预测更保守,对错误类别(如类别1和2)分配的中位概率更低,反映出更谨慎的预测策略。
- QWK优化模型的交叉熵损失最高,表明其概率分布对正确类别的集中程度较低,但因其对类别距离惩罚的更好处理,仍实现了最佳QWK分数。
- 该方法成功学习到在尊重类别顺序的前提下,通过一种简单、可微且无参数的Softmax输出变换,生成有意义的序数类别概率分布。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。