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

使用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

相关产品推荐
方舟 Agent Plan

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

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