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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 05:06:29