大家的AI
机器学习Playground 试玩
加载中…

学习

Ch.04

损失函数深化:类别不平衡与度量学习

想象一场考试:100 道题里 99 道是简单加法,只有 1 道是很难的论述题。加法全对就能拿 99 分,很容易误以为“学得很好”。AI 里这叫类别不平衡(Class Imbalance):只靠多数票,就会漏掉所有关键的少数类(罕见病、不良品、欺诈等)。
本章要重新设计损失函数处理不平衡,并进一步超越“选对标签”,用数据之间的距离学习相似度——走进度量学习(Metric Learning)的世界。
按概念→直观比喻→公式→实践提示掌握加权 CE、Focal loss、Triplet/Contrastive loss,避免“分数高、重点全错”的实务陷阱。
下方 2×2 图展示四种损失如何改变嵌入空间:训练前 → 训练后。
训练前训练后Contrastive Loss训练前训练后Triplet Loss训练前训练后w↑Weighted CE训练前训练后难样本Focal Loss
每格展示训练前(左)→训练后(右)。Contrastive 拉近正样本对、推远负样本对;Triplet 让 anchor–positive 靠近、negative 超出 margin。加权 CE 对少数类 w↑;Focal 让易样本缩小、难样本放大聚焦。

相似的拉近、不同的推远 — 用距离学习

不平衡用权重与 Focal,相似度用 Contrastive 与 Triplet

损失函数进阶:从不平衡与距离中学习

1. 什么是类别不平衡?(掉进多数决陷阱的 AI)
概念: 某些类别的样本数量远远多于或少于其他类别。普通训练会让模型只关注好猜的多数类。
直观比喻: 工厂检测中 99% 为正常品时,模型全预测“正常”也能有 99% 准确率。但真正要抓的是隐藏的 1% 不良品(少数类)。
实践提示: 在损失中加入类权重,或使用聚焦难样本的 Focal loss。
2. 加权交叉熵 (Weighted CE)
概念: 每个类别 ccc 乘以权重 wcw_cwc​,少数类错分时罚得更重。核心公式:L=−wclog⁡(pc)L = - w_c \log(p_c)L=−wc​log(pc​)。
直观比喻: 改评分表:99 道简单题各 1 分,1 道稀有论述题 100 分——学生(AI)就不敢忽略论述题。
实践提示: 权重过大会死记少数类噪声,要慢慢调并查看各类指标。
3. Focal loss — 会的跳过,不会的猛攻
概念: 对已很有把握的简单样本(预测概率 ptp_tpt​ 高)大幅减小损失,让难样本主导训练。核心公式:Lfocal=−(1−pt)γlog⁡(pt)L_{\text{focal}} = - (1-p_t)^\gamma \log(p_t)Lfocal​=−(1−pt​)γlog(pt​),γ\gammaγ 控制专注程度。
直观比喻: 复习时跳过已掌握的章节,把时间全砸在总错的薄弱单元上。
实践提示: 在目标检测等背景(多数)与物体(少数)差距极大时特别有效。
4. 度量学习 — 同类靠近,异类拉开
概念: 不只背答案,而是学距离——猫图彼此近、离狗图远。Triplet loss 用锚点、正例、负例三点,核心公式:L=max⁡(0,d(a,p)−d(a,n)+α)L = \max(0, d(a,p) - d(a,n) + \alpha)L=max(0,d(a,p)−d(a,n)+α)。
直观比喻: 婚礼座位:好友(Positive) 同桌,关系差的人(Negative) 至少隔开安全距离 α\alphaα。
实践提示: 人脸识别、相似商品推荐等需要比较“有多像”的任务广泛使用。

损失函数速览

