1 - 利用线性模型进行二分类

1.1 - 线性模型

这一部分主要是为了证明为什么可以使用线性回归和逻辑斯蒂回归来做二元分类问题。
我们想要将已知的线性模型应用到二分类甚至是多分类的问题。我们已知的线性模型有如下的三种,它们有一个共同的地方是都会利用输入的特征计算一个加权和 s=wTxs=wTx<script type="math/tex" id="MathJax-Element-748">s = w^Tx</script>。


这里写图片描述

线性分类linear classificationlinear classification<script type="math/tex" id="MathJax-Element-749">\text{linear classification}</script>不好解, 因为想要最小化E0/1(w)E0/1(w)<script type="math/tex" id="MathJax-Element-750">E_{0/1}(w)</script>(被分错的点的个数)被证明是一个NPNP<script type="math/tex" id="MathJax-Element-751">\text{NP}</script>难问题。但相比之下,线性回归和逻辑斯蒂回归很方便就可以最小化EinEin<script type="math/tex" id="MathJax-Element-752">E_{in}</script>:线性回归使用平方误差square errorsquare error<script type="math/tex" id="MathJax-Element-753">\text{square error}</script>时有解析解closed-form solution(closed-form solution)<script type="math/tex" id="MathJax-Element-754">\text{(closed-form solution)}</script>;逻辑斯蒂回归由于目标函数是凸函数所以可以使用梯度下降法来求解。所以我们想要做的是:因为linear classificationlinear classification<script type="math/tex" id="MathJax-Element-755">\text{linear classification}</script>不好做,而看起来linear regressionlinear regression<script type="math/tex" id="MathJax-Element-756">\text{linear regression}</script>或者是logistic regressionlogistic regression<script type="math/tex" id="MathJax-Element-757">\text{logistic regression}</script>在最小化EinEin<script type="math/tex" id="MathJax-Element-758">E_{in}</script>这件事情上还是比较简单的。所以我们想要使用线性回归或者是逻辑斯蒂回归作为线性分类的一个替代工具。也就是现在将所有的方法都用在分类问题上,这样的话输出都会被限制在集合{+1,1}{+1,−1}<script type="math/tex" id="MathJax-Element-759">\{+1, -1\}</script>里面。

1.2 - 损失函数(error function)

首先对三种error functionerror function<script type="math/tex" id="MathJax-Element-9539">\text{error function}</script>进行整理, 得到从ss<script type="math/tex" id="MathJax-Element-9540">s</script>到error<script type="math/tex" id="MathJax-Element-9541">error</script>的计算公式, 其中s=wTxs=wTx<script type="math/tex" id="MathJax-Element-9542">s = w^Tx</script>。

  • 线性分类
    h(x)err(h,x,y)err0/1(s,y)=sign(s)=|[h(x)y]|=|[sign(s)y]|=|[sign(ys)1]|(2577)(2578)(2579)(2580)(2581)(2577)h(x)=sign(s)(2578)err(h,x,y)=|[h(x)≠y]|(2579)↓(2580)err0/1(s,y)=|[sign(s)≠y]|(2581)=|[sign(ys)≠1]|
    <script type="math/tex; mode=display" id="MathJax-Element-9543">\begin{align} h(x) &= sign(s) \\ err(h,x,y) &= |[h(x) \ne y]| \\ & \downarrow \\ err_{0/1}(s, y) & = |[sign(s) \ne y]| \\ & = |[sign(ys) \ne 1]| \end{align}</script>
  • 线性回归
    h(x)err(h,x,y)errSQR(s,y)=s=(h(x)y)2=(sy)2=(sy)2y2=(syy2)2=(sy1)2(2582)(2583)(2584)(2585)(2586)(2587)(2588)(2582)h(x)=s(2583)err(h,x,y)=(h(x)−y)2(2584)↓(2585)errSQR(s,y)=(s−y)2(2586)=(s−y)2y2(2587)=(sy−y2)2(2588)=(sy−1)2
    <script type="math/tex; mode=display" id="MathJax-Element-9544">\begin{align} h(x) &= s \\ err(h,x,y) &= (h(x) - y)^2 \\ & \downarrow \\ err_{SQR}(s, y) & = (s - y)^2 \\ & =(s - y)^2y^2 \\ & = (sy - y^2)^2 \\ & = (sy - 1)^2 \end{align}</script>
  • 逻辑斯蒂回归
    h(x)err(h,x,y)errCE(s,y)=11+exp(wTx)=θ(wTx)=θ(s)=lnh(yx)=ln(1+exp(ys))(2589)(2590)(2591)(2592)(2589)h(x)=11+exp(−wTx)=θ(wTx)=θ(s)(2590)err(h,x,y)=−lnh(yx)(2591)↓(2592)errCE(s,y)=ln(1+exp(−ys))
    <script type="math/tex; mode=display" id="MathJax-Element-9545">\begin{align} h(x) &= \frac1{1+exp(-w^Tx)} = \theta(w^Tx)= \theta(s) \\ err(h, x, y)& = -lnh(yx) \\ & \downarrow \\ err_{CE}(s, y) &= ln(1 + exp(-ys)) \end{align}</script>

