Stable Baselines 3:加载PPO模型时设置随机种子的原因探究
关于Stable Baselines3 PPO模型复现性问题的分析
问题背景
我使用Stable Baselines3框架下的PPO模型开展日内VWAP预测的强化学习研究,当前面临模型复现性不佳的问题:即使已固定环境reset方法、观测空间及动作空间的种子,确保环境无随机性,加载训练完成的模型进行测试时,每次得到的结果仍存在差异。
已找到两种解决思路:
- 思路一:加载模型时为torch、numpy等所有涉及随机操作的模块设置随机种子
random_seed=42 # Set seed for reproducibility torch.manual_seed(random_seed) torch.cuda.manual_seed(random_seed) torch.cuda.manual_seed_all(random_seed) # if use multi-GPU torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False np.random.seed(random_seed) random.seed(random_seed)
- 思路二:调用
model.load()方法时指定seed参数
# Load Model model.load("./logs/ppo_vwap_predict_20240919_20240111.zip", env=env, seed=random_seed)
我理解第二种方法的原理,但无法理解为何加载模型时需要为各模块设置随机种子,怀疑这与PPO模型内部的探索机制及动作选择过程(如使用np.choice)有关,恳请相关分析与解答。
问题分析与解答
为什么需要为全局模块设置随机种子?
即使环境的随机性被完全固定,PPO模型在测试阶段仍可能存在未被控制的随机因素,具体原因包括:
- PPO动作采样的随机性:PPO默认会基于策略网络输出的概率分布采样动作(离散空间用
Categorical分布、连续空间用Normal分布),这些采样过程完全依赖PyTorch的随机数生成器。如果不固定PyTorch种子,每次加载模型后的动作采样结果都会不同,直接导致测试轨迹和最终结果不一致。 - 依赖库的隐式随机操作:模型内部或相关工具函数可能调用numpy、Python标准库
random的随机方法(比如某些特征预处理逻辑、策略的辅助计算步骤),这些操作的种子未被固定时,也会引入额外随机性。 - CuDNN的非确定性优化:即使使用CPU推理,CuDNN默认设置可能会选择非确定性算法来提升运算性能(比如卷积操作的算法自动选择)。通过
torch.backends.cudnn.deterministic = True和torch.backends.cudnn.benchmark = False可以强制CuDNN使用确定性算法,消除这部分潜在的随机性。
两种解决思路的互补性
model.load()指定seed:主要负责初始化模型内部的随机生成器(比如与策略探索相关的种子状态),确保模型自身的核心随机环节一致。- 全局模块设置种子:覆盖所有模型外部、依赖库层面的随机操作,填补
model.load()可能未覆盖到的随机源,确保整个测试流程的所有随机环节都被完全固定。
完整代码
import json import datetime import pandas as pd import numpy as np from numpy.random import SeedSequence, default_rng import random import gym import talib as ta import torch from env.ExpectVolumeEnv import ExpectVolumeEnv from env.ExpectVolumeEnvDiscrete import ExpectVolumeEnvDiscrete from stable_baselines3 import PPO from stable_baselines3 import DQN from stable_baselines3.common.callbacks import StopTrainingOnNoModelImprovement, StopTrainingOnRewardThreshold, EvalCallback from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.callbacks import CheckpointCallback from stable_baselines3 import HerReplayBuffer from stable_baselines3.her.goal_selection_strategy import GoalSelectionStrategy import matplotlib.pyplot as plt ''' reference https://github.com/notadamking/Stock-Trading-Environment ''' ''' Data 20XX-XX-XX KOSPI Intraday Data ''' random_seed = 42 # # Set seed for reproducibility # torch.manual_seed(random_seed) # torch.cuda.manual_seed(random_seed) # torch.cuda.manual_seed_all(random_seed) # if use multi-GPU # torch.backends.cudnn.deterministic = True # torch.backends.cudnn.benchmark = False # np.random.seed(random_seed) # random.seed(random_seed) # Load data df = pd.read_csv("data/raw/kospi_minutes/[지수KOSPI계열]일중 시세정보(1분)(주문번호-2499-1)_20240111.csv", encoding='cp949') # DataFrame Preprocessing df = df[df['지수명']=='코스피'] df = df[df['거래시각'] <= '1530'] data_date = str(df['거래일자'].iloc[0]) df = df[['거래시각', '시가', '고가', '저가', '종가', '거래량']] df.columns = ['Time', 'Open', 'High', 'Low', 'Close', 'Volume'] print(df) df = df.astype(float) df = df.reset_index(drop=False) # # Create environment env = ExpectVolumeEnv(df, seed=random_seed) env.action_space.seed(random_seed) env.observation_space.seed(random_seed) # Create model (PPO) model = PPO("MlpPolicy", env, learning_rate=0.00025, batch_size=128, verbose=1, ) # print(help(model.load)) # # Total timesteps / Number of steps per episode = Number of episodes # model.learn(total_timesteps=len(df)*100) # # # Save model # model.save(f"./logs/ppo_vwap_predict_{datetime.datetime.now().strftime('%Y%m%d')}_{data_date}.zip") # Load Model model.load("./logs/ppo_vwap_predict_20240919_20240111.zip", env=env, seed=random_seed) # observation, empty = env.reset(seed=random_seed) observation, empty = env.reset() print("mean: ", df['Close'].mean()) plt.plot(df['Volume'], label=f'{data_date} Market Volume') plt.show() plt.plot(df['Close'], label=f'{data_date} Market Close') plt.show() # Render each environment separately for _ in range(len(df)-1): action, _states = model.predict(observation) observation, reward, terminated, truncated, info = env.step(action) env.render() market_vwap = env.render_plot(data_date=data_date) volume_pattern = pd.read_csv('./data/volume.csv') scaled_mean = volume_pattern['scaled_mean'] proportion = scaled_mean / np.sum(scaled_mean) static_model_vwap = np.sum(df['Close'] * proportion) print(f"Static Model VWAP: {static_model_vwap}") print(f"Static Model VWAP Gap: {market_vwap - static_model_vwap}")
内容的提问来源于stack exchange,提问作者Psyduck_20
相关产品推荐
相关产品推荐

