From b43fbe21709e034ac26875ec84aa890b3ef2522f Mon Sep 17 00:00:00 2001 From: Florian Egger Date: Tue, 26 May 2026 12:43:47 +0200 Subject: [PATCH] initial commit --- README.md | 30 + __pycache__/config.cpython-314.pyc | Bin 0 -> 3016 bytes __pycache__/main.cpython-314.pyc | Bin 0 -> 5951 bytes config.py | 103 +++ main.py | 117 ++++ requirements.txt | 8 + src/data/__pycache__/pipeline.cpython-314.pyc | Bin 0 -> 27417 bytes src/data/pipeline.py | 611 ++++++++++++++++++ src/data/survivorship_bias.py | 89 +++ .../__pycache__/metrics.cpython-314.pyc | Bin 0 -> 4929 bytes src/evaluation/metrics.py | 124 ++++ .../__pycache__/backtester.cpython-314.pyc | Bin 0 -> 8380 bytes .../__pycache__/gnn_model.cpython-314.pyc | Bin 0 -> 2467 bytes .../__pycache__/trainer.cpython-314.pyc | Bin 0 -> 10723 bytes src/models/backtester.py | 177 +++++ src/models/gnn_model.py | 45 ++ src/models/trainer.py | 178 +++++ .../__pycache__/visualization.cpython-314.pyc | Bin 0 -> 6607 bytes src/utils/visualization.py | 137 ++++ stock_gnn.log | 0 20 files changed, 1619 insertions(+) create mode 100644 README.md create mode 100644 __pycache__/config.cpython-314.pyc create mode 100644 __pycache__/main.cpython-314.pyc create mode 100644 config.py create mode 100644 main.py create mode 100644 requirements.txt create mode 100644 src/data/__pycache__/pipeline.cpython-314.pyc create mode 100644 src/data/pipeline.py create mode 100644 src/data/survivorship_bias.py create mode 100644 src/evaluation/__pycache__/metrics.cpython-314.pyc create mode 100644 src/evaluation/metrics.py create mode 100644 src/models/__pycache__/backtester.cpython-314.pyc create mode 100644 src/models/__pycache__/gnn_model.cpython-314.pyc create mode 100644 src/models/__pycache__/trainer.cpython-314.pyc create mode 100644 src/models/backtester.py create mode 100644 src/models/gnn_model.py create mode 100644 src/models/trainer.py create mode 100644 src/utils/__pycache__/visualization.cpython-314.pyc create mode 100644 src/utils/visualization.py create mode 100644 stock_gnn.log diff --git a/README.md b/README.md new file mode 100644 index 0000000..fa2487b --- /dev/null +++ b/README.md @@ -0,0 +1,30 @@ +stock_gnn_project/ +│ +├── data/ +│ ├── raw/ +│ │ ├── news/ # Raw news data +│ │ ├── social_media/ # Raw social media data +│ │ └── ... +│ ├── processed/ +│ │ ├── news_features/ # Processed news features +│ │ ├── social_features/ # Processed social media features +│ │ └── ... +│ └── external/ +│ +├── src/ +│ ├── data/ +│ │ ├── news_processor.py # News data processing +│ │ ├── social_processor.py # Social media processing +│ │ ├── sentiment.py # Sentiment analysis +│ │ └── ... +│ │ +│ ├── models/ +│ │ ├── gnn_model.py # Updated GNN model +│ │ └── ... +│ │ +│ └── utils/ +│ ├── text_processing.py # Text processing utilities +│ └── ... +│ +├── config.py # Updated configuration +└── ... diff --git a/__pycache__/config.cpython-314.pyc b/__pycache__/config.cpython-314.pyc new file mode 100644 index 0000000000000000000000000000000000000000..3993270ed26ef534c4491c656049e7c70fd50302 GIT binary patch literal 3016 zcmb7G&2JmW6<>rN3j1gjr0tIsF&9M)?^w76UimG6_s20S|`}n;# zZ+>rPxF7WQc^P=VeHz~RZ7;+8Xewt9+Y#O_fly^~46+FfF~=Ukma-AsBok()7-Wwy zh`rO(Mw6b6o+;)|CzRNc?FJKb{L|FMn2oxNd7=_+Pc3v1*5K$UZ+GnGOdAZbA%^ti z9K^oCW~~bU8_qD`6Tes&u@t4iyK|S3#@0H^0-MHQ>#hvr( zgSe4**EGd!c##7+?{@kyY4UM7uJ`}M`7Yw-7&Io?NI!BB{}pQYA$NDp0puaUje+dV z%j!SdGuha?UgSON7zv|ZGK732a)lZ$-$VaRX3j-!p%D2w8YD3kCcntp z$p>hN%%BLF1R7+q{i`XD~Vb_=|7*CG9 zi}MTfv-1n^J9bY$nw{R8oi3Wz#rfHV#o5~nmb}oGsdG<_0YiieGTubM|9XSkwc9@k$6+hS70Feu3DyH%T-Z} zMUCJ>NmI)NYsF5)MYEw6Dk`ClC;6gUIgTD_NBQ#hRaFkW#G`z)!369!7 z$!4hSA+@h%6xx$bQhP#>!MBoLhnU1NPuc7vGnpVNamvQ!xJ2z)K1+L6L0FeEF_yY2 z>H$;&>}{5GeGf|AKK4V+`Eo(amuAgd)}Z!7V<%>(?#%;3D??YQr;DI2wV;>vy(7rw zH-I7t0G^D3UQubEDQkKe7xaTNq5b>1p;;?#yZeEW{J9Ppf@h8BW}<-()904{Gtd*g zbg&1rbP#PdcAL5}F6Z}DjHwr6*c=BXQ}$u}yMuhGrBw`4f(4azaQEm2W%u*s zF=cmjtxR1-4K^*LhiqG#mBF3h6b>f@M16=`!(DA4i=;R&%Y5S87Ub6z9*JD~+~nTN z@2Zg1pngFV6ajpS5MSkyOnriw;MY4U4JY`tAS;lrSgoi}6&3L#t;T_(Xn9#t?v4`A4*^$W;rEA`>c{ycuskM%#u6# zx`qgsw)Tf2JmlM3WQ`Y+DTRjKaU_;$7{m1lN2P%CgjBRmII;%Dl>yiVZ_PhC>VZQD zSMc2aQN{cYvRn4?aKq;_&&iFP?q9Tu;`Lb*UzurkZ2F{O(E9|8X_d z3WZTBJT!8#{UYU2k)Eslg7lq9b)z%~ki+d?)Lt z$Ia-y1{;3uWE`$PtbVromG6YSVrN>c)6^P)>!->qcCo8YohDwf3vZm^h_foao?!+@ zVV1E8tDtl(-H6?5Ed8cAh8pZJq#T&2e$w(sow#OhoV)f z6$pZ}KUht+28XMeR$rjH(h7z;UPv4n2T)|J!485FootVTx2kCx;^8%P58T5ezzMj= z?l#zHD>4V5*62s3JaxUnPMBFHOxSIJy>U9i&T8h3&tWa!$4mlb6dOnJsYWPpY<`uR(@s5=9#w|OP%d*Tm5h29)pj%m)e)z$CVvW&977A2E1O| iZ8qEA7|-9C@qan74Zga6V%kH(&Ugb~Mn7{}N&W+*+XsRG literal 0 HcmV?d00001 diff --git a/__pycache__/main.cpython-314.pyc b/__pycache__/main.cpython-314.pyc new file mode 100644 index 0000000000000000000000000000000000000000..ecc6ff49cf66f5991c26fbab93f240dab8d6674c GIT binary patch literal 5951 zcmeGgO>Y~=bto>sKg3T_lx35WB3G0YI(8Vvaumn1Vwp5$m08Dh5y&CdT<968*{Qj~Te1{uE=$~Xlf0iEbIB7trj^0MBA%&FgXh;qy#z+@t6cf+} zQs$&tF{?OZ%96AyRuyJawxnILCmo7Iotsk5q)TzBusP*U4l08xY)N^NUd5}z)|4;l zSNzF<5{RISC^Jo6(#A1@pVj6RVr_tr)0sAoA=aLNILHq;Ug$#l|6?f8o5lxN$IHg< zM74iJ&)0>){s`5l>|>oFq>Nr1VO>XmgHQo-I*KliY8oSuy9xC#=KhW-SdAI5IbSKS zmx=&+&ncDsmDf3i%alZ3D3$rx0N$VHgnU)t6h0^N(t1VO;L3SEx4|n?DKFzdzOo^5 z5}#8lxl4RGe|dwGu7EDr>DASfT>gr}%L*?6BD7SI#EJxZ7V}D}QeM0YQosZxI}C41 z9N3es~AP%3Vcqe6l;meDNY6@8+=TTJ|CT&Dsz2M;5T_eH{al8nJWTg z%_WoQa>a5v4oYH1Y*Kl#S&d1-*2`SEAn=l$*QePE6S4-6|7Ih*Z>>+u^j4;i0Dwhw zEh1>mL@9JV#l55rnsgRr=)R|FV5FKrdmp>29|KQl0?(iRQB9%&EIFR0YL>)-JuDg5 zTvR&S_uq+*=seg%1EbT>m-<;6OuA?dZ==}}>JdU2w+2%vB8t*q_ko`s*lZ(;4v;6I@0r4V@Ipz%GSJ4!jU5qO=u@{Vgf6arq?uDs7t z4sF!eH(fRkY8(^-t?pg5PEihR)N6%xo6#r~qGt%=rYN-QX<qe4esiII-LPKzD*@ z%~~^uDX%tW`zJ)*D}RL3yUQ)J*1st-)JaXn6X!#8WF46~9kB;yZJDD##_G=q>(S3d z=N>umYB>Piul!gybu{zpvuUC$>-?W(-D|H=%@_XB&-38gVjufyAM4p88lNtj!ORlX zpp7&d4_f{RA+lDYSM;r&q8wS@+A@Wy)nD}xg#~hN|920v-b{+hYa{Dt15{fkqTwk- zy>cQpnDwzkE_9?{-6AM`0sR72p|b03evFU}9oWMz8`9<5^Z$JRoK=2CzSDG1T=0Sc-D#d3q6w*eht0TKt40a#jS&frI>?oY%UGRtcdqebR{h9NW zUmMx|>;X2SyU+ThU2AVt+d1sm1d2`6MpEUHQsRVCtyC^X3-H>FioNGKtCn7I>@^eKl8~v{I2}pqK>i5{L}ytPo0Vw9g1H+~`wO zY&s=Y@|Q8Avig-?aJhU{+T`W9ST5Ga*=ku0_)_mdABAXjfkT8f+hXCJ=#o&8`C4SD z*D;U`=z4}|uV<(Q?HR3;h@S0s634J1hYiu%0OZLSPpsF%_9eBoD9K<=sYFlG?&4T2 z42^`cUJ*(aeV?HE>Vg2eZ33BM2G)khgBm}$T?ItHRmrJ{T7ZU@lX$r*C~`b6Z!VLZ zmdXYG-P%0uWum0t5qO0!#G`DNH7awPe1W*Qo08Soz;~p%{f_hlu^DWJ=A+0yTe zPThtQX-=?RgcgwRy6lvypKr zDuVrjR1wRZL=q&K3Ds0m_zfA`OLD0!E2Q7VrggBcV2i|y0+;6{l2VB@E_gsJVB71E zpk)c#Kx|eQRK{j#NMQ|Q0c@qm8^`!`QEp zwLf~>yXBwg_$QnG$&P=f>7Ti?xaE)6S9Y8ew^v%u!}U`-nP{6C+i?%y3bx#l`swY+ z!O!0Ra7q;6boGaXWO1_193&GE{UH4a%t7Xcn#@p}8P>U(XfrSE zcn{pFb`H)r56-u|$LeRc$7eg^^Ud-3`!lWa$c+u`VE`#;(Lr^tt{8|Is) z9Y^pc|LOUTW8#5hV#n$GDE(pj*4$^uKRJG9;*Q+f_hM(?d~@G?Yu_u4k%jvw?^i!R zcmGOb_{>9R`kSA?eZSZkUVi92{k;v1#C~x0_a<4I?eIh=9BYPS55qIRxHGdqxREK? z2u?QU-)uAINs|0G=WkWoOdMce;18p}AN|_tYXqL_1g4vT>22r!M&wW@GS`gE>ETBf z#%6zPH2ci;Gmk*!V224dnP6jR_TDL2sy1_so+VHEA+z!rjxWO2T#g**TrRdBTgYt# z_b~Qua0)5u2}28X|a})EV{J z6YrHHUHJ|Y4z+;FajiOb?~H(s2A|@3*9p0;uun(o)!pt579xbIQWE6&W=XD+r<-b8 zcL_KyadIiYq<+c7Olo!c*VwUK68L2In;v$ZQ{YP7pi^6sS5~_BFl38-F3jbmQ4%_- zf`3NRIRd{yys|DpBvpYwF0I3qe0Gtq!@oyHgTe3xa({~=Um*XtDEuWl^cBj&|1Z(e zFH!6(^g0|;gSCVF50HP`?)_-?!`W@e;cw|R*Z^t(Bz&3L0Xv_9o-TW`}TL!~f!=u^&Vw^W>ZDaWFjP-ZR-y0D#qB{B? DdKmJ6 literal 0 HcmV?d00001 diff --git a/config.py b/config.py new file mode 100644 index 0000000..d3037d0 --- /dev/null +++ b/config.py @@ -0,0 +1,103 @@ +import os +from datetime import datetime + + +class Config: + # Data settings + DATA_DIR = os.path.join(os.path.dirname(__file__), "data") + RAW_DATA_DIR = os.path.join(DATA_DIR, "raw") + PROCESSED_DATA_DIR = os.path.join(DATA_DIR, "processed") + EXTERNAL_DATA_DIR = os.path.join(DATA_DIR, "external") + + # Ensure directories exist + os.makedirs(RAW_DATA_DIR, exist_ok=True) + os.makedirs(PROCESSED_DATA_DIR, exist_ok=True) + os.makedirs(EXTERNAL_DATA_DIR, exist_ok=True) + + # Stock universe settings + INITIAL_TICKERS = [ + "AAPL", + "MSFT", + "GOOGL", + "AMZN", + "META", + "TSLA", + "NVDA", + "JPM", + "V", + "WMT", + "PG", + "DIS", + "NFLX", + "ADBE", + "PYPL", + "INTC", + "CSCO", + "PEP", + "KO", + "XOM", + ] + INDEX_TICKER = "^GSPC" # S&P 500 + DELISTED_TICKERS_FILE = os.path.join(EXTERNAL_DATA_DIR, "delisted_stocks.csv") + + # Date settings + START_DATE = "2010-01-01" + END_DATE = datetime.now().strftime("%Y-%m-%d") + TRAIN_END_DATE = "2020-12-31" + VAL_END_DATE = "2021-12-31" + + # Model settings + MODEL_DIR = os.path.join(os.path.dirname(__file__), "models") + os.makedirs(MODEL_DIR, exist_ok=True) + MODEL_NAME = "stock_gnn" + HIDDEN_CHANNELS = 64 + NUM_HEADS = 8 + DROPOUT = 0.6 + LEARNING_RATE = 0.001 + EPOCHS = 100 + BATCH_SIZE = 32 + LOOKBACK_WINDOW = 30 # Days for feature calculation + + # Backtesting settings + INITIAL_CAPITAL = 100000 + TRANSACTION_COST = 0.001 # 0.1% per trade + + # Evaluation settings + BENCHMARK_TICKER = "^GSPC" + + # News data settings + NEWS_API_KEY = "your_news_api_key" # For NewsAPI or similar + NEWS_SOURCES = ["reuters", "bloomberg", "financial-times", "wsj"] + NEWS_CATEGORIES = ["business", "financial", "economy"] + NEWS_LOOKBACK_DAYS = 7 # Number of days to look back for news + + # Social media settings + TWITTER_BEARER_TOKEN = "your_twitter_bearer_token" + REDDIT_CLIENT_ID = "your_reddit_client_id" + REDDIT_CLIENT_SECRET = "your_reddit_client_secret" + SOCIAL_MEDIA_LOOKBACK_DAYS = 3 # Number of days to look back for social media + + # Sentiment analysis settings + SENTIMENT_MODEL = "vader" # 'vader', 'finbert', or 'custom' + FINBERT_MODEL_PATH = "yiyanghkust/finbert-tone" # HuggingFace model path + + # Alternative data features + NEWS_FEATURES = [ + "sentiment_score", + "mention_count", + "positive_score", + "negative_score", + ] + SOCIAL_FEATURES = [ + "twitter_sentiment", + "reddit_sentiment", + "twitter_volume", + "reddit_volume", + ] + ALTERNATIVE_DATA_WEIGHT = 0.3 # Weight for alternative data in final prediction + + # Database settings for alternative data + ALTERNATIVE_DATA_DB = os.path.join(DATA_DIR, "alternative_data.db") + + +config = Config() diff --git a/main.py b/main.py new file mode 100644 index 0000000..df04aaf --- /dev/null +++ b/main.py @@ -0,0 +1,117 @@ +import logging + +import matplotlib.pyplot as plt +import pandas as pd + +from config import config +from src.data.pipeline import StockDataPipeline +from src.evaluation.metrics import calculate_performance_metrics, compare_to_benchmark +from src.models.backtester import GNNBacktester +from src.models.gnn_model import CorporateActionAwareGNN +from src.models.trainer import GNNTrainer +from src.utils.visualization import plot_performance, plot_trade_log + +# Configure logging +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", + handlers=[logging.FileHandler("stock_gnn.log"), logging.StreamHandler()], +) +logger = logging.getLogger(__name__) + + +def main(): + # Initialize data pipeline + logger.info("Initializing data pipeline") + pipeline = StockDataPipeline() + + # Update all data + logger.info("Updating all data sources") + pipeline.update_all_data() + + # Create datasets + logger.info("Creating training and validation datasets") + train_dataset = pipeline.create_training_dataset() + val_dataset = pipeline.create_validation_dataset() + + logger.info(f"Training dataset size: {len(train_dataset)}") + logger.info(f"Validation dataset size: {len(val_dataset)}") + + # Initialize model + logger.info("Initializing GNN model") + # Get number of features from first data point + num_features = train_dataset[0].x.shape[1] + model = CorporateActionAwareGNN(num_features) + + # Train model + logger.info("Training model") + trainer = GNNTrainer(model) + train_losses, val_losses = trainer.train(train_dataset, val_dataset) + + # Plot training curves + plt.figure(figsize=(10, 5)) + plt.plot(train_losses, label="Training Loss") + plt.plot(val_losses, label="Validation Loss") + plt.title("Training and Validation Loss") + plt.xlabel("Epoch") + plt.ylabel("Loss") + plt.legend() + plt.savefig("training_curves.png") + plt.close() + + # Load best model + trainer.load_model() + + # Run backtest on validation set + logger.info("Running backtest on validation set") + backtester = GNNBacktester(model, pipeline.price_data) + portfolio_values, trade_log = backtester.run_backtest(val_dataset) + + # Get benchmark data + benchmark_data = pipeline.price_data[config.BENCHMARK_TICKER] + benchmark_values = benchmark_data.loc[portfolio_values.index]["Adj Close"] + + # Calculate performance metrics + logger.info("Calculating performance metrics") + portfolio_returns = portfolio_values.pct_change().dropna() + benchmark_returns = benchmark_values.pct_change().dropna() + + metrics = calculate_performance_metrics(portfolio_returns, benchmark_returns) + comparison = compare_to_benchmark(portfolio_values, benchmark_values) + + # Print metrics + logger.info("\nPerformance Metrics:") + for metric, value in metrics.items(): + if isinstance(value, float): + logger.info(f"{metric.replace('_', ' ').title()}: {value:.4f}") + else: + logger.info(f"{metric.replace('_', ' ').title()}: {value}") + + logger.info("\nComparison to Benchmark:") + for metric, value in comparison.items(): + if isinstance(value, float): + logger.info(f"{metric.replace('_', ' ').title()}: {value:.4f}") + else: + logger.info(f"{metric.replace('_', ' ').title()}: {value}") + + # Plot performance + plot_performance(portfolio_values, benchmark_values, "portfolio_performance.png") + + # Plot trade log + plot_trade_log(trade_log, "trade_log.png") + + # Save results + results_df = pd.DataFrame( + { + "date": portfolio_values.index, + "portfolio_value": portfolio_values.values, + "benchmark_value": benchmark_values.values, + } + ) + results_df.to_csv("backtest_results.csv", index=False) + + logger.info("Backtest completed. Results saved to backtest_results.csv") + + +if __name__ == "__main__": + main() diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..cf6a764 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,8 @@ +numpy>=1.21.0 +pandas>=1.3.0 +torch>=1.9.0 +torch-geometric>=2.0.0 +yfinance>=0.1.63 +tqdm>=4.62.0 +matplotlib>=3.4.0 +scikit-learn>=0.24.0 diff --git a/src/data/__pycache__/pipeline.cpython-314.pyc b/src/data/__pycache__/pipeline.cpython-314.pyc new file mode 100644 index 0000000000000000000000000000000000000000..74e6545eba1123da7e2579314c0659df73903928 GIT binary patch literal 27417 zcmeHwZB$#ymEhAS`T_(JLI}_Y-}nOzHXk;&F*aZ@Hef-^-3hT)Mgpr1Nc|-IL3cWy zoSX@r$vKTPlNcx2HO_R8VJCe;CeDPNlXJ%1=}mk#nKQ^BOLyQWK9cjTd@>|rY2r3D@ zDySxJO)wqa($0*Ij9>!WB4g5D6&&9X4yfB zf%3tK=2AhaT@0bjE`d<)p!SOF(p+jpVV6Ou4W^M4IZ1(jLyE$#07#`>389MQsz|Py z}=8w)V%fM+o1qhtBrF}LTkbI>>H_XWHNz(@j~AfpA00VMIb zok8#51rHryRL&8f%gvTD#NtUXbAVC-&&W^!hog+q2k60S>_xRdHmX(rD~!VF^p1Ih zPA6kpH|kYfo(?Jqv_Wu{ikb>0rC-d9+KMNo%L;QezYyO>?)Qqd_|IjSQn)6klqf$tKdC@%H16{pJ<>Tn`bBNIFV;6CY4rYqa-vCxD zfaNh@16YXwmdB3`V5I_B>^FG^2+6}}Dy*@~xH1+X6vdR48pvSFQIwo=n<%P)vQS&p zA&=KlNNVv?O_D;&E2>mX6r2QHuzD^7ig?F{tvJcnA=)=;b$bIB-p`05V3q&E$oo0# zkdP#23J{TnTZfF8JYmMhJR3GQQQz%L?uu`G%w1*e3%Y31zf0cWd8=pCAG|`kdd?N_ z1gdbUXjnw{1CEb)XdELN7a6U^APWLUKH?i5_Rx&fJ2vEFvf8~Po^D^T9UEw)X&+4+ z0E$I;+K8p}SgIfeYV*s-iX}b^!BuKanx?diqUNHAxjJO7UfgqJ$=ov8`o*4TX=S9e zK2%!2*wDRH+B4a@tSO3GiX)aiAHR=9bmOxvyIX3gH*2g%b0lx+MG&T5EB{2goEl3ZmB)Dii#@P@ zVyT3F@hgbSr*PfzI2}*5i?>lu9*5^*cF8u%DLCl(+b)flv&(=b$S1P;To*koVh4Q$ zt?5E2K$|hJU|@xyyHW-?6aY}?#yFIWB0x$>2|}P#VG5tS6PgGNKyZ~>RcIoLypSSq zy5wf{jq1gnP0Nad(RgCSv?pZRvsl}?tT>U7ydz}VvAFBdvf^+;at!)`;sl)!4Uie< z<_G5tsnTOB1P4V|VV2H{8$U-;Zb&VlPK(Zq?4mx=C#h->I_AW`LU=|o1WRMgHR=H_ zAPEF%kT&E)K6=y@WYV3^!4X#=;0EEt=?oxv(N(fqVDOx-u`yo|8I+STr&ugYZ#Y|i|ypUeK8hc z0Tyqic2ecB0>kjN1(a88p$dVV<9~K>F-N2&6DC5WTGT4Z*jJE+8L?O?-^s{Y!!?YX<4O|&INFKHT<3(Tx+;~@Y=!G2th~s5JDV=Ab`udERKx0bW$RtfQFNKYzM%r)b^1PdF~gwrlrdo zs{kAC$FVUt5gVa;@*s&JMCx-ueLs34#mR6S$#7prk|B`%cybi$54>z!5I(L>zi5Ep zkvM=O>;k!xdwNlLN_sdV$dw0~n%@x(;JP7CFA|lc$1Wbo?yHtw}N z2?c@(++0vMvbv-Pct4ChJ}flw{tA|G#&$Q>*bayla6OlU9(v3*!je&WKiM5kq+tmV zQ75vj+{TG3Hqq3&t|FkoR$zM8Ay!W985_A`9S@)+jhi%(J7P_OjG7T&28^5t+yRzK z=d`w+?Ck3YBHiD4tgFr5=WOph*~WqJxXJLyL#H7mYm+G`$aD<#oT(9o3%tjYOQD z5tnA58{|Q5L{=EJjTJV)ivfUVJo}H;1TPL(cEyt@zX)J?Opa%OfWTuGK=3qpLx>dn z|D-zj^%k##TUA8D-elDQc7wE5DW1rE#t($jV|9&;u<~udH%0WSnssV8zOw=iPZ6*BDlW^leWvad6=uXCyTMA&d*QXSJ7m_rKl z^oz@iow47gAw}uTg=Ix;?6)PPu*|5J6}tra%ZhUH3$+gq6nSwRjLkpdANhNXBE)7A zHAvM7bd;5_X+$?BEk+gQ#57_iT@At3Qpu+PDu8s42+ccX{h~({hI9J4!8h{Q_2T6xPvgi72zaQ4AWyH6i^bots}MURg0U4aY7WY-uq*AV zh8X$tOC3_rKa($%ZaNBKgxO?PUf=F0d8W1txoLk!v6g`X?yGunyPovGXlw&Z9MURY zq34&Kc_!5?_S%xFhM#t|UBhYfD=6VC`(+MUZ}IvP90A{`S7lJ7P1N;gpe&L87Ul1d zT40?)U4KJSO3F-xBP%sf!u1JMgRr|Pp!{;dUYSF94hY`%^nNj{?u>y-F0r3i7HX(k zfj|X{Lz!RREQdCqKsgfkig-n$cB1$hmKhO#64N94XzMTtL1Ua0G+`t_C>|w&y|3V5 zbSYp^AMN%82D?ebbtUk=qV&1lrK7t`-R=FAQbra)1tBBxjDgJLr+ve;Cjk0!*Lcw9 zbh)1&4}hrz^cKWaa>3*APne0y&q^3fs_uvnSbz0@;i=4t%USd(7$h2LAIn`;ef=%= ze%x8N5rrcoBZ^5zMmQo3y7xz5ZH2jm&Q=uViaxNWKZ5kpFl9Q+YEw! zMdcp%k9Z;GaWkrbj}C%{9klw4vhDJq$B#;Vx*OXk3?q%*wIf>q$&iSzPzVPC8LEd- zjCn4xtW^P4fpcIBaWf)M06ze?QSh+a`hKjScth?R7Z3;odV|Pb4;2}&s$Wqr8mhh- zm~Z&YmuANnIu`Y9%Zm1g2GexIV&&n8p?T5JJlVb`E=%thJuEDa6jp@_tEN<|decqc z4d1eUchp#LGxJ8~tUPS2zHe-b7Fcig+~}EY4;So>6zqrp7Yp{UDX6?c9x7t23mNO? z&xMWqzu$SUX8GW$u<_K-4d!d-rze&SJJw`Wj^&ApDlp${`f5{DZ;F~qR*YsW9xW`2 zT1%d&qxa+xx$e;cQ&%SJ+!<1 zUi;$i_NCo@i@W;5#r@NYXzhVWZCj|eZK?M7w3Pj@Wn=Yf>g0FxshVc!=;LZC_lW4d zCnrUez7UkaFrq73!}W}78H={&yYhw0cj=!L-*f%A?4EJaVvlJ1mb88U2Qp0 zM2J2M!PY{~7XT`NoR0`Om9p$d6a?Rve!K6xN)p{f!g&dL`7a>hfEc>9goB(q&eb>I ztFV;Z@T~sCn{;-GeWoa)jqrfP~o_(>m zmTK>%s7oNv(GJ08o0myAc%t9rO;f%6mOh?uSMpzu z?6R{Q55|C(4dSi(6%NIwW57y`ifxX;b05nX@FTOWF_3RO2EdB}7a=?ulL9oTwj1Cj zC0+_p38a8jl8%Z<0qOx>rs-F4v$swP(D+r&DrobMK?)#offSHFu_Go0@LNct2LWR> zXh@QU0SnLQQ$y4L4kmWOv@sOz0suxH0G1gHFiN-gqSp-)K!8ab_0VAI3rr{je(%_B z-_THi2JELMN|V_B1`Uq@YL%{We^k)$m^mdehi>u>YoW4R}1L(p)Dkh3h4fSu}15nDsZ*08K=ST$N_ilWBi8GF=d zo2ezSHKnZHFx48-?+EF4%+`nWwZM{1mYW4P3Z~@I+`{WGU4LonrD*x?+b`XE=`UWs zp`Mnm6qnvU_>F_}(xu|MNb%uN@!^GurQ#FQij`CLXo2mqMBFVhO-m!js*tg2_F~vr zKYuxFJaX5#WIXXmMpe|WWl)7hC}2z}qq)}WpSk{-sn0|!YQB{@m-(P#{|(J%P4qq! zEhvH}PA50ZpV`G)Dsz1dEi|MWh{SOL**U^Ufs&_Y6*-E%c`S>riJ!}#=Es| z58ku=VCe3%i~8=T5^+}cV+sOb*8uu{0EygpTTJbg_$}>G5xm~cZYhD+4=HtflkA7; zlJ_W5lR+G&u^=RXm9)A z92)pZOp77xj9+S(Iiy4@0Xz=01n@h79&fP!B%BI*(FKHp?9UX0GV^vqpv<5@g8i9H zFcK+Les(#b%xQ7GLO;)S*HLDLU)C(^x(ale!~}F%IZ+a$%d7%#qwQH9PivFmSX!)4 zmCsRO;)^t4NusdpDq?A?Y75$mnwDfrDNmx5^yd)>c@|3p)2zoVIbc>&gG`?BJAjPNx_3`daaDbk-2C6c{pZo&15YZFt_ zRbBpca9LLzt=>D|@y5$Y=So8Qk{QR6e%B)jr7wA`rgB?=M8GzLNb^t=C(R$ZYq+bp zC;rLadoTRB{$9nRr7xoGzpw2l()=?Z&7*SvPIF6>_{}Eu(K^{%@{*%_WN+<}L7Wf- z$R@3D@$ykl09JrNh#QiC(O!=9MV!a3@C z9()te`$pU#i4AhS#5WRJ+=M(v#b^R*rXa!bDJ!OEL>JgY5}k>%6^z8~H2SGfgyE1? z&nOO#xJJ*pU5A;28!mm)d7-Oh)xi7$R)xpFm3i#OFjQ8dp2g2p?QIlNrS<21wp<;d6GuP z$VkR~mxu^SpTsm|5y@ncwVBg?1lNEb2^{%?+TFi>N`XAor`rR7WG}rij!;7EaeH&L%rd8=a)LaJ`^$6 zgv>Qd`r661mHeWpHfJ$+PpG!zi9}qG0isWVcJd@>d9(%hGt1{Skv&J@|HVB=S1jc> zU;64xUwb)XsSjD|=O>mdN2fHadh6`va8YBVs5w;Byj0Y(CJ{GhPe~)%vXHjywTqGR z1EKN*OWGz7`D%`dLYd|7J;|oDdFVKIf;i5Bg>kX4`EAMFmLGuLWHG-ZqUpS^=|uVp z5bpj&BDR;A^YV5E1(6NvMI5)1KfsJyt0p`aOaGG)~yoT9bDGDu9IWJ@lvQ5y=dXNI2KG zM&AEP!M@+@XVX~$4ZF=fD4M`;&Ero%i{K|F-rS4UFR_c!_~QWck6mV$!(SRPJjm=S zP+z3^rDwFM1S%p*eVGvRlzxVe1e=-X>UJSh`5y9&0rZx{z>|WyB(_Hs`zhF8N}I5_ zN<>Tqv0Be2nh4|$`B|#*`=H=~?Kre9OJxRgB(NBJdV**P1RGfdjH8B|r1+bRBZvYU z8me7eV87ZRX(t3L#HSA#^z|go2VKCtjY5aqN(1@5%ubR>lVP2xuI|8RpE4f-aIv;4DWyOuLiLC5(=Fj|Kw>XZb=paqRi_JmPOY4pmn zXTz%2*maI{yZ+2S>=KA~sSdnU^Fjx-5n#s`#wABMJKG$WzyAowCGBGv7XulWTu0`I z9v2YyDvcAbsg8SsRzEmt26uyU3XMbY{g}&2v_Do?&H^s9hljKVv{+NIiAsY5&N2e}cPRIVNP-TZ_qpF! zwz!|S9vkrmJpHW74Avw#{ROZ1LA^y3c=Rin0^e{Y-~_dncZAh)DLAbj{U?wQIyqQI z6zktlq$DEtsClWA`1gY>7x(+8KZo_o(U=DIFX;u(l>oX`U4&uuxLF|!jdWmK1Epam zRxTw4Tz>GAb<=M60TLV%viNY&Umzf0WE>v}vbH~o*L|7MumV%e7N-^n(~Js#5J3zy zm3*8@P6+-txwATtSb10FY{v2#NyT)t96 z5*cC~xEUN~Fi?8q^+VSV&6t*SyCxN@8vWEzRFgf`5!L9X`lImE9n~17WwD>;8Sz@C zB3m^nUo%qay6e7czGY2mRI9(9c`b8B9@duM*Y1r%l|9#bX4=F0oe_Nv{J*HL0ac;N zeA9BnGHVFu?U_6g%`S*!7lpEmBH2};?5cU`?8SxQ#q6r3?Do4agtI#*k3}=Hr(9pW z94*>~2D|yzc}l>T)0G+Gbnt?7v&Iytgg9t1Y5yyRT~lTjyf_?y$Zp zqOS|->sAYjW{!i(a;7UwI6;IJh}gP&gpP;xIkrzaU_DSxIT!>$y51&43$2AsD6*mJN?GMpWbhTGK~_3>*RTDwMAk3GHd zS=i@*09HVq75Otyt%t}U5QGk#6+=&yDZeN){zYJ_o`|grQExNXv-tbWAN3xI|Keiy zOum;{j@0F@ZC@N!KL(TQH#%}0xwy(=)Im$AX^L5Q7#uQw-t%Csr-CsGU`&oYdzNrM zIlPFq1mEj0Z3+PKg^rI1yGtdz;L->*2&qbR=>`n=lciTY{UNPp{678nz@*uf1A>o4SQ(Vs$Rby zXpR}>h4x=0OA8x0xBD#{OM31%$X9MdV(?58NXrH*CQ%Zwaw8Id6K&wd?Y5R`Z96Q& zRhhWTx7Di@^um$gIFMN9F2RW+;P*lq%jz(v#C|Kcu0Za)4=Xh#)LH(N6F~jrW?ax= z-U=VK;Adus7*=~0tS7z0Y|m~0T7!?@=8!pTWBhnHR`#xQ9KQ?>dR2U6T=Ap@C;pb21t~G*vJs zX7}-!4{?0`3OCG^I~W;Ubu)OL(UX^R*y97$B3zGCg=0+r0P8Elpcn(>MW{=Jy9j94 z7`OuBN73L!KrtEp_Lk0WC+~lshBvEm}8*a}TIC=@X##5N3^23z~gMfC!9!Bm4Bn15`kUTKq27;h>9QC-y81a}NOAOPn zTZ9ZynS$fKXAE?c^dDo2a(pcC!nnsX;bA0r96;$ALo|WUV5ESqH24?>#-JWJwLdn@ zh=RmyNP~Hcnk!;els|lC~VMNb!@mjd&z#FU~VJ)h1(PkQ`;bXP* zKjR=QV|B_gIOR)ps*D5x=~hS#Aeyk}A2a9aT!4G}%K zU^Y&+u9~frJu4NsinQ$5mWb~se4KuJ7e)6@r=cCsxub%;fw_w ze<&25atK)uEx5wgBM*#6;5fwXo?AWh?ctK9NJ&emq~&f=xTNhtLEEaaoO8B(V61 zER7*c<3e`Wa&X~=u;s{-u_bCOnrVVlC=uhnka6EU)Y7=7hGwU**(`HJTN2Wi%sjiS zt%5HrutxIhL;3ad7nkx|rxej#+0is-7Bb=42ewy4p%Xxkfx+79%ApM%9P zg*slNW%}ZmKQmjtoLdvMSZ_$9dDa=jzqibuSk9}5glXxjwk)FE9n$Webw{*&AN7c+ z{2gokA}Y&pU2{z{{i*Nm{PxZ%&64);WD8VRzU%hLt&ym<>{}P+F04tY^1a|qQJ(S3 zjNDhe)1p5afzvSQB>)GdmDiNhHH+rDMN|Dk%aoF<0WRb^fwyvXJv5ikW=717*HlxY zsfLxJowGfWqJ!5?Otnm1g72HI`O5ez-dS`YWI3?VaJO~I(m5?9+VzAKaN{Vn%@}WX z|A%fkH#76>?To+3Ad%NI=1b-;f4Ay^^$2vRs`gIjO6kr>X;Y}Q35KW6{A)^FWuESa z^0rcfh-qPKEz_SDtff(frE8iG1pY)rS&L^H zzxL_b7vT1#JADhr?|ka+=}29Bw7mNEr*D0F{@JDS!!y#A^2*znzj1k9di&+qUyhbk zMM~;IC3W*JEZF9UmP-yrOR6I!^`Vmb`Cz2v;FD&m_|PK>W!eAq1o-=!VH#qS^@-W> zZ%xci;B*~`Xb*<82bZ;nS946)KY8twGlp2^Czo;#z^A0=KFy%?mU!V$&da{1{bXE$6%=TzW84dNfpeblG?m<|pC6)!vYC?|f(2cz6MPox7r7t$|C6X3N5+nuw_` zWU8BYhfPfjU13uPOi*5@Xw{VeYneE&06r$ixTd0V4A(oZbxikt<$28>elvs)Y;pfEK{*S#V|yZq@3H=I?;l z$yJRGuMC*p7tvTBd8Jg&{t=699|@t5jKNh%?Jq)U`9yz{h`eWAB=`=&y0|`26wY~Eett@^KWTJ*nT^IFe7~7 zBQYby*YSrjBcSp>m2rZv^|v-7@HK5tXErn=prPZ45L4N>ov(KbGeUed2`c$LH6tXf zzK_U^kWlBpVlx7d51Zo|7@UM=gshLrj9{2JO;q$A>qXay7gUwt4a-{*Dgxe7?}&@G z!fnBv%AC(UH;#Ax(ZX_IZSeZsmEH7PfWi~`q+&D+tgkxB3_!mP1wMogfc_ztLkj@? zBM8=601j=anv8d8V;g`9#47r~VCDH3{1}5o%|HEjnD+M=yox~xgHa5^7(Bor0zu^t zny8~0Oh}=ir{BR0qM2WT7^88(O?4kiFHiprYx=JkoWkJe7_4IO1_u8PgMYx_7Z^Om z05$M5n8YahpD;l6JN-)x{yPT0!hooe{|CnYCkFou1Jt4ug**MfA@+e($$2qV^7=_& zc7A{vBie>RqE7yWLN8P%*TB_w0%h{vka_PNN2IPJRM)Yj?MzfAAC2gbE$ZRs@&C`1 z$&97yTF0un;@^ZKdF7*0BsZhB4eldIp+}CE?fjS2BCppOuTKFD@3QIvgkDmB?@@btY1Mw9B&xaEW;s+~(kp7(MEJdHefG{mNsGI3@2$&qWf*%eHdk39pUH5`);t4?8S%ORY2R(4Ts-%@AF`xd{?ZNmAz|H%Iu#0M8SAwjOi$UHs=g3Zoo>t+hRQ43= z({o~UwqGO8_D@gBP)YuqxT>6VtJ1)v;Xl{`k`#MyGCCj9sK1Y$&jHs1%mYS95qEz^ z2DWnl6~f}-dVBEjaK$d;QCu1Fg4KM`!ydIx!xAWUFlzYcL%f5>$Um9HNISdRdl)60 zjX6pFoe(+`$GHfB$yj|EV?-n}Vp+vOcCHW8g#d>G8vti=3XD}E(K}SeJCyPrD*fkF z#m}k2Ur^caP`Y=hoS#uS|48MnQfU!Ncc0Qtm*3oVV^^3ek1DkhrRl!XG~KwQEV(Lw PDnG3h Dict: + """ + Get point-in-time data for a stock at a specific date + + Parameters: + ticker: Stock ticker + date: Date as datetime object + + Returns: + Dictionary with point-in-time data + """ + date_str = date.strftime("%Y-%m-%d") + result = { + "ticker": ticker, + "date": date_str, + "price": None, + "sector": None, + "in_index": False, + "index": None, + "upcoming_actions": [], + } + + # Get price data + if ( + ticker in self.price_data + and self.price_data[ticker] is not None + and not self.price_data[ticker].empty + ): + # Find the most recent price before or on the date + price_data = self.price_data[ticker] + idx = price_data.index.get_indexer([date], method="ffill")[0] + if idx >= 0: + result["price"] = price_data.iloc[idx]["Adj Close"] + + # Get sector data + if ticker in self.sector_data: + result["sector"] = self.sector_data[ticker] + + # Check if in index (simplified - would need historical index composition) + for index_ticker, composition in self.index_composition.items(): + # Find the most recent composition before the date + comp_dates = sorted(composition.keys()) + for comp_date in reversed(comp_dates): + if datetime.strptime(comp_date, "%Y-%m-%d") <= date: + if ticker in composition[comp_date]: + result["in_index"] = True + result["index"] = index_ticker + break + + # Get upcoming corporate actions + if ticker in self.corporate_actions: + actions = self.corporate_actions[ticker] + + # Check for upcoming splits (within 30 days) + for action_date, ratio in actions["splits"].items(): + action_date_dt = datetime.strptime(action_date, "%Y-%m-%d") + if date < action_date_dt <= date + timedelta(days=30): + result["upcoming_actions"].append( + { + "type": "split", + "date": action_date, + "ratio": ratio, + "days_until": (action_date_dt - date).days, + } + ) + + # Check for upcoming dividends (within 7 days) + for action_date, amount in actions["dividends"].items(): + action_date_dt = datetime.strptime(action_date, "%Y-%m-%d") + if date < action_date_dt <= date + timedelta(days=7): + result["upcoming_actions"].append( + { + "type": "dividend", + "date": action_date, + "amount": amount, + "days_until": (action_date_dt - date).days, + } + ) + + return result + + def create_training_dataset(self) -> List: + """ + Create a training dataset with proper handling of survivorship bias and corporate actions + + Returns: + List of PyG Data objects + """ + import torch + from torch_geometric.data import Data + + logger.info("Creating training dataset") + dates = pd.date_range(config.START_DATE, config.TRAIN_END_DATE) + dataset = [] + + for date in tqdm(dates, desc="Creating dataset"): + # Get current universe of stocks (those that existed at this date) + current_tickers = [] + for ticker in config.INITIAL_TICKERS + list(self.delisted_tickers): + if ( + ticker in self.price_data + and self.price_data[ticker] is not None + and not self.price_data[ticker].empty + ): + if ( + date >= self.price_data[ticker].index[0] + and date <= self.price_data[ticker].index[-1] + ): + current_tickers.append(ticker) + + # Skip if no stocks available + if not current_tickers: + continue + + # Create node features + node_features = [] + corporate_action_flags = [] + + for ticker in current_tickers: + # Get price data for lookback period + lookback_start = date - timedelta(days=config.LOOKBACK_WINDOW) + price_data = self.price_data[ticker].loc[lookback_start:date] + + if len(price_data) < 5: # Need at least 5 days for meaningful features + # Use default values + features = [ + 0, + 0, + 0, + 0, + 0, + ] # [return, volatility, momentum, volume, price] + else: + returns = price_data["Adj Close"].pct_change().dropna() + features = [ + returns.iloc[-1], # Last return + returns.std(), # Volatility + returns.mean(), # Momentum + np.log(price_data["Volume"].iloc[-1] + 1), # Log volume + price_data["Adj Close"].iloc[-1], # Last price + ] + + node_features.append(features) + + # Get corporate action flag + pit_data = self.get_point_in_time_data(ticker, date) + flag = 0 + if pit_data["upcoming_actions"]: + # Use the type of the soonest upcoming action + soonest = min( + pit_data["upcoming_actions"], key=lambda x: x["days_until"] + ) + if soonest["type"] == "split": + flag = 1 + elif soonest["type"] == "dividend": + flag = 2 + + corporate_action_flags.append(flag) + + # Convert to tensors + x = torch.tensor(node_features, dtype=torch.float) + + # Add corporate action flags as additional features + corporate_action_tensor = torch.tensor( + corporate_action_flags, dtype=torch.float + ).unsqueeze(1) + x = torch.cat([x, corporate_action_tensor], dim=1) + + # Create edges based on sector relationships + edge_index = [] + edge_weight = [] + + for i, ticker1 in enumerate(current_tickers): + for j, ticker2 in enumerate(current_tickers): + if i < j: + # Get sector data + pit1 = self.get_point_in_time_data(ticker1, date) + pit2 = self.get_point_in_time_data(ticker2, date) + + # Create edge if same sector + if ( + pit1["sector"] + and pit2["sector"] + and pit1["sector"] == pit2["sector"] + ): + # Calculate correlation as edge weight + lookback_start = date - timedelta( + days=config.LOOKBACK_WINDOW + ) + returns1 = ( + self.price_data[ticker1] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + returns2 = ( + self.price_data[ticker2] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + + if len(returns1) > 5 and len(returns2) > 5: + corr = returns1.corr(returns2) + if not np.isnan(corr): + edge_index.append([i, j]) + edge_weight.append(corr) + + # Convert to tensors + edge_index = ( + torch.tensor(edge_index, dtype=torch.long).t() + if edge_index + else torch.empty((2, 0), dtype=torch.long) + ) + edge_weight = ( + torch.tensor(edge_weight, dtype=torch.float).unsqueeze(1) + if edge_weight + else torch.empty((0, 1), dtype=torch.float) + ) + + # Create target (next day's return) + y = [] + for ticker in current_tickers: + next_date = date + timedelta(days=1) + if ( + ticker in self.price_data + and self.price_data[ticker] is not None + and next_date in self.price_data[ticker].index + ): + ret = ( + self.price_data[ticker].loc[next_date]["Adj Close"] + / self.price_data[ticker].loc[date]["Adj Close"] + - 1 + ) + y.append(ret) + else: + y.append(0) # Default value + + y = torch.tensor(y, dtype=torch.float).unsqueeze(1) + + # Create Data object + data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y) + data.date = date + data.tickers = current_tickers + + dataset.append(data) + + return dataset + + def create_validation_dataset(self) -> List: + """Create validation dataset (similar to training dataset but for validation period)""" + import torch + from torch_geometric.data import Data + + logger.info("Creating validation dataset") + dates = pd.date_range(config.TRAIN_END_DATE, config.VAL_END_DATE) + dataset = [] + + for date in tqdm(dates, desc="Creating validation dataset"): + # Get current universe of stocks + current_tickers = [] + for ticker in config.INITIAL_TICKERS + list(self.delisted_tickers): + if ( + ticker in self.price_data + and self.price_data[ticker] is not None + and not self.price_data[ticker].empty + ): + if ( + date >= self.price_data[ticker].index[0] + and date <= self.price_data[ticker].index[-1] + ): + current_tickers.append(ticker) + + # Skip if no stocks available + if not current_tickers: + continue + + # Create node features + node_features = [] + corporate_action_flags = [] + + for ticker in current_tickers: + # Get price data for lookback period + lookback_start = date - timedelta(days=config.LOOKBACK_WINDOW) + price_data = self.price_data[ticker].loc[lookback_start:date] + + if len(price_data) < 5: + features = [0, 0, 0, 0, 0] + else: + returns = price_data["Adj Close"].pct_change().dropna() + features = [ + returns.iloc[-1], + returns.std(), + returns.mean(), + np.log(price_data["Volume"].iloc[-1] + 1), + price_data["Adj Close"].iloc[-1], + ] + + node_features.append(features) + + # Get corporate action flag + pit_data = self.get_point_in_time_data(ticker, date) + flag = 0 + if pit_data["upcoming_actions"]: + soonest = min( + pit_data["upcoming_actions"], key=lambda x: x["days_until"] + ) + if soonest["type"] == "split": + flag = 1 + elif soonest["type"] == "dividend": + flag = 2 + + corporate_action_flags.append(flag) + + # Convert to tensors + x = torch.tensor(node_features, dtype=torch.float) + corporate_action_tensor = torch.tensor( + corporate_action_flags, dtype=torch.float + ).unsqueeze(1) + x = torch.cat([x, corporate_action_tensor], dim=1) + + # Create edges based on sector relationships + edge_index = [] + edge_weight = [] + + for i, ticker1 in enumerate(current_tickers): + for j, ticker2 in enumerate(current_tickers): + if i < j: + pit1 = self.get_point_in_time_data(ticker1, date) + pit2 = self.get_point_in_time_data(ticker2, date) + + if ( + pit1["sector"] + and pit2["sector"] + and pit1["sector"] == pit2["sector"] + ): + lookback_start = date - timedelta( + days=config.LOOKBACK_WINDOW + ) + returns1 = ( + self.price_data[ticker1] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + returns2 = ( + self.price_data[ticker2] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + + if len(returns1) > 5 and len(returns2) > 5: + corr = returns1.corr(returns2) + if not np.isnan(corr): + edge_index.append([i, j]) + edge_weight.append(corr) + + edge_index = ( + torch.tensor(edge_index, dtype=torch.long).t() + if edge_index + else torch.empty((2, 0), dtype=torch.long) + ) + edge_weight = ( + torch.tensor(edge_weight, dtype=torch.float).unsqueeze(1) + if edge_weight + else torch.empty((0, 1), dtype=torch.float) + ) + + # Create target + y = [] + for ticker in current_tickers: + next_date = date + timedelta(days=1) + if ( + ticker in self.price_data + and self.price_data[ticker] is not None + and next_date in self.price_data[ticker].index + ): + ret = ( + self.price_data[ticker].loc[next_date]["Adj Close"] + / self.price_data[ticker].loc[date]["Adj Close"] + - 1 + ) + y.append(ret) + else: + y.append(0) + + y = torch.tensor(y, dtype=torch.float).unsqueeze(1) + + # Create Data object + data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y) + data.date = date + data.tickers = current_tickers + + dataset.append(data) + + return dataset diff --git a/src/data/survivorship_bias.py b/src/data/survivorship_bias.py new file mode 100644 index 0000000..c8dca4a --- /dev/null +++ b/src/data/survivorship_bias.py @@ -0,0 +1,89 @@ +import logging +from datetime import datetime +from typing import Dict, List, Set + +logger = logging.getLogger(__name__) + + +class SurvivorshipBiasHandler: + def __init__(self, delisted_tickers: Set[str], index_composition: Dict): + self.delisted_tickers = delisted_tickers + self.index_composition = index_composition + + def get_investable_universe(self, date: datetime) -> List[str]: + """ + Get the investable universe at a specific date + + Parameters: + date: Date for which to get the universe + + Returns: + List of tickers in the investable universe + """ + # Get current index members + current_members = set() + for index_ticker, composition in self.index_composition.items(): + # Find the most recent composition before the date + comp_dates = sorted(composition.keys()) + for comp_date in reversed(comp_dates): + if datetime.strptime(comp_date, "%Y-%m-%d") <= date: + current_members.update(composition[comp_date]) + break + + # Add delisted stocks that were in the index but haven't been delisted yet + investable_universe = list(current_members) + + # Add any delisted stocks that were in the index but are still trading + for ticker in self.delisted_tickers: + if ticker in current_members: + investable_universe.append(ticker) + + return investable_universe + + def filter_available_stocks( + self, tickers: List[str], price_data: Dict, date: datetime + ) -> List[str]: + """ + Filter stocks to only those available at a specific date + + Parameters: + tickers: List of tickers to filter + price_data: Dictionary of price data + date: Date to check availability + + Returns: + List of available tickers + """ + available_tickers = [] + + for ticker in tickers: + if ticker in price_data and not price_data[ticker].empty: + if ( + date >= price_data[ticker].index[0] + and date <= price_data[ticker].index[-1] + ): + available_tickers.append(ticker) + + return available_tickers + + def get_point_in_time_index_membership(self, ticker: str, date: datetime) -> bool: + """ + Check if a stock was in the index at a specific date + + Parameters: + ticker: Stock ticker + date: Date to check + + Returns: + Boolean indicating if the stock was in the index + """ + for index_ticker, composition in self.index_composition.items(): + # Find the most recent composition before the date + comp_dates = sorted(composition.keys()) + for comp_date in reversed(comp_dates): + if datetime.strptime(comp_date, "%Y-%m-%d") <= date: + if ticker in composition[comp_date]: + return True + break + + return False diff --git a/src/evaluation/__pycache__/metrics.cpython-314.pyc b/src/evaluation/__pycache__/metrics.cpython-314.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a24b10656c027ca0c08b7cbc92d4d60ab4809145 GIT binary patch literal 4929 zcmc&&No*U}8GhspR}HDnwkV3GsKu11)vBz;j3T$T8^=3E0~A~+V#nl2o|c*!zZuDv z4+RVK;GCKO0i5nS^x)i-#DH~Z3m0fnUkY$aPbf@Wz(8`zje%f6MvsHd8o+Qo^OpqT3mB6cZe z)9KTEM#&JZswQ&6nE3Q83!gknp+)=jPY}uiErjN5*KJAL0&prVVK+zhK573qd}@I< zrAI07hW>(B(VHa^D@$UI7t@$iRNKYoF9$fwKYURi04$qekLVX5GC}T~83kcv5{p^tiZ0uN1d(vb0N%lNk zz=?JDW-ZdJY;?5e`k&J6;{2W zn*AL+=zqHJn2HBbdVIJ4>OIBgGOIbL1QChDMgb8rT%=koB0z|+5e2hZURaLEFi^PV zb+0$UVisgLEXTcj3RgT*sY?-YDNI%QYs{LsH(9BMaWARWP{>%kZW!L{ ztygChC~CQCwbYxT$gQ$D{x+92N!~RPV&!wZ0##7H&SHs6LrlD=`(&uUya3hL_{prW zx6)b6-pq9l+UD>}ssOS&twx?)H;VF0Q+wZelWdrj7eA_sb&kcd-b4bB`7Fd)g%zzeiX|B>FcaAZ{90YG z0j~-ItUK{ucbmz;ePphFGM8hS+Vj}HF%s@8M4ki4SxT*;p)Fr6wL0i%aUN}(|| zG*%3a-@o=?`J?5n%*Qz`bQ0`dQ0aD%=!_Pk$F*>LYyRPa_IgqaC#!P;FA&?BQ)9DQ z+lc~wU^_Ta3XZ73kz(-Z{ck;(`e0*1lKp)y}=`FPksx5=Xmf`&Q z?WS(6>FDNpt?@+u&FygXH@TnXej%8ic@fk4&aBqZy)pN=A!?cOxAMKI%~@^mWU=qm zC+<@_<~H)^NA89~L#XIFykUFn>e}{qmi#@czo*bUQS?vVp=yMp{|NjF9{yGAj_V!g z+l^&63iT9vCbf>Kt-goR-;Zwl!W+$+ue-$bt4x298Mr6kzfw4Ut~haC8$GWv7arZv zn8b6+{+jnWvWL82#-+LgWiJX1-@Bo994!SW)!<|?I1N#qQ;i(@!kb-ccwA$Sm6&Oj znJzMMSZROf=6BS{gyx?t`QxfTUi8ltd^7nsw#TNny6y+xU;Xuu{?zi9Lw`I3hKvM% zlk^mrs7m*h8OUt@5++FG|NTzD1V-vjwHq}}ussJSymo0;DVX3<4CM(P1(<=Z@**%! zqk;B~&-Vy}z%c^(y9wfl&yz!eyZpO+nIhxyi3i6&JZ_8)xDT{SBCKeiIAx%HyK#I0 zGk*=XPt;I6XjNOS8ruQMd{bR2Wlzwi#d@D0SK_XgqV|(Cu@~TJ`+j&nMO#q?ylALb zkfIi9h;Y8=ST$>at>LBesyKMp%fwN64FQsdk<}xsy~P&`HRPP4%*d=b<>hz`>{i}B zj-BrjLf^)-ByKX z!4Z3fYxW~m;g_L^jNvISi-HMxy!xS)k+(4>s~3N zq%+r9VVTpNSu9Ed;09jEa<_Db7YKaf3`_%1K+G7G5FZ9ZI1Hbd6Hk)y6#3NaI|DAS zh(g-1f(L;q6K19ERtoNcN>m)UpRW0gPlA+j=Kl!;sIwu|9xDaL)xdZ$Fi{H3sDYVc zV75RX+z#~bONW|v;7F-;Ky4i;wN9z6Q-$fXrRmG+^yOk}B7b2!0Op3(!0^4C7I@9D zvx@*lyK$&6^m@@hXB7B>*S6-g*lel)wAz2V*#E|(X087msnxwIedMXzbLULqz>vla zYwqD~ukSsoaNtO>DXMv+&mH!W{-*0# z^{YI_=&qc&ybRz*cPi@=^v0hkxf@BHbymLY9>=+3C7xnNu!+28VTbIxPVB zO{WR)A(YC(k^Fvuj0>`|4DOi~F}s@MX7PDok(*5(hEJKY*=(O8|KCyDXQ<^<)cVxZ dRPqG><_W$(UG((kozI=*f0eJkZ#VDTe*vV^_xJz+ literal 0 HcmV?d00001 diff --git a/src/evaluation/metrics.py b/src/evaluation/metrics.py new file mode 100644 index 0000000..6c38449 --- /dev/null +++ b/src/evaluation/metrics.py @@ -0,0 +1,124 @@ +import logging +from typing import Dict + +import numpy as np +import pandas as pd + +logger = logging.getLogger(__name__) + + +def calculate_performance_metrics( + portfolio_returns: pd.Series, benchmark_returns: pd.Series +) -> Dict: + """ + Calculate performance metrics for a trading strategy + + Parameters: + portfolio_returns: Series of portfolio returns + benchmark_returns: Series of benchmark returns + + Returns: + Dictionary of performance metrics + """ + metrics = {} + + # Total return + metrics["total_return"] = (portfolio_returns + 1).prod() - 1 + + # Annualized return + years = len(portfolio_returns) / 252 # Trading days in a year + metrics["annualized_return"] = (1 + metrics["total_return"]) ** (1 / years) - 1 + + # Volatility + metrics["volatility"] = portfolio_returns.std() * np.sqrt(252) + + # Sharpe ratio (assuming risk-free rate = 0) + metrics["sharpe_ratio"] = metrics["annualized_return"] / metrics["volatility"] + + # Sortino ratio + downside_returns = portfolio_returns[portfolio_returns < 0] + downside_volatility = downside_returns.std() * np.sqrt(252) + metrics["sortino_ratio"] = ( + metrics["annualized_return"] / downside_volatility + if downside_volatility > 0 + else np.inf + ) + + # Maximum drawdown + cumulative_returns = (1 + portfolio_returns).cumprod() + running_max = cumulative_returns.cummax() + drawdown = (cumulative_returns - running_max) / running_max + metrics["max_drawdown"] = drawdown.min() + + # Calmar ratio + metrics["calmar_ratio"] = ( + metrics["annualized_return"] / abs(metrics["max_drawdown"]) + if metrics["max_drawdown"] < 0 + else np.inf + ) + + # Alpha and Beta + if len(benchmark_returns) > 1: + cov = portfolio_returns.cov(benchmark_returns) + var = benchmark_returns.var() + metrics["beta"] = cov / var + + # Alpha = portfolio return - (risk-free rate + beta * (benchmark return - risk-free rate)) + # Assuming risk-free rate = 0 + metrics["alpha"] = metrics["annualized_return"] - metrics["beta"] * ( + (benchmark_returns + 1).prod() ** (252 / len(benchmark_returns)) - 1 + ) + + # Win rate + metrics["win_rate"] = (portfolio_returns > 0).mean() + + # Profit factor + gains = portfolio_returns[portfolio_returns > 0].sum() + losses = -portfolio_returns[portfolio_returns < 0].sum() + metrics["profit_factor"] = gains / losses if losses > 0 else np.inf + + # Return over maximum drawdown + metrics["return_over_max_drawdown"] = ( + metrics["annualized_return"] / abs(metrics["max_drawdown"]) + if metrics["max_drawdown"] < 0 + else np.inf + ) + + return metrics + + +def compare_to_benchmark( + portfolio_values: pd.Series, benchmark_values: pd.Series +) -> Dict: + """ + Compare portfolio performance to benchmark + + Parameters: + portfolio_values: Series of portfolio values + benchmark_values: Series of benchmark values + + Returns: + Dictionary of comparison metrics + """ + # Calculate returns + portfolio_returns = portfolio_values.pct_change().dropna() + benchmark_returns = benchmark_values.pct_change().dropna() + + # Align the returns + common_index = portfolio_returns.index.intersection(benchmark_returns.index) + portfolio_returns = portfolio_returns.loc[common_index] + benchmark_returns = benchmark_returns.loc[common_index] + + # Calculate metrics + metrics = calculate_performance_metrics(portfolio_returns, benchmark_returns) + + # Additional comparison metrics + metrics["benchmark_total_return"] = (benchmark_returns + 1).prod() - 1 + metrics["benchmark_annualized_return"] = ( + 1 + metrics["benchmark_total_return"] + ) ** (252 / len(benchmark_returns)) - 1 + metrics["excess_return"] = ( + metrics["annualized_return"] - metrics["benchmark_annualized_return"] + ) + + return metrics diff --git a/src/models/__pycache__/backtester.cpython-314.pyc b/src/models/__pycache__/backtester.cpython-314.pyc new file mode 100644 index 0000000000000000000000000000000000000000..59ee519384b6bbd9f961c5d2108f9aadef8f292d GIT binary patch literal 8380 zcmd5hZA=?Umep>5V{8))*kEk)5fTS%5(ohjmV^l*VM9C~Lz=zW3@u|C;*7C9-EHzQ zx>=-M=@@A=L-uABn7t81I%z=Ky)dggO(N}SveHSHpLoEr)tS{X+S`*xx|7JkE<5|@ z?p3$j7?0=cbkbc}uCA_n_4?JT_g=mCs@YhoM^Ju0Q*q;OYzX}hzsNmKV@l!EKTq1bFiQD9$;u6V|mz2Enl8RS}eZ-)e*NE_1UJFmM9h_oK$4{OaY<+xwW6;SD%{@v=c#|;ltC6 zn-rA&;SetP%Jtbo`P1@`4}6Iu>ju};={3M-_0P)T=e-u5^AX? zvoAYwXar&1uOOZXAYL(MKVxJQC z@fxv5CHBC&php|f@Hz_d`hXVJr4velpcjInS9ylx88(Djk!1%`LmQK~(*jiwF^ecP zrs(}Ys97&RXd;SaB~M{pb(~2F1QWc(C%zPV-}5?qw)wEbwy=PqS52gtUw{$mNy-> z0#Bnkd|v#fk`Y`p$SuE7ZtXwoP~NzCzh7?o%jMSo^JKPpB>R;1F@UbXKFUe}6lmlb z0R`+OhG}-7Hgtl2y=9BCK!$c$3b(Hf|6q*lLh!W%$z1DEMAYGrD-fcu(_3m zn=AE;8#@G`vuN9ZDr?@-XQ~_DGG(0JO|3m!+iK>NZy7Vv~QEx5M=s#-j90xD4(VuZ^39O3$Y|F-yCz4~F>N8KNGKN|ly@+h+2{{58avIHS< zCMVi}q6GdR-i95e4(yqA7I{%z0e_Mr%;*9+0ERTKq?CRDD=cvp^ylNs0>ss=C}yS9 zW;8Gfvo)h}e@c)rON+KfXIIv zwo;lVp!M=~gmlP(u3?Ueg|dQ0$F^Jx-_0q%Jc^JX3wJ@0RoaZ?maNEejHQTSptLP{ z=Sf^Q`esWW&bXdZ6f+%^UdHQ06j>BPwGVL(i1WhCq36jqB0CuGx zJcGglYsu3Hc~t%)NEGs8DK03KQikw53x?oT=9s_kYsMPBdMsWNAo1^H)|@-RoW~8< z4$CtrwfuESa&LlcvuR2*Y|WE0u%xhJED!~SDh-$p+5x}200ds;HN=fSD4|O8Sf}I@w#{+ZBU0@}qKt8f_07F)%#u0c%*?C#X zr3_y$kr#-8Um+x}`bU&0pb{zg4-vNx=)u(e8-BwYkxxb3e63ktDWwZ&T5?As@X9QGD7KB_qQbpE#ecezzeQ%6SwDhgKGp0Ba8AR&LJJa_*Y&{Aav z6wpx`sP&XNP*RH8a?AE?(xdv(wcKj zL9gxT+CIkKeN1wz4$DjtiO)$#l@|rdUL>fhs7szhRZ>;$;7xrKW?32HUxUMinvqQJ z6s*gR!U`w2Qcm%AfgIADA!|{Xa38w!V|^eIb;+TMD;9O}*BRGTc!r6(vi_W2?+;>U z4Pr3vv_9K;+J!?)S8USdzkSh#9XeNR{00-^IlVq0x_LKcZ`dKuB23%RoeTX84$%6ArF$~^ZzQxYh zo?B=ZqVmi2l2imx1L`2d#46!iMXO3d;XN*>`Eck}h85J@bqI(#!5B)gEEDBvF;3(p zm*3@@m0W&Z&lD47ZpGPN!E?~(XA@CcW~#R*5(@<*T(36=j_sJ_)qW03+D1y3@w29$ zyT18>Ou21gELq;Xt3qYR=hX|26?Mwinl`t7Vs3rb)bqs-7VjUIG!uVqZ63}V(L2Yx>MqkmV>wAj?xu`F4EgJc}F!Ii%huw zFL=&krDZ+1!sp~czOs1JSuE&x@R;cRXZd9RwtQkcU@9a!RN?3ZbL;c}f+92PW4B6B zf*3RGaexujK`{fu;?D>wai86F>>2!~k)(l>VoEG&0ttI5kX&M#tr6kFJTuL))wrz= z#^X$Mg2m3BpozxlDF`tI1w@w)=n;cQLBqY4V3<1$YZGTj6Vs3y&@u32a56?9ZV5_^ zQXs=oo&}c(3FeVNpKrMD{K(*oKKlHN!y^JY#qg{hkEtPkodR)HO)%q$DS?c~;w*Nb zSgif*Q4FF9#UhF6DCf4ZF5KCOJI#gq2pHh=oXD4;w0M&&`Z5zlZk5WtQber`1-a`& zmE2i?Rw#+H3{H?hhKYlT2LcfmDnfGpiI&(;Pey`Mpmrh%`#nr#sU);CAEQO# z4#hZL+9hqAjfEIy0z%0sb2A&<=dFp8;!H!%#?TX!EVjoSW_weUt7thVDYYSAcMB6W z`D0>eDu5dSq3N~-*&XXzYsP$NPW7*rnz{Z=dDXq{yWP-m)GtLABXciqtC7Ps=X+M| zo>zmKFVTzi)AE)z%g4?~&L@?rZhyL)N_JCEPf!`pi3jv5{g6uGi2qX3bLojCOy*adC*)%wI z)#y;wz1pAF-gExUnW?I|ck}Me1%B!E#n(5gn&*{xX*U*ctXMW2-n8RX(s63T(Uo@e zBpp2)j^4DRKk4Y-a16|s?5WV9Iv}FTHh1yU+NR|nJ+1A$tD7f8YxlIgd8_64gMrn7 zwG(UOYaPG5l5Xiuw)AeaoLjfJfj(=)@`<0lv7*|n@~pRAd~#y5Z76fBbwTrU)0VCA zeaA}ErtSFp$-yUmPfuRS96OE&o>w48*LF2(IQD_(zUPD1yGD`c3s>JgyUecy|Mu3F z{n*N>l)Y{K!l#W#GPc@f^8M-+)mqu#Q|XSOWXDkI#8B$!m6X-DLn;ne&R@Y~9A18H zv+8KZ*|L&IIZr+$em#)x8BX>Lr@Dqy&XEP>mff}FTlA&v-lW~TGPCyDhQ0gYm6UyO zn?$zSCF7!Tc_wM~0xI~P@_oaCGGn#l?zGjDw0bhNbt@ffp2rpI_KUlu!d4C3w>g&# zi-wHj@bdIBy>|Jr`r)m|{F4jofvZo>tlLNbh{2whBdZ(MT4`UHNmm?RY5K_bp)Y;x zg{|7UrLo1abZu+0wslRtQQJ8`uw`{DJJ!^xy3Ulf^WlkK&!o>%$+J}I43)Bu0+|)H zKwX7x;Zn-dINzVKSnv7n`qCD6(&ENUnt@fl@pv-Td1c+|0~#wTfyRoOdqa1JGS$vy z^}?;?*Vj(19evpN@Qo*vsk2wtslQn7eRUpncMft$EYTT&loEs7v0&_!~1{WYRyO!XcoR} z*y#taF{qh^tU(O&^nU0rEZK~gV$@YEFNV*`yo{H*7z}kFlUOX9m4`*jcx8R`d|C>y zsPo69_CbHJAmgQ>6k@(1n>V-ynZfKVORI%VlnvOu! zQQ9&$Fh90Dkg8}-8=5~cG(W5F`a<-kSqkR6HKOCY?-8?Sp_e^}t3F&|SJ8eb>{+JYs&D$hd*7R`Kb5RMwOQWi9-f^V;)5aYw`m7b>I0y7c8ZMd|2EfI`lM@%$58D=>? z5{@zu-q2EfVjA=fjtuq<(dYa8gCl)IEKXUoAzZx*l~4hF9=h>1jK+D%CJKsW&^$pDM9U}6Pk3~@Lt{L^3IfLV{~kEZK7lXabY;^pUmiZ`Tt zc6>7=-3#rFlW@ZQF?Y;~SlnhogCnMJbV@v4A>6+m$1h6!9ensnmsRR`Fgg+Bq)-aG z`O-=)aJlQ1E_c0C(I{pU6A|x-1sw?MkQj@x(-{5=u0-m^7#H6(NYRD1M`CU-dmH-j zQaNm`w@HE^o+0CJQSHB=#w2R|9V&l@48KRV-=eN9ojI+m{zO;3(6ym!oK^3sLkgm1 K2jN=W<^Kj3p^k3= literal 0 HcmV?d00001 diff --git a/src/models/__pycache__/gnn_model.cpython-314.pyc b/src/models/__pycache__/gnn_model.cpython-314.pyc new file mode 100644 index 0000000000000000000000000000000000000000..74e4ccb140f2a76fdeaa43e68f868491eb599249 GIT binary patch literal 2467 zcmb_d-EY)J5a0FLzK>kU2T6{Q&`U_0^hy;(K&1lI5>c8X3L!m^d1#1he7@u=#}4Z~ zzFwM#N==%!M7$85pifo*gZ>$;C+OO2}hWBGPKY_8bVyWR4W*8Iq@*W9dYBeyhPVlXfTW=0w!4(8xF)U@RIVG8 z#|?6v@&x49Mdp;xmV9P$XWZsx-y6SXF=t|W+NO{PKcobFxbW>UU?(88LFSbAza*px z)_(E{U4XRb>8GTiNFI}11x^pDa+?ZF6lt%rXqFs{SC|uswuN%B=y;~RV0oV71|o^W z3yxI`M7+rSCBMR#WU-tdi40;szH8NnOB+b?VKhqi6 z8vnr^Yn-D;@YsN0mGU~G+!Uh+^_OXo7;4kr($Q!^0hGlq!f|3Qc37HYOC$6Z%xyb1Wv5k}cvNg}9XQ z;r6J3d4!7GyGavIbB9C)RQpX_8$ zV$ZLMy$`$ltLZ(}bYHczuiDjD-JN;XqW7d$)NfMY*OT4&;z9P5LL zpgGVPX?~f3wRui$kc}B+rVi1@jEh4NIF=Ofs6(vbVJHz0mXuk%UXo_&K#D8`I&!lw zL~Ct~MbZwBw#eT2?^qF@?&0?_GPO3>>fjNg}$0Fvq6$ww)kLG>I27u*MwjX{%En3`S!)P7>7U0Dg9v&EwxD|2ez3@5|MfO$L ztvO(<;9FqJq`Ig7o_<%qul~^XecQ^hYP$W_s= za4jE>Os!6BbgZ3ga{I57CRQib$JZtyLRu2d|yHqvIAr7{b+3;n|Lp~v7wX9s|ZZfFDW z_>^ta;+&N)SGW_v{Entq%ywKCOwq#@syT%XVGwMr@{v1vGVCGPQy4P*0!;MxS5WTf z3cNml;zRZPy%w$G4+2c?TlxYz65MC@!l>u5H(-v~9jT+GiYMQH%N41H%-%+?PU&UN zJN}{rtG97F_Rp|!>Stv+WQHF}myH zlY}FGwB7%DZSKAA-0yqO`Q3BQRbyd+i9mRDyyWWVWrTczA4<@qGh2THnMrbqc$wow zlv74z;9GTEHJ}nzHAqt*R}W}J&45UF2d>J0pLK))L|(UY(tI_3f%OZ5O2(fM(3lc6*sx;x>JS zRfgP*WC*e%Cx*s2w?@+Rh619bIS~>>NqvSBCGFYqSeS!axS-_AB*zP zs3#zXqLH2(EYH0#G!&4>l;iXS$d(7bS=gY9L}a`~R9z|()ukj#eB==0W$Ywu@TwrT ziW*8$Q;LRCv|g>9h&o)7sQ2ncgO7+tuK^}+luR(-S)L6=I6iajCzK7g|y-(sJ_Dbl6EW_Aiz>nHBmlh_@U(4kedOU75b2O`$^18K+vIWmBk$ zr%gjn%{>Y&n?lXZjJ;~$N)W@qc#xIUf!H`mo+``_xr-&O807<3BwdiZ9tv>09XgOK zA;HgHXG3B3a+s4;Py=^R(nez-&yuDm$c{;-7|XL`oXGJ4Z@|)m6Ma42q2og@_`N-6 z`y?Z+9Gc*GNfn7m`jcn+PDBMkG6wh%Gy(yaj3-a^_MPze9`{N{%KH9bh?h(JS~{h>$*WJxOCAs;f4%Er&3-XvKst+*5ZY50S4 z_g`LmIoaGBZ|+St_kYsdfA{cj`x2$zDb1gnQzec&FWr7=HaLH6?pmUx0cfW6^6L4D zxr$`DJ6`TiS*$ZJPQRG6xZ)O9!czNzF1f!uzQ21zO&kqdMCC9+lftic#8El#_=V&B zO8i_euaR5R4~H}9XpB1=r!+q);!$jOKn)54g__Hdc^uSuRR}mq&}IbyX!CV8ocYl4oH))$+-e>rjz@u)v?JjtE4o!YZZMB> zlZG;SD10$eQBH6|LI*??o7%MdHwi?)t8#6qF+I>u)888lr^VZS@4Pk4BqXKIX=Tr1CZEYu zqtLUdGIM%T`A*lBt*hj+ZO&iy)GE|$s+2bsWN9&7A2ZlPh>=vo{i&Q_5KbR53g`1ef%s#b8X{c0()Mr+9l&&G4Y`Al(r?bi&sccw=8APnnm! zD*V?DNP_{nL1ED}Z(DMBTcvR7cX+$Nr^(CP1xkw=GTb7ML5Sy|lFO!EwNHDA$z@&( z@oFA_w?GZ8&+OF|fU)E&z_iVf9j{M4cqA>a#B(S+l9J0x(E0X5s#@YCjYL`dC?Itu z34z&H1$%=3PP}tOT}48S+c2thATc+GCUkwVXyA%dVp=B(o1D~`aiZqc<(yErw2w?! znw&m1oF=&)l9p#9qnxb6>H0jUjt!hit9w)<;1L9Q#9e+5M3NCSx~PAYXM>Vncy*lP zCOAodg%wy)mDu-7?1)d}5&JR%H_N65Eg8;h*KgkhTYRL_KG_Q@@NQ&p~c-<)s$!rX-g z@!soqUr$sWy49buJ8+P9Z!Yw__hzcvH9s~tw%C)XZkbi5>{WLrZcn`X#=@&B_Wdc> zzNG75+;#AQE8%LN)uk$G=8w)DU93%19KO|?vX$LAdHdvhr!X^gdnjq!7q{(OvAI*W z>U-|H?!RbUwKc!{;${UYuivgFJ7A;dR^6%jS1F}mDS&OFS-A+akcE3UleYZkgEOLYO>q4 z?SRFfU}legaO{CO(bTrr8H@&id(8?!z-4?FD<)g2B!y8yUS)v7Iyur zeAc{h@YlY37w=wNIQ8J@-?u=M<(_lP=h)@Tk>x1A%nHj^aa*k_DBL{7kkWqyfWRdlubg z+-u++^GC0Jt8R&j6GbP{cjQFC3p&N9bI5IxR8c`PjIq}^Fp~uyZw5&dW5p|y=4vz) z;ZY}(;ZuXnWDqwA)OUo=CcYOd^g*21@k_B&m-yOPEG3`v6? zV#;|xe+tg?AfG-I;AHdg7f{iC@lXJ+!E;KVYmCX^(KIc$L0+XUqjoD-c<>xm&GVXs z`CL@P89wx6968xhUF5tquOliWSy?U{&X{uf%p>i3b=}qn&iH=8tPF4gSFXVmDo=3jqFJOdwM;Tkv zQTY&Lt}eo^{5h;mwFbDs*`TbdWLRC&gVGTRMMil%LfLO5GrMG>ozOY#;Q1{1EH;^h z2zdT!INV$?(=pvKb9DOX?B1lg<`Z*G<~Fa7Tk3y3esAK>Cm!tm%_|?jva;{=ip9Gw zuslXMNdaN$n;*j#ZmYm4y@8(T zZ5udtHqa0BZ&w(Nwq@I*({^qUi0DFT+W<7n1V!eyt^jF=+yCOzyZz}y0J#3)X#iaR z@HC*K6g{QDTLw_{USs+>zy!E#mh5N7`Rk$UQC_$biaj4 zqQ^Px&{3D`prZ~vRV!%v;8{aeQnCW(!!5uC{jl@bz_lvv3Ppn4YyLoVEEW|)C{0io zVN6mBpm5{ZkM+Sxdmpic9S;bXpc?^w8Q9=PXZ=IAeF;;XV$^*C9p_l2rix&6!Nb^1nZz0e1;(4CW1CCGZ8YS6K)3E$VE|XlF_} zAu!)cn6Kgd4`ABL9AI7G0L`Pp`EtzpL1206eA5GOt0vwzy}*f1@e1b*MXqxIp73_! z9FK%R6B9VR-dLv);{u_PP{4^CWHNa&FDP9mq;xrZA%^>qciy-X0vjImKCnz_pGrq2 zlb8D9g)S4mraGe|&NNiOu#Xzu^sS-pLUY|%b6pVBNHzHPsRT+|G{OZ*Cq#L`97%JH zy9pX5&(SFddE{-$05F5kKAfZ*kKruAP=4}4<)bKRK9OU2pk^DCcp4jH4EmX^FZ8r zAmMCFlpg{Uy{c~E>RQ#oTgLU`ie&NLPm1?0TzFtg6d(CeTq{17vehPS4RKq;;*r(H zo)z2klC6A8PpaEC^`v2d`Fc00sa<$3QGICHP_dyS_KIX_<4S4cRDY_dWUXj_%37PW zHpHzBi-*^&ZJSzRZ`v>rYwfaO&v&s7Hwn*f%L~!}{6eBHN@}kr{eHCL%4((dihNn= zl{?B5^nK=%eZaQ zBVp$e4JNVo>kxrgu#!~OEE~$dD6N|6|J+)-X59<^vQ=CCimfeGwr91hAz5|+{(n+- z05GDYeBt_PWlOTMJzm-V;Ph%`$C|ApRaUupc-7gKbauv_ogeOAb#||nb$@QFOWK;^ zwx&hjnynqYPxk7KGGcWt8=T*TkF=sew$S(eIQWoNB?=d+oTp^EUpYXa=2110ZxhYs zXf*s}5co6PlUj%<2-vhGn*^ediZ?L%7?dSQ;-}A1s=NTiJ0)Z^B^4EU(2A6joYa%s zB28B`9Ds*E%4kIyAp(KQXh_+Pls1Y!0$Nc12yBoyv-%&Q@>0o>Kg^?cLaL@xoeLHx5pw8Wl zAw>fpvYLQ0OMh4*doT_3S4l0A2w#f@=%JHr@DsTf>aG!5Mk5jV9f(JPl8kVuC+JI& zy!KvfMBNZLhb1D<2fb*Lp<8x~9OduACoWky4$+31VVF-z;its-8L9eb;!Y6v7sT=z vDft&t{u$}~+-OZ29iJE-vz-a!-bvlIZjfO*Hwnhu!>V%(Q?f-crbGQVzEmyQ literal 0 HcmV?d00001 diff --git a/src/models/backtester.py b/src/models/backtester.py new file mode 100644 index 0000000..bbd3d06 --- /dev/null +++ b/src/models/backtester.py @@ -0,0 +1,177 @@ +import logging +from datetime import datetime +from typing import Dict, List, Tuple + +import pandas as pd +import torch + +from config import config +from src.models.gnn_model import CorporateActionAwareGNN + +logger = logging.getLogger(__name__) + + +class GNNBacktester: + def __init__( + self, + model: CorporateActionAwareGNN, + price_data: Dict, + initial_capital: float = config.INITIAL_CAPITAL, + ): + self.model = model + self.price_data = price_data + self.initial_capital = initial_capital + self.portfolio_value = initial_capital + self.portfolio = {} # {ticker: shares} + self.trade_log = [] + self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + def run_backtest(self, dataset: List) -> Tuple[pd.Series, pd.DataFrame]: + """ + Run backtest on the given dataset + + Parameters: + dataset: List of PyG Data objects + + Returns: + Tuple of (portfolio_values, trade_log) + """ + portfolio_values = [] + dates = [] + + for data in dataset: + date = data.date + current_tickers = data.tickers + + # Get current prices + current_prices = {} + for ticker in current_tickers: + if ticker in self.price_data and date in self.price_data[ticker].index: + current_prices[ticker] = self.price_data[ticker].loc[date][ + "Adj Close" + ] + + # Calculate current portfolio value + current_value = sum( + shares * current_prices[ticker] + for ticker, shares in self.portfolio.items() + if ticker in current_prices + ) + cash = self.portfolio_value - current_value + current_portfolio_value = current_value + cash + + # Store portfolio value + portfolio_values.append(current_portfolio_value) + dates.append(date) + + # Get predictions from GNN + with torch.no_grad(): + data = data.to(self.device) + predictions = self.model(data).squeeze().cpu().numpy() + + # Create trading signals + signals = {} + for i, ticker in enumerate(current_tickers): + if ticker in current_prices: + # Get corporate action flag (last feature) + corporate_action_flag = data.x[i, -1].item() + + # Buy if predicted return > threshold and no upcoming corporate action + if predictions[i] > 0.005 and corporate_action_flag == 0: + signals[ticker] = "buy" + # Sell if predicted return < threshold or upcoming corporate action + elif predictions[i] < -0.005 or corporate_action_flag > 0: + signals[ticker] = "sell" + + # Execute trades + for ticker, signal in signals.items(): + if signal == "buy" and cash > 0: + # Buy with 10% of available cash + price = current_prices[ticker] + shares_to_buy = int( + (cash * 0.1) / (price * (1 + config.TRANSACTION_COST)) + ) + if shares_to_buy > 0: + cost = shares_to_buy * price * (1 + config.TRANSACTION_COST) + self.portfolio[ticker] = ( + self.portfolio.get(ticker, 0) + shares_to_buy + ) + cash -= cost + self.trade_log.append( + (date, ticker, "buy", shares_to_buy, price) + ) + logger.debug( + f"Bought {shares_to_buy} shares of {ticker} at {price:.2f}" + ) + + elif signal == "sell" and ticker in self.portfolio: + # Sell all shares + shares = self.portfolio.pop(ticker) + proceeds = ( + shares * current_prices[ticker] * (1 - config.TRANSACTION_COST) + ) + cash += proceeds + self.trade_log.append( + (date, ticker, "sell", shares, current_prices[ticker]) + ) + logger.debug( + f"Sold {shares} shares of {ticker} at {current_prices[ticker]:.2f}" + ) + + # Update portfolio value + new_value = sum( + shares * current_prices[ticker] + for ticker, shares in self.portfolio.items() + if ticker in current_prices + ) + self.portfolio_value = new_value + cash + + # Create portfolio value series + portfolio_series = pd.Series(portfolio_values, index=dates) + + # Create trade log DataFrame + if self.trade_log: + trade_log_df = pd.DataFrame(self.trade_log) + trade_log_df.columns = ["date", "ticker", "action", "shares", "price"] + else: + trade_log_df = pd.DataFrame() + trade_log_df.columns = ["date", "ticker", "action", "shares", "price"] + + return portfolio_series, trade_log_df + + def get_portfolio_composition(self, date: datetime) -> Dict[str, float]: + """ + Get portfolio composition at a specific date + + Parameters: + date: Date to get composition for + + Returns: + Dictionary of {ticker: weight} where weight is the percentage of portfolio + """ + # Get current prices + current_prices = {} + for ticker in self.portfolio: + if ticker in self.price_data and date in self.price_data[ticker].index: + current_prices[ticker] = self.price_data[ticker].loc[date]["Adj Close"] + + # Calculate current value + current_value = sum( + shares * current_prices[ticker] + for ticker, shares in self.portfolio.items() + if ticker in current_prices + ) + cash = self.portfolio_value - current_value + + # Calculate weights + composition = {} + for ticker, shares in self.portfolio.items(): + if ticker in current_prices: + composition[ticker] = ( + shares * current_prices[ticker] + ) / self.portfolio_value + + # Add cash + composition["cash"] = cash / self.portfolio_value + + return composition diff --git a/src/models/gnn_model.py b/src/models/gnn_model.py new file mode 100644 index 0000000..59f9f7d --- /dev/null +++ b/src/models/gnn_model.py @@ -0,0 +1,45 @@ +import torch.nn as nn +import torch.nn.functional as F +from torch_geometric.nn import GATConv, LayerNorm + + +class CorporateActionAwareGNN(nn.Module): + def __init__( + self, + num_features: int, + hidden_channels: int = 64, + num_heads: int = 8, + dropout: float = 0.6, + ): + super().__init__() + self.conv1 = GATConv( + num_features, + hidden_channels, + heads=num_heads, + dropout=dropout, + concat=True, + ) + self.norm1 = LayerNorm(hidden_channels * num_heads) + self.conv2 = GATConv( + hidden_channels * num_heads, + hidden_channels, + heads=num_heads, + dropout=dropout, + concat=True, + ) + self.norm2 = LayerNorm(hidden_channels * num_heads) + self.fc = nn.Linear(hidden_channels * num_heads, 1) + self.dropout = nn.Dropout(dropout) + + def forward(self, data): + x, edge_index = data.x, data.edge_index + x = self.conv1(x, edge_index) + x = self.norm1(x) + x = F.elu(x) + x = self.dropout(x) + x = self.conv2(x, edge_index) + x = self.norm2(x) + x = F.elu(x) + x = self.dropout(x) + x = self.fc(x) + return x diff --git a/src/models/trainer.py b/src/models/trainer.py new file mode 100644 index 0000000..0ecdb7a --- /dev/null +++ b/src/models/trainer.py @@ -0,0 +1,178 @@ +import logging +import os +from datetime import datetime +from typing import Dict, List, Set, Tuple + +import torch +import torch.nn as nn + +from config import config +from src.models.gnn_model import CorporateActionAwareGNN + +logger = logging.getLogger(__name__) + + +class GNNTrainer: + def __init__(self, model: CorporateActionAwareGNN): + self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + self.model = model.to(self.device) + self.optimizer = torch.optim.Adam( + self.model.parameters(), lr=config.LEARNING_RATE + ) + self.criterion = nn.MSELoss() + self.model_dir = config.MODEL_DIR + self.model_name = config.MODEL_NAME + + def train( + self, train_dataset: List, val_dataset: List + ) -> Tuple[List[float], List[float]]: + train_losses = [] + val_losses = [] + + for epoch in range(config.EPOCHS): + self.model.train() + total_loss = 0.0 + for data in train_dataset: + data = data.to(self.device) + self.optimizer.zero_grad() + out = self.model(data).squeeze() + if hasattr(data, "y") and data.y is not None: + target = data.y.to(self.device) + if out.dim() == 0: + out = out.unsqueeze(0) + if target.dim() == 0: + target = target.unsqueeze(0) + loss = self.criterion(out, target) + loss.backward() + self.optimizer.step() + total_loss += loss.item() + + avg_train_loss = total_loss / len(train_dataset) if train_dataset else 0.0 + train_losses.append(avg_train_loss) + + self.model.eval() + total_val_loss = 0.0 + with torch.no_grad(): + for data in val_dataset: + data = data.to(self.device) + out = self.model(data).squeeze() + if hasattr(data, "y") and data.y is not None: + target = data.y.to(self.device) + if out.dim() == 0: + out = out.unsqueeze(0) + if target.dim() == 0: + target = target.unsqueeze(0) + loss = self.criterion(out, target) + total_val_loss += loss.item() + + avg_val_loss = total_val_loss / len(val_dataset) if val_dataset else 0.0 + val_losses.append(avg_val_loss) + + logger.info( + f"Epoch {epoch + 1}/{config.EPOCHS}, Train Loss: {avg_train_loss:.4f}, Val Loss: {avg_val_loss:.4f}" + ) + + self.save_model() + return train_losses, val_losses + + def save_model(self): + os.makedirs(self.model_dir, exist_ok=True) + path = os.path.join(self.model_dir, f"{self.model_name}.pt") + torch.save(self.model.state_dict(), path) + logger.info(f"Model saved to {path}") + + def load_model(self): + path = os.path.join(self.model_dir, f"{self.model_name}.pt") + if os.path.exists(path): + self.model.load_state_dict(torch.load(path, map_location=self.device)) + logger.info(f"Model loaded from {path}") + else: + logger.warning(f"No model found at {path}") + + +class SurvivorshipBiasHandler: + def __init__(self, delisted_tickers: Set[str], index_composition: Dict): + self.delisted_tickers = delisted_tickers + self.index_composition = index_composition + + def get_investable_universe(self, date: datetime) -> List[str]: + """ + Get the investable universe at a specific date + + Parameters: + date: Date for which to get the universe + + Returns: + List of tickers in the investable universe + """ + # Get current index members + current_members = set() + for index_ticker, composition in self.index_composition.items(): + # Find the most recent composition before the date + comp_dates = sorted(composition.keys()) + for comp_date in reversed(comp_dates): + if datetime.strptime(comp_date, "%Y-%m-%d") <= date: + current_members.update(composition[comp_date]) + break + + # Add delisted stocks that were in the index but haven't been delisted yet + investable_universe = list(current_members) + + # Add any delisted stocks that were in the index but are still trading + for ticker in self.delisted_tickers: + if ticker in current_members: + investable_universe.append(ticker) + + return investable_universe + + def filter_available_stocks( + self, tickers: List[str], price_data: Dict, date: datetime + ) -> List[str]: + """ + Filter stocks to only those available at a specific date + + Parameters: + tickers: List of tickers to filter + price_data: Dictionary of price data + date: Date to check availability + + Returns: + List of available tickers + """ + available_tickers = [] + + for ticker in tickers: + if ( + ticker in price_data + and price_data[ticker] is not None + and not price_data[ticker].empty + ): + if ( + date >= price_data[ticker].index[0] + and date <= price_data[ticker].index[-1] + ): + available_tickers.append(ticker) + + return available_tickers + + def get_point_in_time_index_membership(self, ticker: str, date: datetime) -> bool: + """ + Check if a stock was in the index at a specific date + + Parameters: + ticker: Stock ticker + date: Date to check + + Returns: + Boolean indicating if the stock was in the index + """ + for index_ticker, composition in self.index_composition.items(): + # Find the most recent composition before the date + comp_dates = sorted(composition.keys()) + for comp_date in reversed(comp_dates): + if datetime.strptime(comp_date, "%Y-%m-%d") <= date: + if ticker in composition[comp_date]: + return True + break + + return False diff --git a/src/utils/__pycache__/visualization.cpython-314.pyc b/src/utils/__pycache__/visualization.cpython-314.pyc new file mode 100644 index 0000000000000000000000000000000000000000..e560fe38cbbac2394e6ed83c201ccf6324e43dc1 GIT binary patch literal 6607 zcmcgwTWl2989uW!d-eLV2HV7tc(7w&ftV5lrf!s!ON=3jYX(@<)WIF^j@J{$yPGq! z25&F1P$aA<7e^rql_E7SmEa+2q$W~&c`%h)^`Q^8utjo$RH>z{cr$e>l(+u>nc1D$ zWx%1TdZa!3pL70m=G^}8`_I3lp*}_+eOhdJ`Q>^-{)UBe5*=Z63WOPQhA6^OGALNW zV899tik8@cb%CSOV9*LVa@dOSx~LUX>J+h!C{i0Snhe3JgK`kcA*URNau~`Hr(6%^ zD3oJPxj_kZk!1WfP9iBX>9C%*nDo4ES}gLtf~Dtksxg!npcQ|Z!tld6ua+R2At`bu zu>UneGEldP92X`8MHqn=I}HcQlKZa&)@T<;q_Xx%s4VSlFH3~%5@F8LZ5*3574p;? z&l`GvbV@afn#r2Rv|M^(Ql)3hHIZ@M&~oaemXuhakYS+{8r3zE25|t>=rJF_7&`T>Nu=#hCFz$ zHp59OLZ5Rs9pp3_uIwESW~Vd6yL&t9%Jqa4#M6m~>>+t(R5=}Y#LXlnDgj08b56d) zy8?fGNs;cOj~iQ)V8y{bmGye&!Y!;?8FWLwy76i~EP-y=SGU1WcaKw9wK}{CH3>$I zxTB^5Tb=ey4_b?7_RTdNp$Iz}aQk|6*Zb+ZGg6{m-jQ3^)=t)yej6J$@|IjNCb^(d zZV|b(CQq4i#inAhVU>a@SsFF_V-hzgi~WdOScS- zh0fcb4^2D&hL+WG875_^o?&5J3Up?{w2?Pq0COULE*YeF$4ZE7+Iger z&|=S_`2)*61Irr+7NiZIZ`^TD-7P76aqz<6LjC3|8<(ZMH>-=-di284g@z}Wq+V=; z56n(2Nv)2&(71U?>U0XPJ@53MJ^9+uKf{eTBXw`ry>3tx9V&pFM$x@&QcI|F~tS?7+L5-NAW<{Y=LKj#Q;U(Zgj#W~eCY%R_K?PNKrG`et194k%(+-iowqzphk3l^xJ z)|ilG;z4K<&Sb}V`o}ZiF9y?jBTp%=h=tIp)F=z8Mqxr_0%g((97dQbMNc9!$j<{Z zY(UlEFphbWJ^}^cil8=Gu%_uIDA7l;J#3~`3pkX@ShX3mPRzP6dmOWF%(h^L2MAEJ z6J<^c60E>vL0p1qB7C`^-@rk?vUl_#v~sxlD+n`WrHMom7hk>j>ddQOL>dsPJ@6lI zxqN=PW!sf!7h85O*F8CN`15$nhmqyhpUi)YfnW%{Ti_I~Qa?NXX8m08^6B?}dga7tp?PybIc4iaPzN|-_Uw|> z?n4o?=a!_l|20EQ@C|q{L%0{-eS-7`ssJhAvV^Pc)!60Xi@J1vvXD1*ymw{B07NF7 zAy{zS$g38oj0dp=;=7sTv%gCoL$-kIW}@X%`h&f5FE1xHuZURj#NJQRAH8%3<$Fm3 z*^86LEpVOj(N+e9tVgUN&;5p#WD0z$&q=M+9EqHg^qa7@W9w*rs#!sJALTBDk!hL8o zkgD8we=zsD@3l~B4+8nHFSw_?Q0nTg0q&fvU2R{0FYjPi)gA3NYz5Ot4~jFo0; zU3$-DtcF@o1Ar|=0N}c+bQ;n2N0y!j$m^HSX?k|T`pLuNF89~U>p}8CFkVGI;lYEd z(*#-&ZsX(bDhNzC&%LKDwmTbn%1|fAGV0T8r|sC59sN&(`rGfnH}nUf%faNAAb{(+p2XW~SODO=UH-7{ zhP-=G-hEx}{qr?aTkP6oKWj`8(IA^U)h5}-Vfd&CM7{vo?L+U0*xY$R3OD)!;VVX4~3WxxQ|x`{EGT%`@7PS_4s`3(Y1VQ-)bP`|QTD2Yxn2gCLrd&6S!5?C;1&D<)cMa%K4h z&Sr=zaLmcL2g$=K%_1-}X-uUPEJ%S5G$uk|nT4himrai^UxV&J`VtIB&tS$OKITBY zo;lA1y*6^w#zDNiAMQ5NvoHWhfU6)t0H?JkfYZF`Mxtvm(RDr1J#);7th{mRQtw>p z-Dj@XcYXNca`$)U_b*F4>H?ns8^uc}F1NosdcCpx+V<;PU%0mYv*RDDH@59x*tY-L z_QkC)EJ;HZ#O_>@x*Vt6j^TU*_)d^*LjotM-;>1QrdU!KN;a~PHC@niS&Fu7V;^D( zOF;Y`S8nKI5NAh;h3tQR3@n%{P8O!=DC&j^DuAcS<1aj}#a}v{;`lDnqc}Fki^|@^ zWIj_gw5KWVAG#PEMOQ>Y5dK28E|RT(Cmmms<}b;*yCI0n-yIMj5OepK0PghNgF+c; H`H24kLCEZ2 literal 0 HcmV?d00001 diff --git a/src/utils/visualization.py b/src/utils/visualization.py new file mode 100644 index 0000000..e4a0f2f --- /dev/null +++ b/src/utils/visualization.py @@ -0,0 +1,137 @@ +from typing import Dict, List, Optional + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + + +def plot_performance( + portfolio_values: pd.Series, + benchmark_values: pd.Series, + filename: Optional[str] = None, +): + """ + Plot portfolio performance vs benchmark + + Parameters: + portfolio_values: Series of portfolio values + benchmark_values: Series of benchmark values + filename: Filename to save plot (optional) + """ + plt.figure(figsize=(12, 6)) + + # Normalize to start at 1 + portfolio_normalized = portfolio_values / portfolio_values.iloc[0] + benchmark_normalized = benchmark_values / benchmark_values.iloc[0] + + plt.plot(portfolio_normalized, label="Portfolio") + plt.plot(benchmark_normalized, label="Benchmark") + + plt.title("Portfolio Performance vs Benchmark") + plt.xlabel("Date") + plt.ylabel("Normalized Value") + plt.legend() + plt.grid(True) + + if filename: + plt.savefig(filename) + plt.close() + else: + plt.show() + + +def plot_trade_log(trade_log: pd.DataFrame, filename: Optional[str] = None): + """ + Plot trade log + + Parameters: + trade_log: DataFrame of trades + filename: Filename to save plot (optional) + """ + if trade_log.empty: + return + + plt.figure(figsize=(12, 6)) + + # Plot buy and sell points + buys = trade_log[trade_log["action"] == "buy"] + sells = trade_log[trade_log["action"] == "sell"] + + plt.scatter( + buys["date"], buys["price"], color="g", label="Buy", marker="^", alpha=0.7 + ) + plt.scatter( + sells["date"], sells["price"], color="r", label="Sell", marker="v", alpha=0.7 + ) + + plt.title("Trade Log") + plt.xlabel("Date") + plt.ylabel("Price") + plt.legend() + plt.grid(True) + + if filename: + plt.savefig(filename) + plt.close() + else: + plt.show() + + +def plot_portfolio_composition( + composition: Dict[str, float], filename: Optional[str] = None +): + """ + Plot portfolio composition + + Parameters: + composition: Dictionary of {ticker: weight} + filename: Filename to save plot (optional) + """ + if not composition: + return + + plt.figure(figsize=(10, 6)) + + # Sort by weight + sorted_composition = sorted(composition.items(), key=lambda x: x[1], reverse=True) + + # Extract tickers and weights + tickers = [item[0] for item in sorted_composition] + weights = [item[1] for item in sorted_composition] + + # Create pie chart + plt.pie(weights, labels=tickers, autopct="%1.1f%%", startangle=140) + plt.title("Portfolio Composition") + + if filename: + plt.savefig(filename) + plt.close() + else: + plt.show() + + +def plot_feature_importance( + importance: np.ndarray, feature_names: List[str], filename: Optional[str] = None +): + """ + Plot feature importance + + Parameters: + importance: Array of feature importance scores + feature_names: List of feature names + filename: Filename to save plot (optional) + """ + plt.figure(figsize=(10, 6)) + + # Sort features by importance + sorted_idx = importance.argsort() + plt.barh(range(len(sorted_idx)), importance[sorted_idx], align="center") + plt.yticks(range(len(sorted_idx)), [feature_names[i] for i in sorted_idx]) + plt.title("Feature Importance") + plt.xlabel("Importance Score") + + if filename: + plt.savefig(filename) + plt.close() + else: + plt.show() diff --git a/stock_gnn.log b/stock_gnn.log new file mode 100644 index 0000000..e69de29