加权 CE — 样本少的类错分罚得更重。
核心公式: L=−wclog⁡(pc)L = -w_c \log(p_c)L=−wc​log(pc​)
符号说明 — pcp_cpc​ 是模型对真实类别 ccc 的预测概率(0~1)。wcw_cwc​ 是类别 ccc 的配分,少数类权重更大。log⁡(pc)\log(p_c)log(pc​) 在概率低(错得多)时变大,前面的负号让训练朝提高 pcp_cpc​ 的方向进行。
权重规则: 总样本 NNN、类别 KKK 时 wc∝N/(Knc)w_c \propto N/(K n_c)wc​∝N/(Knc​) — 类别 ccc 的样本数 ncn_cnc​ 越小,wcw_cwc​ 越大。
数值例: pc=0.2p_c=0.2pc​=0.2, wc=5w_c=5wc​=5 时损失 ≈−5log⁡(0.2)≈8\approx -5\log(0.2) \approx 8≈−5log(0.2)≈8 — 同样错分,加权后罚得更重。
比喻: 普通交叉熵加上按类配分。
Focal loss — 已会做的样本损失变小,聚焦难样本。
核心公式: Lfocal=−(1−pt)γlog⁡(pt)L_{\text{focal}} = -(1-p_t)^\gamma \log(p_t)Lfocal​=−(1−pt​)γlog(pt​)
符号说明 — ptp_tpt​ 是模型对该样本的预测把握度(正确概率)。(1−pt)(1-p_t)(1−pt​) 表示还不确定的程度;(1−pt)γ(1-p_t)^\gamma(1−pt​)γ 在 γ\gammaγ 较大时更强烈地削弱简单样本(ptp_tpt​ 高)的贡献。log⁡(pt)\log(p_t)log(pt​) 与 CE 相同的基础损失骨架。
数值例: pt=0.9p_t=0.9pt​=0.9, γ=2\gamma=2γ=2 时 (1−0.9)2=0.01(1-0.9)^2=0.01(1−0.9)2=0.01 — 损失缩到约 1%,已掌握的题几乎被忽略。pt=0.3p_t=0.3pt​=0.3 时 (0.7)2=0.49(0.7)^2=0.49(0.7)2=0.49,仍有较大惩罚。
比喻: 复习时跳过已掌握的章节,专攻总错的薄弱单元。
Triplet loss — 朋友拉近、陌生人推远。
核心公式: L=max⁡(0, d(a,p)−d(a,n)+α)L = \max(0,\, d(a,p) - d(a,n) + \alpha)L=max(0,d(a,p)−d(a,n)+α)
符号说明 — aaa 是锚点(anchor),ppp 是同身份(positive),nnn 是不同身份(negative)。d(⋅,⋅)d(\cdot,\cdot)d(⋅,⋅) 是两点间的距离(如 L2)。α\alphaα 是 negative 距 anchor 至少要保持的间隔(margin)。max⁡(0,⋅)\max(0,\cdot)max(0,⋅) 表示条件已满足时损失为 0,无需再推。
数值例: d(a,p)=1d(a,p)=1d(a,p)=1, d(a,n)=4d(a,n)=4d(a,n)=4, α=0.5\alpha=0.5α=0.5 → 1−4+0.5=−2.51-4+0.5=-2.51−4+0.5=−2.5 → max⁡(0,−2.5)=0\max(0,-2.5)=0max(0,−2.5)=0(已足够远)。若 d(a,n)=1.2d(a,n)=1.2d(a,n)=1.2 → 仍有 0.3 的惩罚。
比喻: 婚礼座位 — 好友(ppp)坐 anchor 旁,不合的人(nnn)至少隔开 α\alphaα。
Contrastive loss — 同身份(或增强)拉近,不同身份推远。
核心公式 — positive: L+=0.5 d2L_+ = 0.5\,d^2L+​=0.5d2 · negative: L−=0.5 max⁡(0, m−d)2L_- = 0.5\,\max(0,\, m-d)^2L−​=0.5max(0,m−d)2
符号说明 — ddd 是两个嵌入之间的距离。positive 对 ddd 越接近 0 损失越小(拉近)。negative 的 mmm 是最小间隔;当 d≥md \ge md≥m 时 max⁡(0,m−d)=0\max(0,m-d)=0max(0,m−d)=0,无需再推。0.50.50.5 是缩放常数。
数值例: positive 中 d=0.4d=0.4d=0.4 → 0.5×0.16=0.080.5\times0.16=0.080.5×0.16=0.08(近则罚轻)。negative 中 d=0.6d=0.6d=0.6, m=1m=1m=1 → 0.5×(0.4)2=0.080.5\times(0.4)^2=0.080.5×(0.4)2=0.08 — 仍太近,继续推远。
比喻: 同一人的照片聚到一起,不同人的照片至少隔开 mmm。

为什么重要

1. 识破虚假的 100 分成绩单
不平衡数据上的“99% 准确率”可能是错觉。要确认模型是否在做真正重要的事,不能只看准确率,而要用加权 CE或 Focal loss 在损失里写明什么更重要。
2. 损失决定优先级
模型最该重视什么,写在损失的形状里。目标错了,练得再多也解决不了真正的问题。处理不平衡与距离学习时,关键是让损失告诉模型惩罚该落在哪里。
3. 搜索与推荐服务的脊梁
购物 App 的“找相似衣服”、手机的“人脸解锁”时,AI 不是在选客观题答案,而是计算两张图有多像(距离)。度量学习支撑着这些现代 AI 服务。

如何使用

