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

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模型在测试阶段仍可能存在未被控制的随机因素,具体原因包括:

  1. PPO动作采样的随机性:PPO默认会基于策略网络输出的概率分布采样动作(离散空间用Categorical分布、连续空间用Normal分布),这些采样过程完全依赖PyTorch的随机数生成器。如果不固定PyTorch种子,每次加载模型后的动作采样结果都会不同,直接导致测试轨迹和最终结果不一致。
  2. 依赖库的隐式随机操作:模型内部或相关工具函数可能调用numpy、Python标准库random的随机方法(比如某些特征预处理逻辑、策略的辅助计算步骤),这些操作的种子未被固定时,也会引入额外随机性。
  3. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 04:02:36