如何加速Jax神经网络中的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

