Jax中vmap结合lax.scan遇batch维度序列长度不一致问题求解
解决JAX中vmap与动态步数lax.scan冲突的方案
你的问题核心是:lax.scan要求迭代次数是编译时已知的具体值,但vmap后sim_timestep变为追踪数组(tracedArray),无法满足scan的要求。以下是两种可行的替代方案,同时修复你代码中uk索引越界的bug(原代码中uk形状为(10,2,1),但访问了uk[2][0],需调整为(10,3,1)):
方案一:用jax.lax.while_loop替代lax.scan
while_loop支持动态终止条件,不需要编译时确定迭代次数,完美适配batch内不同步数的场景:
from jax import random from jax import lax import jax import jax.numpy as jnp def fwd_dynamics(x_u): # 移除未使用的xs参数 x0, uk = x_u Delta_T = 0.001 lwb = 1.2 psi0 = x0[2][0] v0 = x0[3][0] vdot0 = uk[0][0] delta0 = uk[1][0] thetadot0 = uk[2][0] xdot = jnp.asarray([ [v0 * jnp.cos(psi0)], [v0 * jnp.sin(psi0)], [v0 * jnp.tan(delta0) / lwb], [vdot0], [thetadot0] ]) x_next = x0 + xdot * Delta_T return (x_next, uk) def state_predictor(xk, uk, sim_timestep): sim_steps = jnp.squeeze(sim_timestep) # 将(1,)形状转为标量 # 循环终止条件:当前步数 < 总步数 def cond_fn(carry): x, u, step = carry return step < sim_steps # 循环体:执行一步动力学更新 def body_fn(carry): x, u, step = carry x_next, u_next = fwd_dynamics((x, u)) return (x_next, u_next, step + 1) # 初始状态:初始x、初始u、当前步数0 initial_carry = (xk, uk, 0) final_carry = lax.while_loop(cond_fn, body_fn, initial_carry) return final_carry[0] low = 0 high = 100 key = jax.random.PRNGKey(44) sim_time = jax.random.randint(key, shape=(10, 1), minval=low, maxval=high) xk = jax.random.uniform(key, shape=(10, 5, 1)) uk = jax.random.uniform(key, shape=(10, 3, 1)) # 修复为3维,适配fwd_dynamics的索引 state_predictor_vmap = jax.jit(jax.vmap(state_predictor, in_axes=0, out_axes=0)) x_next = state_predictor_vmap(xk, uk, sim_time) print(x_next.shape) # 输出 (10, 5, 1)
方案二:利用JAX动态形状支持,保留lax.scan
如果你更倾向于使用scan,可以开启JAX的动态形状支持(需JAX版本≥0.4.13),通过jax.jit的dynamic参数允许动态的迭代次数:
from jax import random from jax import lax import jax import jax.numpy as jnp def fwd_dynamics(x_u, _): # 保留原函数结构,忽略未使用的第二个参数 x0, uk = x_u Delta_T = 0.001 lwb = 1.2 psi0 = x0[2][0] v0 = x0[3][0] vdot0 = uk[0][0] delta0 = uk[1][0] thetadot0 = uk[2][0] xdot = jnp.asarray([ [v0 * jnp.cos(psi0)], [v0 * jnp.sin(psi0)], [v0 * jnp.tan(delta0) / lwb], [vdot0], [thetadot0] ]) x_next = x0 + xdot * Delta_T return (x_next, uk), x_next def state_predictor(xk, uk, sim_timestep): sim_steps = jnp.squeeze(sim_timestep) # 使用动态形状的scan,迭代次数由sim_steps决定 (x_next, _), _ = lax.scan(fwd_dynamics, (xk, uk), None, length=sim_steps) return x_next low = 0 high = 100 key = jax.random.PRNGKey(44) sim_time = jax.random.randint(key, shape=(10, 1), minval=low, maxval=high) xk = jax.random.uniform(key, shape=(10, 5, 1)) uk = jax.random.uniform(key, shape=(10, 3, 1)) # 修复索引越界问题 # 开启dynamic=True允许动态长度的scan state_predictor_vmap = jax.jit(jax.vmap(state_predictor, in_axes=0, out_axes=0), dynamic=True) x_next = state_predictor_vmap(xk, uk, sim_time) print(x_next.shape) # 输出 (10, 5, 1)
关键说明
- 索引越界修复:原代码中
uk的维度不匹配fwd_dynamics的索引需求,必须调整为(10,3,1)才能正常运行。 - 动态迭代适配:两种方案都实现了batch内每个元素使用不同的模拟步数,同时兼容vmap和jit的追踪机制。
内容的提问来源于stack exchange,提问作者Prajwal THAKUR
相关产品推荐
相关产品推荐

