developersajidbashir commited on
Commit
0c265a3
·
verified ·
1 Parent(s): fce63f5

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +73 -69
app.py CHANGED
@@ -1,18 +1,17 @@
1
- import gymnasium as gym
2
  import numpy as np
 
3
  import requests
4
  import pandas as pd
5
  from datetime import datetime, timedelta
6
  from stable_baselines3 import PPO
7
  from stable_baselines3.common.vec_env import DummyVecEnv
8
- from gymnasium import spaces
 
9
  import firebase_admin
10
  from firebase_admin import credentials, db
11
  import os
12
- import threading
13
- import time
14
 
15
- # Firebase initialization
16
  cred = credentials.Certificate("credentials.json")
17
  firebase_admin.initialize_app(cred, {"databaseURL": "https://socail-swap-default-rtdb.asia-southeast1.firebasedatabase.app/"})
18
  ref = db.reference()
@@ -21,8 +20,8 @@ buy_signals = []
21
  sell_signals = []
22
 
23
  class TradingEnv(gym.Env):
24
- def __init__(self, data, window_size=50):
25
- super(TradingEnv, self).__init__()
26
  self.data = data
27
  self.window_size = window_size
28
  self.current_step = window_size
@@ -30,11 +29,10 @@ 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, 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
@@ -54,9 +52,9 @@ class TradingEnv(gym.Env):
54
  elif action == 2:
55
  reward = self.data['close'].iloc[self.current_step - 1] - self.data['close'].iloc[self.current_step]
56
 
57
- return self._get_observation(), reward, done, {}, {}
58
-
59
- def fetch_data(symbol='ETH', tsym='USD', start_date='2021-01-01', api_key='YOUR_API_KEY'):
60
  start_date = datetime.strptime(start_date, '%Y-%m-%d')
61
  end_date = datetime.utcnow()
62
  to_ts = int(end_date.timestamp())
