前言
测试时域适应 (Test-Time Adaptation, TTA) 是笔者的第一篇工作 所研究的领域.
笔者认为, 实现域迁移的本质是模型习得了域迁移过程中的不变量.
举个例子, 在自动驾驶中, 模型是基于晴天下数据训练所得; 但是在实际应用中, 模型需要在各种天气状况 (阴天/雨天/雪地等) 下模型识别交通状况并作出反应.
- 域迁移 (domain shift): 晴天 $\rightarrow$ 阴天/雨天/雪地等
- 不变量: 交通状况 (斑马线/交通标识/其他车辆/行人等)
巧合的是, 作者近期在学习对比学习 (Contrastive Learning) 时, 对其本质的想法是非常类似的: 模型习得正样本对中的不变量. 与该思想完全吻合的对比学习相关工作是: CPC v1, 详情见个人博客解读.
幸运的是, 不变量并不需要自己构造, 前人已经告诉我们了——互信息.
这也是 CPC v1 的思想: 最大化互信息.
基于本质的一致性: 习得不变量, 笔者尝试为对比学习和 TTA 的可行性之间搭建桥梁. 笔者首先尝试的是拿方法足够简单有效的 TENT 作为分析对象.
本来笔者还以为是个挺不错的 idea, 结果一调研发现 2020 年左右被”反复”地在多个视角阐述过. 而最早甚至可以追溯到 Bengio 在 NIPS 2004 上的工作. 果然, 大佬的眼光是领先时代几十年的!
因此, 该篇博客也从技术记录变成了论文解读. ╮(╯▽╰)╭
目录
SHOT
损失函数有三项, 其中最小化熵是常见的; SHOT 还引入了另外两项.
\[\begin{align} \mathcal{L}(g_t) = \mathcal{L}_{ent} + \mathcal{L}_{div} -\beta \mathbb{E}_{(x_t,\hat{y}_t)}\sum_{k=1}^{K}\mathbb{1}_{[k = \hat{y}_t]}\log \delta_k(h_t(g_t(x_t))). \end{align}\]$\mathcal{L}_{ent}$ 项
\[\begin{align} \mathcal{L}_{ent}(f_t;\mathcal{X}_t) = -\mathbb{E}_{x_t\in \mathcal{X}_t}\sum_{k=1}^K\delta_k(f_t(x_t))\log \delta_k(f_t(x_t)). \end{align}\]$\mathcal{L}_{\text{div}}$ 项
为了缓解最小化熵导致的预测类别坍缩, 即模型倾向于预测一个类别, 丧失了多样性; 因此作者引入了负熵 $\sum_{k=1}^{K} \hat{p}_k \log \hat{p}_k$.
\[\begin{align} \mathcal{L}_{\text{div}}(\boldsymbol{f}_t; \mathcal{X}_t) = \sum_{k=1}^{K} \hat{p}_k \log \hat{p}_k. \end{align}\]首先我们考虑 KL 散度的定义: $D_{KL}(P\Vert Q)=\sum_k P(k) \log \frac{P(k)}{Q(k)}$. $P=\hat{p}$ (经验类别分布) 和 $Q=\frac{1}{K}\mathbf{1}_{K}$ (均匀分布) 代入:
\[\begin{align} &D_{KL}\left(\hat{p}\|\frac{1}{K}\mathbf{1}_{K}\right) \\ =&\sum_{k=1}^{K}\hat{p}_{k}\log\frac{\hat{p}_{k}}{1/K} \\ =&\sum_{k=1}^{K}\hat{p}_{k}\left[\log\hat{p}_{k}-\log\frac{1}{K}\right] \\ =& \sum_{k=1}^{K} \hat{p}_{k} \log \hat{p}_{k} - \sum_{k=1}^{K} \hat{p}_{k} \log \frac{1}{K} \\ =& \sum_{k=1}^{K} \hat{p}_{k} \log \hat{p}_{k} - \log \frac{1}{K} \cdot \underbrace{\sum_{k=1}^{K} \hat{p}_{k}}_{=1} \\ =& \sum_{k=1}^{K} \hat{p}_{k} \log \hat{p}_{k} + \log K . \end{align}\]移项可得: $\sum_{k=1}^{K} \hat{p}{k} \log \hat{p}{k} = D_{KL}\left(\hat{p}|\frac{1}{K}\mathbf{1}_{K}\right) - \log K$.
最终将负熵作为一个新的损失项, 得到:
\[\begin{align} \mathcal{L}_{\text{div}}(\boldsymbol{f}_t; \mathcal{X}_t) = \sum_{k=1}^{K} \hat{p}_k \log \hat{p}_k = D_{KL}\left(\hat{p}, \frac{1}{K} \mathbf{1}_K\right) - \log K \end{align}\]
Understanding Contrastive Representation Learning
AdaContrast