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

Jax中随机种子的正确使用:DQN环境步骤两种方式差异解析

关于Jax DQN中随机种子生成方式的差异分析与规范说明

问题背景

在基于Jax的DQN强化学习任务中,针对环境步骤采用两种随机种子生成方式时,结果差异显著,需明确差异原因及符合Jax随机种子规范的使用方式。

两种种子生成方式

方式一(Jax官方文档推荐)

rng, step_rng = jax.random.split(rng)
next_state, next_env_state, reward, terminated, info = env.step(step_rng, env_state, action.squeeze(), env_params)

方式二(purejaxrl示例实现)

rng, _rng = jax.random.split(rng)
_rng, step_rng = jax.random.split(_rng)
next_state, next_env_state, reward, terminated, info = env.step(step_rng, env_state, action.squeeze(), env_params)

补充代码片段

环境步骤函数中子环境代码

state, env_state, _, _, rng = runner
q = self.agent_nn.apply(self.agent_params, state)
action = jnp.argmax(q)
rng, rng_step = jax.random.split(rng)
next_state, next_env_state, reward, terminated, info = self.env.step(rng_step, env_state, action, self.env_params)

lax.scan内的随机种子使用代码

def train(rng):

    rng, network_init_rng = jax.random.split(rng)
    network = q_network(env.action_space(env_params).n)
    init_x = jnp.zeros((1, config["STATE_SIZE"]))
    network_params = network.init(network_init_rng, init_x)
    
    training = TrainState.create(apply_fn=network.apply,
                                 params=network_params,
                                 target_params=network_params,
                                 tx=tx)

    rng, _rng = jax.random.split(rng)
    _rng, reset_rng = jax.random.split(_rng)
    state, env_state = env.reset(reset_rng, env_params)

    @jit
    @scan_tqdm(config["TOTAL_STEPS"])
    def _run_step(runner, i_step):

        training, env_state, state, rng, buffer_state, i_episode = runner

        rng, *_rng = jax.random.split(rng, 3)
        random_q_rng, random_number_rng = _rng

        q_state = network.apply(training.params, state)
        random_number = jax.random.uniform(random_number_rng, minval=0, maxval=1, shape=(1,))
        exploitation = jnp.greater(random_number, config["EPS"])
        action = jnp.where(exploitation, jnp.argmax(q_state, 1), random_action)

        rng, _rng = jax.random.split(rng)
        _rng, step_rng = jax.random.split(_rng)
        next_state, next_env_state, reward, terminated, info = env.step(step_rng, env_state, action.squeeze(), env_params)
        
        return runner

    rng, _rng = jax.random.split(rng)
    runner = (training, env_state, state, _rng, buffer_state, 0)
    runner, metrics = lax.scan(_run_step, runner, jnp.arange(config["TOTAL_STEPS"]), config["TOTAL_STEPS"])

    return {"runner": runner}

差异原因分析

Jax的随机数生成依赖确定性的RNG链,每次jax.random.split都会将输入RNG拆分为两个独立的子RNG,后续使用子RNG生成的随机序列完全独立。

两种方式的核心差异在于:

  • 方式一直接将主RNG拆分为rng(新主RNG)和step_rng(环境步骤用RNG),step_rng是主RNG的直接子节点。
  • 方式二先拆分出_rng,再从_rng拆分出step_rng,相当于step_rng是主RNG的孙节点。

由于Jax的RNG拆分是确定性的,两种方式生成的step_rng对应的随机序列完全不同,这会直接影响环境步骤中的随机事件(比如随机奖励、状态转移),最终导致DQN训练结果出现显著差异。

正确使用规范

Jax官方推荐的方式一是更简洁、符合规范的用法,原因如下:

  • 遵循Jax RNG的核心设计原则:每次使用随机数前,通过split从主RNG派生专用子RNG,同时更新主RNG以保证后续RNG链的独立性。
  • 方式二的两次拆分完全冗余,没有额外收益——split一次就足够生成独立的子RNG,多次拆分只会增加不必要的计算,且容易引入RNG管理的混乱。

结合你提供的lax.scan代码来看,当前在_run_step中使用方式二生成step_rng是冗余的,建议修改为官方推荐的方式:

# 替换原方式二的代码
rng, step_rng = jax.random.split(rng)
next_state, next_env_state, reward, terminated, info = env.step(step_rng, env_state, action.squeeze(), env_params)

另外需要注意:在lax.scan这类循环结构中,必须将rng作为runner的一部分传递,确保每次迭代都更新RNG链——你当前的代码已经做到了这一点,这是正确的做法,避免了JAX的"重复使用RNG导致相同随机序列"的问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 12:04:53