通过上面的这些操作,每一个模型的errorerror<script type="math/tex" id="MathJax-Element-9546">\text{error}</script>的计算方法中都出现了ysys<script type="math/tex" id="MathJax-Element-9547">\text{ys}</script>,我们接下来要做的就是看看这些errorerror<script type="math/tex" id="MathJax-Element-9548">\text{error}</script>和ysys<script type="math/tex" id="MathJax-Element-9549">\text{ys}</script>的关系。

为什么要关注ysys<script type="math/tex" id="MathJax-Element-9550">\text{ys}</script>这个量呢?我们简单的看看ysys<script type="math/tex" id="MathJax-Element-9551">\text{ys}</script>的物理意义:对于分类来说,我们希望ysys<script type="math/tex" id="MathJax-Element-9552">\text{ys}</script>越大越好, 首先如果该值是正的, 说明起码分类正确了, 如果这个值还很大,那么说明置信度很高。
可视化一下我们得到了三个error funcitonerror funciton<script type="math/tex" id="MathJax-Element-9553">\text{error funciton}</script>:横轴是ysys<script type="math/tex" id="MathJax-Element-9554">\text{ys}</script>的值,纵轴是errerr<script type="math/tex" id="MathJax-Element-9555">\text{err}</script>的值。

  1. 0/1: err0/1(s,y)=|[sign(sy1)]|0/1: err0/1(s,y)=|[sign(sy≠1)]|<script type="math/tex" id="MathJax-Element-9556">0/1:\ err_{0/1}(s, y) = |[sign(sy \ne 1)]|</script>
  2. sqr: errsqr(s,y)=(ys1)2sqr: errsqr(s,y)=(ys−1)2<script type="math/tex" id="MathJax-Element-9557">sqr:\ err_{sqr}(s, y) = (ys - 1)^2</script>
  3. ce: errce(s,y)=ln(1+exp(ys))ce: errce(s,y)=ln(1+exp(−ys))<script type="math/tex" id="MathJax-Element-9558">ce:\ err_{ce}(s, y) = ln(1+exp(-ys))</script>


三个error function
这里写图片描述

  • 通过比较在x = 1附近的0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9559">\text{0/1 error}</script>和square errorsquare error<script type="math/tex" id="MathJax-Element-9560">\text{square error}</script>可以得到:如果有一个样本在square errorsquare error<script type="math/tex" id="MathJax-Element-9561">\text{square error}</script>上的值很低的话,那么这个样本在0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9562">\text{0/1 error}</script>上也会得到比较低的值。但是如果样本在square errorsquare error<script type="math/tex" id="MathJax-Element-9563">\text{square error}</script>上的值比较大, 很左边或者是很右边, 那么我们是没有办法判别这个样本的0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9564">\text{0/1 error}</script>是小还是大的。所以当能把样本的square errorsquare error<script type="math/tex" id="MathJax-Element-9565">\text{square error}</script>做到很低的时候,可以在一定的程度上保证0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9566">\text{0/1 error}</script>也很低。

  • 通过比较0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9567">\text{0/1 error}</script>和corss entropy errorcorss entropy error<script type="math/tex" id="MathJax-Element-9568">\text{corss entropy error}</script>的函数我们可以得到:corss entropy errorcorss entropy error<script type="math/tex" id="MathJax-Element-9569">\text{corss entropy error}</script>小的时候,0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9570">\text{0/1 error}</script>也是很小的;0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9571">\text{0/1 error}</script>比较小的时候,corss entropy errorcorss entropy error<script type="math/tex" id="MathJax-Element-9572">\text{corss entropy error}</script>也是很小的。

corss entropy errorcorss entropy error<script type="math/tex" id="MathJax-Element-9573">\text{corss entropy error}</script>换底之后得到下面的scaled ce error:errsce(s,y)=log2(1+exp(sy))scaled ce error:errsce(s,y)=log2(1+exp(−sy))<script type="math/tex" id="MathJax-Element-9574">\text{scaled ce error}: err_{sce}(s, y) = log_2(1+exp(-sy))</script>,这样一来,corss entropy errorcorss entropy error<script type="math/tex" id="MathJax-Element-9575">\text{corss entropy error}</script>就一定会是0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9576">0/1\ error</script>的上界。我们得到的新的图像如下:


这里写图片描述

