5.2 量刑标尺:损失函数与优化器选择


5.2 量刑标尺:损失函数与优化器选择

本节摘要:损失函数定义"错在哪里、错多少"——分类族处理离散标签(交叉熵及其加权、焦点变体),回归族处理连续目标(均方、绝对、休伯误差),排序族处理相对顺序(成对损失)。优化器与训练日程决定"怎么改":Adam 族是默认起点,权重衰减、学习率日程与早停构成稳健训练的下半场。本节给出组合配方与完整代码。

损失是任务语义的数学翻译

复盘会的第二个议题从一道选择题开始:同一份欺诈检测的图表示,甲用交叉熵、乙用加权交叉熵、丙用焦点损失,三个模型的能力差异不来自网络结构,全部来自损失函数——损失函数就是任务要求的数学翻译,翻译失真,模型再深也是在错误方向上优化。本节把图任务常用的损失按"分类、回归、排序"三族整理,再讲优化器与训练日程的下半场配置。上一节定了阵型(在哪算),本节定标尺(算什么),下一节诊断顽疾(为什么算不准)。三族损失里图任务的特殊性在于:监督信号往往稀疏(少量节点有标签)、失衡(少数类极稀少)、带结构相关性(相邻标签不独立)——这些特性会把选型从"查表"逼成"推理"。

分类族:交叉熵及其变体

多分类的标准选择是交叉熵:它惩罚"真实类别的预测概率低",配合 softmax 输出构成完整流水线。类别不平衡时按类频反比加权,让少数类样本的梯度话语权回升。焦点损失更进一步:给"已经被分对的易样本"降权,把训练火力集中在难样本上——当负样本海量化(链接预测的负采样场景)时尤其有效。标签平滑是抗过拟合的轻量添加剂:把硬标签软化,防止模型对训练标签过度自信,在小数据图(分子数据集常见)上常有小幅增益。

import torch import torch.nn.functional as F logits = torch.tensor([[2.0, 0.1, -1.0], [0.3, 2.5, -0.2], [1.8, 0.5, 0.9]]) target = torch.tensor([0, 1, 2]) # 基准:普通交叉熵 print("交叉熵:", F.cross_entropy(logits, target).item()) # 类别不平衡:类别权重与类频成反比 counts = torch.tensor([8.0, 3.0, 1.0]) weights = counts.sum() / (3 * counts) print("加权交叉熵:", F.cross_entropy(logits, target, weight=weights).item()) # 焦点损失:易样本降权(调制系数随置信度上升而衰减) def focal_loss(logits, target, gamma=2.0): logp = F.log_softmax(logits, dim=-1) p = logp.exp() logp_t = logp.gather(1, target.unsqueeze(1)).squeeze(1) p_t = p.gather(1, target.unsqueeze(1)).squeeze(1) return (-((1 - p_t) ** gamma) * logp_t).mean() print("焦点损失:", focal_loss(logits, target).item()) # 第一个样本极好分(高置信正确),调制系数把它压到接近零; # 火力自动转移到第三个难样本上——这就是"集中量刑"的含义

回归族:误差度量决定模型的保守程度

回归任务的损失选择塑造模型的"性格"。均方误差对大误差二次惩罚,模型性格是"追逐均值、害怕离群点"——流量预测里偶尔的极端尖峰会把整体预测拉向保守。绝对误差对误差线性惩罚,性格是"抗离群、预测中位数"。休伯误差在小区间用二次、大区间退回线性,兼得两端的温和,是长尾回归目标(溶解度、活性值)的稳妥默认。图回归还有个特性值得注意:目标常呈跨数量级分布,先做对数变换再回归,比在原始尺度上硬扛离群点有效得多。

import torch import torch.nn.functional as F pred = torch.tensor([1.1, 2.0, 3.2, 8.0]) true = torch.tensor([1.0, 2.1, 3.0, 3.5]) # 最后一个是潜在离群(真实尖峰) mse = ((pred - true) ** 2).mean() mae = (pred - true).abs().mean() huber = F.smooth_l1_loss(pred, true) print(f"均方误差 {mse:.4f}|绝对误差 {mae:.4f}|休伯 {huber:.4f}") # 均方误差被离群点主导(单项 4.5 的平方贡献压倒其余), # 绝对误差与休伯温和得多——若那项尖峰值得追,选均方;若是噪声,选休伯

排序族:只关心相对顺序

推荐与补全任务的监督本质常是"正例应排在负例前",而非"正例概率必须高到某个绝对值"。成对损失(贝叶斯个性化排序风格)只约束正负对的分数差:差为正即可,超出边界的部分不再奖励——模型不必为已经排对的样本浪费梯度。这与负采样天然配套:每轮采到的负例不同,模型学到的排序面持续刷新。链接预测实践中,交叉熵与成对损失的差距通常体现在尾部——成对损失对"Top 列表的顺序"更友好,因为它的优化目标与排序指标同构。

优化器与训练日程

优化器的默认答案是 Adam 或 AdamW:自适应学习率让图模型这种"各参数梯度量级悬殊"(聚合权重与任务头梯度不同量级)的场景省心不少。AdamW 的解耦权重衰减比 Adam 的内嵌正则更干净,是当前的稳妥默认。学习率日程与早停构成下半场:余弦退火或阶梯衰减在图任务上普遍有效;早停以验证集指标为准,耐心值设得略宽(图任务的验证曲线常有平台期后二次下降的现象)。权重衰减的取值对图模型格外敏感——它同时压制任务头与聚合层的权重,过强会把表示拉向欠拟合的"全平均"状态,与过平滑症状类似但机理不同,诊断时别混淆。

import torch def make_stuff(): torch.manual_seed(0) w = torch.nn.Parameter(torch.randn(8, 4) * 0.5) x = torch.randn(32, 8) y = (torch.randn(32, 4) > 0).float().argmax(1) return w, x, y w, x, y = make_stuff() opt = torch.optim.AdamW([w], lr=0.01, weight_decay=5e-4) sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=100) best, patience, bad = float("inf"), 15, 0 for step in range(100): opt.zero_grad() loss = torch.nn.functional.cross_entropy(x @ w, y) loss.backward(); opt.step(); sched.step() if loss.item() < best - 1e-4: best, bad = loss.item(), 0 else: bad += 1 if bad >= patience: print(f"第 {step} 轮早停:验证损失 {best:.4f},已连续 {patience} 轮无改善") break print(f"结束:最终学习率 {sched.get_last_lr()[0]:.5f}(余弦退火到近零)")

组合配方速查

任务情境 损失 优化器与日程 备注
均衡多分类 交叉熵 AdamW 加余弦退火 默认起点
极度不平衡 加权交叉熵或焦点 同上,早停以宏 F1 为准 裸准确率无意义
带离群的回归 休伯 同上 目标先对数变换
链接预测 交叉熵或成对损失 同上 负采样口径固定
小数据图 交叉熵加标签平滑 强权重衰减慎用 与过平滑症状区分

⚠️ 常见坑:把训练损失当早停依据。训练损失单调下降与泛化无关,早停必须看验证集;图任务里还要锁死验证掩码,防止结构泄漏让验证指标虚高。

标尺配齐之后,复盘会的第三个议题浮出水面:明明一切配置正确,为什么层数一深、精度反而跳水?下一节诊断图模型的名顽疾——过平滑。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U