@@ -70,7 +68,10 @@ def fetch_data(symbol='ETH', tsym='USD', start_date='2021-01-01', api_key='YOUR_
70
  df = pd.DataFrame(data_points)
71
  df['time'] = pd.to_datetime(df['time'], unit='s')
72
  df.set_index('time', inplace=True)
 
 
73
  df = df[df.index >= start_date]
 
74
  return df[['close']]
75
  else:
76
  print(f"Error fetching data: {data['Message']}")
@@ -80,65 +81,68 @@ def calculate_ema(data, span=20):
80
  data['EMA'] = data['close'].ewm(span=span, adjust=False).mean()
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")
92
- return model
93
-
94
- def run_model(model, env, data):
95
- obs, _ = env.reset()
96
- dates = data.index[50:]
97
- prices = data['close'][50:]
98
- emas = data['EMA'][50:]
99
- actions = []
100
-
101
- for date, price, ema in zip(dates, prices, emas):
102
- action, _ = model.predict(obs)
103
- actions.append(action[0])
104
- obs, _, done, _, _ = env.step(action)
105
- if done:
106
- break
107
-
108
- new_buy_signals = [(date, price, ema) for date, price, ema, action in zip(dates, prices, emas, actions) if action == 1]
109
- new_sell_signals = [(date, price, ema) for date, price, ema, action in zip(dates, prices, emas, actions) if action == 2]
110
-
111
- for signal in new_buy_signals:
112
- if signal[0] not in [s[0] for s in buy_signals] and signal[0] not in [s[0] for s in sell_signals]:
113
- buy_signals.append(signal)
114
-
115
- for signal in new_sell_signals:
116
- if signal[0] not in [s[0] for s in sell_signals] and signal[0] not in [s[0] for s in buy_signals]:
117
- sell_signals.append(signal)
118
-
119
- buy_signals_data = [{'timestamp': signal[0].strftime('%Y-%m-%d %H:%M:%S'), 'type': 'b', 'price': round(signal[1], 2), 'ema': round(signal[2], 2)} for signal in buy_signals]
120
- sell_signals_data = [{'timestamp': signal[0].strftime('%Y-%m-%d %H:%M:%S'), 'type': 's', 'price': round(signal[1], 2), 'ema': round(signal[2], 2)} for signal in sell_signals]
121
-
122
- all_signals_data = buy_signals_data + sell_signals_data
123
-
124
- ref.child('signals').child('data').set(all_signals_data)
125
-
126
- def background_task():
127
  while True:
128
  try:
129
  new_data = fetch_data(start_date=(datetime.utcnow() - timedelta(days=3)).strftime('%Y-%m-%d'))
130
- if new_data is not None:
131
- new_data = calculate_ema(new_data)
132
- if len(new_data) >= 50:
133
- env = DummyVecEnv([lambda: TradingEnv(new_data)])
134
- model = load_or_train_model(env)
135
- run_model(model, env, new_data)
136
- time.sleep(3600)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
137
  except Exception as e:
138
  print(f"An error occurred: {e}")
139
  break
 
 
 
 
 
 
 
 
140
 
141
- if __name__ == "__main__":
142
- background_thread = threading.Thread(target=background_task)
143
- background_thread.start()
144
- background_thread.join()
 
 
 
 
1
+ import gym
2
  import numpy as np
3
+ import matplotlib.pyplot as plt
4
  import requests
5
  import pandas as pd
6
  from datetime import datetime, timedelta
7
  from stable_baselines3 import PPO
8
  from stable_baselines3.common.vec_env import DummyVecEnv
9
+ from gym import spaces
10
+ import time
11
  import firebase_admin
12
  from firebase_admin import credentials, db
13
  import os
 
 
14
 
 
15
  cred = credentials.Certificate("credentials.json")
16
  firebase_admin.initialize_app(cred, {"databaseURL": "https://socail-swap-default-rtdb.asia-southeast1.firebasedatabase.app/"})
17
  ref = db.reference()
 
20
  sell_signals = []
21
 
22
  class TradingEnv(gym.Env):
23
+ def _init_(self, data, window_size=50):
24
+ super(TradingEnv, self)._init_()
25
  self.data = data
26
  self.window_size = window_size
27
  self.current_step = window_size
 
29
  self.observation_space = spaces.Box(
30
  low=0, high=1, shape=(window_size, 2), dtype=np.float32)
31
 
32
+ def reset(self):
 
33
  self.current_step = self.window_size
34
+ return self._get_observation()
35
+
36
  def _get_observation(self):
37
  window_data = self.data[self.current_step-self.window_size:self.current_step]
38
  obs = window_data[['close', 'EMA']].values
 
52
  elif action == 2:
53
  reward = self.data['close'].iloc[self.current_step - 1] - self.data['close'].iloc[self.current_step]
54
 
55
+ return self._get_observation(), reward, done, {}
56
+
57
+ def fetch_data(symbol='ETH', tsym='USD', start_date='2021-01-01', api_key='66bc686cb714fadda1fad0320704c98869d4b31ce7d9d27560c6c574b4d04c54'):
58
  start_date = datetime.strptime(start_date, '%Y-%m-%d')
59
  end_date = datetime.utcnow()
60
  to_ts = int(end_date.timestamp())
 
68
  df = pd.DataFrame(data_points)
69
  df['time'] = pd.to_datetime(df['time'], unit='s')
70
  df.set_index('time', inplace=True)
71
+
72
+ # Filter data based on start_date
73
  df = df[df.index >= start_date]
74
+
75
  return df[['close']]
76
  else:
77
  print(f"Error fetching data: {data['Message']}")
 
81
  data['EMA'] = data['close'].ewm(span=span, adjust=False).mean()
82
  return data
83
 
84
+ def run_model():
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
85
  while True:
86
  try:
87
  new_data = fetch_data(start_date=(datetime.utcnow() - timedelta(days=3)).strftime('%Y-%m-%d'))
88
+ new_data = calculate_ema(new_data)
89
+
90
+ if len(new_data) < 50:
91
+ print("Not enough data to update the environment.")
92
+ time.sleep(3600)
93
+ continue
94
+
95
+ env = DummyVecEnv([lambda: TradingEnv(new_data)])
96
+ model.set_env(env)
97
+
98
+ obs = env.reset()
99
+ dates = new_data.index[50:]
100
+ prices = new_data['close'][50:]
101
+ emas = new_data['EMA'][50:]
102
+ actions = []
103
+
104
+ for date, price, ema in zip(dates, prices, emas):
105
+ action, _ = model.predict(obs)
106
+ actions.append(action[0])
107
+ obs, _, done, _ = env.step(action)
108
+ if done:
109
+ break
110
+
111
+ new_buy_signals = [(date, price, ema) for date, price, ema, action in zip(dates, prices, emas, actions) if action == 1 and date not in [signal[0] for signal in buy_signals]]
112
+ new_sell_signals = [(date, price, ema) for date, price, ema, action in zip(dates, prices, emas, actions) if action == 2 and date not in [signal[0] for signal in sell_signals]]
113
+
114
+ for signal in new_buy_signals:
115
+ if signal[0] not in [s[0] for s in buy_signals] and signal[0] not in [s[0] for s in sell_signals]:
116
+ buy_signals.append(signal)
117
+
118
+
119
+ for signal in new_sell_signals:
120
+ if signal[0] not in [s[0] for s in sell_signals] and signal[0] not in [s[0] for s in buy_signals]:
121
+ sell_signals.append(signal)
122
+
123
+ buy_signals_data = [{'timestamp': signal[0].strftime('%Y-%m-%d %H:%M:%S'), 'type': 'b', 'price': round(signal[1], 2), 'ema': round(signal[2],2)} for signal in buy_signals]
124
+ sell_signals_data = [{'timestamp': signal[0].strftime('%Y-%m-%d %H:%M:%S'), 'type': 's', 'price': round(signal[1], 2), 'ema': round(signal[2],2)} for signal in sell_signals]
125
+
126
+ all_signals_data = buy_signals_data + sell_signals_data
127
+
128
+ ref.child('signals').child('data').set(all_signals_data)
129
+ time.sleep(3600)
130
  except Exception as e:
131
  print(f"An error occurred: {e}")
132
  break
133
+
134
+ if _name_ == "_main_":
135
+ data = fetch_data()
136
+ data = calculate_ema(data)
137
+ if len(data) < 50:
138
+ raise ValueError("Not enough data to fill the window size.")
139
+
140
+ env = DummyVecEnv([lambda: TradingEnv(data)])
141
 
142
+ if os.path.exists("./ppo_trading_agent.zip"):
143
+ model = PPO.load("ppo_trading_agent", env=env)
144
+ else:
145
+ model = PPO('MlpPolicy', env, verbose=1)
146
+ model.learn(total_timesteps=10000)
147
+ print("Ender function")
148
+     run_model()