不确定性校准深度实战:从可靠性图、ECE 到温度缩放与贝叶斯校准的工程全解

一个在 ImageNet 上 85% 准确率的模型,可能给错误预测贴上 0.99 的置信度——这正是现代深度网络普遍存在的"过自信(overconfidence)"问题。准确率衡量的是区分度(separability),而校准(calibration)衡量的是"置信度是否可信"。二者解耦:网络可以很准但并不校准。本文与本站《损失函数》《正则化》《激活函数》《RMSNorm》《相对位置编码》《交叉注意力》《残差连接》《学习率调度》《梯度累积》《词嵌入》共同构成 Transformer 内部机制与训练工程深度解系列,补全"模型输出的可靠性"这一被长期忽视的环节,并与《模型提取与盗取攻击》《成员推断攻击》共享"置信度"这一攻防面。

一、为什么准确率不等于可信

现代过参数化网络(尤其 ResNet、Transformer)在验证集上区分度极高,却普遍过自信:SoftMax 输出的概率与真实正确率严重偏离。Guo 等(2017,On Calibration of Modern Neural Networks)用可靠性图证明——加深/加宽、BatchNorm、权重衰减、更多 epoch 都会推高过自信;而数据增强(Mixup/CutMix)、标签平滑、知识蒸馏这些"让训练更稳"的技巧,恰恰进一步放大过自信。

核心区分:区分度 vs 校准度是可分离的。一个模型可以分类很准但概率不可信。在安全敏感场景(医疗、自动驾驶、风控、模型安全攻防),"模型说 90% 把握"必须真的有 90% 正确率。

二、第一性原理:什么是校准

一个模型是完美校准的,当且仅当对预测置信度 $\hat c$ 的任意取值:

$$\mathbb{P}(\hat Y = Y \mid \hat C = \hat c) = \hat c$$

即"所有被预测为 p 概率的样本,真实正确率也是 p"。理想情况落在图上的对角线(perfect calibration line)。

可靠性图(reliability diagram / calibration curve)是核心诊断工具:把预测置信度分箱(bin),每箱画"平均置信度 vs 经验正确率"。偏离对角线即未校准。

区域 含义 风险
曲线在对角线上方 欠自信(under-confident) 浪费确信度,阈值决策保守
曲线在对角线下方 过自信(over-confident) 高风险误判,OOD 检测失效

三、评估指标

指标 定义 注意
ECE $\sum_{k=1}^B\frac{ B_k }{n}\big\lvert\mathrm{acc}(B_k)-\mathrm{conf}(B_k)\big\rvert$ 等宽分箱,对 bin 数敏感
MCE $\max_k\big\lvert\mathrm{acc}(B_k)-\mathrm{conf}(B_k)\big\rvert$ 最差箱偏差,关注尾部
Class-wise ECE 每个类单独算 ECE 再平均 类别不平衡时比 ECE 更公平
Brier Score $\frac{1}{n}\sum_i(p_i-o_i)^2$ Proper scoring rule,可微可优化
NLL $-\frac{1}{n}\sum_i\log q_{y_i}$ Strictly proper,但惩罚整体分布

ECE 不是 proper scoring rule(可通过对概率做"乱序重排"骗过分箱),NLL/Brier 才是。所以永远同时看 NLL 与 ECE——NLL 降但 ECE 升的背离,说明分布整体更准却更不校准(常见于蒸馏网络)。


import torch, torch.nn.functional as F

def ece(logits, labels, n_bins=15):
    probs = F.softmax(logits, dim=-1)
    conf, pred = probs.max(-1)
    correct = (pred == labels).float()
    edges = torch.linspace(0, 1, n_bins + 1)
    val = 0.0
    for i in range(n_bins):
        lo, hi = edges[i].item(), edges[i + 1].item()
        mask = (conf > lo) & (conf <= hi)
        n = int(mask.sum())
        if n:
            acc = correct[mask].mean()
            val += (n / len(conf)) * (acc - conf[mask].mean()).abs()
    return val.item()

