You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何加速Jax神经网络中的Newton更新步骤?

加速类Newton神经网络更新步骤的优化建议

我正在处理一个可通过小规模神经网络解决的问题,但需要使用类Newton更新方法。以下是可运行示例代码,希望获取加速Newton更新步骤的建议。我重点关注提升代码中更新部分的速度,这是脚本耗时最多的环节,期望专家帮忙大幅提速。

import jax
import equinox as eqx
from jax import numpy as jnp
import matplotlib.pyplot as plt
from jax import flatten_util

key = jax.random.PRNGKey(42)
key, subkey1, subkey2 = jax.random.split(key, 3)
data = jnp.concatenate((jax.random.normal(subkey1, shape=(100, 2)) * 0.1 - 1, jax.random.normal(subkey2, shape=(100, 2)) * 0.1 + 1), axis=0)
labels = jnp.array([0] * 100 + [1] * 100)
plt.scatter(data[:,0], data[:,1])
plt.show()

class MLP(eqx.Module):
    layers: list

    def __init__(self, key):
        key1, key2, key3 = jax.random.split(key, 3)
        self.layers = [
            eqx.nn.Linear(2, 10, key=key1),
            jax.nn.relu,
            eqx.nn.Linear(10, 12, key=key2),
            jax.nn.relu,
            eqx.nn.Linear(12, 2, key=key3),
            jax.nn.log_softmax
        ]

    def __call__(self, x):
        for layer in self.layers:
            x = layer(x)
        return x

def loss_fn(model, ins, ytrue):
    pred_y = eqx.filter_vmap(model)(ins)
    return cross_entropy(ytrue, pred_y)

def cross_entropy(y, pred_y):
    pred_y = jnp.take_along_axis(pred_y, jnp.expand_dims(y, 1), axis=1)
    return -jnp.mean(pred_y)

@eqx.filter_jit
def compute_accuracy(m, x, y):
    pred_y = eqx.filter_vmap(m)(x)
    pred_y = jnp.argmax(pred_y, axis=1)
    return jnp.mean(y == pred_y)

key = jax.random.PRNGKey(42)

@eqx.filter_jit
def step(mlp, xs, ys):
    vals, grads = eqx.filter_value_and_grad(loss_fn)(mlp, xs, ys)
    updates = jax.tree_map(lambda g: -0.1 * g, grads)
    mlp = eqx.apply_updates(mlp, updates)
    return mlp, vals

epochs = 50
batch_size = 100
grad_loss = []
grad_acc = []

key, subkey = jax.random.split(key, 2)
model = MLP(subkey)

for e in range(epochs):
    if e % 20 == 0:
        print(e, "/", epochs)
    
    key, subkey = jax.random.split(key, 2)
    inds = jax.random.randint(subkey, minval=0, maxval=len(data), shape=(batch_size,))
    inputs = data[inds]
    ls = labels[inds]
    
    model, loss = step(model, inputs, ls)
    grad_loss.append(loss)
    grad_acc.append(compute_accuracy(model, data, labels))

@eqx.filter_jit
def loss_h(arrs, static, ins, ytrue, uf):
    arrs = uf(arrs)
    model = eqx.combine(arrs, static)
    pred_y = eqx.filter_vmap(model)(ins)
    return cross_entropy(ytrue, pred_y)

@eqx.filter_jit
def step_h(mlp, xs, ys):
    vals, grads = eqx.filter_value_and_grad(loss_fn)(mlp, xs, ys)
    a, s = eqx.partition(mlp, eqx.is_inexact_array)
    flat_a, unflat_a = flatten_util.ravel_pytree(a)
    h = jax.hessian(loss_h)(flat_a, s, xs, ys, unflat_a)
    g_flat, unflat = flatten_util.ravel_pytree(grads)
    updates = unflat(-1 * jnp.linalg.pinv(h) @ g_flat)
    mlp = eqx.apply_updates(mlp, updates)
    return mlp, vals

key = jax.random.PRNGKey(1)
key, subkey = jax.random.split(key, 2)
model_h = MLP(subkey)

epochs = 50
batch_size = 100
h_loss = []
h_acc = []

for e in range(epochs):
    if e % 20 == 0:
        print(e, "/", epochs)
    
    key, subkey = jax.random.split(key, 2)
    inds = jax.random.randint(subkey, minval=0, maxval=len(data), shape=(batch_size,))
    inputs = data[inds]
    ls = labels[inds]
    
    model_h, loss = step_h(model_h, inputs, ls)
    h_loss.append(loss_fn(model_h, data, labels))
    h_acc.append(compute_accuracy(model_h, data, labels))

fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 4))

ax1.plot(h_loss, label="Newton")
ax2.plot(h_acc, label="Newton")

ax1.plot(grad_loss, label="Grad")
ax2.plot(grad_acc, label="Grad")

plt.legend()
plt.show()

核心优化方向

原代码中Newton步骤的性能瓶颈主要在于完整Hessian矩阵的计算和伪逆求解,以下是针对性的提速方案:

1. 用Hessian-vector乘积替代完整Hessian计算

