如何基于JAX批量计算神经网络对输入的各类导数?
JAX中神经网络多输入导数计算实现方案
问题说明
现有神经网络net(x, t),输入为d维变量x(批量输入形状(b, d))和标量t(批量输入形状(b, 1)),输出为长度d的向量(批量输出形状(b, d))。需完成以下计算:
- 神经网络输出对
t的导数d out/dt,形状(batch, d); - 神经网络输出对
x的导数d out/dx; - 神经网络输出的散度对
x的梯度,形状(batch, d)。
已知PyTorch实现方式,但刚接触JAX,需要对应实现方案。
基础示例代码
import jaxlib import jax from jax import numpy as jnp import flax.linen as nn from flax.training import train_state class NN(nn.Module): hid_dim : int # 隐藏层神经元数量 output_dim : int # 输出神经元数量 @nn.compact def __call__(self, x, t): out = jnp.hstack((x, t)) out = nn.tanh(nn.Dense(features=self.hid_dim)(out)) out = nn.tanh(nn.Dense(features=self.hid_dim)(out)) out = nn.Dense(features=self.output_dim)(out) return out d = 3 batch_size = 10 net = NN(hid_dim=100, output_dim=d) rng_nn, rng_inp1, rng_inp2 = jax.random.split(jax.random.PRNGKey(100), 3) inp_x = jax.random.normal(rng_inp1, (1, d)) # 单样本输入x,形状(1, d) inp_t = jax.random.normal(rng_inp2, (1, 1)) # 单样本输入t,形状(1, 1) params_net = net.init(rng_nn, inp_x, inp_t) x = jax.random.normal(rng_inp2, (batch_size, d)) # 批量输入x,形状(batch_size, d) t = jax.random.normal(rng_inp1, (batch_size, 1)) # 批量输入t,形状(batch_size, 1) out_net = net.apply(params_net, x, t) optimizer = optax.adam(1e-3) model_state = train_state.TrainState.create(apply_fn=net.apply, params= params_net, tx=optimizer)
用户初步尝试代码
def find_derivatives(net, params, X, t): d_dt = lambda net, params, X, t: jax.jvp(lambda time: net(params, X, t), (t, ), (jnp.ones_like(t), )) d_dx = lambda net, params, X, t: jax.jvp(lambda X: net(params, X, t), (X, ), (jnp.ones_like(X), )) out_f, df_dt = d_dt(net.apply, params, X, t) d_ddx = lambda net, params, X, t: d_dx(lambda params, X, t: d_dx(net, params, X, t)[1], params, X, t) df_dx, df_ddx = d_ddx(net.apply, params, X, t) return out_f, df_dt, df_dx, df_ddx out_f, df_dt, df_dx, df_ddx = find_derivatives(net, params_net, x, t)
正确实现方案
JAX中处理向量输出的导数,可通过jax.jacfwd/jax.jacrev计算Jacobian矩阵,再基于此推导所需的导数和散度梯度,以下是分步实现:
1. 计算输出对t的导数d out/dt
使用jax.jacfwd对t求导,提取批量中每个样本自身的导数:
def df_dt(params, x, t): # 定义仅对t求导的函数 f_t = lambda t_val: net.apply(params, x, t_val) # 计算Jacobian,形状(b, d, b, 1) jac = jax.jacfwd(f_t)(t) # 提取每个样本的导数,得到(b, d) return jac[jnp.arange(batch_size), :, jnp.arange(batch_size), 0]
2. 计算输出对x的导数d out/dx
输出对x的Jacobian为(b, d, b, d),提取每个样本对应的d×d Jacobian矩阵:
def df_dx(params, x, t): f_x = lambda x_val: net.apply(params, x_val, t) jac = jax.jacfwd(f_x)(x) # 提取每个样本的Jacobian矩阵,形状(b, d, d) return jac[jnp.arange(batch_size), :, jnp.arange(batch_size), :]
3. 计算输出散度对x的梯度
先计算输出的散度(Jacobian矩阵的迹),再对x求梯度:
def grad_div_f(params, x, t): def div_fn(x_val): # 计算每个样本的散度 jac = jax.jacfwd(lambda x: net.apply(params, x, t))(x_val) div = jnp.trace(jac[jnp.arange(x_val.shape[0]), :, jnp.arange(x_val.shape[0]), :], axis1=1, axis2=2) # 返回散度均值(grad要求函数输出为标量) return jnp.mean(div) # 对x求梯度,形状(b, d) return jax.grad(div_fn)(x)
整合所有导数计算
def compute_all_derivatives(params, x, t): f = net.apply(params, x, t) df_dt_val = df_dt(params, x, t) df_dx_val = df_dx(params, x, t) grad_div_val = grad_div_f(params, x, t) return f, df_dt_val, df_dx_val, grad_div_val # 调用示例 f, df_dt_val, df_dx_val, grad_div_val = compute_all_derivatives(params_net, x, t)
关键说明
jax.jacfwd适合输入维度与输出维度相近的场景,这里x是d维、输出也是d维,正向与反向模式效率相近;- 批量处理时需注意Jacobian的维度提取,确保仅保留每个样本自身的导数关联;
- 若需要每个样本单独的散度梯度,可结合
jax.vmap对单个样本求导后批量处理。
内容的提问来源于stack exchange,提问作者Michaela
相关产品推荐
相关产品推荐

