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

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)

关键说明

  1. 索引越界修复:原代码中uk的维度不匹配fwd_dynamics的索引需求,必须调整为(10,3,1)才能正常运行。
  2. 动态迭代适配:两种方案都实现了batch内每个元素使用不同的模拟步数,同时兼容vmap和jit的追踪机制。

内容的提问来源于stack exchange,提问作者Prajwal THAKUR

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 10:12:04