Tamed Subgradient ULA:非光滑非凸场景下的Langevin采样算法解析
2026/8/28 12:33:15 网站建设 项目流程

Langevin 类算法这两年又热起来了,原因很直接:它能把优化和采样统一到同一个框架里。而 The Tamed Subgradient Unadjusted Langevin Algorithm 这个名字虽然长,翻译过来就三件事:用次梯度处理不可导函数,用 taming 技巧抑制无界梯度导致的数值爆炸,再把理论分析放到非凸场景下。这三个点恰好是很多真实机器学习问题绕不开的困难:目标函数带 L1 正则,网络里有 ReLU,损失面非凸,梯度在远处还可能长得很快。

我不打算复述论文,而是按这个标题里的技术路线,把它拆成能看懂、能动手跑一跑的内容。适合谁看?正在做贝叶斯推断、研究非光滑优化算法、或者想理解 Langevin 采样为何需要这么多变种的工程师和研究生。看完你能搞清楚:ULA 的假设边界在哪,taming 到底改了哪个项,次梯度版本和普通版本的差别在哪,以及自己写一个最小实现需要注意哪些参数和调试点。

1. 一个名字拆成三份:Tamed、Subgradient、Unadjusted 各自解决什么问题

先说结论:这个算法不是又一个全新的采样器,而是在标准 ULA 上做了三个针对性改造,应对三类机器学习里非常常见的现实问题。

1.1 ULA 的基本更新式其实非常朴素

标准 ULA 的迭代式长这样:

X_{k+1} = X_k - η∇f(X_k) + √(2η)Z_k,其中 Z_k ~ N(0, I)。

理解它不需要太深的数学。把它拆成三部分看:

  • X_k - η∇f(X_k):这是沿着负梯度方向走一小步,跟梯度下降完全一样,作用是让状态往目标函数的低值区域移动。
  • +√(2η)Z_k:这是往状态里加一个高斯噪声,作用是让状态不会死死钉在某个局部极小点,而是有能力在目标分布周围游走。
  • η 是步长:它同时控制下降速度、噪声幅度和离散化误差。

连续时间的版本是 Langevin 扩散 dX_t = -∇f(X_t)dt + √2 dW_t。它的不变分布是 π(x) ∝ exp(-f(x))。这意味着,如果你把 x 反复迭代足够多步,最后得到的样本分布会接近这个 π,恰好在 f 较小的地方概率密度更高。这个性质让 Langevin 类算法既能做优化,也能做采样。

问题在于,上面这些都是建立在非常理想的条件上:f 要可导,梯度不能太野,步长要足够小。现实不是这个样子的。

1.2 普通 ULA 会在三类场景下失控

第一类场景是 f 不可导。最典型的例子就是 f(x) = |x|,以及机器学习里最常见的 L1 正则项 λ|x|。在 x = 0 处没有传统导数,你只能退到次梯度。这个变化看起来不大,但对理论分析影响很大,因为很多用来证明收敛的式子都依赖梯度的连续性。

第二类场景是梯度或次梯度的范数无界增长。比如 f(x) = |x|^4,在 x 很大的时候,次梯度大约按 x³ 增长。欧拉格式的核心假设是步长内梯度变化不能太大,如果梯度本身快速增长,一步的更新量可能会非常大,链就直接飞到远处。轻则震荡剧烈,重则直接出 NaN。

第三类场景是 f 非凸。神经网络、混合模型、带强正则的非线性模型基本都是非凸的。非凸不是某一小块的局部形状问题,而是全局结构复杂:有多个局部极小点、鞍点、平坦区域。这种结构会让 Markov 链在模式之间迁移很慢,也会让理论证明变得困难。

TSULA 的三个关键词就是分别针对这三类问题:Subgradient 解决不可导,Tamed 解决无界次梯度,beyond Convexity 表示理论分析覆盖非凸场景。所以在读这篇论文之前,先把这三个坑列出来,后面就顺了。

1.3 Unadjusted 是效率选择,不是缺陷

Unadjusted 指这个算法不做 Metropolis-Hastings 接受修正。对比一下:

  • ULA / TSULA:每一步就是一个简单的迭代,没有额外计算。
  • MALA:每一步计算候选点后,还要计算接受概率,并抛一个随机数决定是否接受。这样能保证样本精确来自目标分布,但计算成本高,而且在次梯度场景下接受率可能很低。

为什么这里用 Unadjusted?因为很多实际任务要的不是精确的 MCMC 样本,而是后验均值、损失面附近的代表性样本。在这些任务里,只要步长足够小,ULA 的离散化偏差是可以接受的。尤其是在大数据场景下,一步迭代的代价远大于多跑几步去补偿偏差。

当然代价也存在。Unadjusted 不保证严格收敛到 π,它收敛到的是 π 的一个近似分布,偏差量级和步长 η 相关。理解这一点很重要:使用 TSULA 时,最终结果是有偏的,这个偏差可以靠小步长压小,但完全的纠偏需要回到带接受修正的框架。