上述是通过函数图像的直观理解,接下来开始利用以上提到的error functionerror function<script type="math/tex" id="MathJax-Element-9577">\text{error function}</script>来说明为什么linear regressionlinear regression<script type="math/tex" id="MathJax-Element-9578">linear\ regression</script> 或者是logistics regressionlogistics regression<script type="math/tex" id="MathJax-Element-9579">logistics\ regression</script>可以是一个好的二分类的方法。

如果想做要计算分类的错误的话,它的上限如下:

err0/1(y,s)errsce(s,y)=1ln2errce(s,y)err0/1(y,s)≤errsce(s,y)=1ln2errce(s,y)
<script type="math/tex; mode=display" id="MathJax-Element-9580">err_{0/1}(y, s) \le err_{sce}(s, y) = \frac1{ln2}err_{ce}(s, y)</script>
那么我们就可以得到:(err是E的平均)
E0/1in(w)ESCEin(w)=1ln2ECEin(w)Ein0/1(w)≤EinSCE(w)=1ln2EinCE(w)
<script type="math/tex; mode=display" id="MathJax-Element-9581">E_{in}^{0/1}(w) \le E_{in}^{SCE}(w) = \frac1{ln2}E_{in}^{CE}(w)</script>
同样也可以得到:
E0/1out(w)ESCEout(w)=1ln2ECEout(w)Eout0/1(w)≤EoutSCE(w)=1ln2EoutCE(w)
<script type="math/tex; mode=display" id="MathJax-Element-9582">E_{out}^{0/1}(w) \le E_{out}^{SCE}(w) = \frac1{ln2}E_{out}^{CE}(w)</script>

根据VC维的理论我们可以得到:

E0/1outE0/1in+ΩESCEin(w)+ΩEout0/1≤Ein0/1+Ω≤EinSCE(w)+Ω
<script type="math/tex; mode=display" id="MathJax-Element-9583">E_{out}^{0/1} \le E_{in}^{0/1} + \Omega \le E_{in}^{SCE}(w) + \Omega</script>

所以如果我们能够把逻辑斯蒂回归中的cross entropy errorcross entropy error<script type="math/tex" id="MathJax-Element-9584">cross\ entropy\ error</script>做到最小的话,从某种角度上来说也就是把0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9585">0/1\ error</script>做到了最小。或者说把square errorsquare error<script type="math/tex" id="MathJax-Element-9586">square\ error</script>最小化也是在间接的最小化0/1 error0/1 error<script type="math/tex" id="MathJax-Element-9587">0/1\ error</script>, 只不过从图中可以看得出来,square errorsquare error<script type="math/tex" id="MathJax-Element-9588">square\ error</script>是0/1error0/1error<script type="math/tex" id="MathJax-Element-9589">0/1 error</script>的一个更为宽松的上界,这就表示也可以把linear regressionlinear regression<script type="math/tex" id="MathJax-Element-9590">linear\ regression</script>用在linear classificationlinear classification<script type="math/tex" id="MathJax-Element-9591">linear\ classification</script>上。我们称之为regression for classificationregression for classification。<script type="math/tex" id="MathJax-Element-9592">regression\ for\ classification。</script>

Regression for ClassificationRegression for Classification<script type="math/tex" id="MathJax-Element-9593">\text{Regression for Classification}</script>

  1. run logisitic/linear regression on D with yn{+1,1} to get wregrun logisitic/linear regression on D with yn∈{+1,−1} to get wreg<script type="math/tex" id="MathJax-Element-9594">\text{run logisitic/linear regression on D with }y_n \in \{+1, -1\} \text{ to get } w_{reg}</script>
  2. return g(x)=sign(wTregx)return g(x)=sign(wregTx)<script type="math/tex" id="MathJax-Element-9595">\text{return } g(x) = sign(w^T_{reg}x)</script>

1.3 - 小结

  • 如果我们有一个二分类的问题,linear regressionlinear regression<script type="math/tex" id="MathJax-Element-9882">linear\ regression</script>计算非常的方便,但是太为宽松(从图中可以得到),所以我们通常是使用linear regressionlinear regression<script type="math/tex" id="MathJax-Element-9883">linear\ regression</script>的结果作为PLA/pocket/logistic regressionPLA/pocket/logistic regression<script type="math/tex" id="MathJax-Element-9884">PLA/pocket/logistic\ regression</script>的初始的w0w0<script type="math/tex" id="MathJax-Element-9885">w_0</script>值。linear regression sometimes used to set w0 for PLA/pocket/logistic regressionlinear regression sometimes used to set w0 for PLA/pocket/logistic regression<script type="math/tex" id="MathJax-Element-9886">\text{linear regression sometimes used to set } w_0 \text{ for PLA/pocket/logistic regression}</script>
  • 对于二分类这个其实比较困难的问题,在每一轮中logistic regressionlogistic regression<script type="math/tex" id="MathJax-Element-9887">logistic\ regression</script>和pocketpocket<script type="math/tex" id="MathJax-Element-9888">pocket</script>的时间复杂度其实是差不多的,所以我们通常选用logistic regressionlogistic regression<script type="math/tex" id="MathJax-Element-9889">logistic\ regression</script>而不使用pocketpocket<script type="math/tex" id="MathJax-Element-9890">pocket</script>来做binary classificationbinary classification<script type="math/tex" id="MathJax-Element-9891">binary\ classification</script>, 由于它比较好的optimizationoptimization<script type="math/tex" id="MathJax-Element-9892">optimization</script>性质

