developersajidbashir commited on
Commit
21f0d1d
·
verified ·
1 Parent(s): 126b84c

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +9 -6
app.py CHANGED
@@ -30,10 +30,11 @@ class TradingEnv(gym.Env):
30
  self.observation_space = spaces.Box(
31
  low=0, high=1, shape=(window_size, 2), dtype=np.float32)
32
 
33
- def reset(self):
 
34
  self.current_step = self.window_size
35
- return self._get_observation()
36
-
37
  def _get_observation(self):
38
  window_data = self.data[self.current_step-self.window_size:self.current_step]
39
  obs = window_data[['close', 'EMA']].values
@@ -80,9 +81,11 @@ def calculate_ema(data, span=20):
80
  return data
81
 
82
  def load_or_train_model(env):
83
- if os.path.exists("./ppo_trading_agent.zip"):
84
- model = PPO.load("ppo_trading_agent", env=env)
85
- else:
 
 
86
  model = PPO('MlpPolicy', env, verbose=1)
87
  model.learn(total_timesteps=10000)
88
  model.save("ppo_trading_agent")
 
30
  self.observation_space = spaces.Box(
31
  low=0, high=1, shape=(window_size, 2), dtype=np.float32)
32
 
33
+ def reset(self, seed=None, options=None):
34
+ super().reset(seed=seed)
35
  self.current_step = self.window_size
36
+ return self._get_observation(), {}
37
+
38
  def _get_observation(self):
39
  window_data = self.data[self.current_step-self.window_size:self.current_step]
40
  obs = window_data[['close', 'EMA']].values
 
81
  return data
82
 
83
  def load_or_train_model(env):
84
+ custom_objects = {"clip_range": 0.2, "lr_schedule": 0.0003} # Update with your actual values or leave as is
85
+ try:
86
+ model = PPO.load("ppo_trading_agent", env=env, custom_objects=custom_objects)
87
+ except Exception as e:
88
+ print(f"Loading model failed: {e}. Training a new model.")
89
  model = PPO('MlpPolicy', env, verbose=1)
90
  model.learn(total_timesteps=10000)
91
  model.save("ppo_trading_agent")