使用JAX的vmap时遇到PyTree相关错误,求排查解决
问题分析与解决
错误的核心是vmap的输入参数结构和in_axes指定不匹配:
grad(loss)返回的函数需要两个参数:params和batch- 你指定的
in_axes=(None, 0)表示第一个参数(params)不做批量映射,第二个参数(batch)沿轴0做批量映射 - 但调用时只传了
batch一个参数,导致JAX无法匹配参数树结构,抛出错误
另外,原代码中batch的构造也有问题:gt_inputs是(1,20)形状,转置+squeeze后得到的结构不符合loss函数对batch的预期(loss期望batch是(inputs, targets)的二元组)。
修正后的代码
import jax import jax.numpy as jnp import numpy as np from jax import grad, vmap def predict(params, input): y = jnp.sin(params * input) return y def loss(params, batch): inputs, targets = batch predictions = predict(params, inputs) return jnp.sum((predictions - targets)**2) # 生成数据:让inputs和targets都是(20,)形状的数组 gt_inputs = np.random.random(20) gt_targets = jnp.sin(2.3 * gt_inputs) # 构造batch:每个元素是(input, target),形状为(20, 2) batch = jnp.stack([gt_inputs, gt_targets], axis=1) param_init = 0.2 # 正确调用vmap:传入两个参数,对应grad(loss)的两个输入 grads = vmap(grad(loss), in_axes=(None, 0))(param_init, batch)
关键修正点
- 参数传递:调用vmap包裹的函数时,必须传入
grad(loss)所需的全部参数(param_init和batch),对应in_axes的两个指定项 - batch结构:调整数据生成方式,让
batch的每个元素是单个(input, target)对,确保loss函数能正确解包 - 简化代码:移除了未用到的导入(如
random、hessian等),让代码更简洁
内容的提问来源于stack exchange,提问作者q than a
相关产品推荐
相关产品推荐

