如何用JAX的vmap批量计算二维神经网络的雅可比矩阵?
用JAX批量计算二维神经网络近似函数的一阶/二阶导数
你当前代码存在两处关键问题:
- 重复定义
u时,第二个lambda仅传入x参数,遗漏了y,导致后续求导逻辑失效 - 二维输入场景下,
vmap与jacfwd的组合方式未匹配输入的批量维度结构
以下是修正后的完整实现,确保输出形状符合(10,1)的预期:
1. 修正网络函数与基础定义
首先补全网络函数中遗漏的偏置项,并正确定义近似函数u:
import jax import jax.numpy as jnp from jax import jnn, vmap, jacfwd def f(params, inputs_x, inputs_y): inputs = jnp.concatenate((inputs_x, inputs_y), axis=1) for w, b in params: outputs = jnp.dot(inputs, w) + b # 补全原代码遗漏的偏置项b inputs = jnn.swish(outputs) return outputs # 正确定义u:接收x、y两个输入参数 u = lambda x, y: f(params, x, y)
2. 初始化参数与批量输入
按照[2,5,1]的层结构初始化参数,并生成10批测试输入:
# 初始化参数PyTree:对应[2→5]、[5→1]两层的权重与偏置 params = [ (jnp.ones((2, 5)), jnp.zeros(5)), (jnp.ones((5, 1)), jnp.zeros(1)) ] # 生成形状为(10,1)的批量输入 x = jnp.ones((10, 1)) y = jnp.ones((10, 1))
3. 批量计算一阶导数
通过jacfwd计算单样本导数,再用vmap批量处理所有样本:
# du/dx:对x求导,批量处理 u_x = vmap(jacfwd(u, argnums=0), in_axes=(0, 0))(x, y) # du/dy:对y求导,批量处理 u_y = vmap(jacfwd(u, argnums=1), in_axes=(0, 0))(x, y) # 验证输出形状:均为(10,1) print(u_x.shape) # 输出 (10, 1) print(u_y.shape) # 输出 (10, 1)
4. 批量计算二阶导数
对一阶导数函数再次求导,结合vmap实现批量处理:
# d²u/dx²:嵌套求导后批量处理 u_xx = vmap(lambda x_i, y_i: jacfwd(lambda x_val: jacfwd(u, argnums=0)(x_val, y_i))(x_i), in_axes=(0, 0))(x, y) # d²u/dy²:同理实现 u_yy = vmap(lambda x_i, y_i: jacfwd(lambda y_val: jacfwd(u, argnums=1)(x_i, y_val))(y_i), in_axes=(0, 0))(x, y) # 验证输出形状:均为(10,1) print(u_xx.shape) # 输出 (10, 1) print(u_yy.shape) # 输出 (10, 1)
内容的提问来源于stack exchange,提问作者Sumanta Roy
相关产品推荐
相关产品推荐