def brier_score(logits, labels):
    probs = F.softmax(logits, dim=-1)
    return ((probs.gather(1, labels.unsqueeze(1)) - 1) ** 2).mean().item()

def class_wise_ece(logits, labels, n_bins=15):
    out = []
    for c in labels.unique():
        m = (labels == c)
        out.append(ece(logits[m], labels[m], n_bins))
    return sum(out) / len(out)

四、校准方法谱系

4.1 参数化缩放(后处理,不打乱预测)

在冻结模型、独立验证集上学习一个从 logits 到 calibrated 概率的映射,是当前工业界首选,因为不改动训练、不损失区分度:

方法 形式 参数量 特点
Platt Scaling $\sigma(a z + b)$ 2 二分类经典
Isotonic Regression 保序阶梯函数 $O(n)$ 非参数,强但易过拟合小验证集
Vector Scaling $W z + b$ $d+1$ 每类独立缩放
Matrix Scaling $W z$ $d^2$ 全矩阵,需大验证集
Temperature Scaling $\mathrm{Softmax}(z/T)$ 1 单标量,保留 argmax 与 AUC

温度缩放(Temperature Scaling, Guo 2017)是性价比之王:仅一个标量 $T>0$,把 logits 除以 $T$ 再 SoftMax。在独立验证集上最小化 NLL 拟合 $T$。它不改变预测类别、不改变 AUC,只把概率整体"软化/收紧",却把过自信压回校准线。


from torch.optim import LBFGS

def fit_temperature(logits, labels, lr=0.05, max_iter=200):
    """在独立验证集上最小化 NLL 拟合单标量温度 T(>0)。"""
    T = torch.nn.Parameter(torch.ones(1))
    opt = LBFGS([T], lr=lr, max_iter=max_iter)
    def closure():
        opt.zero_grad()
        loss = F.cross_entropy(logits / T.clamp_min(1e-3), labels)
        loss.backward()
        return loss
    opt.step(closure)
    return T.item()           # 过自信网络通常 T>1(软化)

# 推理:先 eval(),全程 fp32
# q_calibrated = torch.softmax(logits / T, dim=-1)

经验:T 通常 > 1(模型过自信,需把分布"摊平")。关键约束:拟合 T 必须在 .eval() 模式下、用 fp32 logits——否则 BN/Dropout 噪声污染 logits,且低精度会溢出。

4.2 训练期方法(及其反作用)

  • 标签平滑:把目标从 one-hot 软化,反而让模型更过自信(平滑目标与校准目标冲突)——与本站《损失函数》呼应。
  • 数据增强(Mixup/CutMix):软标签让网络对边界样本过自信。
  • 知识蒸馏:拟合教师软分布,学生天然过自信。
  • 结论:若最终要校准,训练期就该直接用温度缩放或标签平滑的校准友好变体,而非事后补救。

4.3 真实不确定性建模(认知 + 偶然)

温度缩放置信度"看起来"校准,但仍是点估计。要量化"模型不知道什么",需建模不确定性:

  • 偶然不确定性(aleatoric):数据本身噪声,不可约。
  • 认知不确定性(epistemic):模型参数不确定,数据越多越小。
  • MC Dropout(在 .eval() 下多次前向取方差)、Deep Ensemble(多模型方差)、变分推断、Conformal Prediction(给出带覆盖保证的预测集 $\{y: \hat\mu(x)+\hat\sigma(x)\cdot t_\alpha\}$,不假定分布)。

五、生产陷阱清单(必读)

