developersajidbashir commited on
Commit
e9db3fa
·
verified ·
1 Parent(s): 84a35b6

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +57 -64
app.py CHANGED
@@ -1,17 +1,18 @@
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()
@@ -53,8 +54,8 @@ class TradingEnv(gym.Env):
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,10 +69,7 @@ def fetch_data(symbol='ETH', tsym='USD', start_date='2021-01-01', api_key='66bc6
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,68 +79,63 @@ def calculate_ema(data, span=20):
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()
 
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()
 
54
  reward = self.data['close'].iloc[self.current_step - 1] - self.data['close'].iloc[self.current_step]
55
 
56
  return self._get_observation(), reward, done, {}
57
+
58
+ def fetch_data(symbol='ETH', tsym='USD', start_date='2021-01-01', api_key='YOUR_API_KEY'):
59
  start_date = datetime.strptime(start_date, '%Y-%m-%d')
60
  end_date = datetime.utcnow()
61
  to_ts = int(end_date.timestamp())
 
69
  df = pd.DataFrame(data_points)
70
  df['time'] = pd.to_datetime(df['time'], unit='s')
71
  df.set_index('time', inplace=True)
 
 
72
  df = df[df.index >= start_date]
 
73
  return df[['close']]
74
  else:
75
  print(f"Error fetching data: {data['Message']}")
 
79
  data['EMA'] = data['close'].ewm(span=span, adjust=False).mean()
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")
89
+ return model
90
+
91
+ def run_model(model, env, data):
92
+ obs = env.reset()
93
+ dates = data.index[50:]
94
+ prices = data['close'][50:]
95
+ emas = data['EMA'][50:]
96
+ actions = []
 
 
 
 
97
 
98
+ for date, price, ema in zip(dates, prices, emas):
99
+ action, _ = model.predict(obs)
100
+ actions.append(action[0])
101
+ obs, _, done, _ = env.step(action)
102
+ if done:
103
+ break
104
 
105
+ new_buy_signals = [(date, price, ema) for date, price, ema, action in zip(dates, prices, emas, actions) if action == 1]
106
+ new_sell_signals = [(date, price, ema) for date, price, ema, action in zip(dates, prices, emas, actions) if action == 2]
107
 
108
+ for signal in new_buy_signals:
109
+ if signal[0] not in [s[0] for s in buy_signals] and signal[0] not in [s[0] for s in sell_signals]:
110
+ buy_signals.append(signal)
111
+
112
+ for signal in new_sell_signals:
113
+ if signal[0] not in [s[0] for s in sell_signals] and signal[0] not in [s[0] for s in buy_signals]:
114
+ sell_signals.append(signal)
115
+
116
+ 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]
117
+ 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]
118
 
119
+ all_signals_data = buy_signals_data + sell_signals_data
 
 
 
 
 
120
 
121
+ ref.child('signals').child('data').set(all_signals_data)
122
 
123
+ def background_task():
124
+ while True:
125
+ try:
126
+ new_data = fetch_data(start_date=(datetime.utcnow() - timedelta(days=3)).strftime('%Y-%m-%d'))
127
+ if new_data is not None:
128
+ new_data = calculate_ema(new_data)
129
+ if len(new_data) >= 50:
130
+ env = DummyVecEnv([lambda: TradingEnv(new_data)])
131
+ model = load_or_train_model(env)
132
+ run_model(model, env, new_data)
133
+ time.sleep(3600)
134
  except Exception as e:
135
  print(f"An error occurred: {e}")
136
  break
 
 
 
 
 
 
137
 
138
+ if __name__ == "__main__":
139
+ background_thread = threading.Thread(target=background_task)
140
+ background_thread.start()
141
+ background_thread.join()