2. 为什么“非光滑 + 非凸 + 次梯度无界”必须一起处理

上一节列了三个问题,这一节讲它们为什么偏偏要放在同一个算法里解决,以及各自背后的数学直觉。

2.1 次梯度在工程上怎么获取

数学上,次梯度是集合值映射。对凸函数 f 和点 x,∂f(x) 包含了所有满足 f(y) ≥ f(x) + ⟨g, y-x⟩ 的向量 g。在光滑点,这个集合只有一个元素,就是梯度;在不可导点,可能是一段区间。比如 |x| 在 0 处的次梯度是 [-1, 1] 区间。

但工程框架不会真的返回一个区间。PyTorch、TensorFlow 或自己写 JAX 代码时,你会得到某个具体的子梯度向量。不同库的选择不完全一样,比如 ReLU 在 0 处可能返回 0,也可能返回 1。这个具体选择对很多算法影响不大,但在边界分析时需要知道:你的次梯度实现并不唯一,不同的自动微分框架可能有细微差别。

写自己的优化器时,最容易出错的地方就是不可导点的返回值不是合法次梯度。比如对 f(x) = |x|,如果你手写在 x = 0 处返回了一个很大的数,那整个链都会被带偏。

2.2 taming 的核心思想:给次梯度加一个自适应缩放

taming 并不是一个很新的技巧,它最早源于随机微分方程的数值格式研究。标准欧拉方法在漂移项不满足 Lipschitz 条件时,数值解可能不收敛。解决办法很简单:把漂移项除以一个随漂移大小增长的量。

常见的 taming 形式有两种:

T(g) = g / (1 + η·‖g‖)

T(g) = g / max(1, τ·‖g‖)

两种都是连续缩放。g 很小时,分母约等于 1,T(g) ≈ g,算法退化成普通更新;g 很大时,分母随 ‖g‖ 线性增长,T(g) 的范数趋向一个常数或一个缓慢增长的函数。这样每步更新的范数就被限制住了,数值上不会一步跳出稳定区域。

这里有一个容易被忽略的细节:taming 不是把次梯度裁到一个固定阈值。它保留方向并按比例缩小,不做生硬截断。这个连续性很重要,因为生硬截断会创造一个不连续的方向场,而连续性在理论分析中很关键。

2.3 beyond convexity 之后,理论分析依赖什么假设

说到 beyond convexity,很多人会以为这意味着完全不要求凸性。不是的。非凸问题的分析通常要依赖更宽松的替代假设,用来保证 Markov 链不会永远漂走。常见的假设包括:

  • 目标函数满足耗散条件:当 x 的范数很大时,⟨∇f(x), x⟩ 应当为正且足够大,相当于说目标函数在远处像一个深井,把链往回拽。
  • 次梯度增长有界:比如 ‖∂f(x)‖ 随 ‖x‖ 的增长速度不高于某个多项式量级。
  • 局部 Lipschitz 或弱光滑性:保证相邻点的次梯度不会跳变得太离谱。

在这些假设下,即使 f 全局非凸,也能证明 ULA 变体的样本分布按某种概率度量收敛到目标分布附近。注意这里说的是“附近”,因为 Unadjusted 本身有偏差,加上非凸的限制,精确收敛的保证会更弱。具体用了哪种弱凸性定义、收敛速率是多少,要以论文原版为准。实际项目里要验证这些假设并不容易,耗散条件可以通过观察长链是否爆掉来经验性判断。如果链总是飘到很远的地方,要么是假设不满足,要么是步长开太大了。

3. 手写一个最小 TSULA:从伪代码到 Python

理论讲完,直接进入可以运行的版本。我会用一个具体的非凸非光滑函数做例子,把每一步都写出来。

3.1 算法流程和关键约定

我采用的 taming 约定如下:

  1. 初始化 X_0。
  2. 计算次梯度 g = ∂f(X_k)。
  3. 计算驯化方向 t = g / max(1, τ·‖g‖)。
  4. 更新 X_{k+1} = X_k - η·t + √(2η)·Z_k,其中 Z_k ~ N(0, I)。

注意不同论文里 τ 的位置和形式可能不同。有的把 τ 写成与 η 合并,有的用 1 + τ‖g‖ 做分母。这里用 max(1, τ‖g‖) 是为了让代码和解释都更直观。

3.2 一个可直接运行的实现

测试函数选 f(x) = (x² - 1)² + 0.5|x|。它有两个极小值点,分别在 x ≈ ±1 附近,函数在 x = 0 处有一个尖点,整体是非凸、非光滑的。

import numpy as np def subgradient_f(x): # f(x) = (x^2 - 1)^2 + 0.5 * |x| # 在 x=0 处返回 0,它是 |x| 的一个合法次梯度 return 4.0 * x * (x * x - 1.0) + 0.5 * np.sign(x)

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询