Flax神经网络输出对输入Hessian向量积计算及报错解决
问题根因
嵌套vmap+grad时函数签名和输入维度不匹配:第一次vmap封装后的批量一阶导函数,会自动将批量输入沿第0维拆分为单样本送入求导逻辑;如果直接对这个批量函数再次嵌套grad+vmap,会将输入重复拆分,最终传入Dense层的是0维标量,调用jnp.shape(inputs)[-1]时空元组无对应索引,触发越界报错。
正确实现方案
核心逻辑是先定义单样本维度的求导逻辑,再统一用vmap做批量映射,不要对已经做过vmap的批量函数重复套vmap求导。
可运行的完整代码如下:
import jax import jax.numpy as jnp import flax.linen as nn from jax import jit, grad, vmap from typing import Sequence # 原网络定义保持不变 class MLP(nn.Module): features: Sequence[int] @nn.compact def __call__(self, x): for feat in self.features[:-1]: x = nn.tanh(nn.Dense(feat)(x)) x = nn.Dense(self.features[-1])(x) return x # 初始化逻辑 model = MLP([20, 20, 20, 20, 20, 1]) batch = jnp.ones((32, 3)) # 用于初始化的哑输入 params = model.init(jax.random.PRNGKey(0), batch) X = jnp.ones((32, 3)) # -------------------------- # 1. 先写单样本版本的计算逻辑:输入x形状为(3,),无批量维度 # -------------------------- @jit def u_single(params, x): # 单样本前向,输出标量 u = model.apply(params, x) return jnp.squeeze(u) # 单样本一阶导:输出形状(3,),对应每个输入维度的一阶偏导 du_dx_single = grad(u_single, argnums=1) # 单样本二阶导:输出形状(3,3),对应单样本的海塞矩阵 d2u_dx2_single = jax.jacfwd(du_dx_single, argnums=1) # -------------------------- # 2. 统一用vmap映射到批量维度 # -------------------------- u_batch = vmap(u_single, in_axes=(None, 0), out_axes=0) du_dx_batch = vmap(du_dx_single, in_axes=(None, 0), out_axes=0) d2u_dx2_batch = vmap(d2u_dx2_single, in_axes=(None, 0), out_axes=0) # 运行测试 u_val = u_batch(params, X) # 输出形状(32,),和原前向结果一致 u_X = du_dx_batch(params, X) # 输出形状(32, 3),批量一阶导 u_XX = d2u_dx2_batch(params, X) # 输出形状(32, 3, 3),批量二阶导(海塞矩阵)
补充说明
- 计算二阶导时如果不需要完整海塞矩阵,只需要海塞向量积(HVP),可以用
jax.jvp/jax.vjp组合实现,内存占用和计算速度会远优于显式计算完整海塞,适合高维输入场景。 - 求导时优先对单样本逻辑封装,再做批量映射,能避免绝大多数维度不匹配的问题,也更便于调试。
内容的提问来源于stack exchange,提问作者Vignesh Gopakumar
相关产品推荐
相关产品推荐