① 解决不平衡 — 按顺序诊断
症状: 常见类全对,稀有类全错。
步骤:
1. 看分布 — 各类数量差多少?
2. 换损失 — 试 加权 CE 或 Focal loss。
3. 复查训练设置 — 损失变了,学习率、batch 等也要一起看。
② 度量学习 — 凑对训练
度量学习需要把数据准备成“组合”。
- Triplet: (基准照、同一人、不同人) 三个一组。
- Contrastive: (原图、增强原图) 作 positive 拉近,其余推远。
让模型学距离,使相似样本聚在一起。
③ 扔掉太简单的题 (Hard Negative Mining)
只学苹果 vs 汽车太容易,模型会自满。故意加入很像但标签不同的 Hard Negative,像做难模考一样提升实战力。
④ 选尺子 — L2 距离 vs 余弦相似度
- L2: 直线距离,大小和位置都重要。
- 余弦: 箭头方向是否一致,适合语义或“气质”。
实务: 忽视少数→权重/Focal;只做简单题→提高 Focal γ\gammaγ;检索差→加 Hard negative。

总结

一句话: 本章学习了如何通过设计损失函数处理类别不平衡,以及如何用嵌入距离学习相似度。
四种核心损失各有分工。加权 CE 用逆频率 wcw_cwc​ 重罚少数类误分;Focal loss 结合 α\alphaα 与 (1−pt)γ(1-p_t)^\gamma(1−pt​)γ 削弱简单样本的贡献;Triplet loss 用 margin α\alphaα 让正样本靠近锚点、负样本远离;Contrastive loss 通过拉近/推远正负对整理嵌入空间。
实务上,若模型只预测多数类,应改用加权 CE 或 Focal,并查看各类 F1。若总在简单样本上打转,可调 Focal 的 γ\gammaγ,但过大可能不稳定。若 Triplet 损失接近 0,往往说明 negative 太简单,需要 hard negative mining。人脸认证、相似检索则要关注度量嵌入以及 L2 与余弦距离的选择。
调参时先确认类别分布与指标,再选损失并调整学习率等设置。不要一次改多项,逐项对比效果更稳妥。

解题方法说明

本章题目可分为两条线:类别不平衡分类与基于距离的相似度学习。不平衡时,只猜多数类也能得到很高的准确率,仅用普通交叉熵容易漏掉稀有类。加权 CE 为每类设置 wcw_cwc​,加大对少数类误分的惩罚;Focal loss 用 (1−pt)γ(1-p_t)^\gamma(1−pt​)γ 降低简单样本的影响,让难样本主导训练。度量学习则不在标签上“选对答案”,而在嵌入空间里用距离学“有多像”。Triplet loss 用 anchor·positive·negative 三点最小化 L=max⁡(0,d(a,p)−d(a,n)+α)L=\max(0,d(a,p)-d(a,n)+\alpha)L=max(0,d(a,p)−d(a,n)+α);Contrastive loss 拉近正样本对、推远负样本对。题干若写“只有少数类总错”,先想权重/Focal;若写“negative 太简单”,先想 hard mining;若是人脸验证、相似检索,先想度量嵌入。
定义题要先想损失在“更重地惩罚什么”。例如“加权 CE 为何给少数类更大的 wcw_cwc​?”
① 跳过反传与机制无关,
③ 固定 batch 也不是加权目的。核心是提高少数类误分代价,选
②,答案 2。

应用题先读数据设定。“欺诈 1%、正常 99%”这类极端不平衡下,只保留 CE(①)或删除少数类(③)都不合理,应优先尝试加权 CE 或 Focal loss(②)。→ 答案 2

计算题按公式逐步代入。wB=N/(K⋅nB)w_B=N/(K\cdot n_B)wB​=N/(K⋅nB​),N=1000N=1000N=1000, K=2K=2K=2, nB=100n_B=100nB​=100 时得 1000/(2⋅100)=51000/(2\cdot100)=51000/(2⋅100)=5,答案
②。
定义例 — “Focal 中 (1−pt)γ(1-p_t)^\gamma(1−pt​)γ 的作用?”它降低简单样本的损失权重,让训练聚焦难样本,不是
① 相同损失或
③ 提高学习率。→ 答案 2

判断例 — “Triplet 需要 anchor、positive、negative。”正确,因为靠三点距离关系学习。→ 答案 1

应用例 — “做人脸验证嵌入”更适合用 Triplet/Contrastive 学距离,而非普通分类。→ 答案 1

选择例 — 极端不平衡(如检测背景极多)常优先 Focal(②) 而非仅 CE(①)。→ 答案 2

概念例 — d(a,p)=1d(a,p)=1d(a,p)=1, d(a,n)=4d(a,n)=4d(a,n)=4, α=0.5\alpha=0.5α=0.5 时 max⁡(0,1−4+0.5)=0\max(0,1-4+0.5)=0max(0,1−4+0.5)=0,已满足 margin,无额外惩罚。→ 答案 3

计算例 — L2 距离 (0,0)(0,0)(0,0)–(3,4)(3,4)(3,4) 为 9+16=5\sqrt{9+16}=59+16​=5。→ 答案
②