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

JAX中嵌套vmap至pmap时遭遇追踪错误的技术求助

问题分析与解决方案

错误原因

你遇到的追踪错误核心在于:pmap在追踪run_trajectory时,timings参数的具体值无法在Python层面被获取——大概率是run_trajectory里用了timings的属性(比如timings.step_count这类Python数值)做原生Python控制流(比如for循环次数、if条件判断),而JAX的追踪机制无法处理依赖追踪值的Python逻辑。

另外,嵌套pmap(vmap(...))的写法本身没问题,但如果输入轴划分和函数逻辑不匹配,也会放大这类追踪问题。


解决步骤

1. 重构run_trajectory中的Python控制流

如果run_trajectory里有依赖timings具体值的Python循环/条件,必须替换为JAX原生的可追踪操作:

  • 把Python for循环改成jax.lax.fori_loop
  • 把Python if/else改成jax.lax.cond

示例重构:
原错误写法(依赖Python数值的循环):

def run_trajectory(sim_state, timings, lambda_val):
    # timings.num_steps是Python整数,追踪时无法获取具体值
    for _ in range(timings.num_steps):
        sim_state = update_sim(sim_state, lambda_val)
    return sim_state

修改为JAX可追踪写法:

import jax

def run_trajectory(sim_state, timings, lambda_val):
    def step_fn(step_idx, current_state):
        return update_sim(current_state, lambda_val)
    
    # 用jax.lax.fori_loop替代Python循环,timings.num_steps可传入JAX数组
    final_state = jax.lax.fori_loop(0, timings.num_steps, step_fn, sim_state)
    return final_state

2. 确保timings是JAX兼容类型

如果timings是自定义Python类,改成用JAX数组或PyTree存储(比如把timings.num_steps转成jax.numpy.array),避免传递无法被JAX追踪的原生Python对象。

3. 调整并行嵌套逻辑(可选但更规范)

外层用pmap做设备间并行,内层用vmap做单设备内的批量并行,确保输入轴划分清晰:
假设总仿真数N = 设备数D × 单设备批量数M,先把输入reshape为(D, M, ...)的形状:

# 先reshape输入,拆分设备轴和单设备批量轴
D = jax.device_count()  # 获取可用GPU数量
M = sim_state.shape[0] // D  # 单设备处理的仿真数

reshaped_sim_state = sim_state.reshape((D, M) + sim_state.shape[1:])
reshaped_lambda_array = lambda_array.reshape((D, M) + lambda_array.shape[1:])

# 先定义单设备批量运行的vmap函数,再用pmap分发到多设备
batch_run = jax.vmap(run_trajectory, in_axes=(0, None, 0))
traj_state = jax.pmap(batch_run, in_axes=(0, None, 0))(reshaped_sim_state, timings, reshaped_lambda_array)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 21:32:59