2 - 随机梯度下降算法(SGD)

PLAPLA<script type="math/tex" id="MathJax-Element-10061">PLA</script>算法和gradient descentgradient descent<script type="math/tex" id="MathJax-Element-10062">gradient\ descent</script>算法都可以看成是iterative optimizationiterative optimization<script type="math/tex" id="MathJax-Element-10063">iterative\ optimization</script>(一步一步的接近最佳的wOPTwOPT<script type="math/tex" id="MathJax-Element-10064">w_{OPT}</script>)。不同的是:PLAPLA<script type="math/tex" id="MathJax-Element-10065">PLA</script>算法每一次只看一个数据点来更新来ww<script type="math/tex" id="MathJax-Element-10066">w</script>(如果这个点被划分错误了);而在logistic regression<script type="math/tex" id="MathJax-Element-10067">logistic\ regression</script>中的gradient descentgradient descent<script type="math/tex" id="MathJax-Element-10068">gradient\ descent</script>算法每一次对ww<script type="math/tex" id="MathJax-Element-10069">w</script>的更新都要扫描所有的样本。这一小节我们想要做的是优化梯度下降算法。让logistic regression<script type="math/tex" id="MathJax-Element-10070">logistic\ regression</script>和PLAPLA<script type="math/tex" id="MathJax-Element-10071">PLA</script>一样快。目前来说,logistic regressionlogistic regression<script type="math/tex" id="MathJax-Element-10072">logistic\ regression</script>是和pocketpocket<script type="math/tex" id="MathJax-Element-10073">pocket</script>算法差不多快的。

2.1 - 逻辑斯蒂回归中的梯度下降算法

逻辑斯蒂回归使用梯度下降算法更新权重的规则:

wt+1wt+η1Nn=1Nθ(ynwTtxn)(ynxn)Ein(wt)wt+1⟵wt+η1N∑n=1Nθ(−ynwtTxn)(ynxn)⏟−▽Ein(wt)
<script type="math/tex; mode=display" id="MathJax-Element-10099">w_{t+1} \longleftarrow w_{t} + \eta\underbrace{\frac1N\sum_{n=1}^{N}\theta(-y_nw^T_tx_n)(y_nx_n)}_{-\triangledown E_{in}(w_t)}</script>

我们看到在更新的时候有1NNn=11N∑n=1N<script type="math/tex" id="MathJax-Element-10100">\frac1N\sum_{n=1}^{N}</script>这样的操作,那么我们可以随机的抽取一个样本替代这个加和取平均的操作(比如这里有1w个数字,我们把这1w个数字求和取平均的值应该差不多等于我们只抽取几个数字来求和取平均的值,应该也差不多等于我们随机抽取一个数字的值)。只在这一个样本点上求取梯度值,即只针对一个点的errorerror<script type="math/tex" id="MathJax-Element-10101">\text{error}</script>求取偏微分,这个值称为随机梯度stochastic gradientstochastic gradient<script type="math/tex" id="MathJax-Element-10102">\text{stochastic gradient}</script>。这样的话整体的梯度可以看成是这个随机梯度的期望值。

2.2 - 随机梯度下降算法(stochastic gradient descent)

使用随机的梯度取代真正的梯度:


wt+1wt+ηθ(ynwTtxn)(ynxn)errin(wt, xn, yn)(1)(1)wt+1⟵wt+ηθ(−ynwtTxn)(ynxn)⏟−▽errin(wt, xn, yn)
<script type="math/tex; mode=display" id="MathJax-Element-15460">w_{t+1} \longleftarrow w_{t} + \eta\underbrace{\theta(-y_nw^T_tx_n)(y_nx_n)}_{-\triangledown err_{in}(w_t,\ x_n,\ y_n)} \tag1</script>

算法的优点:

  • 如果迭代足够多次的话, 真实的梯度和随机的梯度应该是会差不多接近的。
  • 简单,不再需要对所有的点计算梯度,而只计算一个点的梯度。在大数据的背景下这样的方式会很有用。
  • 当资料本身就是一笔一笔的来的时候(online learningonline learning<script type="math/tex" id="MathJax-Element-15461">online\ learning</script>), stochastic gradient descentstochastic gradient descent<script type="math/tex" id="MathJax-Element-15462">stochastic\ gradient\ descent</script>会很适合这样的场景。

