强制间隔投影(Margin-Enforcing Projection)
TL;DR · AI 摘要
Title: Enforcing Projection) - 科学空间 Scientific Spaces URL Source: Markdown Content: 这篇文章我们介绍一个数学运算,名为“强制间隔投影(Margin-Enfo...
核心要点
- 主题聚焦:强制间隔投影(Margin-Enforcing Projection)
- 来源:科学空间,建议结合原文判断细节。
- AI 分析暂不可用,本条为保底评分与摘要。
这篇文章我们介绍一个数学运算,名为“强制间隔投影(Margin-Enforcing Projection,MEP)”,它将向量按指定方式划分为两部分,然后要求这两部分的间隔至少为$m$(Margin)。
背景简介[#](https://spaces.ac.cn/archives/11784#%E8%83%8C%E6%99%AF%E7%AE%80%E4%BB%8B)
我们知道,分类任务是要选出正确的类别,所以训练目标通常是“正类分数大于负类分数”就行。但在某些场景下,我们不仅希望正类分数要超过负类分数,还希望至少超过指定间隔$m > 0$。
这些场景主要分两种。第一种是希望分类结果更为稳健,不轻易受到随机噪声的干扰,尤其是低精度推理场景,所以我们希望预测结果的间隔能更显著一些;第二种是根本不用分类模型,而是想着通过分类来学习特征,最终场景是用特征去检索,此时如果不设置Margin,那么边界处的检索结果就很容易出错。
实际上,这两个场景在早些年就已经非常成熟,尤其是场景二,典型代表就是人脸识别特征模型的训练,所以本文其实算是“挖坟”了。该问题的标准思路是设计各种Margin Loss,比如Hinge Loss、Margin Softmax、AM-Softmax等,我们之前在《基于GRU和AM-Softmax的句子相似度模型》和《从三角不等式到Margin Softmax》中也略有介绍。
本文的设想是,如果我们有某种运算,能够将预测分数$\mathbf{\mathit{x}}$投影成符合间隔要求的分数$\mathbf{\mathit{y}}$,那么直接以它为目标给模型学就是了,比如最小化$\parallel \mathbf{\mathit{y}} - \mathbf{\mathit{x}} \parallel_{2}^{2}$。而这个投影运算,便是接下来要研究的问题。
数学定义[#](https://spaces.ac.cn/archives/11784#%E6%95%B0%E5%AD%A6%E5%AE%9A%E4%B9%89)
记$\mathbf{\mathit{x}} = \left(\right. x_{1} , x_{2} , \hdots , x_{n} \left.\right) , \mathbf{\mathit{z}} = \left(\right. z_{1} , z_{2} , \hdots , z_{n} \left.\right)$,$\mathbf{\mathit{x}}$代表模型的预测分数,简单起见,我们假设前$k$个分数代表得分,后$n - k$个代表负类得分,我们定义
$$ (\text{1}) \mathcal{P}_{m} \left(\right. \mathbf{\mathit{x}} \left.\right) \triangleq \underset{\mathbf{\mathit{z}} \in \mathbb{R}^{n}}{\text{argmin}} d \left(\right. \mathbf{\mathit{z}} , \mathbf{\mathit{x}} \left.\right) \text{s}.\text{t}. min \left(\right. \mathbf{\mathit{z}}_{\leq k} \left.\right) - max \left(\right. \mathbf{\mathit{z}}_{> k} \left.\right) \geq m $$
其中$\mathbf{\mathit{z}}_{\leq k} = \left(\right. z_{1} , \hdots , z_{k} \left.\right) , \mathbf{\mathit{z}}_{> k} = \left(\right. z_{k + 1} , \hdots , z_{n} \left.\right)$,$k$是大于0、小于$n$的整数,$m > 0$是给定的间隔,$d \left(\right. \mathbf{\mathit{z}} , \mathbf{\mathit{x}} \left.\right)$是要最小化的距离函数,我们有L1距离和L2距离两种选择,接下来会分别讨论。 这个定义很直观,就是在满足间隔要求的前提下,找跟预测分数最接近的向量,这符合投影运算的一般思想,所以我们称之为“强制间隔投影(Margin-Enforcing Projection,MEP)”。以这样的投影结果作为学习目标,理论上能让模型往最轻松、最快捷的路径走,同时也能在达到目标后及时“刹车”,避免过度训练。
共同结果[#](https://spaces.ac.cn/archives/11784#%E5%85%B1%E5%90%8C%E7%BB%93%E6%9E%9C)
若$\mathbf{\mathit{x}}$本就满足$min \left(\right. \mathbf{\mathit{x}}_{\leq k} \left.\right) - max \left(\right. \mathbf{\mathit{x}}_{> k} \left.\right) \geq m$,那么显然$\mathbf{\mathit{z}}^{*} = \mathbf{\mathit{x}}$,这是平凡的。不失一般性,下面都假设$min \left(\right. \mathbf{\mathit{x}}_{\leq k} \left.\right) - max \left(\right. \mathbf{\mathit{x}}_{> k} \left.\right) < m$。不难看出,不管$d$选择L1还是L2距离,最优解$\mathbf{\mathit{z}}^{*}$都必然在
$$ (\text{2}) min \left(\right. \mathbf{\mathit{z}}_{\leq k}^{*} \left.\right) - max \left(\right. \mathbf{\mathit{z}}_{> k}^{*} \left.\right) = m $$
取到,否则总可以让某些$z_{i}$更靠近$x_{i}$来降低目标值。基于这个观察,设$min \left(\right. \mathbf{\mathit{z}}_{\leq k}^{*} \left.\right) = ℓ$,那么$max \left(\right. \mathbf{\mathit{z}}_{> k}^{*} \left.\right) = ℓ - m$,那么可以得到
$$ (\text{3}) \mathbf{\mathit{z}}_{\leq k}^{*} = max \left(\right. \mathbf{\mathit{x}}_{\leq k} , ℓ \left.\right) , \mathbf{\mathit{z}}_{> k}^{*} = min \left(\right. \mathbf{\mathit{x}}_{> k} , ℓ - m \left.\right) $$
此时
$$ (\text{4}) \begin{matrix}\mathbf{\mathit{z}}^{*} - \mathbf{\mathit{x}} = & \left[\right. max \left(\right. \mathbf{\mathit{x}}_{\leq k} , ℓ \left.\right) - \mathbf{\mathit{x}}_{\leq k} , min \left(\right. \mathbf{\mathit{x}}_{> k} , ℓ - m \left.\right) - \mathbf{\mathit{x}}_{> k} \left]\right. \\ = & \left[\right. max \left(\right. ℓ - \mathbf{\mathit{x}}_{\leq k} , 0 \left.\right) , min \left(\right. ℓ - m - \mathbf{\mathit{x}}_{> k} , 0 \left.\right) \left]\right. \\ = & \left[\right. max \left(\right. ℓ - \mathbf{\mathit{x}}_{\leq k} , 0 \left.\right) , - max \left(\right. \mathbf{\mathit{x}}_{> k} + m - ℓ , 0 \left.\right) \left]\right.\end{matrix} $$
接下来就是要根据不同的$d$来求解$ℓ$。
距离之二[#](https://spaces.ac.cn/archives/11784#%E8%B7%9D%E7%A6%BB%E4%B9%8B%E4%BA%8C)
先考虑L2距离,此时我们有
$$ (\text{5}) \parallel \mathbf{\mathit{z}}^{*} - \mathbf{\mathit{x}} \parallel_{2}^{2} = \sum_{i = 1}^{k} max \left(\right. ℓ - x_{i} , 0 \left.\right)^{2} + \sum_{j = k + 1}^{n} max \left(\right. x_{j} + m - ℓ , 0 \left.\right)^{2} \triangleq f \left(\right. ℓ \left.\right) $$
求导得
$$ (\text{6}) f^{'} \left(\right. ℓ \left.\right) = 2 \sum_{i = 1}^{k} max \left(\right. ℓ - x_{i} , 0 \left.\right) - 2 \sum_{j = k + 1}^{n} max \left(\right. x_{j} + m - ℓ , 0 \left.\right) \\ (\text{7}) f^{''} \left(\right. ℓ \left.\right) = 2 \# \left{\right. ℓ > x_{i} \left.\right} + 2 \# \left{\right. ℓ < x_{j} + m \left.\right} $$
其中$\#$是计数函数,约定$1 \leq i \leq k < j \leq n$。显然$f^{''} \left(\right. ℓ \left.\right) \geq 0$,但还可以加强到$f^{''} \left(\right. ℓ \left.\right) > 0$。 这是因为$f^{''} \left(\right. ℓ \left.\right) = 0$意味着$\# \left{\right. ℓ > x_{i} \left.\right} = 0$且$\# \left{\right. ℓ < x_{j} + m \left.\right} = 0$,即同时成立$min \left(\right. \mathbf{\mathit{x}}_{\leq k} \left.\right) \geq ℓ$和$max \left(\right. \mathbf{\mathit{x}}_{> k} \left.\right) \leq ℓ - m$,这跟假设$min \left(\right. \mathbf{\mathit{x}}_{\leq k} \left.\right) - max \left(\right. \mathbf{\mathit{x}}_{> k} \left.\right) < m$矛盾。所以$f^{''} \left(\right. ℓ \left.\right) > 0$,即$f \left(\right. ℓ \left.\right)$是严格凸的,加上$f \left(\right. ℓ \left.\right)$和$f^{'} \left(\right. ℓ \left.\right)$的连续性,以及$f^{'} \left(\right. - \infty \left.\right) = - \infty$和$f^{'} \left(\right. \infty \left.\right) = \infty$,可以得出$f \left(\right. ℓ \left.\right)$的最小值点只有一个,并且必然在$f^{'} \left(\right. ℓ \left.\right) = 0$处取到。
由于$f^{'} \left(\right. ℓ \left.\right)$是分段线性函数$max \left(\right. x , 0 \left.\right)$的复合,所以$f^{'} \left(\right. ℓ \left.\right)$也是$ℓ$的分段线性函数,边界点是全体$x_{i}$和$x_{j} + m$。为了求解$f^{'} \left(\right. ℓ \left.\right) = 0$,我们先将边界点$\left{\right. x_{1} , \hdots , x_{k} , x_{k + 1} + m , \hdots , x_{n} + m \left.\right}$从小到大排列,得到$n - 1$个区间,在单个区间内,$f^{'} \left(\right. ℓ \left.\right)$是一条直线,遍历所有区间$\left[\right. a , b \left]\right.$,找到$f^{'} \left(\right. a \left.\right) \leq 0$且$f^{'} \left(\right. b \left.\right) \geq 0$的区间,在该区间内求直线的零点即可。
距离之一[#](https://spaces.ac.cn/archives/11784#%E8%B7%9D%E7%A6%BB%E4%B9%8B%E4%B8%80)
接着考虑L1距离,此时我们有
$$ (\text{8}) \parallel \mathbf{\mathit{z}}^{*} - \mathbf{\mathit{x}} \parallel_{1} = \sum_{i = 1}^{k} max \left(\right. ℓ - x_{i} , 0 \left.\right) + \sum_{j = k + 1}^{n} max \left(\right. x_{j} + m - ℓ , 0 \left.\right) \triangleq g \left(\right. ℓ \left.\right) $$
显然,$g \left(\right. ℓ \left.\right)$本身就是分段线性函数,这种函数的最小值只能在边界点取到,所以最朴素的解法就是遍历所有边界点$\left{\right. x_{1} , \hdots , x_{k} , x_{k + 1} + m , \hdots , x_{n} + m \left.\right}$,取让$g \left(\right. ℓ \left.\right)$最小者,复杂度为$\mathcal{O} \left(\right. n^{2} \left.\right)$。但我们还可以更进一步简化,首先求导得
$$ (\text{9}) \begin{matrix}g^{'} \left(\right. ℓ \left.\right) = & \# \left{\right. ℓ > x_{i} \left.\right} - \# \left{\right. ℓ < x_{j} + m \left.\right} \\ = & \# \left{\right. ℓ > x_{i} \left.\right} + \# \left{\right. ℓ \geq x_{j} + m \left.\right} - \left(\right. n - k \left.\right)\end{matrix} $$
第二个等号用了恒等式$\# \left{\right. ℓ < x_{j} + m \left.\right} + \# \left{\right. ℓ \geq x_{j} + m \left.\right} = n - k$,现在容易看出,$g^{'} \left(\right. ℓ \left.\right)$是从负到正单调递增的。当然,$g^{'} \left(\right. ℓ \left.\right)$不是连续的,我们没法保证找到$g^{'} \left(\right. ℓ \left.\right) = 0$的点。 不过,如果我们将$\# \left{\right. ℓ > x_{i} \left.\right}$改为$\# \left{\right. ℓ \geq x_{i} \left.\right}$(不连续的边界点的导数可以视为任意的,所以这种调整是允许的),那么$g^{'} \left(\right. ℓ \left.\right) = 0$的含义正好是“小于等于$ℓ$的边界点刚好有$n - k$个”,如果进一步假设所有边界点两两不同,那么$l$正好是全体边界点从小到大排序后的第$n - k$个!所以,L1场景下,其实一次排序就可以得到$ℓ$,相当漂亮。
参考实现[#](https://spaces.ac.cn/archives/11784#%E5%8F%82%E8%80%83%E5%AE%9E%E7%8E%B0)
两个版本的MEP参考实现如下:
import jax
import jax.numpy as jnp
@jax.jit
def l2mep(inputs, mask, margin):
x, m = inputs, margin
u = jnp.where(mask, x, x + m).sort()[:, None]
v = jnp.where(mask, jnp.fmax(u - x, 0), -jnp.fmax(x + m - u, 0)).sum(axis=1)
i = ((v[:-1] < 0) & (v[1:] >= 0)).argmax()
l = (u[i + 1] * v[i] - u[i] * v[i + 1]) / (v[i] - v[i + 1])
return jnp.where(mask, jnp.fmax(x, l), jnp.fmin(x, l - m))
@jax.jit
def l1mep(inputs, mask, margin):
x, m = inputs, margin
u = jnp.where(mask, x, x + m).sort(axis=-1)
l = jnp.take_along_axis(u, (~mask).sum(axis=-1, keepdims=True) - 1)
return jnp.where(mask, jnp.fmax(x, l), jnp.fmin(x, l - m))这里稍加说明一下,实际场景下,正类通常不会简单地排列在前$k$个,数目、位置都可能不确定,所以在上述实现中,我们用一个mask向量来标记正负类。
l2mep当前实现只支持1d的输入,如果需要支持batch维度,用jax.vmap包装一层即可。这个算法通过遍历所有区间来找零点区间,每步复杂度是$\mathcal{O} \left(\right. n \left.\right)$,总复杂度是$\mathcal{O} \left(\right. n^{2} \left.\right)$。如果想要提高效率,可以改用二分法来找零点区间,这样能降低到$\mathcal{O} \left(\right. n log n \left.\right)$
l1mep由于算法简单,所以现在的写法就已经支持任意的batch维度,它的核心运算只有排序,复杂度是$\mathcal{O} \left(\right. n log n \left.\right)$,不管从简洁还是速度看都已经相当理想,如果没有别的考虑,实践中推荐用L1版。
文章小结[#](https://spaces.ac.cn/archives/11784#%E6%96%87%E7%AB%A0%E5%B0%8F%E7%BB%93)
本文介绍了“强制间隔投影(Margin-Enforcing Projection,MEP)”这一运算:给定一个分数向量,把它投影到满足“正类最小值至少比负类最大值大$m$”的最近向量上。这可以为传统的Margin Learning提供一些新的思路。
_转载到请包括本文地址:[https://spaces.ac.cn/archives/11784](https://spaces.ac.cn/archives/11784 "强制间隔投影(Margin-Enforcing Projection)")_
_更详细的转载事宜请参考:_[《科学空间FAQ》](https://spaces.ac.cn/archives/6508#%E6%96%87%E7%AB%A0%E5%A6%82%E4%BD%95%E8%BD%AC%E8%BD%BD/%E5%BC%95%E7%94%A8 "《科学空间FAQ》")