尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Tamed SGULA:驯化次梯度Langevin采样算法解决非凸非光滑难题

Tamed SGULA:驯化次梯度Langevin采样算法解决非凸非光滑难题 在贝叶斯推断、概率机器学习和采样型生成模型的实现中经常要面对同一个问题已知目标分布只差一个归一化常数即 (p(x) \propto \exp(-f(x)))需要生成一批服从该分布的样本。当 (f) 是光滑凸函数时梯度型 Langevin 采样器已经足够好用但当目标函数同时带有非光滑项和非凸结构时比如稀疏贝叶斯后验、带 L1 正则的鲁棒估计或者带惩罚项的非凸统计模型直接使用基于梯度的采样器会出现轨迹发散、偏差变大和混合变慢三类问题。Tamed Subgradient Unadjusted Langevin AlgorithmTamed SGULA即驯化次梯度非修正 Langevin 算法就是研究这类非凸非光滑采样问题的代表算法之一。这篇文章从 Langevin 采样的基本链路讲起说明为什么要引入次梯度、为什么需要驯化、算法怎么实现以及在使用和复现时容易踩到哪些坑。1. 先厘清三条技术主线Langevin、Unadjusted 与次梯度1.1 Langevin 动力学是采样器的源头通俗理解Langevin 动力学是一组带有随机扰动的微分方程它让粒子在高维空间中沿着“势能下降”方向和“噪声扩散”方向同时运动。只要噪声的强度与温度参数匹配粒子长期运行后所在位置的分布会收敛到 Boltzmann 分布也就是与 (\exp(-f(x))) 成比例的密度对应的分布。技术定义设目标密度为 (\pi(x) \propto \exp(-f(x)))其中 (f) 称为势能函数。在维纳过程扰动下的梯度下降系统写作 Langevin SDE[ dX_t -\nabla f(X_t) dt \sqrt{2} dB_t ]其中 (B_t) 是标准布朗运动(\sqrt{2}) 是保证稳态分布为 (\pi) 的关键系数。通过 Fokker-Planck 方程可以验证该 SDE 的唯一不变测度正是目标分布 (\pi)。从 MCMC 视角看Langevin 动力学把采样问题转化成模拟随机微分方程的问题这也是后面所有离散化算法的源头。这里容易误解的是Langevin 动力学只有在长时间极限下才收敛到 (\pi)任何有限时间窗口内都存在时间离散化带来的偏差。离散化算法真正要做的是在计算代价和偏差之间取平衡。1.2 Euler-Maruyama 离散化得到 ULA把上面的连续时间 SDE 用 Euler-Maruyama 格式离散化步长为 (\eta)得到[ X_{k1} X_k - \eta \nabla f(X_k) \sqrt{2\eta} Z_k, \quad Z_k \sim N(0, I_d) ]这就是 Unadjusted Langevin AlgorithmULA。所谓 Unadjusted指算法不带 Metropolis-Hastings 拒绝步骤而是直接把离散链当作近似采样器。这样做的好处是每一步只需要一次梯度求值适合高维和大数据场景代价是离散化误差没有被修正链的稳态分布与目标 (\pi) 之间存在 (O(\eta)) 量级的偏差。实际项目里如果没法做拒绝校正就要通过调小 (\eta) 来控制这个偏差。1.3 从梯度到次梯度非光滑目标只能这样走如果 (f) 不是处处可微比如包含 (|x|)、ReLU、分位数损失等项(\nabla f) 在部分点不存在。此时“梯度下降”的替换物是次梯度对凸函数来说(g \in \partial f(x)) 满足 (f(y) \ge f(x) \langle g, y-x\rangle)。对非凸函数次梯度通常指 Clarke 次微分或其他广义梯度性质更复杂但算法层面仍可以选取一个方向 (g) 来替代 (\nabla f)。把 ULA 中的 (\nabla f) 换成 (g \in \partial f(x))就得到 Subgradient Langevin AlgorithmSGULA。这一步看起来只是把符号换掉实际上引入了两个新问题次梯度的选取方式会影响随机噪声的方差次梯度在非凸区域的增长可能接近无界直接带入 Euler 格式会导致轨迹在少数几步内飞出去。2. 非凸非光滑场景为什么难处理2.1 次梯度无界是离散化发散的根因考虑目标函数 (f(x) (x^2 - 1)^2 / 2 \lambda |x|)。在 (x) 离原点较远时光滑部分的导数约为 (2x^3)增长速度超过线性。ULA 更新公式里有一项 (-\eta \nabla f(x))当当前点位于远离原点的区域时这一步会让点跳向更远的位置之后导数更大形成正反馈。理论上连续时间 SDE 中噪声和漂移之间的平衡可以维持稳定但离散化后每个大步长内“漂移过大”会导致轨迹爆炸。数学上这对应着漂移系数的线性增长条件被破坏即通常要求的 (|\nabla f(x)| \le C(1 |x|)) 不成立。当 (\nabla f) 换成次梯度 (g) 后问题更明显次梯度在不可微点往往有取值区间且可能比两侧导数更不规整一些解析表达式还会在数值上产生很大的瞬时值。若直接做 (x - \eta g)数值稳定性几乎没有保证。2.2 非凸性破坏全局收缩凸函数的好处是在强凸条件下链与稳态之间的 Wasserstein 距离每一步都会按常数比例收缩。非凸函数不存在全局强凸分析时必须引入“在无穷远处耗散”之类的条件例如 (\langle \nabla f(x), x\rangle \ge a|x|^2 - b)说明势能在远处足够“陡”能把粒子拉回来。但这些条件仍不足以保证全局快速混合只能保证链不会发散并且在一个较弱的度量下收敛到目标附近。2.3 直接套 ULA 的三个典型失败模式实际观察中非凸非光滑目标上直接使用 ULA 或 SGULA 会出现以下现象。第一轨迹爆炸若干步内数值溢出日志里出现 inf 或 nan。第二稳态偏差链不爆炸但长期停留在某个势阱附近样本分布明显偏离 (\pi)。第三估计方差过大由于高杠杆点偶尔出现样本均值或方差的 Monte Carlo 估计剧烈抖动。这三类问题正是 Tamed SGULA 想解决的。3. Tamed Subgradient ULA 的算法设计3.1 Taming 的思想让漂移项的增长速度被压制“驯化”一词来自随机微分方程的数值格式研究。原始思想来自处理超线性系数方程当漂移项 (b(x)) 增长太快时用带归一化的形式替换 (b(x))使得无论 (b(x)) 多大更新量都被限制在可控范围同时保留原系统的主要方向信息。对采样问题Tamed SGULA 的更新可以写作[ X_{k1} X_k - \frac{\eta g_k}{1 \eta |g_k|} \sqrt{2\eta} Z_k ]其中 (g_k \in \partial f(X_k))。当 (|g_k|) 很小分母接近 1迭代退化为普通 SGULA当 (|g_k|) 很大漂移项被压制到接近 ±1粒子不会因为单步漂移过大而飞出。另一种常见形式是使用 (1 |g_k|) 做分母性质类似。两种变体都不改变目标分布的大致方向但都会额外引入 (O(\eta)) 量级的偏差理论分析需要把它们计入误差项。为什么不让分母变化太剧烈因为算法最终要模拟稳态为 (\pi) 的扩散过程驯化只应限制极端漂移不能改变小梯度区域的动力学。如果分母设计不合理链的稳态分布会明显偏离目标。3.2 完整算法伪代码下面给出带烧入期和采样期的完整流程输入: 势能 f次梯度 oracle g ∂f步长 η迭代总数 N烧入期 B 初始化: 从某个初始分布 μ0 中采样 X0 for k 0 to N - 1 do 计算 g_k ∈ ∂f(X_k) X_{k1} X_k - η g_k / (1 η |g_k|) sqrt(2η) Z_k, Z_k ~ N(0, I_d) end for 输出: {X_{B1}, ..., X_N}对多维情形(|g_k|) 表示欧氏范数。实际实现中计算分母时要注意不要对整个向量做逐元素的绝对值后直接相加而应使用范数如果目标函数在维度间完全独立也可以按维度分别驯化但要明确这种改动会影响理论假设的匹配。3.3 关键参数速查表参数含义常见起点调小调大(\eta)离散化步长0.001 ~ 0.01偏差更小混合更慢混合更快偏差和发散风险增大分母中的常数驯化强度1 或与 (\eta) 配套驯化效果减弱漂移被过度压制稳态偏差增大(B) 烧入期丢弃前多少个样本总迭代的 10%~30%初始分布影响残留浪费计算量采样间隔每多少步保留一个样本1 或 10~100链内相关性高信息损失更多这里的取值只是给初学者的起点不能当作最优配置。真正做实验时要根据目标函数的尺度、维度、步长和收敛诊断来调整。4. 收敛性分析理论证明通常需要哪些条件4.1 非凸场景下的常见假设优化和采样理论中的分析本质上要回答两个问题链会不会爆炸链离目标分布有多远。围绕这两个问题Tamed SGULA 这类工作的证明通常包含以下假设。势能函数 (f) 有下界避免稳态密度退化。耗散条件存在 (a0,\ b\ge0)使得 (\langle g, x\rangle \ge a|x|^2 - b) 对所有 (g \in \partial f(x)) 成立。这保证链在无穷远处会被拉回。次梯度增长条件存在 (C0) 和 (p\ge0)使 (|g| \le C(1 |x|^p))。如果 (p) 太大普通 SGULA 的爆炸风险高而驯化格式可以在较宽的 (p) 范围内保持稳定这正是 “tamed” 的价值。若需要量化收敛速度还会要求目标分布满足 Poincaré 不等式或 log-Sobolev 不等式这些条件描述的是 (\pi) 尾部衰减和函数空间的谱性质而不是 (f) 的凸性。注意这些条件中没有任何一条要求 (f) 是凸函数因此分析框架天然覆盖非凸场景。4.2 典型结论的形式这类证明得到的界通常写成[ W_2(\mu_k, \pi) \le C(1 - c\eta)^k W_2(\mu_0, \pi) C\eta^{q} ]第一项反映初始分布的记忆在一定条件下按几何速度衰减第二项是离散化偏差(q) 通常为 1 或 1/2取决于光滑性和驯化的具体形式。这种“收缩项加偏差项”的结构在非凸问题的采样分析中很常见。具体常数 (C)、(c)、(q) 的数值与假设的精确形式强相关因此复现论文实验时必须对照原文的假设和参数设置不能把一个函数族的参数直接搬到另一个目标函数上。4.3 三种算法的适用边界对比算法对 (f) 的要求漂移计算稳态偏差主要风险ULA光滑梯度温和增长(\nabla f)(O(\eta))步长稍大就发散SGULA非光滑次梯度增长温和任选 (g \in \partial f)(O(\eta)) 次梯度方差远端瞬时值过大导致爆发Tamed SGULA非光滑次梯度增长可较快驯化后的 (g)(O(\eta)) 驯化偏差驯化过强导致稳态失真这张表也解释了为什么算法名称中 “Subgradient” 和 “Tamed” 同时出现前者解决可微性不足后者解决稳定性不足二者面向的是同一类非凸非光滑目标函数的两个不同困难。5. Python 数值实验从实现到验证5.1 实验目标与测试目标函数下面用一个一维非凸非光滑目标函数验证算法行为。取势能[ f(x) \frac{(x^2 - 1)^2}{2} \lambda |x| ]其中 (\lambda|x|) 是非光滑项((x^2-1)^2/2) 是双势阱结构因此 (f) 既非凸也非处处可微。目标密度 (\pi(x) \propto \exp(-f(x))) 在 (x \pm 1) 附近有两个峰(x0) 附近有低密度区域。代码目标是观察三点普通 SGULA 在大步长下是否发散Tamed SGULA 是否更稳定样本直方图是否与目标密度形状一致。5.2 实现代码先定义势能函数的次梯度 oracle在 (x0) 处取 (\text{sign}(0)0)属于次梯度集合中的一个合法选择。import numpy as np def subgrad_f(x, lam0.8): # f(x) (x^2 - 1)^2 / 2 lam * |x| smooth_grad 2.0 * x * (x * x - 1.0) subgrad lam * np.sign(x) return smooth_grad subgrad def sgula_step(x, eta, grad_fn): z np.random.randn() return x - eta * grad_fn(x) np.sqrt(2.0 * eta) * z def tamed_sgula_step(x, eta, grad_fn): g grad_fn(x) drift eta * g / (1.0 eta * abs(g)) z np.random.randn() return x - drift np.sqrt(2.0 * eta) * z def run_chain(step_fn, x0, eta, n_iter, burn_in, grad_fn): x x0 samples [] for k in range(n_iter): x step_fn(x, eta, grad_fn) if k burn_in: samples.append(x) return np.array(samples) np.random.seed(0) x0 0.0 eta 0.05 samples_tamed run_chain( tamed_sgula_step, x0, eta, n_iter200000, burn_in5000, grad_fnsubgrad_f )这里把步长故意设为 0.05对尺度较大的非凸势能已经偏大。普通 SGULA 在这种配置下容易出现轨迹长时间卡在远端甚至数值溢出Tamed 版本则能把漂移压制住使链保持在有限范围内。5.3 结果验证画直方图并对比理论密度为了判断样本是否近似服从目标分布可以在网格上数值归一化目标密度import matplotlib.pyplot as plt xs np.linspace(-3.0, 3.0, 1000) dens np.exp(-((xs * xs - 1.0) ** 2 / 2.0 0.8 * np.abs(xs))) dens / np.trapezoid(dens, xs) # 旧版 NumPy 可用 np.trapz plt.hist(samples_tamed, bins120, densityTrue, alpha0.5, labelTamed SGULA) plt.plot(xs, dens, r-, labeltarget density) plt.legend() plt.show()验证时不要只看图形是否接近。可以计算样本均值、标准差以及样本落在峰区间的比例与数值积分得到的理论值比较。更严谨的情况下使用多个独立链计算不同链之间的差异或者用能量距离、Wasserstein 距离与参考样本做比较。6. 常见问题与排查链路6.1 样本路径发散出现 inf 或 nan现象运行几步后输出 inf或绘图时坐标轴突然出现极大值。可能原因步长过大次梯度瞬时值过大分母中的驯化项没有生效比如把范数写成了逐元素绝对值后相加导致分母偏小。检查方式打印前几十步的 (x) 和 (g)如果 (|g|) 达到 (10^8) 量级检查势能函数是否在数值上溢确认步长是否符合目标函数尺度。处理建议先减小 (\eta) 一个数量级再用 (1 |g|) 形式的驯化分母最后检查势能表达式避免类似 ((x^2-1)^2) 展开成 (x^4 - 2x^2 1) 时在精度上被严重抵消。6.2 链不爆炸但样本分布长期不收敛现象直方图明显偏向一侧或者多链对比时链之间的均值差异大。可能原因烧入期不足初始点附近的势阱让链没有完全混合步长太小导致移动缓慢目标多峰分布而链没有足够的跳跃能力穿过低密度区域。检查方式画轨迹图看是否存在长时间滞留比较不同初值得到的两条链的样本均值观察自相关时间是否过长。处理建议增加烧入期适当增大步长或者使用退火、并行多链等方案对真实目标不能期望单链在超多峰情况下快速混合这是采样算法的固有限制而不是驯化的问题。6.3 次梯度选择影响结果现象同一份数据、同一步长不同次梯度 oracle 得到的估计不同。可能原因在不可微点上次梯度是一个集合不同选择对应不同的噪声注入方式对于非凸函数广义次梯度的定义也可能影响分析假设。检查方式检查代码中 (\text{sign}(0))、ReLU 在 0 处的导数等特殊分支统计不可微点被访问的频率如果频率很低影响通常较小。处理建议保持固定的、可解释的选择规则在复现论文时严格对齐原文的次梯度定义如果不可微点访问频繁考虑使用 Huber 等光滑化近似。6.4 推荐排查顺序按以下顺序排查先确认目标函数和次梯度 oracle 正确再确认维度、范数和指数正确然后检查步长和驯化参数接着看烧入期与采样间隔最后观察日志中的异常数值和轨迹图。7. 工程实践建议与后续学习路径7.1 学习环境与生产环境的差别学习阶段用一维或低维例子验证算法语义可以固定随机种子直接打印轨迹和直方图步长可以选得大一点看失败模式。生产环境则要额外处理目标函数和次梯度的数值稳定性日志记录每次更新的范数便于快速定位发散针对多维参数的分块采样自动调节步长的方案以及定期用多链诊断检查混合质量。采样算法在生产系统里通常不是单独运行而是作为更大推理管线的一部分因此还要考虑容错、断点续跑和结果版本管理。7.2 可复用检查清单在发布或上线前按下面清单逐项确认目标密度 (\exp(-f)) 不需要计算归一化常数但 (f) 的表达式在数值上稳定。次梯度 oracle 在不可微点有明确定义且实现与论文一致。步长 (\eta) 从 0.001 开始逐步放大不超过目标函数尺度的倒数。驯化分母的范数计算正确多维情形使用欧氏范数。有足够长的烧入期单链轨迹图和自相关图均已检查。
返回列表