SBX训练RL算法跨设备结果不一致问题及代码排查请求
SBX强化学习跨设备训练结果不一致问题排查
问题背景
使用SBX训练强化学习算法时,发现HPC集群与个人电脑(PC)上的训练结果不一致,但同一设备内多次运行结果始终稳定。目前已将所有计算限制在CPU上运行,编写了最小可复现代码验证,需排查代码中是否存在随机种子遗漏等问题。
测试代码
# simple_sbx_test.py import jax import numpy as np import random import os import gymnasium as gym from sbx import DQN from stable_baselines3.common.callbacks import EvalCallback from stable_baselines3.common.vec_env import DummyVecEnv def set_seed(seed): """Set seed for reproducibility.""" os.environ['PYTHONHASHSEED'] = str(seed) random.seed(seed) np.random.seed(seed) def make_env(env_name, seed): """Create environment with fixed seed""" def _init(): env = gym.make(env_name) env.reset(seed=seed) return env return _init def main(): # Fixed seeds AGENT_SEED = 42 ENV_SEED = 123 EVAL_SEED = 456 set_seed(AGENT_SEED) print("=== Simple SBX DQN Cross-Platform Test (JAX) ===") print(f"JAX: {jax.__version__}") print(f"NumPy: {np.__version__}") print(f"JAX devices: {jax.devices()}") print(f"Agent seed: {AGENT_SEED}, Env seed: {ENV_SEED}, Eval seed: {EVAL_SEED}") print("-" * 50) # Create environments train_env = DummyVecEnv([make_env("CartPole-v1", ENV_SEED)]) eval_env = DummyVecEnv([make_env("CartPole-v1", EVAL_SEED)]) # Create model model = DQN( "MlpPolicy", train_env, learning_rate=1e-3, buffer_size=10000, learning_starts=1000, batch_size=32, gamma=0.99, train_freq=4, target_update_interval=1000, exploration_initial_eps=1.0, exploration_final_eps=0.05, exploration_fraction=0.1, verbose=0, seed=AGENT_SEED ) # Print initial model parameters (JAX uses params instead of weights) if hasattr(model, 'qf') and hasattr(model.qf, 'params'): print("Initial parameters available") # JAX parameters are nested dictionaries, harder to inspect directly print(" Model initialized successfully") # Evaluation callback eval_callback = EvalCallback( eval_env, best_model_save_path=None, log_path=None, eval_freq=2000, n_eval_episodes=10, deterministic=True, render=False, verbose=1 # Enable to see evaluation results ) # Train print("\nTraining...") model.learn(total_timesteps=10000, callback=eval_callback) print("Training completed") # Final evaluation print("\nFinal evaluation:") rewards = [] for i in range(10): obs = eval_env.reset() total_reward = 0 done = False while not done: action, _ = model.predict(obs, deterministic=True) obs, reward, done, info = eval_env.step(action) total_reward += reward[0] rewards.append(total_reward) print(f"Episode {i + 1}: {total_reward}") print(f"\nFinal Results:") print(f"Mean reward: {np.mean(rewards):.2f}") print(f"Std reward: {np.std(rewards):.2f}") print(f"All rewards: {rewards}") if __name__ == "__main__": main()
PC端运行结果
Final evaluation: Episode 1: 208.0 Episode 2: 237.0 Episode 3: 200.0 Episode 4: 242.0 Episode 5: 206.0 Episode 6: 334.0 Episode 7: 278.0 Episode 8: 235.0 Episode 9: 248.0 Episode 10: 206.0
HPC集群端运行结果
Final evaluation: Episode 1: 201.0 Episode 2: 256.0 Episode 3: 193.0 Episode 4: 218.0 Episode 5: 192.0 Episode 6: 326.0 Episode 7: 239.0 Episode 8: 226.0 Episode 9: 237.0 Episode 10: 201.0
排查建议
补充JAX随机种子设置:SBX基于JAX实现,当前
set_seed函数未设置JAX的全局种子,需添加:jax.random.PRNGKey(seed) jax.config.update('jax_platform_name', 'cpu') # 确保强制使用CPU jax.config.update('jax_deterministic', True) # 强制开启确定性计算强化环境种子设置:当前仅调用
env.reset(seed=seed),需同时设置动作空间和观测空间的种子,修改make_env函数:def make_env(env_name, seed): def _init(): env = gym.make(env_name) env.reset(seed=seed) env.action_space.seed(seed) env.observation_space.seed(seed) return env return _init同步依赖版本:跨设备结果不一致可能源于JAX、SBX、Gymnasium等依赖包版本差异,需确保PC和HPC上所有依赖版本完全相同,可通过
pip freeze导出依赖列表并同步。验证向量环境种子传递:尝试给DummyVecEnv手动设置种子:
train_env.seed(ENV_SEED) eval_env.seed(EVAL_SEED)统一浮点数精度:不同CPU架构可能存在浮点数计算微小差异,可在JAX中强制使用单精度浮点数,或开启确定性模式减少此类差异。
内容的提问来源于stack exchange,提问作者desert_ranger
相关产品推荐
相关产品推荐

