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

如何基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 07:27:07