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
相关产品推荐
相关产品推荐

