Team Ai
Apppublic

developersajidbashir/socialmodel

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py160 linesDownload Raw Back to root
1import threading2import gym3import numpy as np4import requests5import pandas as pd6from datetime import datetime, timedelta7from stable_baselines3 import PPO8from stable_baselines3.common.vec_env import DummyVecEnv9from gym import spaces10import time11import firebase_admin12from firebase_admin import credentials, db13import os14import gradio as gr15 16cred = credentials.Certificate("credentials.json")17firebase_admin.initialize_app(cred, {"databaseURL": "https://socail-swap-default-rtdb.asia-southeast1.firebasedatabase.app/"})18ref = db.reference()19buy_signals = []20sell_signals = []21hold_signals = []22 23class TradingEnv(gym.Env):24    def __init__(self, data, window_size=50):25        super(TradingEnv, self).__init__()26        super(TradingEnv, self).__init__()27        self.data = data28        self.window_size = window_size29        self.current_step = window_size30        self.action_space = spaces.Discrete(3)31        self.observation_space = spaces.Box(32            low=0, high=1, shape=(window_size, 2), dtype=np.float32)33        34    def reset(self):35        self.current_step = self.window_size36        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']].values41        obs = (obs - obs.min()) / (obs.max() - obs.min())42        return obs43    44    def step(self, action):45        reward = 046        done = False47        self.current_step += 148 49        if self.current_step >= len(self.data):50            done = True51        else:52            if action == 1:  53                reward = self.data['close'].iloc[self.current_step] - self.data['close'].iloc[self.current_step - 1]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    59def fetch_data(symbol='ETH', tsym='USD', start_date='2021-01-01', api_key='66bc686cb714fadda1fad0320704c98869d4b31ce7d9d27560c6c574b4d04c54'):60    start_date = datetime.strptime(start_date, '%Y-%m-%d')61    end_date = datetime.utcnow()62    to_ts = int(end_date.timestamp())63    64    url = f'https://min-api.cryptocompare.com/data/v2/histohour?fsym={symbol}&tsym={tsym}&toTs={to_ts}&api_key={api_key}'65    response = requests.get(url)66    data = response.json()67    68    if data['Response'] == 'Success':69        data_points = data['Data']['Data']70        df = pd.DataFrame(data_points)71        df['time'] = pd.to_datetime(df['time'], unit='s')72        df.set_index('time', inplace=True)73        74        # Filter data based on start_date75        df = df[df.index >= start_date]76        77        return df[['close']]78    else:79        print(f"Error fetching data: {data['Message']}")80        return None81 82def calculate_ema(data, span=20):83    data['EMA'] = data['close'].ewm(span=span, adjust=False).mean()84    return data85 86def run_model():87    while True:88        try:89            new_data = fetch_data(start_date=(datetime.utcnow() - timedelta(days=3)).strftime('%Y-%m-%d'))90            new_data = calculate_ema(new_data)91            92            if len(new_data) < 50:93                print("Not enough data to update the environment.")94                time.sleep(3600)95                continue96            97            env = DummyVecEnv([lambda: TradingEnv(new_data)])98            model.set_env(env)99            100            obs = env.reset()101            dates = new_data.index[50:]102            prices = new_data['close'][50:]103            emas = new_data['EMA'][50:]104            actions = []105 106            for date, price, ema in zip(dates, prices, emas):107                action, _ = model.predict(obs)108                actions.append(action[0])109                obs, _, done, _ = env.step(action)110                if done:111                    break112 113            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]]114            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]]115            new_hold_signals = [(date, price, ema) for date, price, ema, action in zip(dates, prices, emas, actions) if action == 0 and date not in [signal[0] for signal in hold_signals]]116            117            for signal in new_buy_signals:118                if signal[0] not in [s[0] for s in buy_signals] and signal[0] not in [s[0] for s in sell_signals] and signal[0] not in [s[0] for s in hold_signals]:119                    buy_signals.append(signal)120                121 122            for signal in new_sell_signals:123                if signal[0] not in [s[0] for s in sell_signals] and signal[0] not in [s[0] for s in buy_signals] and signal[0] not in [s[0] for s in hold_signals]:124                    sell_signals.append(signal)125 126            for signal in new_hold_signals:127                if signal[0] not in [s[0] for s in hold_signals] and signal[0] not in [s[0] for s in buy_signals] and signal[0] not in [s[0] for s in sell_signals]:128                    hold_signals.append(signal)129            130            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]131            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]132            hold_signals_data = [{'timestamp': signal[0].strftime('%Y-%m-%d %H:%M:%S'), 'type': 'h', 'price': round(signal[1], 2), 'ema': round(signal[2], 2)} for signal in hold_signals]133            all_signals_data = buy_signals_data + sell_signals_data + hold_signals_data134            all_signals_data.sort(key=lambda x: x['timestamp'], reverse=True)135 136            ref.child('signals').child('data').set(all_signals_data)137            time.sleep(3590) 138        except Exception as e:139            print(f"An error occurred: {e}")140            break141 142def start_model_in_background():143    thread = threading.Thread(target=run_model)144    thread.daemon = True145    thread.start()146 147def dummy_interface():148    return "Model is running in the background."149    150if __name__ == "__main__":151    data = fetch_data()152    data = calculate_ema(data)153    if len(data) < 50:154        raise ValueError("Not enough data to fill the window size.")155 156    env = DummyVecEnv([lambda: TradingEnv(data)])157 158    model = PPO.load("ppo_trading_agent", env=env)159    start_model_in_background()160    gr.Interface(fn=dummy_interface, inputs=[], outputs="text").launch()