# 陷阱 后果 修复
1 仅看 accuracy 忽略校准 高风险场景误判 安全/医疗场景强制报 ECE
2 在训练集算 ECE 数据泄漏,虚假 0 ECE 必须用独立验证集
3 拟合 T 时未 .eval() BN/Dropout 噪声污染 logits 校准前 model.eval()
4 T 用 fp16/bf16 计算 低精度溢出,T 估计失真 全程 fp32
5 ECE bin 数随意(如 10 vs 20) ECE 波动大不可比 固定 n_bins=15,报告等样本分箱
6 类别不平衡只看整体 ECE 多数类掩盖少数类失准 用 class-wise ECE
7 以为温度缩放能改决策 T 不改变 argmax 需改阈值时单独重标定
8 NLL↓ 就当校准好了 NLL 与 ECE 可背离 二者同报
9 蒸馏/增强后的网络直接上线 天然过自信 上线前先测 ECE 或训练期校准
10 用校准概率做 OOD 检测 T 偏低使 OOD 分数失效 先温度缩放再取 OOD 分数
11 多标签/多任务直接套单 T 各头尺度不一失效 每头独立 T 或 per-class
12 在线/流式场景一次标定不更新 数据漂移致校准退化 定期在近期窗口重拟 T

六、可复现工具箱


import torch, torch.nn.functional as F
from torch.optim import LBFGS

class CalibrationKit:
    """置信度校准工程工具箱:可靠性图、ECE/MCE、Brier、温度缩放拟合。"""
    @staticmethod
    def reliability_curve(logits, labels, n_bins=15):
        probs = F.softmax(logits, dim=-1)
        conf, pred = probs.max(-1)
        correct = (pred == labels).float()
        edges = torch.linspace(0, 1, n_bins + 1)
        xs, ys = [], []
        for i in range(n_bins):
            lo, hi = edges[i].item(), edges[i + 1].item()
            mask = (conf > lo) & (conf <= hi)
            if int(mask.sum()):
                xs.append(conf[mask].mean().item())
                ys.append(correct[mask].mean().item())
        return xs, ys

    @staticmethod
    def ece(logits, labels, n_bins=15):
        probs = F.softmax(logits, dim=-1)
        conf, pred = probs.max(-1)
        correct = (pred == labels).float()
        edges = torch.linspace(0, 1, n_bins + 1)
        val = 0.0
        for i in range(n_bins):
            lo, hi = edges[i].item(), edges[i + 1].item()
            mask = (conf > lo) & (conf <= hi)
            n = int(mask.sum())
            if n:
                val += (n / len(conf)) * (correct[mask].mean() - conf[mask].mean()).abs()
        return val.item()

    @classmethod
    def fit_temperature(cls, logits, labels, lr=0.05, max_iter=200):
        assert not logits.requires_grad or True
        T = torch.nn.Parameter(torch.ones(1))
        opt = LBFGS([T], lr=lr, max_iter=max_iter)
        def closure():
            opt.zero_grad()
            loss = F.cross_entropy(logits / T.clamp_min(1e-3), labels)
            loss.backward()
            return loss
        opt.step(closure)
        return T.item()

# 用法:
# model.eval()
# T = CalibrationKit.fit_temperature(val_logits, val_labels)   # 通常 > 1
# calibrated = torch.softmax(test_logits / T, dim=-1)
# print('ECE before/after:', CalibrationKit.ece(test_logits, y),
#       CalibrationKit.ece(test_logits / T, y))

七、与系列其他文章的桥接

  • 损失函数:标签平滑、知识蒸馏都让网络更过自信——本文是它们的"下游校准补丁";理解 ECE 能解释为何平滑后验证指标好看但置信度失真。
  • 正则化:Dropout 既是正则又是 MC Dropout 不确定性估计的来源;数据增强的隐式正则代价是过自信,需本文方法补救。
  • 模型安全(成员推断 / 模型提取):成员推断攻击高度依赖模型输出的细粒度置信度;若模型经温度缩放校准,黑盒攻击的信噪比下降,直接影响《成员推断攻击》《模型提取与盗取攻击》的攻击成功率——校准因此是一把"双刃剑"。
  • LLM 推理:解码时的 temperature 参数正是温度缩放的工程化身——低温更确定、高温更多样,本质同本文 $z/T$ 的软化逻辑。

校准不是"锦上添花",而是把"模型说有 90% 把握"变成"真的有 90% 把握"的最后一道工序。在一切高风险决策之前,先问一句:这个概率,校准了吗?

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿
网站二维码

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部