Skip to content

算法Algorithm

SAGA 梯度表法

SAGA · SAGA gradient table method

逐项维护历史梯度表,以一次新分量查询形成条件无偏方向,证明表误差与迭代误差共同收缩,并核验更新顺序和存储成本。

形式陈述 ​

设 f(x)=n−1∑i=1nfi(x),每个分量可按索引重复查询。普通有限和SGD只记当前点,SAGA还记住每个分量最近一次算到的梯度。给定固定初值 x0、步长 η>0 与更新数 T,初始化

αi0=∇fi(x0),α¯0=1n∑iαi0.

初始化需 n 次分量查询。第 t 步先独立均匀抽 It∈{1,…,n},再按下面的次序执行:

  1. 新查询 g=∇fIt(xt),保存旧表项 a=αItt
  2. 用尚未修改的表形成 vt=g−a+α¯t,更新 xt+1=xt−ηvt
  3. 只替换 αItt+1=g,其余表项不变;更新均值 α¯t+1=α¯t+(g−a)/n

输出末点 xT。状态不变量是 α¯t=n−1∑iαit;每个表项来自该分量最近被访问时的查询点,不要求所有表项来自同一个位置。必须先用旧项校正,再替换它。

条件于当前点和整张旧表,

(1)E[vt∣Ft]=1n∑i[∇fi(xt)−αit]+α¯t=∇f(xt).

表项可以陈旧,式(1)仍然精确成立。若均值没有与表同步维护,抵消就会失效。

一个自足的保守收敛保证 ​

假设每个 fi 在 Rd 上凸、有 L-Lipschitz梯度,L>0;平均目标 f 为$\mu$-强凸,μ>0。记最优点 x∗,定义

At=1n∑i‖αit−∇fi(x∗)‖2,Φt=‖xt−x∗‖2+4nη2At.

取 η=1/(6L)、q=1−min{μ/(6L),1/(2n)},则

(2)EΦT≤qTΦ0,E‖xT−x∗‖2≤qTΦ0.

这里给出方便逐行验证的常数,不声称是SAGA的最锐步长或速率。当前点误差与表误差必须合起来分析,单看目标一次是否下降不足以判定进度。

直觉

SVRG把许多历史梯度统一放在一个快照上,SAGA则让每一格表在被访问时单独变新。旧表的平均方向可能不等于当前整梯度,但“新梯度减去被替换的旧项”会恰好修正它的平均偏差。

先校正,后替换,始终维护表均值

这是一种存储换查询的方式:每步只新增一个分量梯度,却要保留整个梯度表。SVRG不缓存时每阶段需要整梯度与每个内步两次查询;SAGA省去了周期性刷新阶段,增加了长期表状态的责任。

例子与边界

为什么表误差也会消失 ​

写 Δt=f(xt)−f(x∗)。对凸光滑分量,下降引理作用于

fi(x)−fi(x∗)−⟨∇fi(x∗),x−x∗⟩≥0

得到

(3)1n∑i‖∇fi(xt)−∇fi(x∗)‖2≤2LΔt.

设 bi=αit−∇fi(x∗)。方向可以写成当前梯度差减去 bI−EbI,利用 ‖u+w‖2≤2‖u‖2+2‖w‖2 与中心化不增二阶矩,有

(4)E[‖vt‖2∣Ft]≤4LΔt+2At.

每步随机更新一格,故表误差还有独立的精确更新机制:旧误差平均保留 1−1/n,新的一格受式(3)控制,因而

(5)E[At+1∣Ft]≤(1−1/n)At+2LnΔt.

另一方面,式(1)、强凸一阶下界和式(4)给距离递推

(6)E[‖xt+1−x∗‖2∣Ft]≤(1−μη)‖xt−x∗‖2−2η(1−2Lη)Δt+2η2At.

把式(5)乘 4nη2 后加到式(6)。目标差的系数为 −2η+12Lη2=0;表误差系数为 4nη2(1−1/(2n))。其余距离系数为 1−μη,所以两部分都至多乘上 q,即得式(2)。这个证明只需要平均目标强凸,不要求每个分量强凸。

初始化时式(3)还给 A0≤2LΔ0。因此若已知 ‖x0−x∗‖2≤R2、Δ0≤D0,可用

Φ0≤R2+2n9LD0

作为可计算的初始上界。表项并非额外独立样本,它们只是同一次算法历史的状态;上面的条件证明没有使用它们之间的独立性。

与SVRG相同的两分量数据 ​

取 f1(x)=(x−1)2/2、f2(x)=3(x+1)2/2、x0=0。此时 L=3,μ=2,x∗=−1/2,取定理步长 η=1/18。初始化表为 (−1,3),均值1。

第一步无论抽谁,新的分量梯度都与其旧项相同,方向为1,故 x1=−1/18,表暂未改变。第二步若抽分量2,新梯度为 17/6,方向

v1=17/6−3+1=5/6,

所以 x2=−11/108。新表为 (−1,17/6),均值 11/12,这时已不是当前整梯度 2x2+1=43/54。

再检查第三步的两种方向:抽分量1得到 22/27,抽分量2得到 7/9,平均恰为 43/54,方差 1/2916。陈旧表仍保持条件无偏;更新前后的具体表内容足以逐项复核。

前两步连同初始化一共4次分量查询,而无缓存SVRG的两个内步连同两项快照为6次。这里比较的是查询账,不是宣称二者在这两个不同步长和输出协议下具有相同精度。SAGA还长期保存两个梯度;一般 n,d 大时这份空间不能忽略。

更新顺序是算法的一部分 ​

如果先把本轮表项改成 g,再仍用旧均值计算 g−αIt+α¯,差分被错误消掉,方向只剩旧均值。上述第二步的旧均值为1,但当前真梯度为 8/9,所以这个实现已经有偏。若先更新均值再算方向,同样需要重推公式,不能靠变量名称相同认定实现正确。

若只更新表项、不更新均值缓存,则不变量也会断开。稳健的实现可定期核对缓存均值与表平均的残差;浮点累计误差须纳入容差,不能以精确代数的式(1)掩盖实现漂移。

非凸分量会破坏式(3)的非负凸余项证明;非均匀抽样则要重新加权校正项及表更新的分析。来自不可重访新样本流的梯度也不能凭空赋予固定索引并充当同一有限和分量。

推论与应用

查询、存储与返回保证 ​

使用完整初始化时,T步总共 n+T 次梯度查询;一般梯度表需 O(nd) 空间,当前点、表均值和临时向量另需 O(d)。每步维护均值只改一项,向量算术为 O(d),不应每次重新加整张表。在线性预测等特殊模型中可保存产生梯度的标量系数而降低空间,需按模型重新说明。

由平均函数光滑性,f(x)−f(x∗)≤L‖x−x∗‖2/2,所以式(2)也给期望目标差界。若要达到 ε>0,可用 LqTΦ0/2≤ε 选择预算;总查询阶为

O(n+(n+L/μ)log⁡LΦ0ε)

(初始已经合格时不必继续更新)。这是固定预算期望承诺;高概率或实际停止仍需相应转换、证书或验证。

输出应带末点、查询数和完整表状态的校验情况。若初始化未完成、梯度非有限、缓存均值不一致或存储不足,应报告失败或暂停状态。删掉表后继续运行会变成另一个算法。

参考资料
关系图谱12 个相邻概念 · 2 类关系

拖动节点调整位置。

显示关系

显示:依赖

  1. 前置三跳
  2. 前置二跳
  3. 前置一跳
  4. 当前条目
  5. 后续一跳
  6. 后续二跳
  7. 后续三跳
文字版关系按与当前条目的最短距离分组
类型化关系