在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
相关产品推荐
相关产品推荐