算法的缺点:

  • 算法在性质上可能会不稳定。特别是当步长很大的时候。

公式(1)(1)<script type="math/tex" id="MathJax-Element-15463">(1)</script>和我们之前看到过的PLAPLA<script type="math/tex" id="MathJax-Element-15464">PLA</script>算法的更新规则非常的相似。

  • PLAPLA<script type="math/tex" id="MathJax-Element-15465">PLA</script>:

    wt+1wt+1|[ynsign(wTtxn)]| (ynxn)wt+1⟵wt+1⋅|[yn≠sign(wtTxn)]| (ynxn)
    <script type="math/tex; mode=display" id="MathJax-Element-15466">w_{t+1} \longleftarrow w_{t} + 1 \centerdot |[y_n \ne sign(w^T_tx_n)]|\ (y_nx_n)</script>

    其中如果<script type="math/tex" id="MathJax-Element-15467">\bigcirc</script>成立,|[]|=1|[◯]|=1<script type="math/tex" id="MathJax-Element-15468">|[\bigcirc]| = 1</script>, 否则|[]|=0|[◯]|=0<script type="math/tex" id="MathJax-Element-15469">|[\bigcirc]| = 0</script>

  • SGD logistic regressionSGD logistic regression<script type="math/tex" id="MathJax-Element-15470">SGD\ logistic\ regression</script>:

    wt+1wt+ηθ(ynwTtxn)(ynxn)wt+1⟵wt+η⋅θ(−ynwtTxn)(ynxn)
    <script type="math/tex; mode=display" id="MathJax-Element-15471">w_{t+1} \longleftarrow w_{t} +\eta\centerdot\theta(-y_nw^T_tx_n)(y_nx_n)</script>

SGD logistic regression 'soft' PLASGD logistic regression ≈'soft' PLA<script type="math/tex" id="MathJax-Element-15472">\text{SGD logistic regression }\approx \text{'soft' PLA}</script>
即使用SGDSGD<script type="math/tex" id="MathJax-Element-15473">SGD</script>算法的logistic regressionlogistic regression<script type="math/tex" id="MathJax-Element-15474">logistic\ regression</script>可以看成是‘软’的PLAPLA<script type="math/tex" id="MathJax-Element-15475">PLA</script>,因为不同于PLAPLA<script type="math/tex" id="MathJax-Element-15476">PLA</script>,不再只是看错了没有, 而是看看错了多少, 错的多就更新的多一点, 错的少就更新的少一点。

  • ynwTtxn0ynwtTxn≪0<script type="math/tex" id="MathJax-Element-15477">y_nw^T_tx_n \ll 0</script>,说明样本(xn,yn)(xn,yn)<script type="math/tex" id="MathJax-Element-15478">(x_n, y_n)</script>被划分错了, 此时θ(ynwTtxn)1θ(−ynwtTxn)≈1<script type="math/tex" id="MathJax-Element-15479">\theta(-y_nw^T_tx_n)\approx 1</script>,也就是说,错的多,更新的多。
  • ynwTtxn0ynwtTxn≫0<script type="math/tex" id="MathJax-Element-15480">y_nw^T_tx_n \gg 0</script>,说明样本(xn,yn)(xn,yn)<script type="math/tex" id="MathJax-Element-15481">(x_n, y_n)</script>被划分对了, 此时θ(ynwTtxn)0θ(−ynwtTxn)≈0<script type="math/tex" id="MathJax-Element-15482">\theta(-y_nw^T_tx_n)\approx 0</script>,也就是说,错的少,基本不用更新。

通过比对上面的两个式子我们可以得到

  1. 逻辑斯蒂回归使用随机梯度下降算法大概相当于是一个softsoft<script type="math/tex" id="MathJax-Element-15483">soft</script>的PLAPLA<script type="math/tex" id="MathJax-Element-15484">PLA</script>算法。在PLAPLA<script type="math/tex" id="MathJax-Element-15485">PLA</script>算法中|[ynsign(wTtxn)]||[yn≠sign(wtTxn)]|<script type="math/tex" id="MathJax-Element-15486">|[y_n \ne sign(w^T_tx_n)]|</script>的取值是0或者1,而在SGD logistic regressionSGD logistic regression<script type="math/tex" id="MathJax-Element-15487">SGD\ logistic\ regression</script>中θ(ynwTtxn)θ(−ynwtTxn)<script type="math/tex" id="MathJax-Element-15488">\theta(-y_nw^T_tx_n)</script>的取值在0和1之间。
  2. η=1η=1<script type="math/tex" id="MathJax-Element-15489">\eta = 1</script>并且wTtxnwtTxn<script type="math/tex" id="MathJax-Element-15490">w^T_tx_n</script>非常大的时候(这是θ=0/1θ=0/1<script type="math/tex" id="MathJax-Element-15491">\theta=0/1</script>),PLAPLA<script type="math/tex" id="MathJax-Element-15492">PLA</script>算法和SGD logistic regressionSGD logistic regression<script type="math/tex" id="MathJax-Element-15493">SGD\ logistic\ regression</script>算法几乎是一样的。

