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

在JAX中获取雅可比矩阵的预期维度问题

问题分析与解决方法

你的核心问题在于**jacobian和vmap的调用顺序错误**,以及对函数输入输出维度的匹配逻辑理解偏差,以下是具体拆解和解决方案:

维度逻辑梳理

先明确你的函数行为:

  • 当输入单个参数向量v_params(形状(2,))时,reparameterize依赖全局的eps(形状(3,)),输出是形状(3,)的theta。此时单个参数组的雅可比矩阵形状应为(3, 2)——输出的每个维度对输入的每个维度求导。
  • 当用vmap(reparameterize)处理形状(3,2)的v_params时,相当于对v_params的每一行(共3组参数)单独应用函数,最终输出形状是(3, 3)(3组输出,每组是(3,)的向量)。

此时调用jacobian(vmap(reparameterize))(v_params),计算的是整个批量输出对整个批量输入的全量雅可比,JAX会返回形状(3,3,3,2)——对应输出的每个批量元素的每个维度,对输入的每个批量元素的每个维度的导数。这显然不是你需要的“每个参数组对应自身的雅可比”。

正确解决方案

1. 调换vmap与jacobian的顺序

你需要先对单个参数组计算雅可比,再用vmap批量应用到所有参数行上:

jax.vmap(jax.jacobian(reparameterize))(v_params)

该代码返回形状为(3, 3, 2)的结果——对应3组参数,每组参数的雅可比矩阵是(3,2),完全符合“遍历所有eps并堆叠雅可比”的手动实现逻辑。

2. 调整函数设计(若需匹配单元素输出的雅可比)

如果你期望最终得到形状(3,2)的结果,说明你可能希望每个参数组对应一个eps元素(而非全局共享(3,)的eps),可以修改函数让eps与参数维度对齐:

def reparameterize(v_params, eps):
    theta = v_params[0] + jnp.exp(v_params[1]) * eps
    return theta

# 将eps调整为(3,),与v_params的批量维度匹配
key, _ = random.split(key)
eps = random.normal(key, shape=(3,))
# 用vmap同时映射参数和eps,仅对参数求雅可比
jax.vmap(jax.jacobian(reparameterize, argnums=0))(v_params, eps)

此时返回形状为(3,1,2),可通过jnp.squeeze()压缩为(3,2),完全符合你最初的期望。

小维度场景的特殊情况解释

当eps是(1,)、v_params是(2,)时,reparameterize的输出是(1,),JAX会自动挤压雅可比矩阵的单维度,所以jacobian(reparameterize)(v_params)返回(2,)(本质是(1,2)被压缩后的结果)。但批量场景下JAX不会自动压缩维度,因此你看到了完整的多维度输出。

内容的提问来源于stack exchange,提问作者hasco641

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 22:25:25