完整Hessian是O(n²)复杂度的矩阵(n为参数数量),即使小规模网络也会带来冗余计算。改用Hessian-vector乘积(HVP)结合共轭梯度法,可以在O(n)复杂度下求解H⁻¹g,避免生成完整Hessian矩阵。

修改后的step_h函数:

@eqx.filter_jit
def step_h(mlp, xs, ys):
    vals, grads = eqx.filter_value_and_grad(loss_fn)(mlp, xs, ys)
    # 将梯度展平为向量
    g_flat, unflat_g = flatten_util.ravel_pytree(grads)
    
    # 定义Hessian-vector乘积函数
    def hvp(v):
        v_tree = unflat_g(v)
        # 通过VJP+JVP计算Hessian与向量的乘积
        _, grad_vjp = eqx.filter_vjp(loss_fn, mlp, xs, ys)
        loss_vjp = grad_vjp(v_tree)[0]
        _, hv = eqx.filter_value_and_grad(lambda m: loss_vjp)(mlp)
        hv_flat, _ = flatten_util.ravel_pytree(hv)
        return hv_flat
    
    # 用共轭梯度求解 H^{-1} * g
    delta_flat, _ = jax.scipy.sparse.linalg.cg(hvp, g_flat, maxiter=10)
    updates = unflat_g(-delta_flat)
    mlp = eqx.apply_updates(mlp, updates)
    return mlp, vals

2. 正则化Hessian并改用线性求解替代伪逆

jnp.linalg.pinv依赖SVD分解,计算成本远高于普通线性求解。通过给Hessian添加小的对角正则项(Tikhonov正则化),可以确保矩阵可逆,然后用jnp.linalg.solve替代伪逆:

修改原step_h中的求解部分:

h = jax.hessian(loss_h)(flat_a, s, xs, ys, unflat_a)
# 添加Tikhonov正则化,避免矩阵奇异
h_reg = h + 1e-4 * jnp.eye(h.shape[0])
# 用solve替代pinv,速度提升明显
delta_flat = jnp.linalg.solve(h_reg, g_flat)
updates = unflat(-delta_flat)

3. 减少Hessian计算频率

Newton方法不需要每轮迭代都重新计算Hessian。可以每隔k轮更新一次Hessian,中间复用之前的Hessian进行更新,大幅降低计算开销:

key, subkey = jax.random.split(key, 2)
model_h = MLP(subkey)
h_reg = None
unflat = None

for e in range(epochs):
    if e % 20 == 0:
        print(e, "/", epochs)
    
    key, subkey = jax.random.split(key, 2)
    inds = jax.random.randint(subkey, minval=0, maxval=len(data), shape=(batch_size,))
    inputs = data[inds]
    ls = labels[inds]
    
    vals, grads = eqx.filter_value_and_grad(loss_fn)(model_h, inputs, ls)
    g_flat, unflat = flatten_util.ravel_pytree(grads)
    
    # 每5轮更新一次Hessian
    if e % 5 == 0:
        a, s = eqx.partition(model_h, eqx.is_inexact_array)
        flat_a, unflat_a = flatten_util.ravel_pytree(a)
        h = jax.hessian(loss_h)(flat_a, s, inputs, ls, unflat_a)
        h_reg = h + 1e-4 * jnp.eye(h.shape[0])
    
    delta_flat = jnp.linalg.solve(h_reg, g_flat)
    updates = unflat(-delta_flat)
    model_h = eqx.apply_updates(model_h, updates)
    
    h_loss.append(loss_fn(model_h, data, labels))
    h_acc.append(compute_accuracy(model_h, data, labels))

4. 用Equinox原生工具简化Hessian计算

原代码手动处理参数的partition和flatten,代码冗余且容易出错。改用eqx.filter_hessian可以直接处理模型的Hessian,自动区分静态/动态参数:

@eqx.filter_jit
def step_h(mlp, xs, ys):
    vals, grads = eqx.filter_value_and_grad(loss_fn)(mlp, xs, ys)
    # 直接计算模型参数的Hessian树
    hessian_tree = eqx.filter_hessian(loss_fn)(mlp, xs, ys)
    # 将梯度和Hessian展平为向量/矩阵
    g_flat, unflat_g = flatten_util.ravel_pytree(grads)
    h_flat, _ = flatten_util.ravel_pytree(hessian_tree)
    h = h_flat.reshape(g_flat.shape[0], g_flat.shape[0])
    
    h_reg = h + 1e-4 * jnp.eye(h.shape[0])
    delta_flat = jnp.linalg.solve(h_reg, g_flat)
    updates = unflat_g(-delta_flat)
    mlp = eqx.apply_updates(mlp, updates)
    return mlp, vals

5. 最大化JAX编译优化

确保所有计算逻辑都被eqx.filter_jit包裹,避免jit边界外的冗余操作。同时可以尝试开启JAX的XLA优化(默认已开启),对于GPU/TPU环境,HVP和线性求解会自动并行加速。

内容的提问来源于stack exchange,提问作者Baba Yara Fahiz

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 04:34:51