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

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

排查建议

  1. 补充JAX随机种子设置:SBX基于JAX实现,当前set_seed函数未设置JAX的全局种子,需添加:

    jax.random.PRNGKey(seed)
    jax.config.update('jax_platform_name', 'cpu')  # 确保强制使用CPU
    jax.config.update('jax_deterministic', True)  # 强制开启确定性计算
    
  2. 强化环境种子设置:当前仅调用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
    
  3. 同步依赖版本:跨设备结果不一致可能源于JAX、SBX、Gymnasium等依赖包版本差异,需确保PC和HPC上所有依赖版本完全相同,可通过pip freeze导出依赖列表并同步。

  4. 验证向量环境种子传递:尝试给DummyVecEnv手动设置种子:

    train_env.seed(ENV_SEED)
    eval_env.seed(EVAL_SEED)
    
  5. 统一浮点数精度:不同CPU架构可能存在浮点数计算微小差异,可在JAX中强制使用单精度浮点数,或开启确定性模式减少此类差异。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 10:44:56