海森矩阵在机器学习中的应用

从牛顿法到 HVP:二阶信息在深度学习里到底还剩多少位置

海森矩阵是函数二阶偏导数构成的矩阵:

$$ \boldsymbol H= \begin{bmatrix} \frac{\partial^2 f}{\partial x_1^2} & \frac{\partial^2 f}{\partial x_1\partial x_2} & \dots\\ \frac{\partial^2 f}{\partial x_2\partial x_1} & \frac{\partial^2 f}{\partial x_2^2} & \dots\\ \vdots & \vdots & \ddots \end{bmatrix} $$

它刻画的是损失函数的曲率。一阶导数是梯度,给出斜率;二阶的海森则告诉你曲面弯曲的程度——同样的斜率,走在一条平缓的沟里和走在一堵陡壁上,该迈的步子完全不一样。


一、牛顿法(二阶优化)

梯度下降只用一阶梯度,沿斜率方向走一个固定步长;牛顿法则用海森做二阶近似,直接解出这一步该走多远:

$$\boldsymbol x_{t+1}=\boldsymbol x_t - H^{-1}\nabla f(\boldsymbol x_t)$$

优点:利用了曲率信息,可以显著加速收敛。靠近极小点时,其收敛速度远快于梯度下降。

缺点:

  1. $H$ 的维度等于参数数量。大模型参数千万级,海森矩阵内存直接爆炸;
  2. 需要求逆,计算代价极高;
  3. 海森非正定时,更新方向不一定是下降方向。

实际深度学习中几乎不用原始牛顿法,而是衍生出了各种拟牛顿法


二、拟牛顿法(L-BFGS)

思路是不直接计算真实海森,而是迭代地近似海森的逆矩阵。


三、高斯-牛顿法 Gauss-Newton

针对最小二乘损失(回归问题),用 $H \approx J^T J$ 近似海森,其中 $J$ 是雅可比矩阵。这样就完全规避了二阶导数的计算。

应用:非线性最小二乘、SLAM、传统非线性回归。


四、分析损失曲面与临界点判断

海森的特征值可以判断一个临界点的性质:

  1. 特征值全部 > 0:正定,局部极小点
  2. 特征值全部 < 0:负定,局部极大点;
  3. 正负特征值同时存在:鞍点

深度学习的损失曲面上有大量鞍点——梯度几乎为 0,但并不是极值。梯度下降很容易卡在鞍点附近,而海森特征值可以把它们识别出来。

一个有意思的现象

高维空间里鞍点极多,局部极小反而很少。而且很多极小点的海森特征值有大量接近 0 的方向,对应的是平坦极小——这类解的泛化性通常更好。


五、二阶优化在深度学习中的变种

5.1 自然梯度 Natural Gradient

用 Fisher 信息矩阵(等价于期望海森)代替海森,在概率模型和变分推断中使用。计算成本依然很高。

5.2 海森对角近似

只保留海森的对角线元素,忽略交叉二阶导数。

AdaGrad、RMSprop、Adam 都可以理解为对角海森的粗糙近似——用梯度平方来模拟二阶曲率信息,不需要显式构造海森矩阵。

注意:Adam 并没有真的去算海森,它只是启发式地模仿了二阶缩放的行为。

5.3 海森向量乘积 HVP(Hessian-vector product)

不构造完整的 $H$ 矩阵,而是直接实现 $Hv$,从而避免存储整个海森。

用途:计算海森特征值、对抗扰动、模型敏感性分析、二阶信任域优化(TRPO)。强化学习中 TRPO 的核心就是海森-向量乘积。


六、模型可解释性:用海森识别重要参数

损失对参数的海森对角元大小,代表该参数对损失变化的敏感程度:

这可以用于参数重要性筛选和剪枝分析。


七、贝叶斯机器学习

拉普拉斯近似:在 MAP 点用海森矩阵近似后验协方差,$\Sigma \approx H^{-1}$。

也就是拿到极大后验点之后,用二阶曲率去估计参数的后验分布。小模型的贝叶斯推断中很常用。


八、局限性总结

为什么深度学习很少直接用海森?三条原因:

  1. 参数是 N 维,海森大小就是 $N \times N$。百万参数的模型,海森矩阵有 $10^{12}$ 个元素,内存完全放不下;
  2. 每一步求逆的代价是 $O(N^3)$,无法用来训练大网络;
  3. 非凸的神经网络损失下,海森经常不定,牛顿方向不一定是下降方向。

实践现状:

小数据集、传统机器学习(如逻辑回归):L-BFGS 这类拟牛顿法表现优秀;

深度神经网络:几乎不使用完整海森,只使用海森-向量乘积、对角近似等轻量化变体;主流依然是一阶优化器 Adam / SGD。


九、代码示例:计算海森矩阵与 HVP

下面用 PyTorch 的自动微分实现,不需要手动求导。完整海森矩阵只适合低维;海森向量乘积 HVP 不生成完整矩阵,是高维场景的必备手段

import torch

# 定义损失函数 f(x,y) = x^2 + 2*y^2 + x*y
def loss_fn(params):
    x, y = params
    return x**2 + 2 * y**2 + x * y

# ---------------------- 1. 显式计算完整海森矩阵(仅适合低维) ----------------------
def compute_hessian(loss, params):
    n = params.numel()
    hessian = torch.zeros((n, n))
    grad = torch.autograd.grad(loss, params, create_graph=True)[0]
    for i in range(n):
        hessian[i] = torch.autograd.grad(grad[i], params, retain_graph=True)[0]
    return hessian

params = torch.tensor([1.0, 2.0], requires_grad=True)
loss_val = loss_fn(params)
H = compute_hessian(loss_val, params)
print("完整海森矩阵:")
print(H)

# ---------------------- 2. 海森-向量乘积 HVP: H@v,不构造完整矩阵 ----------------------
def hessian_vector_product(loss, params, vec):
    """
    计算 H * vec
    :param loss: 标量损失
    :param params: 模型参数张量
    :param vec: 要相乘的向量,shape 和 params 一致
    :return: Hv 结果向量
    """
    grad = torch.autograd.grad(loss, params, create_graph=True)[0]
    grad_dot_v = torch.sum(grad * vec)
    hvp = torch.autograd.grad(grad_dot_v, params)[0]
    return hvp

v = torch.tensor([1., 0.])
hv_result = hessian_vector_product(loss_val, params, v)
print("\nHVP H @ v 结果向量:")
print(hv_result)

# 验证:用完整海森矩阵做矩阵乘法,和 HVP 结果对比
verify = H @ v
print("完整矩阵乘法验证结果:")
print(verify)

输出结果:

完整海森矩阵:
tensor([[2., 1.],
        [1., 4.]])

HVP H @ v 结果向量:
tensor([2., 1.])
完整矩阵乘法验证结果:
tensor([2., 1.])

代码要点:

1. 完整海森只适合参数很少的场景;网络参数上万、上百万时,绝对不能生成完整的 $N \times N$ 矩阵;

2. HVP 是深度学习二阶方法的标准实现手段,TRPO、二阶敏感性分析全部基于此,而且只需要 $O(N)$ 内存。


十、对比简表

方法是否显式海森适用场景
梯度下降 SGD深度网络大模型
牛顿法完整 H低维小优化问题
L-BFGS 拟牛顿近似 H 逆传统回归,中小凸问题
Gauss-Newton$J^TJ$ 近似 H非线性最小二乘
TRPOH-vector product强化学习
拉普拉斯近似H 逆作为协方差贝叶斯小模型
← 返回首页