现在有两个问题:

  1. 如何决定停止条件:因为不再扫描所有的点,所以不知道是不是已经梯度为0了,所以一般的停止条件是根据预先设定的迭代次数
  2. 如何设置步长ηη<script type="math/tex" id="MathJax-Element-15494">\eta</script>:如果X的范围不是很糟糕的话,0.1还不错,不过这只是个经验值,会有专门的方法来帮助我们选择参数,我们之后也会介绍到。
  3. 如果把SGDSGD<script type="math/tex" id="MathJax-Element-15495">SGD</script>使用在linear regressionlinear regression<script type="math/tex" id="MathJax-Element-15496">linear\ regression</script>上, 更新的方向如下:
    2(ynwTtxn)xn2(yn−wtTxn)xn
    <script type="math/tex; mode=display" id="MathJax-Element-15497">2(y_n - w^T_tx_n)x_n</script>
    SGD linear regressionSGD linear regression<script type="math/tex" id="MathJax-Element-15498">\text{SGD linear regression}</script>
    wt+1=wt+η(ynwTtxn)xnwt+1=wt+η(yn−wtTxn)xn
    <script type="math/tex; mode=display" id="MathJax-Element-15499">w_{t+1} = w_{t} + \eta(y_n - w^T_tx_n)x_n</script>

    也是朝着xx<script type="math/tex" id="MathJax-Element-15500">x</script>的方更新,真实值和预测值的差异越大,更新的越大。

3 - 利用逻辑斯蒂回归解决多分类问题

之前的内容只能用来解决是非题,即分类问题,现在我们想要做的是多选题。即判断一个样本属于多个类别中的哪一个,接下来我们想要做的是将二元分类的方法延伸到可以帮助我们求解多分类的问题。

具体的思想就是,将多元分类转化为多个二元分类问题, 每一次只识别一个类型。


这里写图片描述

我们可以把这个问题转换为一个二分类的问题,即首先我们关心数据是 <script type="math/tex" id="MathJax-Element-15519">\square</script> 或不是 <script type="math/tex" id="MathJax-Element-15520">\square</script> ,是的话其label记为1, 不是的话其label记为0。 (对应到我们熟悉的二分类问题就是将 <script type="math/tex" id="MathJax-Element-15521">\square</script> 记为 <script type="math/tex" id="MathJax-Element-15522">\bigcirc</script>, 将其他的记为 ××<script type="math/tex" id="MathJax-Element-15523">\times</script>), 得到如下图所示的一个新的二分类问题,这个时候就可以使用binary classificationbinary classification<script type="math/tex" id="MathJax-Element-15524">binary\ classification</script>算法来解决。


这里写图片描述

使用相同的方法解决以下的三个二分类问题,即针对每一个类别做一个该类别和其他的类别的二元分类问题,把该类和其他的类别分开。


这里写图片描述

通过这样的步骤我们就得到了如下的四个分类器:


这里写图片描述

将4次二元分类的结果综合起来得到如下图所示的多分类器:这个多远的分类器告诉我们一些结果:

  • 在黑色圈起来的部分中的数据点可以很明确的知道自己属于哪一个类别
  • 在红色区域中的数据点会同时属于两个类别
  • 在蓝色区域中的数据点没有类别。


这里写图片描述

缺点:所以这样的分类器对于蓝色区域和红色区域中的样本是不确定的。

3.1 - 利用逻辑斯蒂回归进行改进

同样是上面的任务, 同样是每一次划分一个类别,但是不再是给出“是”或者“不是”的结果,而是使用logistic regressionlogistic regression<script type="math/tex" id="MathJax-Element-15569">\text{logistic regression}</script>给出一个模糊的边界。即给出每个数据点属于某个类别的概率(下图中颜色越蓝表示概率越接近1(正例),颜色越红表明概率越接近0(负例),颜色越白表明概率越接近0.5(难以辨别))。同样是进行4次二元分类分别得到如下的结果:


这里写图片描述

综合上述的结果,得到最后的一个多元的分类器如下:因为在每一个数据点(这个例子中是二维平面上的每一个点),上述的四个二元分类器都会给出该点属于某一个类别的一个概率,选取4个概率中值最高的分类器对该数据点的划分结果作为该点的类别就可以了。


这里写图片描述

分类器的表达式:

