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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 09:10:40