g(x)=argmaxky θ(wT[k]x)g(x)=argmaxk∈y θ(w[k]Tx)
<script type="math/tex; mode=display" id="MathJax-Element-15570">g(x) = argmax_{k \in y}\ \theta(w^T_{[k]}x)</script>
在上述的例子中k = 4, 或者不需要使用θθ<script type="math/tex" id="MathJax-Element-15571">\theta</script>, 因为该函数是单调的。
g(x)=argmaxky wT[k]xg(x)=argmaxk∈y w[k]Tx
<script type="math/tex; mode=display" id="MathJax-Element-15572">g(x) = argmax_{k \in y}\ w^T_{[k]}x</script>

3.2 - 多元分类算法OVA Decomposition

One-Versus-All(OVA) Decomposition

  1. 对于每一个类别kk<script type="math/tex" id="MathJax-Element-15857">k</script>,通过在数据集Dk<script type="math/tex" id="MathJax-Element-15858">D_{k}</script>上使用逻辑斯蒂回归得到wkwk<script type="math/tex" id="MathJax-Element-15859">w_{k}</script>,
    其中Dk={(xn,yn=2|[yn=k]|1)}Nn=1Dk={(xn,yn′=2|[yn=k]|−1)}n=1N<script type="math/tex" id="MathJax-Element-15860">D_{k} = \{(x_n, y_n' = 2|[y_n = k]| - 1)\}_{n=1}^N</script>(是这个类的其yy<script type="math/tex" id="MathJax-Element-15861">y</script>为1<script type="math/tex" id="MathJax-Element-15862">1</script>,不是这个类的其yy<script type="math/tex" id="MathJax-Element-15863">y</script>为1<script type="math/tex" id="MathJax-Element-15864">-1</script>)
  2. 得到多元分类器 :g(x)=argmaxky  wTkxg(x)=argmaxk∈y  wkTx<script type="math/tex" id="MathJax-Element-15865">g(x) = argmax_{k\in y} \ \ w^T_kx</script>

DecompositionDecomposition<script type="math/tex" id="MathJax-Element-15866">\text{Decomposition}</script>:是指将原来的多分类问题拆成多个二分类问题
OVAOVA<script type="math/tex" id="MathJax-Element-15867">\text{OVA}</script>:是指每次对一个类别和其他的所有的类别进行分类

优点

  • 有效率,可以使用任何和逻辑斯蒂回归相似的算法来进行计算,只要是分类的结果输出的也是可以比大小的实数值就行。
  • 当面对的是一个K分类的问题的时候,我们其实要做的是KK<script type="math/tex" id="MathJax-Element-15868">K</script>个logistic regression<script type="math/tex" id="MathJax-Element-15869">logistic\ regression</script>,每一个logistic regressionlogistic regression<script type="math/tex" id="MathJax-Element-15870">logistic\ regression</script>所用的资料和原来的资料其实并没有差太远,只是将labellabel<script type="math/tex" id="MathJax-Element-15871">label</script>换掉了而已。值得一提的是, 这KK<script type="math/tex" id="MathJax-Element-15872">K</script>个不同logistic regression<script type="math/tex" id="MathJax-Element-15873">logistic\ regression</script>问题我们是可以分布在KK<script type="math/tex" id="MathJax-Element-15874">K</script>个不同的机器上运行的。所以这是一个很容易并行处理的算法。

缺点

  • K<script type="math/tex" id="MathJax-Element-15875">K</script>很大,即类别很多的时候,会有数据不平衡的问题,

4 - 利用二元分类解决多分类问题

上一小节讲到使用OVA(One-Versus-All)OVA(One-Versus-All)<script type="math/tex" id="MathJax-Element-16197">\text{OVA(One-Versus-All)}</script>进行多元分类的时候,当KK<script type="math/tex" id="MathJax-Element-16198">K</script>的数量很大的时候,很有可能会出现数据不平衡的问题。 也就是只有很少量的label<script type="math/tex" id="MathJax-Element-16199">label</script>为11<script type="math/tex" id="MathJax-Element-16200">1</script>的数据, 其余的都是label<script type="math/tex" id="MathJax-Element-16201">label</script>为00<script type="math/tex" id="MathJax-Element-16202">0</script>的数据。这样就会造成分类的结果不好。 为了避免这种情况的发生,我们采用的方法是:每一次只选择两个类别进行分类(例如对于猫, 狗,汽车这样的三分类问题,分别对(猫, 狗),(猫, 汽车),(狗, 汽车)做一次二元分类;而OVA<script type="math/tex" id="MathJax-Element-16203">OVA</script>在每一次分类中所有的类都参与, 同样是上面的问题,OVAOVA<script type="math/tex" id="MathJax-Element-16204">\text{OVA}</script>要做的是(猫, (狗,汽车)), (狗, (猫, 汽车)), (汽车, (猫, 狗))这样的三个二分类问题)。在如下的四个类别的分类问题中,按照上述的思路,一共需要进行6次二元分类问题的求解。(每次选两个类进行划分C24=4!2!(42)!=6C42=4!2!(4−2)!=6<script type="math/tex" id="MathJax-Element-16205">C_4^2 = \frac{4!}{2!(4-2)!}=6</script>)


这里写图片描述


6次二元分类问题的求解结果
这里写图片描述

将这些二元分类的结果进行合并得到如下最终的4元分类器,


这里写图片描述

得到了分类器之后怎么对新样本进行分类?
对于一个新的样本,分别使用这个6个分类器进行计算,计算得到的结果中,类别最多的类就该样本的类。具体的细节是,对于任意一个给定的数据点,分别使用上面得到的6个二元的分类器来进行类别的划分,选择6个划分结果中出现次数最多的作为该类的类别。例如6个分类器的分类结果分别是:,,,,,◻,◻,◻,◊,★,★<script type="math/tex" id="MathJax-Element-16206">\square,\square,\square,\lozenge,\bigstar,\bigstar</script>, 那么我们判断该点的最终的类别是 <script type="math/tex" id="MathJax-Element-16207">\square</script>。

4.1 - 多元分类算法OVO

One-versus-one(OVO) DecompositionOne-versus-one(OVO) Decomposition<script type="math/tex" id="MathJax-Element-16797">\text{One-versus-one(OVO) Decomposition}</script>

  1. 对于数据集中的任意两个类别k,lk,l<script type="math/tex" id="MathJax-Element-16798">k, l</script>, 通过在数据集Dk,lDk,l<script type="math/tex" id="MathJax-Element-16799">D_{k, l}</script>上使用任一种二元分类算法来得到wk,lwk,l<script type="math/tex" id="MathJax-Element-16800">w_{k,l}</script>。
    其中数据集如下:
    Dk,l={(xn,yn=2|[yn=k]|1):yn=k or yn=l}Dk,l={(xn,yn′=2|[yn=k]|−1):yn=k or yn=l}<script type="math/tex" id="MathJax-Element-16801">D_{k,l} = \{(x_n, y_n' = 2|[y_n = k]| - 1): y_n = k\ or\ y_n = l\}</script>。(是kk<script type="math/tex" id="MathJax-Element-16802">k</script>类的样本其label<script type="math/tex" id="MathJax-Element-16803">label</script>设置为11<script type="math/tex" id="MathJax-Element-16804">1</script>,是l<script type="math/tex" id="MathJax-Element-16805">l</script>类的样本其labellabel<script type="math/tex" id="MathJax-Element-16806">label</script>设置为00<script type="math/tex" id="MathJax-Element-16807">0</script>)
  2. 在得到了所有的分类器之后,通过一个投票函数来决定一个新的数据x<script type="math/tex" id="MathJax-Element-16808">x</script>的属类:g(x)=vote(wTk,lx)g(x)=vote(wk,lTx)<script type="math/tex" id="MathJax-Element-16809">g(x) = vote(w^T_{k,l}x)</script>

优点:

  • 虽然如果共有4个类别,却要做C42=6C24=6<script type="math/tex" id="MathJax-Element-16810">C_2^4 = 6</script>次二元分类,但是每一次二元分类涉及到的数据量很少(只包括两个类的数据)。
    • 可以使用任意的binary classificationbinary classification<script type="math/tex" id="MathJax-Element-16811">binary\ classification</script>的方法。
    • 可以并行计算

缺点:

  • 预测的时间比较长:因为要使用得到的6个分类器来进行投票决定。
  • 需要更多的存储:同样是因为有更多的ww<script type="math/tex" id="MathJax-Element-16812">w</script>,所以需要占用更多的空间。

5 - 小结

首先声明了三个linear model<script type="math/tex" id="MathJax-Element-16822">linear\ model</script>都可以用来做binary classificationbinary classification<script type="math/tex" id="MathJax-Element-16823">binary\ classification</script>。然后将解决logistics regressionlogistics regression<script type="math/tex" id="MathJax-Element-16824">logistics\ regression</script>的方法从GDGD<script type="math/tex" id="MathJax-Element-16825">GD</script>算法改进到了SGDSGD<script type="math/tex" id="MathJax-Element-16826">SGD</script>算法,然后发现这样的话logistics regressionlogistics regression<script type="math/tex" id="MathJax-Element-16827">logistics\ regression</script>和PLAPLA<script type="math/tex" id="MathJax-Element-16828">PLA</script>算法看起来原理差不多。第三部分和第四部分给出了两种不同的做多类别分类的方法OVAOVA<script type="math/tex" id="MathJax-Element-16829">OVA</script>和OVOOVO<script type="math/tex" id="MathJax-Element-16830">OVO</script>。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