Compare commits

...

58 Commits

Author SHA1 Message Date
leo2650 809f46fe36 to discard 2025-11-04 18:02:38 +00:00
leo2650 413abafe0f My First Commit 2025-11-04 17:55:08 +00:00
oleg 5d46c1e32c . 2025-10-27 18:46:26 -04:00
oleg 889f7ba1c3 . 2025-10-27 18:46:14 -04:00
oleg 1515b2d077 . 2025-10-27 18:39:51 -04:00
oleg b4ae3e715d . 2025-10-27 18:36:26 -04:00
Cryptoval Trading Technologies 6f845d32c6 . 2025-07-25 22:13:49 +00:00
Cryptoval Trading Technologies a04e8878fb lg_changes 2025-07-25 22:11:49 +00:00
Oleg Sheynin 71822c64b0 progress 2025-07-25 20:39:59 +00:00
Oleg Sheynin c2f701e3a2 progress 2025-07-25 20:20:23 +00:00
Oleg Sheynin 21a473a4c2 fix close position trades 2025-07-25 18:21:52 +00:00
Oleg Sheynin 98a15d301a bug fix - multiple dates 2025-07-25 07:04:44 +00:00
Oleg Sheynin bcf4447cb6 bug fixes 2025-07-25 06:39:17 +00:00
Oleg Sheynin 1af35000ab cleaning 2025-07-25 01:28:59 +00:00
Oleg Sheynin 2c08b6f1a9 intermarket fix for weekends 2025-07-25 00:47:19 +00:00
Oleg Sheynin 24f1f82d1f fixes to notebook 2025-07-24 22:45:21 +00:00
Oleg Sheynin af0a6f62a9 progress 2025-07-24 21:09:13 +00:00
Oleg Sheynin a7b4777f76 bug fix 2025-07-24 07:44:33 +00:00
Oleg Sheynin e30b0df4db progress and result.py fixes 2025-07-24 06:51:46 +00:00
Oleg Sheynin 577fb5c109 notebook progress 2025-07-23 03:32:43 +00:00
Oleg Sheynin e0138907be progress 2025-07-23 02:56:00 +00:00
Oleg Sheynin b7292c11f3 notebook cleaning 2025-07-23 02:11:02 +00:00
Oleg Sheynin aac8b9dc50 fixes 2025-07-22 18:04:23 +00:00
Oleg Sheynin 9bb36dddd7 notebook fixes 2025-07-22 17:42:14 +00:00
Oleg Sheynin 31eb9f800c bug fix 2025-07-22 17:25:16 +00:00
Oleg Sheynin 0e83142d0a progress: added zscore fit 2025-07-22 00:20:14 +00:00
Oleg Sheynin b87b40a6ed progress 2025-07-21 05:15:33 +00:00
Oleg Sheynin 28386cdf12 fix trading pair, loading scripts 2025-07-20 18:11:45 +00:00
Oleg Sheynin fb3dc68a1d minor: rename 2025-07-19 01:49:46 +00:00
Oleg Sheynin c776c95d69 progress: stop signals 2025-07-19 01:04:09 +00:00
Oleg Sheynin ca9fff8d88 progress 2025-07-18 23:13:11 +00:00
Oleg Sheynin 705330a9f7 progress 2025-07-18 22:51:29 +00:00
Oleg Sheynin 2272a31765 cointegration test initial 2025-07-17 00:19:49 +00:00
Oleg Sheynin facf7fb0c6 Using timezone for trading session 2025-07-16 18:17:34 +00:00
Oleg Sheynin 9c34d935bd added close position and trade session 2025-07-16 18:06:33 +00:00
Oleg Sheynin 20f150a6b7 progress 2025-07-16 03:21:14 +00:00
Oleg Sheynin d46bcb64d6 progress 2025-07-16 03:06:16 +00:00
Oleg Sheynin 26659ede12 fix pair market data 2025-07-16 02:32:16 +00:00
Oleg Sheynin e9995312a0 minor 2025-07-15 22:26:10 +00:00
Oleg Sheynin a46c8a7576 minor 2025-07-15 20:36:15 +00:00
Oleg Sheynin fe2ebbb27f fixed 2025-07-15 20:32:11 +00:00
Oleg Sheynin ddd9f4adb9 progress 2025-07-15 19:29:26 +00:00
Oleg Sheynin 4bc947cf07 progress 2025-07-15 19:24:18 +00:00
Oleg Sheynin 51944b3a2f progress 2025-07-15 04:14:57 +00:00
Oleg Sheynin bff1c54b48 fix 2025-07-15 03:57:18 +00:00
Oleg Sheynin 9c91f37bcc progress 2025-07-15 03:37:29 +00:00
Oleg Sheynin 76547e1176 progress 2025-07-15 02:52:10 +00:00
Oleg Sheynin 80cf1b60ef outstanding positions bug fix 2025-07-15 00:10:46 +00:00
Oleg Sheynin 94ffb32f50 progress 2025-07-14 22:42:08 +00:00
Oleg Sheynin 967c01c367 progress 2025-07-14 22:26:52 +00:00
Oleg Sheynin 747ca05b16 progress 2025-07-14 21:56:28 +00:00
Oleg Sheynin 30ae95a808 minor 2025-07-14 19:07:28 +00:00
Oleg Sheynin bcba183768 cleaned up sliding notebook 2025-07-14 19:05:19 +00:00
Oleg Sheynin cc0072dcc8 minor change 2025-07-14 05:19:14 +00:00
Oleg Sheynin 35a1cd748e notebook changes 2025-07-14 05:15:47 +00:00
Oleg Sheynin 3b003c7811 progress 2025-07-14 00:41:46 +00:00
Oleg Sheynin b24285802a sliding fit fix 2025-07-13 22:33:48 +00:00
Oleg Sheynin 48f18f7b4f progress, sliding model - buggy 2025-07-12 03:17:12 +00:00
29 changed files with 23002 additions and 4267 deletions
+1 -2
View File
@@ -1,11 +1,10 @@
# SpecStory explanation file # SpecStory explanation file
__pycache__/ __pycache__/
__OLD__/ __OLD__/
.specstory/
.history/ .history/
.cursorindexingignore .cursorindexingignore
data data
.vscode/ ####.vscode/
cvttpy cvttpy
# SpecStory explanation file # SpecStory explanation file
.specstory/.what-is-this.md .specstory/.what-is-this.md
+1
View File
@@ -11,6 +11,7 @@ The enhanced `pt_backtest.py` script now supports multi-day and multi-instrument
- Support for wildcard patterns in configuration files - Support for wildcard patterns in configuration files
- CLI override for data file specification - CLI override for data file specification
### 2. Dynamic Instrument Selection ### 2. Dynamic Instrument Selection
- Auto-detection of instruments from database - Auto-detection of instruments from database
- CLI override for instrument specification - CLI override for instrument specification
+1 -4
View File
@@ -38,15 +38,12 @@ CONFIG = EQT_CONFIG # For equity data
``` ```
Each configuration dictionary specifies: Each configuration dictionary specifies:
- `security_type`: "CRYPTO" or "EQUITY".
- `data_directory`: Path to the data files. - `data_directory`: Path to the data files.
- `datafiles`: A list of database files to process. You can comment/uncomment specific files to include/exclude them from the backtest. - `datafiles`: A list of database files to process. You can comment/uncomment specific files to include/exclude them from the backtest.
- `db_table_name`: The name of the table within the SQLite database. - `db_table_name`: The name of the table within the SQLite database.
- `instruments`: A list of symbols to consider for forming trading pairs. - `instruments`: A list of symbols to consider for forming trading pairs.
- `trading_hours`: Defines the session start and end times, crucial for equity markets. - `trading_hours`: Defines the session start and end times, crucial for equity markets.
- `price_column`: The column in the data to be used as the price (e.g., "close"). - `stat_model_price`: The column in the data to be used as the price (e.g., "close").
- `min_required_points`: Minimum data points needed for statistical calculations.
- `zero_threshold`: A small value to handle potential division by zero.
- `dis-equilibrium_open_trshld`: The threshold (in standard deviations) of the dis-equilibrium for opening a trade. - `dis-equilibrium_open_trshld`: The threshold (in standard deviations) of the dis-equilibrium for opening a trade.
- `dis-equilibrium_close_trshld`: The threshold (in standard deviations) of the dis-equilibrium for closing an open trade. - `dis-equilibrium_close_trshld`: The threshold (in standard deviations) of the dis-equilibrium for closing an open trade.
- `training_minutes`: The length of the rolling window (in minutes) used to train the model (e.g., calculate cointegration, mean, and standard deviation of the dis-equilibrium). - `training_minutes`: The length of the rolling window (in minutes) used to train the model (e.g., calculate cointegration, mean, and standard deviation of the dis-equilibrium).
-33
View File
@@ -1,33 +0,0 @@
{
"security_type": "CRYPTO",
"data_directory": "./data/crypto",
"datafiles": [
"2025*.mktdata.ohlcv.db"
],
"db_table_name": "md_1min_bars",
"exchange_id": "BNBSPOT",
"instrument_id_pfx": "PAIR-",
# "instruments": [
# "BTC-USDT",
# "BCH-USDT",
# "ETH-USDT",
# "LTC-USDT",
# "XRP-USDT",
# "ADA-USDT",
# "SOL-USDT",
# "DOT-USDT"
# ],
"trading_hours": {
"begin_session": "00:00:00",
"end_session": "23:59:00",
"timezone": "UTC"
},
"price_column": "close",
"min_required_points": 30,
"zero_threshold": 1e-10,
"dis-equilibrium_open_trshld": 2.0,
"dis-equilibrium_close_trshld": 0.5,
"training_minutes": 120,
"funding_per_pair": 2000.0,
"fit_method_class": "pt_trading.fit_methods.StaticFit"
}
+5 -3
View File
@@ -2,7 +2,7 @@
"security_type": "EQUITY", "security_type": "EQUITY",
"data_directory": "./data/equity", "data_directory": "./data/equity",
"datafiles": [ "datafiles": [
"202506*.mktdata.ohlcv.db", "20250618.mktdata.ohlcv.db",
], ],
"db_table_name": "md_1min_bars", "db_table_name": "md_1min_bars",
"exchange_id": "ALPACA", "exchange_id": "ALPACA",
@@ -19,7 +19,9 @@
"dis-equilibrium_close_trshld": 1.0, "dis-equilibrium_close_trshld": 1.0,
"training_minutes": 120, "training_minutes": 120,
"funding_per_pair": 2000.0, "funding_per_pair": 2000.0,
"fit_method_class": "pt_trading.fit_methods.SlidingFit", # "fit_method_class": "pt_trading.sliding_fit.SlidingFit",
"exclude_instruments": ["CAN"] "fit_method_class": "pt_trading.static_fit.StaticFit",
"exclude_instruments": ["CAN"],
"close_outstanding_positions": false
} }
+26
View File
@@ -0,0 +1,26 @@
{
"security_type": "EQUITY",
"data_directory": "./data/equity",
"datafiles": [
"20250602.mktdata.ohlcv.db",
],
"db_table_name": "md_1min_bars",
"exchange_id": "ALPACA",
"instrument_id_pfx": "STOCK-",
"trading_hours": {
"begin_session": "9:30:00",
"end_session": "16:00:00",
"timezone": "America/New_York"
},
"price_column": "close",
"min_required_points": 30,
"zero_threshold": 1e-10,
"dis-equilibrium_open_trshld": 2.0,
"dis-equilibrium_close_trshld": 1.0,
"training_minutes": 120,
"funding_per_pair": 2000.0,
"fit_method_class": "pt_trading.fit_methods.StaticFit",
"exclude_instruments": ["CAN"]
}
# "fit_method_class": "pt_trading.fit_methods.SlidingFit",
# "fit_method_class": "pt_trading.fit_methods.StaticFit",
+43
View File
@@ -0,0 +1,43 @@
{
"market_data_loading": {
"CRYPTO": {
"data_directory": "./data/crypto",
"db_table_name": "md_1min_bars",
"instrument_id_pfx": "PAIR-",
},
"EQUITY": {
"data_directory": "./data/equity",
"db_table_name": "md_1min_bars",
"instrument_id_pfx": "STOCK-",
}
},
# ====== Funding ======
"funding_per_pair": 2000.0,
# ====== Trading Parameters ======
"stat_model_price": "close", # "vwap"
"execution_price": {
"column": "vwap",
"shift": 1,
},
"dis-equilibrium_open_trshld": 2.0,
"dis-equilibrium_close_trshld": 1.0,
"training_minutes": 120,
"fit_method_class": "pt_trading.vecm_rolling_fit.VECMRollingFit",
# ====== Stop Conditions ======
"stop_close_conditions": {
"profit": 2.0,
"loss": -0.5
}
# ====== End of Session Closeout ======
"close_outstanding_positions": true,
# "close_outstanding_positions": false,
"trading_hours": {
"timezone": "America/New_York",
"begin_session": "9:30:00",
"end_session": "18:30:00",
}
}
+42
View File
@@ -0,0 +1,42 @@
{
"market_data_loading": {
"CRYPTO": {
"data_directory": "./data/crypto",
"db_table_name": "md_1min_bars",
"instrument_id_pfx": "PAIR-",
},
"EQUITY": {
"data_directory": "./data/equity",
"db_table_name": "md_1min_bars",
"instrument_id_pfx": "STOCK-",
}
},
# ====== Funding ======
"funding_per_pair": 2000.0,
# ====== Trading Parameters ======
"stat_model_price": "close",
"execution_price": {
"column": "vwap",
"shift": 1,
},
"dis-equilibrium_open_trshld": 2.0,
"dis-equilibrium_close_trshld": 0.5,
"training_minutes": 120,
"fit_method_class": "pt_trading.z-score_rolling_fit.ZScoreRollingFit",
# ====== Stop Conditions ======
"stop_close_conditions": {
"profit": 2.0,
"loss": -0.5
}
# ====== End of Session Closeout ======
"close_outstanding_positions": true,
# "close_outstanding_positions": false,
"trading_hours": {
"timezone": "America/New_York",
"begin_session": "9:30:00",
"end_session": "18:30:00",
}
}
+115
View File
@@ -0,0 +1,115 @@
07.11.2025
pairs_trading/configuration <---- directory for config
equity_lg.cfg <-------- copy of equity.cfg
How to run a Program: TRIANGLEsquare ----> triangle EQUITY backtest
Results are in > results (timestamp table for all runs)
table "...timestamp... .pt_backtest_results.equity.db"
going to table using sqlite
> sqlite3 '/home/coder/results/20250721_175750.pt_backtest_results.equity.db'
sqlite> .databases
main: /home/coder/results/20250717_180122.pt_backtest_results.equity.db r/w
sqlite> .tables
config outstanding_positions pt_bt_results
sqlite> PRAGMA table_info('pt_bt_results');
0|date|DATE|0||0
1|pair|TEXT|0||0
2|symbol|TEXT|0||0
3|open_time|DATETIME|0||0
4|open_side|TEXT|0||0
5|open_price|REAL|0||0
6|open_quantity|INTEGER|0||0
7|open_disequilibrium|REAL|0||0
8|close_time|DATETIME|0||0
9|close_side|TEXT|0||0
10|close_price|REAL|0||0
11|close_quantity|INTEGER|0||0
12|close_disequilibrium|REAL|0||0
13|symbol_return|REAL|0||0
14|pair_return|REAL|0||0
select count(*) as cnt from pt_bt_results;
8
select * from pt_bt_results;
select
date, close_time, pair, symbol, symbol_return, pair_return
from pt_bt_results ;
select date, sum(symbol_return) as daily_return
from pt_bt_results where date = '2025-06-18' group by date;
.quit
sqlite3 '/home/coder/results/20250717_172435.pt_backtest_results.equity.db'
sqlite> select date, sum(symbol_return) as daily_return
from pt_bt_results group by date;
2025-06-02|1.29845390060828
...
2025-06-18|-43.5084977104115 <========== ????? ==========>
2025-06-20|11.8605547517183
select
date, close_time, pair, symbol, symbol_return, pair_return
from pt_bt_results ;
select date, close_time, pair, symbol, symbol_return, pair_return
from pt_bt_results where date = '2025-06-18';
./scripts/load_equity_pair_intraday.sh -A NVDA -B QQQ -d 20250701 -T ./intraday_md
to inspect exactly what sources, formats, and processing steps you can open the script with:
head -n 50 ./scripts/load_equity_pair_intraday.sh
✓ Data file found: /home/coder/pairs_trading/data/crypto/20250605.mktdata.ohlcv.db
sqlite3 '/home/coder/results/20250722_201930.pt_backtest_results.crypto.db'
sqlite3 '/home/coder/results/xxxxxxxx_yyyyyy.pt_backtest_results.pseudo.db'
11111111
=== At your terminal, run these commands:
sqlite3 '/home/coder/results/20250722_201930.pt_backtest_results.crypto.db'
=== Then inside the SQLite prompt:
.mode csv
.headers on
.output results_20250722.csv
SELECT * FROM pt_bt_results;
.output stdout
.quit
cd /home/coder/
# === mode csv formats output as CSV
# === headers on includes column names
# === output my_table.csv directs output to that file
# === Run your SELECT query, then revert output
# === Open my_table.csv in Excel directly
# ======== Using scp (Secure Copy)
# === On your local machine, open a terminal and run:
scp cvtt@953f6e8df266:/home/coder/results_20250722.csv ~/Downloads/
# ===== convert cvs pandas dataframe ====== -->
import pandas as pd
# Replace with the actual path to your CSV file
file_path = '/home/coder/results_20250722.csv'
# Read the CSV file into a DataFrame
df = pd.read_csv(file_path)
# Show the first few rows
print(df.head())
+52
View File
@@ -0,0 +1,52 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from enum import Enum
from typing import Dict, Optional, cast
import pandas as pd
from pt_trading.results import BacktestResult
from pt_trading.trading_pair import TradingPair
NanoPerMin = 1e9
class PairsTradingFitMethod(ABC):
TRADES_COLUMNS = [
"time",
"symbol",
"side",
"action",
"price",
"disequilibrium",
"scaled_disequilibrium",
"signed_scaled_disequilibrium",
"pair",
]
@staticmethod
def create(config: Dict) -> PairsTradingFitMethod:
import importlib
fit_method_class_name = config.get("fit_method_class", None)
assert fit_method_class_name is not None
module_name, class_name = fit_method_class_name.rsplit(".", 1)
module = importlib.import_module(module_name)
fit_method = getattr(module, class_name)()
return cast(PairsTradingFitMethod, fit_method)
@abstractmethod
def run_pair(
self, pair: TradingPair, bt_result: BacktestResult
) -> Optional[pd.DataFrame]: ...
@abstractmethod
def reset(self) -> None: ...
@abstractmethod
def create_trading_pair(
self,
config: Dict,
market_data: pd.DataFrame,
symbol_a: str,
symbol_b: str,
) -> TradingPair: ...
-419
View File
@@ -1,419 +0,0 @@
from abc import ABC, abstractmethod
from enum import Enum
from typing import Dict, Optional, cast
import pandas as pd # type: ignore[import]
from pt_trading.results import BacktestResult
from pt_trading.trading_pair import TradingPair
NanoPerMin = 1e9
class PairsTradingFitMethod(ABC):
TRADES_COLUMNS = [
"time",
"action",
"symbol",
"price",
"disequilibrium",
"scaled_disequilibrium",
"pair",
]
@abstractmethod
def run_pair(self, config: Dict, pair: TradingPair, bt_result: BacktestResult) -> Optional[pd.DataFrame]:
...
@abstractmethod
def reset(self):
...
class StaticFit(PairsTradingFitMethod):
def run_pair(self, config: Dict, pair: TradingPair, bt_result: BacktestResult) -> Optional[pd.DataFrame]: # abstractmethod
pair.get_datasets(training_minutes=config["training_minutes"])
try:
is_cointegrated = pair.train_pair()
if not is_cointegrated:
print(f"{pair} IS NOT COINTEGRATED")
return None
except Exception as e:
print(f"{pair}: Training failed: {str(e)}")
return None
try:
pair.predict()
except Exception as e:
print(f"{pair}: Prediction failed: {str(e)}")
return None
pair_trades = self.create_trading_signals(pair=pair, config=config, result=bt_result)
return pair_trades
def create_trading_signals(self, pair: TradingPair, config: Dict, result: BacktestResult) -> pd.DataFrame:
beta = pair.vecm_fit_.beta # type: ignore
colname_a, colname_b = pair.colnames()
predicted_df = pair.predicted_df_
open_threshold = config["dis-equilibrium_open_trshld"]
close_threshold = config["dis-equilibrium_close_trshld"]
# Iterate through the testing dataset to find the first trading opportunity
open_row_index = None
for row_idx in range(len(predicted_df)):
curr_disequilibrium = predicted_df["scaled_disequilibrium"][row_idx]
# Check if current row has sufficient disequilibrium (not near-zero)
if curr_disequilibrium >= open_threshold:
open_row_index = row_idx
break
# If no row with sufficient disequilibrium found, skip this pair
if open_row_index is None:
print(f"{pair}: Insufficient disequilibrium in testing dataset. Skipping.")
return pd.DataFrame()
# Look for close signal starting from the open position
trading_signals_df = (
predicted_df["scaled_disequilibrium"][open_row_index:] < close_threshold
)
# Adjust indices to account for the offset from open_row_index
close_row_index = None
for idx, value in trading_signals_df.items():
if value:
close_row_index = idx
break
open_row = predicted_df.loc[open_row_index]
open_tstamp = open_row["tstamp"]
open_disequilibrium = open_row["disequilibrium"]
open_scaled_disequilibrium = open_row["scaled_disequilibrium"]
open_px_a = open_row[f"{colname_a}"]
open_px_b = open_row[f"{colname_b}"]
abs_beta = abs(beta[1])
pred_px_b = predicted_df.loc[open_row_index][f"{colname_b}_pred"]
pred_px_a = predicted_df.loc[open_row_index][f"{colname_a}_pred"]
if pred_px_b * abs_beta - pred_px_a > 0:
open_side_a = "BUY"
open_side_b = "SELL"
close_side_a = "SELL"
close_side_b = "BUY"
else:
open_side_b = "BUY"
open_side_a = "SELL"
close_side_b = "SELL"
close_side_a = "BUY"
# If no close signal found, print position and unrealized PnL
if close_row_index is None:
last_row_index = len(predicted_df) - 1
# Use the new method from BacktestResult to handle outstanding positions
result.handle_outstanding_position(
pair=pair,
pair_result_df=predicted_df,
last_row_index=last_row_index,
open_side_a=open_side_a,
open_side_b=open_side_b,
open_px_a=open_px_a,
open_px_b=open_px_b,
open_tstamp=open_tstamp,
)
# Return only open trades (no close trades)
trd_signal_tuples = [
(
open_tstamp,
open_side_a,
pair.symbol_a_,
open_px_a,
open_disequilibrium,
open_scaled_disequilibrium,
pair,
),
(
open_tstamp,
open_side_b,
pair.symbol_b_,
open_px_b,
open_disequilibrium,
open_scaled_disequilibrium,
pair,
),
]
else:
# Close signal found - create complete trade
close_row = predicted_df.loc[close_row_index]
close_tstamp = close_row["tstamp"]
close_disequilibrium = close_row["disequilibrium"]
close_scaled_disequilibrium = close_row["scaled_disequilibrium"]
close_px_a = close_row[f"{colname_a}"]
close_px_b = close_row[f"{colname_b}"]
print(f"{pair}: Close signal found at index {close_row_index}")
trd_signal_tuples = [
(
open_tstamp,
open_side_a,
pair.symbol_a_,
open_px_a,
open_disequilibrium,
open_scaled_disequilibrium,
pair,
),
(
open_tstamp,
open_side_b,
pair.symbol_b_,
open_px_b,
open_disequilibrium,
open_scaled_disequilibrium,
pair,
),
(
close_tstamp,
close_side_a,
pair.symbol_a_,
close_px_a,
close_disequilibrium,
close_scaled_disequilibrium,
pair,
),
(
close_tstamp,
close_side_b,
pair.symbol_b_,
close_px_b,
close_disequilibrium,
close_scaled_disequilibrium,
pair,
),
]
# Add tuples to data frame
return pd.DataFrame(
trd_signal_tuples,
columns=self.TRADES_COLUMNS, # type: ignore
)
def reset(self) -> None:
pass
class PairState(Enum):
INITIAL = 1
OPEN = 2
CLOSED = 3
class SlidingFit(PairsTradingFitMethod):
def __init__(self) -> None:
super().__init__()
self.curr_training_start_idx_ = 0
def run_pair(self, config: Dict, pair: TradingPair, bt_result: BacktestResult) -> Optional[pd.DataFrame]:
print(f"***{pair}*** STARTING....")
pair.user_data_['state'] = PairState.INITIAL
pair.user_data_["trades"] = pd.DataFrame(columns=self.TRADES_COLUMNS) # type: ignore
pair.user_data_["is_cointegrated"] = False
open_threshold = config["dis-equilibrium_open_trshld"]
close_threshold = config["dis-equilibrium_open_trshld"]
training_minutes = config["training_minutes"]
while True:
print(self.curr_training_start_idx_, end='\r')
pair.get_datasets(
training_minutes=training_minutes,
training_start_index=self.curr_training_start_idx_,
testing_size=1
)
if len(pair.training_df_) < training_minutes:
print(f"{pair}: {self.curr_training_start_idx_} Not enough training data. Completing the job.")
if pair.user_data_["state"] == PairState.OPEN:
print(f"{pair}: {self.curr_training_start_idx_} Position is not closed.")
# outstanding positions
# last_row_index = self.curr_training_start_idx_ + training_minutes
bt_result.handle_outstanding_position(
pair=pair,
pair_result_df=pair.predicted_df_,
last_row_index=0,
open_side_a=pair.user_data_["open_side_a"],
open_side_b=pair.user_data_["open_side_b"],
open_px_a=pair.user_data_["open_px_a"],
open_px_b=pair.user_data_["open_px_b"],
open_tstamp=pair.user_data_["open_tstamp"],
)
break
try:
is_cointegrated = pair.train_pair()
except Exception as e:
raise RuntimeError(f"{pair}: Training failed: {str(e)}") from e
if pair.user_data_["is_cointegrated"] != is_cointegrated:
pair.user_data_["is_cointegrated"] = is_cointegrated
if not is_cointegrated:
if pair.user_data_["state"] == PairState.OPEN:
print(f"{pair} {self.curr_training_start_idx_} LOST COINTEGRATION. Consider closing positions...")
else:
print(f"{pair} {self.curr_training_start_idx_} IS NOT COINTEGRATED. Moving on")
else:
print('*' * 80)
print(f"Pair {pair} ({self.curr_training_start_idx_}) IS COINTEGRATED")
print('*' * 80)
if not is_cointegrated:
self.curr_training_start_idx_ += 1
continue
try:
pair.predict()
except Exception as e:
raise RuntimeError(f"{pair}: Prediction failed: {str(e)}") from e
if pair.user_data_["state"] == PairState.INITIAL:
open_trades = self._get_open_trades(pair, open_threshold=open_threshold)
if open_trades is not None:
pair.user_data_["trades"] = open_trades
pair.user_data_["state"] = PairState.OPEN
elif pair.user_data_["state"] == PairState.OPEN:
close_trades = self._get_close_trades(pair, close_threshold=close_threshold)
if close_trades is not None:
pair.user_data_["trades"] = pd.concat([pair.user_data_["trades"], close_trades], ignore_index=True)
pair.user_data_["state"] = PairState.CLOSED
break
self.curr_training_start_idx_ += 1
print(f"***{pair}*** FINISHED ... {len(pair.user_data_['trades'])}")
return pair.user_data_["trades"]
def _get_open_trades(self, pair: TradingPair, open_threshold: float) -> Optional[pd.DataFrame]:
colname_a, colname_b = pair.colnames()
predicted_df = pair.predicted_df_
# Check if we have any data to work with
if len(predicted_df) == 0:
return None
open_row = predicted_df.iloc[0]
open_tstamp = open_row["tstamp"]
open_disequilibrium = open_row["disequilibrium"]
open_scaled_disequilibrium = open_row["scaled_disequilibrium"]
open_px_a = open_row[f"{colname_a}"]
open_px_b = open_row[f"{colname_b}"]
if open_scaled_disequilibrium < open_threshold:
return None
# creating the trades
if open_disequilibrium > 0:
open_side_a = "SELL"
open_side_b = "BUY"
close_side_a = "BUY"
close_side_b = "SELL"
else:
open_side_a = "BUY"
open_side_b = "SELL"
close_side_a = "SELL"
close_side_b = "BUY"
# save closing sides
pair.user_data_["open_side_a"] = open_side_a
pair.user_data_["open_side_b"] = open_side_b
pair.user_data_["open_px_a"] = open_px_a
pair.user_data_["open_px_b"] = open_px_b
pair.user_data_["open_tstamp"] = open_tstamp
pair.user_data_["close_side_a"] = close_side_a
pair.user_data_["close_side_b"] = close_side_b
# create opening trades
trd_signal_tuples = [
(
open_tstamp,
open_side_a,
pair.symbol_a_,
open_px_a,
open_disequilibrium,
open_scaled_disequilibrium,
pair,
),
(
open_tstamp,
open_side_b,
pair.symbol_b_,
open_px_b,
open_disequilibrium,
open_scaled_disequilibrium,
pair,
),
]
return pd.DataFrame(
trd_signal_tuples,
columns=self.TRADES_COLUMNS, # type: ignore
)
def _get_close_trades(self, pair: TradingPair, close_threshold: float) -> Optional[pd.DataFrame]:
colname_a, colname_b = pair.colnames()
# Check if we have any data to work with
if len(pair.predicted_df_) == 0:
return None
close_row = pair.predicted_df_.iloc[0]
close_tstamp = close_row["tstamp"]
close_disequilibrium = close_row["disequilibrium"]
close_scaled_disequilibrium = close_row["scaled_disequilibrium"]
close_px_a = close_row[f"{colname_a}"]
close_px_b = close_row[f"{colname_b}"]
close_side_a = pair.user_data_["close_side_a"]
close_side_b = pair.user_data_["close_side_b"]
if close_scaled_disequilibrium > close_threshold:
return None
trd_signal_tuples = [
(
close_tstamp,
close_side_a,
pair.symbol_a_,
close_px_a,
close_disequilibrium,
close_scaled_disequilibrium,
pair,
),
(
close_tstamp,
close_side_b,
pair.symbol_b_,
close_px_b,
close_disequilibrium,
close_scaled_disequilibrium,
pair,
),
]
# Add tuples to data frame
return pd.DataFrame(
trd_signal_tuples,
columns=self.TRADES_COLUMNS, # type: ignore
)
def reset(self):
self.curr_training_start_idx_ = 0
+338 -372
View File
@@ -1,28 +1,30 @@
from typing import Any, Dict, List
import pandas as pd
import sqlite3
import os import os
from datetime import datetime, date import sqlite3
from datetime import date, datetime
from typing import Any, Dict, List, Optional, Tuple
import pandas as pd
from pt_trading.trading_pair import TradingPair
# Recommended replacement adapters and converters for Python 3.12+ # Recommended replacement adapters and converters for Python 3.12+
# From: https://docs.python.org/3/library/sqlite3.html#sqlite3-adapter-converter-recipes # From: https://docs.python.org/3/library/sqlite3.html#sqlite3-adapter-converter-recipes
def adapt_date_iso(val): def adapt_date_iso(val: date) -> str:
"""Adapt datetime.date to ISO 8601 date.""" """Adapt datetime.date to ISO 8601 date."""
return val.isoformat() return val.isoformat()
def adapt_datetime_iso(val): def adapt_datetime_iso(val: datetime) -> str:
"""Adapt datetime.datetime to timezone-naive ISO 8601 date.""" """Adapt datetime.datetime to timezone-naive ISO 8601 date."""
return val.isoformat() return val.isoformat()
def convert_date(val): def convert_date(val: bytes) -> date:
"""Convert ISO 8601 date to datetime.date object.""" """Convert ISO 8601 date to datetime.date object."""
return datetime.fromisoformat(val.decode()).date() return datetime.fromisoformat(val.decode()).date()
def convert_datetime(val): def convert_datetime(val: bytes) -> datetime:
"""Convert ISO 8601 datetime to datetime.datetime object.""" """Convert ISO 8601 datetime to datetime.datetime object."""
return datetime.fromisoformat(val.decode()) return datetime.fromisoformat(val.decode())
@@ -39,6 +41,12 @@ def create_result_database(db_path: str) -> None:
Create the SQLite database and required tables if they don't exist. Create the SQLite database and required tables if they don't exist.
""" """
try: try:
# Create directory if it doesn't exist
db_dir = os.path.dirname(db_path)
if db_dir and not os.path.exists(db_dir):
os.makedirs(db_dir, exist_ok=True)
print(f"Created directory: {db_dir}")
conn = sqlite3.connect(db_path) conn = sqlite3.connect(db_path)
cursor = conn.cursor() cursor = conn.cursor()
@@ -60,7 +68,8 @@ def create_result_database(db_path: str) -> None:
close_quantity INTEGER, close_quantity INTEGER,
close_disequilibrium REAL, close_disequilibrium REAL,
symbol_return REAL, symbol_return REAL,
pair_return REAL pair_return REAL,
close_condition TEXT
) )
""" """
) )
@@ -112,8 +121,8 @@ def store_config_in_database(
config_file_path: str, config_file_path: str,
config: Dict, config: Dict,
fit_method_class: str, fit_method_class: str,
datafiles: List[str], datafiles: List[Tuple[str, str]],
instruments: List[str], instruments: List[Dict[str, str]],
) -> None: ) -> None:
""" """
Store configuration information in the database for reference. Store configuration information in the database for reference.
@@ -131,8 +140,13 @@ def store_config_in_database(
config_json = json.dumps(config, indent=2, default=str) config_json = json.dumps(config, indent=2, default=str)
# Convert lists to comma-separated strings for storage # Convert lists to comma-separated strings for storage
datafiles_str = ", ".join(datafiles) datafiles_str = ", ".join([f"{datafile}" for _, datafile in datafiles])
instruments_str = ", ".join(instruments) instruments_str = ", ".join(
[
f"{inst['symbol']}:{inst['instrument_type']}:{inst['exchange_id']}"
for inst in instruments
]
)
# Insert configuration record # Insert configuration record
cursor.execute( cursor.execute(
@@ -163,251 +177,23 @@ def store_config_in_database(
traceback.print_exc() traceback.print_exc()
def store_results_in_database( def convert_timestamp(timestamp: Any) -> Optional[datetime]:
db_path: str, datafile: str, bt_result: "BacktestResult" """Convert pandas Timestamp to Python datetime object for SQLite compatibility."""
) -> None: if timestamp is None:
""" return None
Store backtest results in the SQLite database. if isinstance(timestamp, pd.Timestamp):
""" return timestamp.to_pydatetime()
if db_path.upper() == "NONE": elif isinstance(timestamp, datetime):
return
def convert_timestamp(timestamp):
"""Convert pandas Timestamp to Python datetime object for SQLite compatibility."""
if timestamp is None:
return None
if hasattr(timestamp, "to_pydatetime"):
return timestamp.to_pydatetime()
return timestamp return timestamp
elif isinstance(timestamp, date):
return datetime.combine(timestamp, datetime.min.time())
elif isinstance(timestamp, str):
return datetime.strptime(timestamp, "%Y-%m-%d %H:%M:%S")
elif isinstance(timestamp, int):
return datetime.fromtimestamp(timestamp)
else:
raise ValueError(f"Unsupported timestamp type: {type(timestamp)}")
try:
# Extract date from datafile name (assuming format like 20250528.mktdata.ohlcv.db)
filename = os.path.basename(datafile)
date_str = filename.split(".")[0] # Extract date part
# Convert to proper date format
try:
date_obj = datetime.strptime(date_str, "%Y%m%d").date()
except ValueError:
# If date parsing fails, use current date
date_obj = datetime.now().date()
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
# Process each trade from bt_result
trades = bt_result.get_trades()
for pair_name, symbols in trades.items():
# Calculate pair return for this pair
pair_return = 0.0
pair_trades = []
# First pass: collect all trades and calculate returns
for symbol, symbol_trades in symbols.items():
if len(symbol_trades) == 0: # No trades for this symbol
print(
f"Warning: No trades found for symbol {symbol} in pair {pair_name}"
)
continue
elif len(symbol_trades) >= 2: # Completed trades (entry + exit)
# Handle both old and new tuple formats
if len(symbol_trades[0]) == 2: # Old format: (action, price)
entry_action, entry_price = symbol_trades[0]
exit_action, exit_price = symbol_trades[1]
open_disequilibrium = 0.0 # Fallback for old format
open_scaled_disequilibrium = 0.0
close_disequilibrium = 0.0
close_scaled_disequilibrium = 0.0
open_time = datetime.now()
close_time = datetime.now()
else: # New format: (action, price, disequilibrium, scaled_disequilibrium, timestamp)
(
entry_action,
entry_price,
open_disequilibrium,
open_scaled_disequilibrium,
open_time,
) = symbol_trades[0]
(
exit_action,
exit_price,
close_disequilibrium,
close_scaled_disequilibrium,
close_time,
) = symbol_trades[1]
# Handle None values
open_disequilibrium = (
open_disequilibrium
if open_disequilibrium is not None
else 0.0
)
open_scaled_disequilibrium = (
open_scaled_disequilibrium
if open_scaled_disequilibrium is not None
else 0.0
)
close_disequilibrium = (
close_disequilibrium
if close_disequilibrium is not None
else 0.0
)
close_scaled_disequilibrium = (
close_scaled_disequilibrium
if close_scaled_disequilibrium is not None
else 0.0
)
# Convert pandas Timestamps to Python datetime objects
open_time = convert_timestamp(open_time) or datetime.now()
close_time = convert_timestamp(close_time) or datetime.now()
# Calculate actual share quantities based on funding per pair
# Split funding equally between the two positions
funding_per_position = bt_result.config["funding_per_pair"] / 2
shares = funding_per_position / entry_price
# Calculate symbol return
symbol_return = 0.0
if entry_action == "BUY" and exit_action == "SELL":
symbol_return = (exit_price - entry_price) / entry_price * 100
elif entry_action == "SELL" and exit_action == "BUY":
symbol_return = (entry_price - exit_price) / entry_price * 100
pair_return += symbol_return
pair_trades.append(
{
"symbol": symbol,
"entry_action": entry_action,
"entry_price": entry_price,
"exit_action": exit_action,
"exit_price": exit_price,
"symbol_return": symbol_return,
"open_disequilibrium": open_disequilibrium,
"open_scaled_disequilibrium": open_scaled_disequilibrium,
"close_disequilibrium": close_disequilibrium,
"close_scaled_disequilibrium": close_scaled_disequilibrium,
"open_time": open_time,
"close_time": close_time,
"shares": shares,
"is_completed": True,
}
)
# Skip one-sided trades - they will be handled by outstanding_positions table
elif len(symbol_trades) == 1:
print(
f"Skipping one-sided trade for {symbol} in pair {pair_name} - will be stored in outstanding_positions table"
)
continue
else:
# This should not happen, but handle unexpected cases
print(
f"Warning: Unexpected number of trades ({len(symbol_trades)}) for symbol {symbol} in pair {pair_name}"
)
continue
# Second pass: insert completed trade records into database
for trade in pair_trades:
# Only store completed trades in pt_bt_results table
cursor.execute(
"""
INSERT INTO pt_bt_results (
date, pair, symbol, open_time, open_side, open_price,
open_quantity, open_disequilibrium, close_time, close_side,
close_price, close_quantity, close_disequilibrium,
symbol_return, pair_return
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
date_obj,
pair_name,
trade["symbol"],
trade["open_time"],
trade["entry_action"],
trade["entry_price"],
trade["shares"],
trade["open_scaled_disequilibrium"],
trade["close_time"],
trade["exit_action"],
trade["exit_price"],
trade["shares"],
trade["close_scaled_disequilibrium"],
trade["symbol_return"],
pair_return,
),
)
# Store outstanding positions in separate table
outstanding_positions = bt_result.get_outstanding_positions()
for pos in outstanding_positions:
# Calculate position quantity (negative for SELL positions)
position_qty_a = (
pos["shares_a"] if pos["side_a"] == "BUY" else -pos["shares_a"]
)
position_qty_b = (
pos["shares_b"] if pos["side_b"] == "BUY" else -pos["shares_b"]
)
# Calculate unrealized returns
# For symbol A: (current_price - open_price) / open_price * 100 * position_direction
unrealized_return_a = (
(pos["current_px_a"] - pos["open_px_a"]) / pos["open_px_a"] * 100
) * (1 if pos["side_a"] == "BUY" else -1)
unrealized_return_b = (
(pos["current_px_b"] - pos["open_px_b"]) / pos["open_px_b"] * 100
) * (1 if pos["side_b"] == "BUY" else -1)
# Store outstanding position for symbol A
cursor.execute(
"""
INSERT INTO outstanding_positions (
date, pair, symbol, position_quantity, last_price, unrealized_return, open_price, open_side
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
date_obj,
pos["pair"],
pos["symbol_a"],
position_qty_a,
pos["current_px_a"],
unrealized_return_a,
pos["open_px_a"],
pos["side_a"],
),
)
# Store outstanding position for symbol B
cursor.execute(
"""
INSERT INTO outstanding_positions (
date, pair, symbol, position_quantity, last_price, unrealized_return, open_price, open_side
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
date_obj,
pos["pair"],
pos["symbol_b"],
position_qty_b,
pos["current_px_b"],
unrealized_return_b,
pos["open_px_b"],
pos["side_b"],
),
)
conn.commit()
conn.close()
except Exception as e:
print(f"Error storing results in database: {str(e)}")
import traceback
traceback.print_exc()
class BacktestResult: class BacktestResult:
@@ -420,17 +206,20 @@ class BacktestResult:
self.trades: Dict[str, Dict[str, Any]] = {} self.trades: Dict[str, Dict[str, Any]] = {}
self.total_realized_pnl = 0.0 self.total_realized_pnl = 0.0
self.outstanding_positions: List[Dict[str, Any]] = [] self.outstanding_positions: List[Dict[str, Any]] = []
self.pairs_trades_: Dict[str, List[Dict[str, Any]]] = {}
def add_trade( def add_trade(
self, self,
pair_nm, pair_nm: str,
symbol, symbol: str,
action, side: str,
price, action: str,
disequilibrium=None, price: Any,
scaled_disequilibrium=None, disequilibrium: Optional[float] = None,
timestamp=None, scaled_disequilibrium: Optional[float] = None,
): timestamp: Optional[datetime] = None,
status: Optional[str] = None,
) -> None:
"""Add a trade to the results tracking.""" """Add a trade to the results tracking."""
pair_nm = str(pair_nm) pair_nm = str(pair_nm)
@@ -439,14 +228,23 @@ class BacktestResult:
if symbol not in self.trades[pair_nm]: if symbol not in self.trades[pair_nm]:
self.trades[pair_nm][symbol] = [] self.trades[pair_nm][symbol] = []
self.trades[pair_nm][symbol].append( self.trades[pair_nm][symbol].append(
(action, price, disequilibrium, scaled_disequilibrium, timestamp) {
"symbol": symbol,
"side": side,
"action": action,
"price": price,
"disequilibrium": disequilibrium,
"scaled_disequilibrium": scaled_disequilibrium,
"timestamp": timestamp,
"status": status,
}
) )
def add_outstanding_position(self, position: Dict[str, Any]): def add_outstanding_position(self, position: Dict[str, Any]) -> None:
"""Add an outstanding position to tracking.""" """Add an outstanding position to tracking."""
self.outstanding_positions.append(position) self.outstanding_positions.append(position)
def add_realized_pnl(self, realized_pnl: float): def add_realized_pnl(self, realized_pnl: float) -> None:
"""Add realized PnL to the total.""" """Add realized PnL to the total."""
self.total_realized_pnl += realized_pnl self.total_realized_pnl += realized_pnl
@@ -462,36 +260,44 @@ class BacktestResult:
"""Get all trades.""" """Get all trades."""
return self.trades return self.trades
def clear_trades(self): def clear_trades(self) -> None:
"""Clear all trades (used when processing new files).""" """Clear all trades (used when processing new files)."""
self.trades.clear() self.trades.clear()
def collect_single_day_results(self, result): def collect_single_day_results(self, pairs_trades: List[pd.DataFrame]) -> None:
"""Collect and process single day trading results.""" """Collect and process single day trading results."""
if result is None: result = pd.concat(pairs_trades, ignore_index=True)
return result["time"] = pd.to_datetime(result["time"])
result = result.set_index("time").sort_index()
print("\n -------------- Suggested Trades ") print("\n -------------- Suggested Trades ")
print(result) print(result)
for row in result.itertuples(): for row in result.itertuples():
side = row.side
action = row.action action = row.action
symbol = row.symbol symbol = row.symbol
price = row.price price = row.price
disequilibrium = getattr(row, "disequilibrium", None) disequilibrium = getattr(row, "disequilibrium", None)
scaled_disequilibrium = getattr(row, "scaled_disequilibrium", None) scaled_disequilibrium = getattr(row, "scaled_disequilibrium", None)
timestamp = getattr(row, "time", None) if hasattr(row, "time"):
timestamp = getattr(row, "time")
else:
timestamp = convert_timestamp(row.Index)
status = row.status
self.add_trade( self.add_trade(
pair_nm=row.pair, pair_nm=str(row.pair),
action=action, symbol=str(symbol),
symbol=symbol, side=str(side),
price=price, action=str(action),
price=float(str(price)),
disequilibrium=disequilibrium, disequilibrium=disequilibrium,
scaled_disequilibrium=scaled_disequilibrium, scaled_disequilibrium=scaled_disequilibrium,
timestamp=timestamp, timestamp=timestamp,
status=str(status) if status is not None else "?",
) )
def print_single_day_results(self): def print_single_day_results(self) -> None:
"""Print single day results summary.""" """Print single day results summary."""
for pair, symbols in self.trades.items(): for pair, symbols in self.trades.items():
print(f"\n--- {pair} ---") print(f"\n--- {pair} ---")
@@ -501,7 +307,7 @@ class BacktestResult:
side, price = trade_data[:2] side, price = trade_data[:2]
print(f"{symbol} {side} at ${price}") print(f"{symbol} {side} at ${price}")
def print_results_summary(self, all_results): def print_results_summary(self, all_results: Dict[str, Dict[str, Any]]) -> None:
"""Print summary of all processed files.""" """Print summary of all processed files."""
print("\n====== Summary of All Processed Files ======") print("\n====== Summary of All Processed Files ======")
for filename, data in all_results.items(): for filename, data in all_results.items():
@@ -512,105 +318,137 @@ class BacktestResult:
) )
print(f"{filename}: {trade_count} trades") print(f"{filename}: {trade_count} trades")
def calculate_returns(self, all_results: Dict): def calculate_returns(self, all_results: Dict[str, Dict[str, Any]]) -> None:
"""Calculate and print returns by day and pair.""" """Calculate and print returns by day and pair."""
def _symbol_return(trade1_side: str, trade1_px: float, trade2_side: str, trade2_px: float) -> float:
if trade1_side == "BUY" and trade2_side == "SELL":
return (trade2_px - trade1_px) / trade1_px * 100
elif trade1_side == "SELL" and trade2_side == "BUY":
return (trade1_px - trade2_px) / trade1_px * 100
else:
return 0
print("\n====== Returns By Day and Pair ======") print("\n====== Returns By Day and Pair ======")
trades = []
for filename, data in all_results.items(): for filename, data in all_results.items():
day_return = 0 pairs = list(data["trades"].keys())
for pair in pairs:
self.pairs_trades_[pair] = []
trades_dict = data["trades"][pair]
for symbol in trades_dict.keys():
trades.extend(trades_dict[symbol])
trades = sorted(trades, key=lambda x: (x["timestamp"], x["symbol"]))
print(f"\n--- {filename} ---") print(f"\n--- {filename} ---")
# Process each pair self.outstanding_positions = data["outstanding_positions"]
for pair, symbols in data["trades"].items():
pair_return = 0
pair_trades = []
# Calculate individual symbol returns in the pair day_return = 0.0
for symbol, trades in symbols.items(): for idx in range(0, len(trades), 4):
if len(trades) >= 2: # Need at least entry and exit symbol_a = trades[idx]["symbol"]
# Get entry and exit trades - handle both old and new tuple formats trade_a_1 = trades[idx]
if len(trades[0]) == 2: # Old format: (action, price) trade_a_2 = trades[idx + 2]
entry_action, entry_price = trades[0]
exit_action, exit_price = trades[1]
open_disequilibrium = None
open_scaled_disequilibrium = None
close_disequilibrium = None
close_scaled_disequilibrium = None
else: # New format: (action, price, disequilibrium, scaled_disequilibrium, timestamp)
entry_action, entry_price = trades[0][:2]
exit_action, exit_price = trades[1][:2]
open_disequilibrium = (
trades[0][2] if len(trades[0]) > 2 else None
)
open_scaled_disequilibrium = (
trades[0][3] if len(trades[0]) > 3 else None
)
close_disequilibrium = (
trades[1][2] if len(trades[1]) > 2 else None
)
close_scaled_disequilibrium = (
trades[1][3] if len(trades[1]) > 3 else None
)
# Calculate return based on action symbol_b = trades[idx + 1]["symbol"]
symbol_return = 0 trade_b_1 = trades[idx + 1]
if entry_action == "BUY" and exit_action == "SELL": trade_b_2 = trades[idx + 3]
# Long position
symbol_return = (
(exit_price - entry_price) / entry_price * 100
)
elif entry_action == "SELL" and exit_action == "BUY":
# Short position
symbol_return = (
(entry_price - exit_price) / entry_price * 100
)
pair_trades.append( symbol_return = 0
( assert (
symbol, trade_a_1["timestamp"] < trade_a_2["timestamp"]
entry_action, ), f"Trade 1: {trade_a_1['timestamp']} is not less than Trade 2: {trade_a_2['timestamp']}"
entry_price, assert (
exit_action, trade_a_1["action"] == "OPEN" and trade_a_2["action"] == "CLOSE"
exit_price, ), f"Trade 1: {trade_a_1['action']} and Trade 2: {trade_a_2['action']} are the same"
symbol_return,
open_scaled_disequilibrium, # Calculate return based on action combination
close_scaled_disequilibrium, trade_return = 0
) symbol_a_return = _symbol_return(trade_a_1["side"], trade_a_1["price"], trade_a_2["side"], trade_a_2["price"])
symbol_b_return = _symbol_return(trade_b_1["side"], trade_b_1["price"], trade_b_2["side"], trade_b_2["price"])
pair_return = symbol_a_return + symbol_b_return
self.pairs_trades_[pair].append(
{
"symbol": symbol_a,
"open_side": trade_a_1["side"],
"open_action": trade_a_1["action"],
"open_price": trade_a_1["price"],
"close_side": trade_a_2["side"],
"close_action": trade_a_2["action"],
"close_price": trade_a_2["price"],
"symbol_return": symbol_a_return,
"open_disequilibrium": trade_a_1["disequilibrium"],
"open_scaled_disequilibrium": trade_a_1["scaled_disequilibrium"],
"close_disequilibrium": trade_a_2["disequilibrium"],
"close_scaled_disequilibrium": trade_a_2["scaled_disequilibrium"],
"open_time": trade_a_1["timestamp"],
"close_time": trade_a_2["timestamp"],
"shares": self.config["funding_per_pair"] / 2 / trade_a_1["price"],
"is_completed": True,
"close_condition": trade_a_2["status"],
"pair_return": pair_return
}
)
self.pairs_trades_[pair].append(
{
"symbol": symbol_b,
"open_side": trade_b_1["side"],
"open_action": trade_b_1["action"],
"open_price": trade_b_1["price"],
"close_side": trade_b_2["side"],
"close_action": trade_b_2["action"],
"close_price": trade_b_2["price"],
"symbol_return": symbol_b_return,
"open_disequilibrium": trade_b_1["disequilibrium"],
"open_scaled_disequilibrium": trade_b_1["scaled_disequilibrium"],
"close_disequilibrium": trade_b_2["disequilibrium"],
"close_scaled_disequilibrium": trade_b_2["scaled_disequilibrium"],
"open_time": trade_b_1["timestamp"],
"close_time": trade_b_2["timestamp"],
"shares": self.config["funding_per_pair"] / 2 / trade_b_1["price"],
"is_completed": True,
"close_condition": trade_b_2["status"],
"pair_return": pair_return
}
)
# Print pair returns with disequilibrium information
day_return = 0.0
if pair in self.pairs_trades_:
print(f"{pair}:")
pair_return = 0.0
for trd in self.pairs_trades_[pair]:
disequil_info = ""
if (
trd["open_scaled_disequilibrium"] is not None
and trd["open_scaled_disequilibrium"] is not None
):
disequil_info = (
f' | Open Dis-eq: {trd["open_scaled_disequilibrium"]:.2f},'
f' Close Dis-eq: {trd["close_scaled_disequilibrium"]:.2f}'
) )
pair_return += symbol_return
# Print pair returns with disequilibrium information print(
if pair_trades: f' {trd["open_time"].time()}-{trd["close_time"].time()} {trd["symbol"]}: '
print(f" {pair}:") f' {trd["open_side"]} @ ${trd["open_price"]:.2f},'
for ( f' {trd["close_side"]} @ ${trd["close_price"]:.2f},'
symbol, f' Return: {trd["symbol_return"]:.2f}%{disequil_info}'
entry_action, )
entry_price, pair_return += trd["symbol_return"]
exit_action,
exit_price,
symbol_return,
open_scaled_disequilibrium,
close_scaled_disequilibrium,
) in pair_trades:
disequil_info = ""
if (
open_scaled_disequilibrium is not None
and close_scaled_disequilibrium is not None
):
disequil_info = f" | Open Dis-eq: {open_scaled_disequilibrium:.2f}, Close Dis-eq: {close_scaled_disequilibrium:.2f}"
print( print(f" Pair Total Return: {pair_return:.2f}%")
f" {symbol}: {entry_action} @ ${entry_price:.2f}, {exit_action} @ ${exit_price:.2f}, Return: {symbol_return:.2f}%{disequil_info}" day_return += pair_return
)
print(f" Pair Total Return: {pair_return:.2f}%")
day_return += pair_return
# Print day total return and add to global realized PnL # Print day total return and add to global realized PnL
if day_return != 0: if day_return != 0:
print(f" Day Total Return: {day_return:.2f}%") print(f" Day Total Return: {day_return:.2f}%")
self.add_realized_pnl(day_return) self.add_realized_pnl(day_return)
def print_outstanding_positions(self): def print_outstanding_positions(self) -> None:
"""Print all outstanding positions with share quantities and current values.""" """Print all outstanding positions with share quantities and current values."""
if not self.get_outstanding_positions(): if not self.get_outstanding_positions():
print("\n====== NO OUTSTANDING POSITIONS ======") print("\n====== NO OUTSTANDING POSITIONS ======")
@@ -684,22 +522,22 @@ class BacktestResult:
print(f"{'TOTAL OUTSTANDING VALUE':<80} ${total_value:<12.2f}") print(f"{'TOTAL OUTSTANDING VALUE':<80} ${total_value:<12.2f}")
def print_grand_totals(self): def print_grand_totals(self) -> None:
"""Print grand totals across all pairs.""" """Print grand totals across all pairs."""
print(f"\n====== GRAND TOTALS ACROSS ALL PAIRS ======") print(f"\n====== GRAND TOTALS ACROSS ALL PAIRS ======")
print(f"Total Realized PnL: {self.get_total_realized_pnl():.2f}%") print(f"Total Realized PnL: {self.get_total_realized_pnl():.2f}%")
def handle_outstanding_position( def handle_outstanding_position(
self, self,
pair, pair: TradingPair,
pair_result_df, pair_result_df: pd.DataFrame,
last_row_index, last_row_index: int,
open_side_a, open_side_a: str,
open_side_b, open_side_b: str,
open_px_a, open_px_a: float,
open_px_b, open_px_b: float,
open_tstamp, open_tstamp: datetime,
): ) -> Tuple[float, float, float]:
""" """
Handle calculation and tracking of outstanding positions when no close signal is found. Handle calculation and tracking of outstanding positions when no close signal is found.
@@ -716,7 +554,7 @@ class BacktestResult:
last_row = pair_result_df.loc[last_row_index] last_row = pair_result_df.loc[last_row_index]
last_tstamp = last_row["tstamp"] last_tstamp = last_row["tstamp"]
colname_a, colname_b = pair.colnames() colname_a, colname_b = pair.exec_prices_colnames()
last_px_a = last_row[colname_a] last_px_a = last_row[colname_a]
last_px_b = last_row[colname_b] last_px_b = last_row[colname_b]
@@ -727,8 +565,8 @@ class BacktestResult:
shares_b = funding_per_position / open_px_b shares_b = funding_per_position / open_px_b
# Calculate current position values (shares * current price) # Calculate current position values (shares * current price)
current_value_a = shares_a * last_px_a current_value_a = shares_a * last_px_a * (-1 if open_side_a == "SELL" else 1)
current_value_b = shares_b * last_px_b current_value_b = shares_b * last_px_b * (-1 if open_side_b == "SELL" else 1)
total_current_value = current_value_a + current_value_b total_current_value = current_value_a + current_value_b
# Get disequilibrium information # Get disequilibrium information
@@ -775,3 +613,131 @@ class BacktestResult:
) )
return current_value_a, current_value_b, total_current_value return current_value_a, current_value_b, total_current_value
def store_results_in_database(
self, db_path: str, day: str
) -> None:
"""
Store backtest results in the SQLite database.
"""
if db_path.upper() == "NONE":
return
try:
# Extract date from datafile name (assuming format like 20250528.mktdata.ohlcv.db)
date_str = day
# Convert to proper date format
try:
date_obj = datetime.strptime(date_str, "%Y%m%d").date()
except ValueError:
# If date parsing fails, use current date
date_obj = datetime.now().date()
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
# Process each trade from bt_result
trades = self.get_trades()
for pair_name, _ in trades.items():
# Second pass: insert completed trade records into database
for trade_pair in sorted(self.pairs_trades_[pair_name], key=lambda x: x["open_time"]):
# Only store completed trades in pt_bt_results table
cursor.execute(
"""
INSERT INTO pt_bt_results (
date, pair, symbol, open_time, open_side, open_price,
open_quantity, open_disequilibrium, close_time, close_side,
close_price, close_quantity, close_disequilibrium,
symbol_return, pair_return, close_condition
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
date_obj,
pair_name,
trade_pair["symbol"],
trade_pair["open_time"],
trade_pair["open_side"],
trade_pair["open_price"],
trade_pair["shares"],
trade_pair["open_scaled_disequilibrium"],
trade_pair["close_time"],
trade_pair["close_side"],
trade_pair["close_price"],
trade_pair["shares"],
trade_pair["close_scaled_disequilibrium"],
trade_pair["symbol_return"],
trade_pair["pair_return"],
trade_pair["close_condition"]
),
)
# Store outstanding positions in separate table
outstanding_positions = self.get_outstanding_positions()
for pos in outstanding_positions:
# Calculate position quantity (negative for SELL positions)
position_qty_a = (
pos["shares_a"] if pos["side_a"] == "BUY" else -pos["shares_a"]
)
position_qty_b = (
pos["shares_b"] if pos["side_b"] == "BUY" else -pos["shares_b"]
)
# Calculate unrealized returns
# For symbol A: (current_price - open_price) / open_price * 100 * position_direction
unrealized_return_a = (
(pos["current_px_a"] - pos["open_px_a"]) / pos["open_px_a"] * 100
) * (1 if pos["side_a"] == "BUY" else -1)
unrealized_return_b = (
(pos["current_px_b"] - pos["open_px_b"]) / pos["open_px_b"] * 100
) * (1 if pos["side_b"] == "BUY" else -1)
# Store outstanding position for symbol A
cursor.execute(
"""
INSERT INTO outstanding_positions (
date, pair, symbol, position_quantity, last_price, unrealized_return, open_price, open_side
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
date_obj,
pos["pair"],
pos["symbol_a"],
position_qty_a,
pos["current_px_a"],
unrealized_return_a,
pos["open_px_a"],
pos["side_a"],
),
)
# Store outstanding position for symbol B
cursor.execute(
"""
INSERT INTO outstanding_positions (
date, pair, symbol, position_quantity, last_price, unrealized_return, open_price, open_side
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
date_obj,
pos["pair"],
pos["symbol_b"],
position_qty_b,
pos["current_px_b"],
unrealized_return_b,
pos["open_px_b"],
pos["side_b"],
),
)
conn.commit()
conn.close()
except Exception as e:
print(f"Error storing results in database: {str(e)}")
import traceback
traceback.print_exc()
+317
View File
@@ -0,0 +1,317 @@
from abc import ABC, abstractmethod
from enum import Enum
from typing import Any, Dict, Optional, cast
import pandas as pd # type: ignore[import]
from pt_trading.fit_method import PairsTradingFitMethod
from pt_trading.results import BacktestResult
from pt_trading.trading_pair import PairState, TradingPair
from statsmodels.tsa.vector_ar.vecm import VECM, VECMResults
NanoPerMin = 1e9
class RollingFit(PairsTradingFitMethod):
"""
N O T E:
=========
- This class remains to be abstract
- The following methods are to be implemented in the subclass:
- create_trading_pair()
=========
"""
def __init__(self) -> None:
super().__init__()
def run_pair(
self, pair: TradingPair, bt_result: BacktestResult
) -> Optional[pd.DataFrame]:
print(f"***{pair}*** STARTING....")
config = pair.config_
curr_training_start_idx = pair.get_begin_index()
end_index = pair.get_end_index()
pair.user_data_["state"] = PairState.INITIAL
# Initialize trades DataFrame with proper dtypes to avoid concatenation warnings
pair.user_data_["trades"] = pd.DataFrame(columns=self.TRADES_COLUMNS).astype(
{
"time": "datetime64[ns]",
"symbol": "string",
"side": "string",
"action": "string",
"price": "float64",
"disequilibrium": "float64",
"scaled_disequilibrium": "float64",
"pair": "object",
}
)
training_minutes = config["training_minutes"]
curr_predicted_row_idx = 0
while True:
print(curr_training_start_idx, end="\r")
pair.get_datasets(
training_minutes=training_minutes,
training_start_index=curr_training_start_idx,
testing_size=1,
)
if len(pair.training_df_) < training_minutes:
print(
f"{pair}: current offset={curr_training_start_idx}"
f" * Training data length={len(pair.training_df_)} < {training_minutes}"
" * Not enough training data. Completing the job."
)
break
try:
# ================================ PREDICTION ================================
self.pair_predict_result_ = pair.predict()
except Exception as e:
raise RuntimeError(
f"{pair}: TrainingPrediction failed: {str(e)}"
) from e
# break
curr_training_start_idx += 1
if curr_training_start_idx > end_index:
break
curr_predicted_row_idx += 1
self._create_trading_signals(pair, config, bt_result)
print(f"***{pair}*** FINISHED *** Num Trades:{len(pair.user_data_['trades'])}")
return pair.get_trades()
def _create_trading_signals(
self, pair: TradingPair, config: Dict, bt_result: BacktestResult
) -> None:
predicted_df = self.pair_predict_result_
assert predicted_df is not None
open_threshold = config["dis-equilibrium_open_trshld"]
close_threshold = config["dis-equilibrium_close_trshld"]
for curr_predicted_row_idx in range(len(predicted_df)):
pred_row = predicted_df.iloc[curr_predicted_row_idx]
scaled_disequilibrium = pred_row["scaled_disequilibrium"]
if pair.user_data_["state"] in [
PairState.INITIAL,
PairState.CLOSE,
PairState.CLOSE_POSITION,
PairState.CLOSE_STOP_LOSS,
PairState.CLOSE_STOP_PROFIT,
]:
if scaled_disequilibrium >= open_threshold:
open_trades = self._get_open_trades(
pair, row=pred_row, open_threshold=open_threshold
)
if open_trades is not None:
open_trades["status"] = PairState.OPEN.name
print(f"OPEN TRADES:\n{open_trades}")
pair.add_trades(open_trades)
pair.user_data_["state"] = PairState.OPEN
pair.on_open_trades(open_trades)
elif pair.user_data_["state"] == PairState.OPEN:
if scaled_disequilibrium <= close_threshold:
close_trades = self._get_close_trades(
pair, row=pred_row, close_threshold=close_threshold
)
if close_trades is not None:
close_trades["status"] = PairState.CLOSE.name
print(f"CLOSE TRADES:\n{close_trades}")
pair.add_trades(close_trades)
pair.user_data_["state"] = PairState.CLOSE
pair.on_close_trades(close_trades)
elif pair.to_stop_close_conditions(predicted_row=pred_row):
close_trades = self._get_close_trades(
pair, row=pred_row, close_threshold=close_threshold
)
if close_trades is not None:
close_trades["status"] = pair.user_data_[
"stop_close_state"
].name
print(f"STOP CLOSE TRADES:\n{close_trades}")
pair.add_trades(close_trades)
pair.user_data_["state"] = pair.user_data_["stop_close_state"]
pair.on_close_trades(close_trades)
# Outstanding positions
if pair.user_data_["state"] == PairState.OPEN:
print(f"{pair}: *** Position is NOT CLOSED. ***")
# outstanding positions
if config["close_outstanding_positions"]:
close_position_row = pd.Series(pair.market_data_.iloc[-2])
close_position_row["disequilibrium"] = 0.0
close_position_row["scaled_disequilibrium"] = 0.0
close_position_row["signed_scaled_disequilibrium"] = 0.0
close_position_trades = self._get_close_trades(
pair=pair, row=close_position_row, close_threshold=close_threshold
)
if close_position_trades is not None:
close_position_trades["status"] = PairState.CLOSE_POSITION.name
print(f"CLOSE_POSITION TRADES:\n{close_position_trades}")
pair.add_trades(close_position_trades)
pair.user_data_["state"] = PairState.CLOSE_POSITION
pair.on_close_trades(close_position_trades)
else:
if predicted_df is not None:
bt_result.handle_outstanding_position(
pair=pair,
pair_result_df=predicted_df,
last_row_index=0,
open_side_a=pair.user_data_["open_side_a"],
open_side_b=pair.user_data_["open_side_b"],
open_px_a=pair.user_data_["open_px_a"],
open_px_b=pair.user_data_["open_px_b"],
open_tstamp=pair.user_data_["open_tstamp"],
)
def _get_open_trades(
self, pair: TradingPair, row: pd.Series, open_threshold: float
) -> Optional[pd.DataFrame]:
colname_a, colname_b = pair.exec_prices_colnames()
open_row = row
open_tstamp = open_row["tstamp"]
open_disequilibrium = open_row["disequilibrium"]
open_scaled_disequilibrium = open_row["scaled_disequilibrium"]
signed_scaled_disequilibrium = open_row["signed_scaled_disequilibrium"]
open_px_a = open_row[f"{colname_a}"]
open_px_b = open_row[f"{colname_b}"]
# creating the trades
# use outer single quotes so we can reference DataFrame keys with double quotes inside
print(f'OPEN_TRADES: {open_tstamp} open_scaled_disequilibrium={open_scaled_disequilibrium}')
if open_disequilibrium > 0:
open_side_a = "SELL"
open_side_b = "BUY"
close_side_a = "BUY"
close_side_b = "SELL"
else:
open_side_a = "BUY"
open_side_b = "SELL"
close_side_a = "SELL"
close_side_b = "BUY"
# save closing sides
pair.user_data_["open_side_a"] = open_side_a
pair.user_data_["open_side_b"] = open_side_b
pair.user_data_["open_px_a"] = open_px_a
pair.user_data_["open_px_b"] = open_px_b
pair.user_data_["open_tstamp"] = open_tstamp
pair.user_data_["close_side_a"] = close_side_a
pair.user_data_["close_side_b"] = close_side_b
# create opening trades
trd_signal_tuples = [
(
open_tstamp,
pair.symbol_a_,
open_side_a,
"OPEN",
open_px_a,
open_disequilibrium,
open_scaled_disequilibrium,
signed_scaled_disequilibrium,
pair,
),
(
open_tstamp,
pair.symbol_b_,
open_side_b,
"OPEN",
open_px_b,
open_disequilibrium,
open_scaled_disequilibrium,
signed_scaled_disequilibrium,
pair,
),
]
# Create DataFrame with explicit dtypes to avoid concatenation warnings
df = pd.DataFrame(trd_signal_tuples, columns=self.TRADES_COLUMNS)
# Ensure consistent dtypes
return df.astype(
{
"time": "datetime64[ns]",
"action": "string",
"symbol": "string",
"price": "float64",
"disequilibrium": "float64",
"scaled_disequilibrium": "float64",
"signed_scaled_disequilibrium": "float64",
"pair": "object",
}
)
def _get_close_trades(
self, pair: TradingPair, row: pd.Series, close_threshold: float
) -> Optional[pd.DataFrame]:
colname_a, colname_b = pair.exec_prices_colnames()
close_row = row
close_tstamp = close_row["tstamp"]
close_disequilibrium = close_row["disequilibrium"]
close_scaled_disequilibrium = close_row["scaled_disequilibrium"]
signed_scaled_disequilibrium = close_row["signed_scaled_disequilibrium"]
close_px_a = close_row[f"{colname_a}"]
close_px_b = close_row[f"{colname_b}"]
close_side_a = pair.user_data_["close_side_a"]
close_side_b = pair.user_data_["close_side_b"]
trd_signal_tuples = [
(
close_tstamp,
pair.symbol_a_,
close_side_a,
"CLOSE",
close_px_a,
close_disequilibrium,
close_scaled_disequilibrium,
signed_scaled_disequilibrium,
pair,
),
(
close_tstamp,
pair.symbol_b_,
close_side_b,
"CLOSE",
close_px_b,
close_disequilibrium,
close_scaled_disequilibrium,
signed_scaled_disequilibrium,
pair,
),
]
# Add tuples to data frame with explicit dtypes to avoid concatenation warnings
df = pd.DataFrame(
trd_signal_tuples,
columns=self.TRADES_COLUMNS,
)
# Ensure consistent dtypes
return df.astype(
{
"time": "datetime64[ns]",
"action": "string",
"symbol": "string",
"price": "float64",
"disequilibrium": "float64",
"scaled_disequilibrium": "float64",
"signed_scaled_disequilibrium": "float64",
"pair": "object",
}
)
def reset(self) -> None:
curr_training_start_idx = 0
+272 -100
View File
@@ -1,13 +1,79 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from enum import Enum
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
import pandas as pd # type:ignore import pandas as pd # type:ignore
from statsmodels.tsa.vector_ar.vecm import VECM, VECMResults # type:ignore
class TradingPair: class PairState(Enum):
INITIAL = 1
OPEN = 2
CLOSE = 3
CLOSE_POSITION = 4
CLOSE_STOP_LOSS = 5
CLOSE_STOP_PROFIT = 6
class CointegrationData:
EG_PVALUE_THRESHOLD = 0.05
tstamp_: pd.Timestamp
pair_: str
eg_pvalue_: float
johansen_lr1_: float
johansen_cvt_: float
eg_is_cointegrated_: bool
johansen_is_cointegrated_: bool
def __init__(self, pair: TradingPair):
training_df = pair.training_df_
assert training_df is not None
from statsmodels.tsa.vector_ar.vecm import coint_johansen
df = training_df[pair.colnames()].reset_index(drop=True)
# Run Johansen cointegration test
result = coint_johansen(df, det_order=0, k_ar_diff=1)
self.johansen_lr1_ = result.lr1[0]
self.johansen_cvt_ = result.cvt[0, 1]
self.johansen_is_cointegrated_ = self.johansen_lr1_ > self.johansen_cvt_
# Run Engle-Granger cointegration test
from statsmodels.tsa.stattools import coint # type: ignore
col1, col2 = pair.colnames()
assert training_df is not None
series1 = training_df[col1].reset_index(drop=True)
series2 = training_df[col2].reset_index(drop=True)
self.eg_pvalue_ = float(coint(series1, series2)[1])
self.eg_is_cointegrated_ = bool(self.eg_pvalue_ < self.EG_PVALUE_THRESHOLD)
self.tstamp_ = training_df.index[-1]
self.pair_ = pair.name()
def to_dict(self) -> Dict[str, Any]:
return {
"tstamp": self.tstamp_,
"pair": self.pair_,
"eg_pvalue": self.eg_pvalue_,
"johansen_lr1": self.johansen_lr1_,
"johansen_cvt": self.johansen_cvt_,
"eg_is_cointegrated": self.eg_is_cointegrated_,
"johansen_is_cointegrated": self.johansen_is_cointegrated_,
}
def __repr__(self) -> str:
return f"CointegrationData(tstamp={self.tstamp_}, pair={self.pair_}, eg_pvalue={self.eg_pvalue_}, johansen_lr1={self.johansen_lr1_}, johansen_cvt={self.johansen_cvt_}, eg_is_cointegrated={self.eg_is_cointegrated_}, johansen_is_cointegrated={self.johansen_is_cointegrated_})"
class TradingPair(ABC):
market_data_: pd.DataFrame market_data_: pd.DataFrame
symbol_a_: str symbol_a_: str
symbol_b_: str symbol_b_: str
price_column_: str stat_model_price_: str
training_mu_: float training_mu_: float
training_std_: float training_std_: float
@@ -15,27 +81,81 @@ class TradingPair:
training_df_: pd.DataFrame training_df_: pd.DataFrame
testing_df_: pd.DataFrame testing_df_: pd.DataFrame
vecm_fit_: VECMResults
user_data_: Dict[str, Any] user_data_: Dict[str, Any]
# predicted_df_: Optional[pd.DataFrame]
def __init__( def __init__(
self, market_data: pd.DataFrame, symbol_a: str, symbol_b: str, price_column: str self,
config: Dict[str, Any],
market_data: pd.DataFrame,
symbol_a: str,
symbol_b: str,
): ):
self.symbol_a_ = symbol_a self.symbol_a_ = symbol_a
self.symbol_b_ = symbol_b self.symbol_b_ = symbol_b
self.price_column_ = price_column self.stat_model_price_ = config["stat_model_price"]
self.user_data_ = {}
self.predicted_df_ = None
self.config_ = config
self._set_market_data(market_data)
def _set_market_data(self, market_data: pd.DataFrame) -> None:
self.market_data_ = pd.DataFrame( self.market_data_ = pd.DataFrame(
self._transform_dataframe(market_data)[["tstamp"] + self.colnames()] self._transform_dataframe(market_data)[["tstamp"] + self.colnames()]
) )
self.market_data_ = self.market_data_.dropna().reset_index(drop=True)
self.market_data_["tstamp"] = pd.to_datetime(self.market_data_["tstamp"])
self.market_data_ = self.market_data_.sort_values("tstamp")
self._set_execution_price_data()
pass
self.user_data_ = {} def _set_execution_price_data(self) -> None:
if "execution_price" not in self.config_:
self.market_data_[f"exec_price_{self.symbol_a_}"] = self.market_data_[f"{self.stat_model_price_}_{self.symbol_a_}"]
self.market_data_[f"exec_price_{self.symbol_b_}"] = self.market_data_[f"{self.stat_model_price_}_{self.symbol_b_}"]
return
execution_price_column = self.config_["execution_price"]["column"]
execution_price_shift = self.config_["execution_price"]["shift"]
self.market_data_[f"exec_price_{self.symbol_a_}"] = self.market_data_[f"{self.stat_model_price_}_{self.symbol_a_}"].shift(-execution_price_shift)
self.market_data_[f"exec_price_{self.symbol_b_}"] = self.market_data_[f"{self.stat_model_price_}_{self.symbol_b_}"].shift(-execution_price_shift)
self.market_data_ = self.market_data_.dropna().reset_index(drop=True)
def get_begin_index(self) -> int:
if "trading_hours" not in self.config_:
return 0
assert "timezone" in self.config_["trading_hours"]
assert "begin_session" in self.config_["trading_hours"]
start_time = (
pd.to_datetime(self.config_["trading_hours"]["begin_session"])
.tz_localize(self.config_["trading_hours"]["timezone"])
.time()
)
mask = self.market_data_["tstamp"].dt.time >= start_time
return int(self.market_data_.index[mask].min())
def get_end_index(self) -> int:
if "trading_hours" not in self.config_:
return 0
assert "timezone" in self.config_["trading_hours"]
assert "end_session" in self.config_["trading_hours"]
end_time = (
pd.to_datetime(self.config_["trading_hours"]["end_session"])
.tz_localize(self.config_["trading_hours"]["timezone"])
.time()
)
mask = self.market_data_["tstamp"].dt.time <= end_time
return int(self.market_data_.index[mask].max())
def _transform_dataframe(self, df: pd.DataFrame) -> pd.DataFrame: def _transform_dataframe(self, df: pd.DataFrame) -> pd.DataFrame:
# Select only the columns we need # Select only the columns we need
df_selected: pd.DataFrame = pd.DataFrame( df_selected: pd.DataFrame = pd.DataFrame(
df[["tstamp", "symbol", self.price_column_]] df[["tstamp", "symbol", self.stat_model_price_]]
) )
# Start with unique timestamps # Start with unique timestamps
@@ -53,13 +173,13 @@ class TradingPair:
) )
# Create column name like "close-COIN" # Create column name like "close-COIN"
new_price_column = f"{self.price_column_}_{symbol}" new_price_column = f"{self.stat_model_price_}_{symbol}"
# Create temporary dataframe with timestamp and price # Create temporary dataframe with timestamp and price
temp_df = pd.DataFrame( temp_df = pd.DataFrame(
{ {
"tstamp": df_symbol["tstamp"], "tstamp": df_symbol["tstamp"],
new_price_column: df_symbol[self.price_column_], new_price_column: df_symbol[self.stat_model_price_],
} }
) )
@@ -69,7 +189,7 @@ class TradingPair:
drop=True drop=True
) # do not dropna() since irrelevant symbol would affect dataset ) # do not dropna() since irrelevant symbol would affect dataset
return result_df return result_df.dropna()
def get_datasets( def get_datasets(
self, self,
@@ -80,7 +200,7 @@ class TradingPair:
testing_start_index = training_start_index + training_minutes testing_start_index = training_start_index + training_minutes
self.training_df_ = self.market_data_.iloc[ self.training_df_ = self.market_data_.iloc[
training_start_index:testing_start_index, : training_start_index:testing_start_index, :training_minutes
].copy() ].copy()
assert self.training_df_ is not None assert self.training_df_ is not None
self.training_df_ = self.training_df_.dropna().reset_index(drop=True) self.training_df_ = self.training_df_.dropna().reset_index(drop=True)
@@ -97,112 +217,164 @@ class TradingPair:
def colnames(self) -> List[str]: def colnames(self) -> List[str]:
return [ return [
f"{self.price_column_}_{self.symbol_a_}", f"{self.stat_model_price_}_{self.symbol_a_}",
f"{self.price_column_}_{self.symbol_b_}", f"{self.stat_model_price_}_{self.symbol_b_}",
] ]
def fit_VECM(self): def exec_prices_colnames(self) -> List[str]:
assert self.training_df_ is not None return [
vecm_df = self.training_df_[self.colnames()].reset_index(drop=True) f"exec_price_{self.symbol_a_}",
vecm_model = VECM(vecm_df, coint_rank=1) f"exec_price_{self.symbol_b_}",
vecm_fit = vecm_model.fit() ]
assert vecm_fit is not None def add_trades(self, trades: pd.DataFrame) -> None:
if self.user_data_["trades"] is None or len(self.user_data_["trades"]) == 0:
# If trades is empty or None, just assign the new trades directly
self.user_data_["trades"] = trades.copy()
else:
# Ensure both DataFrames have the same columns and dtypes before concatenation
existing_trades = self.user_data_["trades"]
# URGENT check beta and alpha # If existing trades is empty, just assign the new trades
if len(existing_trades) == 0:
self.user_data_["trades"] = trades.copy()
else:
# Ensure both DataFrames have the same columns
if set(existing_trades.columns) != set(trades.columns):
# Add missing columns to trades with appropriate default values
for col in existing_trades.columns:
if col not in trades.columns:
if col == "time":
trades[col] = pd.Timestamp.now()
elif col in ["action", "symbol"]:
trades[col] = ""
elif col in [
"price",
"disequilibrium",
"scaled_disequilibrium",
]:
trades[col] = 0.0
elif col == "pair":
trades[col] = None
else:
trades[col] = None
# Check if the model converged properly # Concatenate with explicit dtypes to avoid warnings
if not hasattr(vecm_fit, "beta") or vecm_fit.beta is None: self.user_data_["trades"] = pd.concat(
print(f"{self}: VECM model failed to converge properly") [existing_trades, trades], ignore_index=True, copy=False
)
self.vecm_fit_ = vecm_fit def get_trades(self) -> pd.DataFrame:
# print(f"{self}: beta={self.vecm_fit_.beta} alpha={self.vecm_fit_.alpha}" ) return (
# print(f"{self}: {self.vecm_fit_.summary()}") self.user_data_["trades"] if "trades" in self.user_data_ else pd.DataFrame()
pass
def check_cointegration_johansen(self):
assert self.training_df_ is not None
from statsmodels.tsa.vector_ar.vecm import coint_johansen
df = self.training_df_[self.colnames()].reset_index(drop=True)
result = coint_johansen(df, det_order=0, k_ar_diff=1)
print(
f"{self}: lr1={result.lr1[0]} > cvt={result.cvt[0, 1]}? {result.lr1[0] > result.cvt[0, 1]}"
) )
is_cointegrated = result.lr1[0] > result.cvt[0, 1]
return is_cointegrated def cointegration_check(self) -> Optional[pd.DataFrame]:
print(f"***{self}*** STARTING....")
config = self.config_
def check_cointegration_engle_granger(self): curr_training_start_idx = 0
from statsmodels.tsa.stattools import coint
col1, col2 = self.colnames() COINTEGRATION_DATA_COLUMNS = {
assert self.training_df_ is not None "tstamp": "datetime64[ns]",
series1 = self.training_df_[col1].reset_index(drop=True) "pair": "string",
series2 = self.training_df_[col2].reset_index(drop=True) "eg_pvalue": "float64",
"johansen_lr1": "float64",
"johansen_cvt": "float64",
"eg_is_cointegrated": "bool",
"johansen_is_cointegrated": "bool",
}
# Initialize trades DataFrame with proper dtypes to avoid concatenation warnings
result: pd.DataFrame = pd.DataFrame(
columns=[col for col in COINTEGRATION_DATA_COLUMNS.keys()]
) # .astype(COINTEGRATION_DATA_COLUMNS)
# Run Engle-Granger cointegration test training_minutes = config["training_minutes"]
pvalue = coint(series1, series2)[1] while True:
# Define cointegration if p-value < 0.05 (i.e., reject null of no cointegration) print(curr_training_start_idx, end="\r")
is_cointegrated = pvalue < 0.05 self.get_datasets(
print(f"{self}: is_cointegrated={is_cointegrated} pvalue={pvalue}") training_minutes=training_minutes,
return is_cointegrated training_start_index=curr_training_start_idx,
testing_size=1,
)
def train_pair(self) -> bool: if len(self.training_df_) < training_minutes:
is_cointegrated_johansen = self.check_cointegration_johansen() print(
is_cointegrated_engle_granger = self.check_cointegration_engle_granger() f"{self}: current offset={curr_training_start_idx}"
if not is_cointegrated_johansen and not is_cointegrated_engle_granger: f" * Training data length={len(self.training_df_)} < {training_minutes}"
" * Not enough training data. Completing the job."
)
break
new_row = pd.Series(CointegrationData(self).to_dict())
result.loc[len(result)] = new_row
curr_training_start_idx += 1
return result
def to_stop_close_conditions(self, predicted_row: pd.Series) -> bool:
config = self.config_
if (
"stop_close_conditions" not in config
or config["stop_close_conditions"] is None
):
return False return False
pass if "profit" in config["stop_close_conditions"]:
current_return = self._current_return(predicted_row)
#
# print(f"time={predicted_row['tstamp']} current_return={current_return}")
#
if current_return >= config["stop_close_conditions"]["profit"]:
print(f"STOP PROFIT: {current_return}")
self.user_data_["stop_close_state"] = PairState.CLOSE_STOP_PROFIT
return True
if "loss" in config["stop_close_conditions"]:
if current_return <= config["stop_close_conditions"]["loss"]:
print(f"STOP LOSS: {current_return}")
self.user_data_["stop_close_state"] = PairState.CLOSE_STOP_LOSS
return True
return False
# print('*' * 80 + '\n' + f"**************** {self} IS COINTEGRATED ****************\n" + '*' * 80) def on_open_trades(self, trades: pd.DataFrame) -> None:
self.fit_VECM() if "close_trades" in self.user_data_:
assert self.training_df_ is not None and self.vecm_fit_ is not None del self.user_data_["close_trades"]
diseq_series = self.training_df_[self.colnames()] @ self.vecm_fit_.beta self.user_data_["open_trades"] = trades
print(diseq_series.shape)
self.training_mu_ = float(diseq_series[0].mean())
self.training_std_ = float(diseq_series[0].std())
self.training_df_["dis-equilibrium"] = ( def on_close_trades(self, trades: pd.DataFrame) -> None:
self.training_df_[self.colnames()] @ self.vecm_fit_.beta del self.user_data_["open_trades"]
) self.user_data_["close_trades"] = trades
# Normalize the dis-equilibrium
self.training_df_["scaled_dis-equilibrium"] = (
diseq_series - self.training_mu_
) / self.training_std_
return True def _current_return(self, predicted_row: pd.Series) -> float:
if "open_trades" in self.user_data_:
open_trades = self.user_data_["open_trades"]
if len(open_trades) == 0:
return 0.0
def predict(self) -> pd.DataFrame: def _single_instrument_return(symbol: str) -> float:
assert self.testing_df_ is not None instrument_open_trades = open_trades[open_trades["symbol"] == symbol]
assert self.vecm_fit_ is not None instrument_open_price = instrument_open_trades["price"].iloc[0]
predicted_prices = self.vecm_fit_.predict(steps=len(self.testing_df_))
# Convert prediction to a DataFrame for readability sign = -1 if instrument_open_trades["side"].iloc[0] == "SELL" else 1
# predicted_df = instrument_price = predicted_row[f"{self.stat_model_price_}_{symbol}"]
instrument_return = (
sign
* (instrument_price - instrument_open_price)
/ instrument_open_price
)
return float(instrument_return) * 100.0
self.predicted_df_ = pd.merge( instrument_a_return = _single_instrument_return(self.symbol_a_)
self.testing_df_.reset_index(drop=True), instrument_b_return = _single_instrument_return(self.symbol_b_)
pd.DataFrame( return instrument_a_return + instrument_b_return
predicted_prices, columns=pd.Index(self.colnames()), dtype=float return 0.0
),
left_index=True,
right_index=True,
suffixes=("", "_pred"),
).dropna()
self.predicted_df_["disequilibrium"] = (
self.predicted_df_[self.colnames()] @ self.vecm_fit_.beta
)
self.predicted_df_["scaled_disequilibrium"] = (
abs(self.predicted_df_["disequilibrium"] - self.training_mu_)
/ self.training_std_
)
# Reset index to ensure proper indexing
self.predicted_df_ = self.predicted_df_.reset_index()
return self.predicted_df_
def __repr__(self) -> str: def __repr__(self) -> str:
return self.name()
def name(self) -> str:
return f"{self.symbol_a_} & {self.symbol_b_}" return f"{self.symbol_a_} & {self.symbol_b_}"
# return f"{self.symbol_a_} & {self.symbol_b_}"
@abstractmethod
def predict(self) -> pd.DataFrame: ...
# @abstractmethod
# def predicted_df(self) -> Optional[pd.DataFrame]: ...
+193
View File
@@ -0,0 +1,193 @@
# original script moved to vecm_rolling_fit_01.py
# 09.09.25 Added GARCH model - predicting volatility
# Rule of thumb:
# alpha + beta ≈ 1 → strong volatility clustering, persistence.
# If much lower → volatility mean reverts quickly.
# If > 1 → model is unstable / non-stationary (bad).
# the VECM disequilibrium (mean reversion signal) and
# the GARCH volatility forecast (risk measure).
# combine them → e.g., only enter trades when:
# high_volatility = 1 → persistence > 0.95 or volatility > 2 (rule of thumb: unstable / risky regime).
# high_volatility = 0 → stable regime.
# VECM disequilibrium z-score > threshold and
# GARCH-forecasted volatility is not too high (avoid noise-driven signals).
# This creates a volatility-adjusted pairs trading strategy, more robust than plain VECM
# now pair_predict_result_ DataFrame includes:
# disequilibrium, scaled_disequilibrium, z-scores, garch_alpha, garch_beta, garch_persistence (α+β rule-of-thumb)
# garch_vol_forecast (1-step volatility forecast)
# Would you like me to also add a warning flag column
# (e.g., "high_volatility" = 1 if persistence > 0.95 or vol_forecast > threshold)
# so you can easily detect unstable regimes?
# VECM/GARCH
# vecm_rolling_fit.py:
from typing import Any, Dict, Optional, cast
import numpy as np
import pandas as pd
from typing import Any, Dict, Optional
from pt_trading.results import BacktestResult
from pt_trading.rolling_window_fit import RollingFit
from pt_trading.trading_pair import TradingPair
from statsmodels.tsa.vector_ar.vecm import VECM, VECMResults
from arch import arch_model
NanoPerMin = 1e9
class VECMTradingPair(TradingPair):
vecm_fit_: Optional[VECMResults]
pair_predict_result_: Optional[pd.DataFrame]
def __init__(
self,
config: Dict[str, Any],
market_data: pd.DataFrame,
symbol_a: str,
symbol_b: str,
):
super().__init__(config, market_data, symbol_a, symbol_b)
self.vecm_fit_ = None
self.pair_predict_result_ = None
self.garch_fit_ = None
self.sigma_spread_forecast_ = None
self.garch_alpha_ = None
self.garch_beta_ = None
self.garch_persistence_ = None
self.high_volatility_flag_ = None
def _train_pair(self) -> None:
self._fit_VECM()
assert self.vecm_fit_ is not None
diseq_series = self.training_df_[self.colnames()] @ self.vecm_fit_.beta
self.training_mu_ = float(diseq_series[0].mean())
self.training_std_ = float(diseq_series[0].std())
self.training_df_["disequilibrium"] = diseq_series
self.training_df_["scaled_disequilibrium"] = (
diseq_series - self.training_mu_
) / self.training_std_
def _fit_VECM(self) -> None:
assert self.training_df_ is not None
vecm_df = self.training_df_[self.colnames()].reset_index(drop=True)
vecm_model = VECM(vecm_df, coint_rank=1)
vecm_fit = vecm_model.fit()
self.vecm_fit_ = vecm_fit
# Error Correction Term (spread)
ect_series = (vecm_df @ vecm_fit.beta).iloc[:, 0]
# Difference the spread for stationarity
dz = ect_series.diff().dropna()
if len(dz) < 30:
print("Not enough data for GARCH fitting.")
return
# Rescale if variance too small
if dz.std() < 0.1:
dz = dz * 1000
# print("Scale check:", dz.std())
try:
garch = arch_model(dz, vol="GARCH", p=1, q=1, mean="Zero", dist="normal")
garch_fit = garch.fit(disp="off")
self.garch_fit_ = garch_fit
# Extract parameters
params = garch_fit.params
self.garch_alpha_ = params.get("alpha[1]", np.nan)
self.garch_beta_ = params.get("beta[1]", np.nan)
self.garch_persistence_ = self.garch_alpha_ + self.garch_beta_
# print (f"GARCH α: {self.garch_alpha_:.4f}, β: {self.garch_beta_:.4f}, "
# f"α+β (persistence): {self.garch_persistence_:.4f}")
# One-step-ahead volatility forecast
forecast = garch_fit.forecast(horizon=1)
sigma_next = np.sqrt(forecast.variance.iloc[-1, 0])
self.sigma_spread_forecast_ = float(sigma_next)
# print("GARCH sigma forecast:", self.sigma_spread_forecast_)
# Rule of thumb: persistence close to 1 or large volatility forecast
self.high_volatility_flag_ = int(
(self.garch_persistence_ is not None and self.garch_persistence_ > 0.95)
or (self.sigma_spread_forecast_ is not None and self.sigma_spread_forecast_ > 2)
)
except Exception as e:
print(f"GARCH fit failed: {e}")
self.garch_fit_ = None
self.sigma_spread_forecast_ = None
self.high_volatility_flag_ = None
def predict(self) -> pd.DataFrame:
self._train_pair()
assert self.testing_df_ is not None
assert self.vecm_fit_ is not None
# VECM predictions
predicted_prices = self.vecm_fit_.predict(steps=len(self.testing_df_))
predicted_df = pd.merge(
self.testing_df_.reset_index(drop=True),
pd.DataFrame(predicted_prices, columns=pd.Index(self.colnames()), dtype=float),
left_index=True,
right_index=True,
suffixes=("", "_pred"),
).dropna()
# Disequilibrium and z-scores
predicted_df["disequilibrium"] = (
predicted_df[self.colnames()] @ self.vecm_fit_.beta
)
predicted_df["signed_scaled_disequilibrium"] = (
predicted_df["disequilibrium"] - self.training_mu_
) / self.training_std_
predicted_df["scaled_disequilibrium"] = abs(
predicted_df["signed_scaled_disequilibrium"]
)
# Add GARCH parameters + volatility forecast
predicted_df["garch_alpha"] = self.garch_alpha_
predicted_df["garch_beta"] = self.garch_beta_
predicted_df["garch_persistence"] = self.garch_persistence_
predicted_df["garch_vol_forecast"] = self.sigma_spread_forecast_
predicted_df["high_volatility"] = self.high_volatility_flag_
# Save results
if self.pair_predict_result_ is None:
self.pair_predict_result_ = predicted_df
else:
self.pair_predict_result_ = pd.concat(
[self.pair_predict_result_, predicted_df], ignore_index=True
)
return self.pair_predict_result_
class VECMRollingFit(RollingFit):
def __init__(self) -> None:
super().__init__()
def create_trading_pair(
self,
config: Dict,
market_data: pd.DataFrame,
symbol_a: str,
symbol_b: str,
) -> TradingPair:
return VECMTradingPair(
config=config,
market_data=market_data,
symbol_a = symbol_a,
symbol_b = symbol_b,
)
+124
View File
@@ -0,0 +1,124 @@
from typing import Any, Dict, Optional
import pandas as pd
import statsmodels.api as sm
from pt_trading.rolling_window_fit import RollingFit
from pt_trading.trading_pair import TradingPair
NanoPerMin = 1e9
class ZScoreTradingPair(TradingPair):
"""TradingPair implementation that fits a hedge ratio with OLS and
computes a standardized spread (z-score).
The class stores training spread mean/std and hedge ratio so the model
can be applied to testing data consistently.
"""
zscore_model_: Optional[sm.regression.linear_model.RegressionResultsWrapper]
pair_predict_result_: Optional[pd.DataFrame]
zscore_df_: Optional[pd.Series]
hedge_ratio_: Optional[float]
spread_mean_: Optional[float]
spread_std_: Optional[float]
def __init__(
self,
config: Dict[str, Any],
market_data: pd.DataFrame,
symbol_a: str,
symbol_b: str,
):
super().__init__(config, market_data, symbol_a, symbol_b)
self.zscore_model_ = None
self.pair_predict_result_ = None
self.zscore_df_ = None
self.hedge_ratio_ = None
self.spread_mean_ = None
self.spread_std_ = None
def _fit_zscore(self) -> None:
"""Fit OLS on the training window and compute training z-score."""
assert self.training_df_ is not None
# Extract price series for the two symbols from the training frame.
px_df = self.training_df_[self.colnames()]
symbol_a_px = px_df.iloc[:, 0]
symbol_b_px = px_df.iloc[:, 1]
# Align indexes and fit OLS: symbol_a ~ const + symbol_b
symbol_a_px, symbol_b_px = symbol_a_px.align(symbol_b_px, join="inner")
X = sm.add_constant(symbol_b_px)
self.zscore_model_ = sm.OLS(symbol_a_px, X).fit()
# Hedge ratio is the slope on symbol_b
params = self.zscore_model_.params
self.hedge_ratio_ = float(params.iloc[1]) if len(params) > 1 else 0.0
# Training spread and its standardized z-score
spread = symbol_a_px - self.hedge_ratio_ * symbol_b_px
self.spread_mean_ = float(spread.mean())
self.spread_std_ = float(spread.std(ddof=0)) if spread.std(ddof=0) != 0 else 1.0
self.zscore_df_ = (spread - self.spread_mean_) / self.spread_std_
def predict(self) -> pd.DataFrame:
"""Apply fitted hedge ratio to the testing frame and return a
dataframe with canonical columns:
- disequilibrium: signed z-score
- scaled_disequilibrium: absolute z-score
- signed_scaled_disequilibrium: same as disequilibrium (keeps sign)
"""
# Fit on training window
self._fit_zscore()
assert self.zscore_df_ is not None
assert self.hedge_ratio_ is not None
assert self.spread_mean_ is not None and self.spread_std_ is not None
# Keep training columns for inspection
self.training_df_["disequilibrium"] = self.zscore_df_
self.training_df_["scaled_disequilibrium"] = self.zscore_df_.abs()
# Apply model to testing frame
assert self.testing_df_ is not None
test_df = self.testing_df_.copy()
px_test = test_df[self.colnames()]
a_test = px_test.iloc[:, 0]
b_test = px_test.iloc[:, 1]
a_test, b_test = a_test.align(b_test, join="inner")
# Compute test spread and standardize using training mean/std
test_spread = a_test - self.hedge_ratio_ * b_test
test_zscore = (test_spread - self.spread_mean_) / self.spread_std_
# Attach canonical columns
# Align back to test_df index if needed
test_zscore = test_zscore.reindex(test_df.index)
test_df["disequilibrium"] = test_zscore
test_df["signed_scaled_disequilibrium"] = test_zscore
test_df["scaled_disequilibrium"] = test_zscore.abs()
# Reset index and accumulate results across windows
test_df = test_df.reset_index(drop=True)
if self.pair_predict_result_ is None:
self.pair_predict_result_ = test_df
else:
self.pair_predict_result_ = pd.concat(
[self.pair_predict_result_, test_df], ignore_index=True
)
self.pair_predict_result_ = self.pair_predict_result_.reset_index(drop=True)
return self.pair_predict_result_.dropna()
class ZScoreRollingFit(RollingFit):
def __init__(self) -> None:
super().__init__()
def create_trading_pair(
self, config: Dict, market_data: pd.DataFrame, symbol_a: str, symbol_b: str
) -> TradingPair:
return ZScoreTradingPair(
config=config, market_data=market_data, symbol_a=symbol_a, symbol_b=symbol_b
)
+81 -68
View File
@@ -1,10 +1,17 @@
from __future__ import annotations
import sqlite3 import sqlite3
from typing import Dict, List, cast from typing import Dict, List, cast
import pandas as pd import pandas as pd
def load_sqlite_to_dataframe(db_path:str, query:str) -> pd.DataFrame:
df: pd.DataFrame = pd.DataFrame()
import os
if not os.path.exists(db_path):
print(f"WARNING: database file {db_path} does not exist")
return df
def load_sqlite_to_dataframe(db_path, query):
try: try:
conn = sqlite3.connect(db_path) conn = sqlite3.connect(db_path)
@@ -21,13 +28,14 @@ def load_sqlite_to_dataframe(db_path, query):
conn.close() conn.close()
def convert_time_to_UTC(value: str, timezone: str) -> str: def convert_time_to_UTC(value: str, timezone: str, extra_minutes: int = 0) -> str:
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from datetime import datetime from datetime import datetime, timedelta
# Parse it to naive datetime object # Parse it to naive datetime object
local_dt = datetime.strptime(value, "%Y-%m-%d %H:%M:%S") local_dt = datetime.strptime(value, "%Y-%m-%d %H:%M:%S")
local_dt = local_dt + timedelta(minutes=extra_minutes)
zinfo = ZoneInfo(timezone) zinfo = ZoneInfo(timezone)
result: datetime = local_dt.replace(tzinfo=zinfo).astimezone(ZoneInfo("UTC")) result: datetime = local_dt.replace(tzinfo=zinfo).astimezone(ZoneInfo("UTC"))
@@ -35,25 +43,28 @@ def convert_time_to_UTC(value: str, timezone: str) -> str:
return result.strftime("%Y-%m-%d %H:%M:%S") return result.strftime("%Y-%m-%d %H:%M:%S")
def load_market_data(datafile: str, config: Dict) -> pd.DataFrame: def load_market_data(
from tools.data_loader import load_sqlite_to_dataframe datafile: str,
instruments: List[Dict[str, str]],
db_table_name: str,
trading_hours: Dict = {},
extra_minutes: int = 0,
) -> pd.DataFrame:
instrument_ids = [ insts = [
'"' + config["instrument_id_pfx"] + instrument + '"' '"' + instrument["instrument_id_pfx"] + instrument["symbol"] + '"'
for instrument in config["instruments"] for instrument in instruments
] ]
security_type = config["security_type"] instrument_ids = list(set(insts))
exchange_id = config["exchange_id"] exchange_ids = list(
set(['"' + instrument["exchange_id"] + '"' for instrument in instruments])
)
query = "select" query = "select"
if security_type == "CRYPTO": query += " tstamp"
query += " strftime('%Y-%m-%d %H:%M:%S', tstamp_ns/1000000000, 'unixepoch') as tstamp" query += ", tstamp_ns as time_ns"
query += ", tstamp as time_ns"
else:
query += " tstamp"
query += ", tstamp_ns as time_ns"
query += f", substr(instrument_id, {len(config['instrument_id_pfx']) + 1}) as symbol" query += f", substr(instrument_id, instr(instrument_id, '-') + 1) as symbol"
query += ", open" query += ", open"
query += ", high" query += ", high"
query += ", low" query += ", low"
@@ -62,74 +73,76 @@ def load_market_data(datafile: str, config: Dict) -> pd.DataFrame:
query += ", num_trades" query += ", num_trades"
query += ", vwap" query += ", vwap"
query += f" from {config['db_table_name']}" query += f" from {db_table_name}"
query += f" where exchange_id ='{exchange_id}'" query += f" where exchange_id in ({','.join(exchange_ids)})"
query += f" and instrument_id in ({','.join(instrument_ids)})" query += f" and instrument_id in ({','.join(instrument_ids)})"
df = load_sqlite_to_dataframe(db_path=datafile, query=query) df = load_sqlite_to_dataframe(db_path=datafile, query=query)
# Trading Hours # Trading Hours
date_str = df["tstamp"][0][0:10] if len(df) > 0 and len(trading_hours) > 0:
trading_hours = config["trading_hours"] date_str = df["tstamp"][0][0:10]
start_time = convert_time_to_UTC( start_time = convert_time_to_UTC(
f"{date_str} {trading_hours['begin_session']}", trading_hours["timezone"] f"{date_str} {trading_hours['begin_session']}", trading_hours["timezone"]
) )
end_time = convert_time_to_UTC( end_time = convert_time_to_UTC(
f"{date_str} {trading_hours['end_session']}", trading_hours["timezone"] f"{date_str} {trading_hours['end_session']}", trading_hours["timezone"], extra_minutes=extra_minutes # to get execution price
) )
# Perform boolean selection # Perform boolean selection
df = df[(df["tstamp"] >= start_time) & (df["tstamp"] <= end_time)] df = df[(df["tstamp"] >= start_time) & (df["tstamp"] <= end_time)]
df["tstamp"] = pd.to_datetime(df["tstamp"]) df["tstamp"] = pd.to_datetime(df["tstamp"])
return cast(pd.DataFrame, df) return cast(pd.DataFrame, df)
def get_available_instruments_from_db(datafile: str, config: Dict) -> List[str]: # def get_available_instruments_from_db(datafile: str, config: Dict) -> List[str]:
""" # """
Auto-detect available instruments from the database by querying distinct instrument_id values. # Auto-detect available instruments from the database by querying distinct instrument_id values.
Returns instruments without the configured prefix. # Returns instruments without the configured prefix.
""" # """
try: # try:
conn = sqlite3.connect(datafile) # conn = sqlite3.connect(datafile)
# Build exclusion list with full instrument_ids # # Build exclusion list with full instrument_ids
exclude_instruments = config.get("exclude_instruments", []) # exclude_instruments = config.get("exclude_instruments", [])
prefix = config.get("instrument_id_pfx", "") # prefix = config.get("instrument_id_pfx", "")
exclude_instrument_ids = [f"{prefix}{inst}" for inst in exclude_instruments] # exclude_instrument_ids = [f"{prefix}{inst}" for inst in exclude_instruments]
# Query to get distinct instrument_ids # # Query to get distinct instrument_ids
query = f""" # query = f"""
SELECT DISTINCT instrument_id # SELECT DISTINCT instrument_id
FROM {config['db_table_name']} # FROM {config['db_table_name']}
WHERE exchange_id = ? # WHERE exchange_id = ?
""" # """
# Add exclusion clause if there are instruments to exclude # # Add exclusion clause if there are instruments to exclude
if exclude_instrument_ids: # if exclude_instrument_ids:
placeholders = ','.join(['?' for _ in exclude_instrument_ids]) # placeholders = ",".join(["?" for _ in exclude_instrument_ids])
query += f" AND instrument_id NOT IN ({placeholders})" # query += f" AND instrument_id NOT IN ({placeholders})"
cursor = conn.execute(query, (config["exchange_id"],) + tuple(exclude_instrument_ids)) # cursor = conn.execute(
else: # query, (config["exchange_id"],) + tuple(exclude_instrument_ids)
cursor = conn.execute(query, (config["exchange_id"],)) # )
instrument_ids = [row[0] for row in cursor.fetchall()] # else:
conn.close() # cursor = conn.execute(query, (config["exchange_id"],))
# instrument_ids = [row[0] for row in cursor.fetchall()]
# conn.close()
# Remove the configured prefix to get instrument symbols # # Remove the configured prefix to get instrument symbols
instruments = [] # instruments = []
for instrument_id in instrument_ids: # for instrument_id in instrument_ids:
if instrument_id.startswith(prefix): # if instrument_id.startswith(prefix):
symbol = instrument_id[len(prefix) :] # symbol = instrument_id[len(prefix) :]
instruments.append(symbol) # instruments.append(symbol)
else: # else:
instruments.append(instrument_id) # instruments.append(instrument_id)
return sorted(instruments) # return sorted(instruments)
except Exception as e: # except Exception as e:
print(f"Error auto-detecting instruments from {datafile}: {str(e)}") # print(f"Error auto-detecting instruments from {datafile}: {str(e)}")
return [] # return []
# if __name__ == "__main__": # if __name__ == "__main__":
+114 -106
View File
@@ -24,11 +24,14 @@ hjson>=3.0.2
html5lib>=1.1 html5lib>=1.1
httplib2>=0.20.2 httplib2>=0.20.2
idna>=3.3 idna>=3.3
ipython>=8.18.1
ipywidgets>=8.1.1
ifaddr>=0.1.7 ifaddr>=0.1.7
IMDbPY>=2021.4.18 IMDbPY>=2021.4.18
ipykernel>=6.29.5 ipykernel>=6.29.5
jeepney>=0.7.1 jeepney>=0.7.1
jsonschema>=3.2.0 jsonschema>=3.2.0
jupyter>=1.0.0
keyring>=23.5.0 keyring>=23.5.0
launchpadlib>=1.10.16 launchpadlib>=1.10.16
lazr.restfulclient>=0.14.4 lazr.restfulclient>=0.14.4
@@ -42,19 +45,23 @@ more-itertools>=8.10.0
multidict>=6.0.4 multidict>=6.0.4
mypy>=0.942 mypy>=0.942
mypy-extensions>=0.4.3 mypy-extensions>=0.4.3
nbformat>=5.10.2
netaddr>=0.8.0 netaddr>=0.8.0
######### netifaces>=0.11.0 ######### netifaces>=0.11.0
numpy>=1.26.4,<2.3.0
oauthlib>=3.2.0 oauthlib>=3.2.0
packaging>=23.1 packaging>=23.1
pandas>=2.2.3
pathspec>=0.11.1 pathspec>=0.11.1
pexpect>=4.8.0 pexpect>=4.8.0
Pillow>=9.0.1 Pillow>=9.0.1
platformdirs>=3.2.0 platformdirs>=3.2.0
plotly>=5.19.0
protobuf>=3.12.4 protobuf>=3.12.4
psutil>=5.9.0 psutil>=5.9.0
ptyprocess>=0.7.0 ptyprocess>=0.7.0
pycurl>=7.44.1 pycurl>=7.44.1
pyelftools>=0.27 # pyelftools>=0.27
Pygments>=2.11.2 Pygments>=2.11.2
pyparsing>=2.4.7 pyparsing>=2.4.7
pyrsistent>=0.18.1 pyrsistent>=0.18.1
@@ -62,11 +69,12 @@ python-debian>=0.1.43 #+ubuntu1.1
python-dotenv>=0.19.2 python-dotenv>=0.19.2
python-magic>=0.4.24 python-magic>=0.4.24
python-xlib>=0.29 python-xlib>=0.29
pyxdg>=0.27 # pyxdg>=0.27
PyYAML>=6.0 PyYAML>=6.0
reportlab>=3.6.8 reportlab>=3.6.8
requests>=2.25.1 requests>=2.25.1
requests-file>=1.5.1 requests-file>=1.5.1
scipy<1.13.0
seaborn>=0.13.2 seaborn>=0.13.2
SecretStorage>=3.3.1 SecretStorage>=3.3.1
setproctitle>=1.2.2 setproctitle>=1.2.2
@@ -74,113 +82,113 @@ six>=1.16.0
soupsieve>=2.3.1 soupsieve>=2.3.1
ssh-import-id>=5.11 ssh-import-id>=5.11
statsmodels>=0.14.4 statsmodels>=0.14.4
texttable>=1.6.4 # texttable>=1.6.4
tldextract>=3.1.2 tldextract>=3.1.2
tomli>=1.2.2 tomli>=1.2.2
######## typed-ast>=1.4.3 ######## typed-ast>=1.4.3
types-aiofiles>=0.1 # types-aiofiles>=0.1
types-annoy>=1.17 # types-annoy>=1.17
types-appdirs>=1.4 # types-appdirs>=1.4
types-atomicwrites>=1.4 # types-atomicwrites>=1.4
types-aws-xray-sdk>=2.8 # types-aws-xray-sdk>=2.8
types-babel>=2.9 # types-babel>=2.9
types-backports-abc>=0.5 # types-backports-abc>=0.5
types-backports.ssl-match-hostname>=3.7 # types-backports.ssl-match-hostname>=3.7
types-beautifulsoup4>=4.10 # types-beautifulsoup4>=4.10
types-bleach>=4.1 # types-bleach>=4.1
types-boto>=2.49 # types-boto>=2.49
types-braintree>=4.11 # types-braintree>=4.11
types-cachetools>=4.2 # types-cachetools>=4.2
types-caldav>=0.8 # types-caldav>=0.8
types-certifi>=2020.4 # types-certifi>=2020.4
types-characteristic>=14.3 # types-characteristic>=14.3
types-chardet>=4.0 # types-chardet>=4.0
types-click>=7.1 # types-click>=7.1
types-click-spinner>=0.1 # types-click-spinner>=0.1
types-colorama>=0.4 # types-colorama>=0.4
types-commonmark>=0.9 # types-commonmark>=0.9
types-contextvars>=0.1 # types-contextvars>=0.1
types-croniter>=1.0 # types-croniter>=1.0
types-cryptography>=3.3 # types-cryptography>=3.3
types-dataclasses>=0.1 # types-dataclasses>=0.1
types-dateparser>=1.0 # types-dateparser>=1.0
types-DateTimeRange>=0.1 # types-DateTimeRange>=0.1
types-decorator>=0.1 # types-decorator>=0.1
types-Deprecated>=1.2 # types-Deprecated>=1.2
types-docopt>=0.6 # types-docopt>=0.6
types-docutils>=0.17 # types-docutils>=0.17
types-editdistance>=0.5 # types-editdistance>=0.5
types-emoji>=1.2 # types-emoji>=1.2
types-entrypoints>=0.3 # types-entrypoints>=0.3
types-enum34>=1.1 # types-enum34>=1.1
types-filelock>=3.2 # types-filelock>=3.2
types-first>=2.0 # types-first>=2.0
types-Flask>=1.1 # types-Flask>=1.1
types-freezegun>=1.1 # types-freezegun>=1.1
types-frozendict>=0.1 # types-frozendict>=0.1
types-futures>=3.3 # types-futures>=3.3
types-html5lib>=1.1 # types-html5lib>=1.1
types-httplib2>=0.19 # types-httplib2>=0.19
types-humanfriendly>=9.2 # types-humanfriendly>=9.2
types-ipaddress>=1.0 # types-ipaddress>=1.0
types-itsdangerous>=1.1 # types-itsdangerous>=1.1
types-JACK-Client>=0.1 # types-JACK-Client>=0.1
types-Jinja2>=2.11 # types-Jinja2>=2.11
types-jmespath>=0.10 # types-jmespath>=0.10
types-jsonschema>=3.2 # types-jsonschema>=3.2
types-Markdown>=3.3 # types-Markdown>=3.3
types-MarkupSafe>=1.1 # types-MarkupSafe>=1.1
types-mock>=4.0 # types-mock>=4.0
types-mypy-extensions>=0.4 # types-mypy-extensions>=0.4
types-mysqlclient>=2.0 # types-mysqlclient>=2.0
types-oauthlib>=3.1 # types-oauthlib>=3.1
types-orjson>=3.6 # types-orjson>=3.6
types-paramiko>=2.7 # types-paramiko>=2.7
types-Pillow>=8.3 # types-Pillow>=8.3
types-polib>=1.1 # types-polib>=1.1
types-prettytable>=2.1 # types-prettytable>=2.1
types-protobuf>=3.17 # types-protobuf>=3.17
types-psutil>=5.8 # types-psutil>=5.8
types-psycopg2>=2.9 # types-psycopg2>=2.9
types-pyaudio>=0.2 # types-pyaudio>=0.2
types-pycurl>=0.1 # types-pycurl>=0.1
types-pyfarmhash>=0.2 # types-pyfarmhash>=0.2
types-Pygments>=2.9 # types-Pygments>=2.9
types-PyMySQL>=1.0 # types-PyMySQL>=1.0
types-pyOpenSSL>=20.0 # types-pyOpenSSL>=20.0
types-pyRFC3339>=0.1 # types-pyRFC3339>=0.1
types-pysftp>=0.2 # types-pysftp>=0.2
types-pytest-lazy-fixture>=0.6 # types-pytest-lazy-fixture>=0.6
types-python-dateutil>=2.8 # types-python-dateutil>=2.8
types-python-gflags>=3.1 # types-python-gflags>=3.1
types-python-nmap>=0.6 # types-python-nmap>=0.6
types-python-slugify>=5.0 # types-python-slugify>=5.0
types-pytz>=2021.1 # types-pytz>=2021.1
types-pyvmomi>=7.0 # types-pyvmomi>=7.0
types-PyYAML>=5.4 # types-PyYAML>=5.4
types-redis>=3.5 # types-redis>=3.5
types-requests>=2.25 # types-requests>=2.25
types-retry>=0.9 # types-retry>=0.9
types-selenium>=3.141 # types-selenium>=3.141
types-Send2Trash>=1.8 # types-Send2Trash>=1.8
types-setuptools>=57.4 # types-setuptools>=57.4
types-simplejson>=3.17 # types-simplejson>=3.17
types-singledispatch>=3.7 # types-singledispatch>=3.7
types-six>=1.16 # types-six>=1.16
types-slumber>=0.7 # types-slumber>=0.7
types-stripe>=2.59 # types-stripe>=2.59
types-tabulate>=0.8 # types-tabulate>=0.8
types-termcolor>=1.1 # types-termcolor>=1.1
types-toml>=0.10 # types-toml>=0.10
types-toposort>=1.6 # types-toposort>=1.6
types-ttkthemes>=3.2 # types-ttkthemes>=3.2
types-typed-ast>=1.4 # types-typed-ast>=1.4
types-tzlocal>=0.1 # types-tzlocal>=0.1
types-ujson>=0.1 # types-ujson>=0.1
types-vobject>=0.9 # types-vobject>=0.9
types-waitress>=0.1 # types-waitress>=0.1
types-Werkzeug>=1.0 #types-Werkzeug>=1.0
types-xxhash>=2.0 #types-xxhash>=2.0
typing-extensions>=3.10.0.2 typing-extensions>=3.10.0.2
Unidecode>=1.3.3 Unidecode>=1.3.3
urllib3>=1.26.5 urllib3>=1.26.5
+127
View File
@@ -0,0 +1,127 @@
import argparse
import glob
import importlib
import os
from datetime import date, datetime
from typing import Any, Dict, List, Optional
import pandas as pd
from tools.config import expand_filename, load_config
from tools.data_loader import get_available_instruments_from_db
from pt_trading.results import (
BacktestResult,
create_result_database,
store_config_in_database,
store_results_in_database,
)
from pt_trading.fit_method import PairsTradingFitMethod
from pt_trading.trading_pair import TradingPair
from research.research_tools import create_pairs, resolve_datafiles
def main() -> None:
parser = argparse.ArgumentParser(description="Run pairs trading backtest.")
parser.add_argument(
"--config", type=str, required=True, help="Path to the configuration file."
)
parser.add_argument(
"--datafile",
type=str,
required=False,
help="Market data file to process.",
)
parser.add_argument(
"--instruments",
type=str,
required=False,
help = "Comma-separated list of instrument symbols (e.g., COIN,GBTC). If not provided, auto-detects from database.",
)
args = parser.parse_args()
config: Dict = load_config(args.config)
# Resolve data files (CLI takes priority over config)
datafile = resolve_datafiles(config, args.datafile)[0]
if not datafile:
print("No data files found to process.")
return
print(f"Found {datafile} data files to process:")
# # Create result database if needed
# if args.result_db.upper() != "NONE":
# args.result_db = expand_filename(args.result_db)
# create_result_database(args.result_db)
# # Initialize a dictionary to store all trade results
# all_results: Dict[str, Dict[str, Any]] = {}
# # Store configuration in database for reference
# if args.result_db.upper() != "NONE":
# # Get list of all instruments for storage
# all_instruments = []
# for datafile in datafiles:
# if args.instruments:
# file_instruments = [
# inst.strip() for inst in args.instruments.split(",")
# ]
# else:
# file_instruments = get_available_instruments_from_db(datafile, config)
# all_instruments.extend(file_instruments)
# # Remove duplicates while preserving order
# unique_instruments = list(dict.fromkeys(all_instruments))
# store_config_in_database(
# db_path=args.result_db,
# config_file_path=args.config,
# config=config,
# fit_method_class=fit_method_class_name,
# datafiles=datafiles,
# instruments=unique_instruments,
# )
# Process each data file
stat_model_price = config["stat_model_price"]
print(f"\n====== Processing {os.path.basename(datafile)} ======")
# Determine instruments to use
if args.instruments:
# Use CLI-specified instruments
instruments = [inst.strip() for inst in args.instruments.split(",")]
print(f"Using CLI-specified instruments: {instruments}")
else:
# Auto-detect instruments from database
instruments = get_available_instruments_from_db(datafile, config)
print(f"Auto-detected instruments: {instruments}")
if not instruments:
print(f"No instruments found in {datafile}...")
return
# Process data for this file
try:
cointegration_data: pd.DataFrame = pd.DataFrame()
for pair in create_pairs(datafile, stat_model_price, config, instruments):
cointegration_data = pd.concat([cointegration_data, pair.cointegration_check()])
pd.set_option('display.width', 400)
pd.set_option('display.max_colwidth', None)
pd.set_option('display.max_columns', None)
with pd.option_context('display.max_rows', None, 'display.max_columns', None):
print(f"cointegration_data:\n{cointegration_data}")
except Exception as err:
print(f"Error processing {datafile}: {str(err)}")
import traceback
traceback.print_exc()
if __name__ == "__main__":
main()
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
-771
View File
@@ -1,771 +0,0 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Pairs Trading Visualization Notebook\n",
"\n",
"This notebook allows you to visualize pairs trading strategies on individual instrument pairs.\n",
"You can examine the relationship between two instruments, their dis-equilibrium, and trading signals."
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### 🎯 Key Features:\n",
"\n",
"1. **Interactive Configuration**: \n",
" - Easy switching between CRYPTO and EQUITY configurations\n",
" - Simple parameter adjustment for thresholds and training periods\n",
"\n",
"2. **Single Pair Focus**: \n",
" - Instead of running multiple pairs, focuses on one pair at a time\n",
" - Allows deep analysis of the relationship between two instruments\n",
"\n",
"3. **Step-by-Step Visualization**:\n",
" - **Raw price data**: Individual prices, normalized comparison, and price ratios\n",
" - **Training analysis**: Cointegration testing and VECM model fitting\n",
" - **Dis-equilibrium visualization**: Both raw and scaled dis-equilibrium with threshold lines\n",
" - **Strategy execution**: Trading signal generation and visualization\n",
" - **Prediction analysis**: Actual vs predicted prices with trading signals overlaid\n",
"\n",
"4. **Rich Analytics**:\n",
" - Cointegration status and VECM model details\n",
" - Statistical summaries for all stages\n",
" - Threshold crossing analysis\n",
" - Trading signal breakdown\n",
"\n",
"5. **Interactive Experimentation**:\n",
" - Easy parameter modification\n",
" - Re-run capabilities for different configurations\n",
" - Support for both StaticFitStrategy and SlidingFitStrategy\n",
"\n",
"### 🚀 How to Use:\n",
"\n",
"1. **Start Jupyter**:\n",
" ```bash\n",
" cd src/notebooks\n",
" jupyter notebook pairs_trading_visualization.ipynb\n",
" ```\n",
"\n",
"2. **Customize Your Analysis**:\n",
" - Change `SYMBOL_A` and `SYMBOL_B` to your desired trading pair\n",
" - Switch between `CRYPTO_CONFIG` and `EQT_CONFIG`\n",
" - Only **StaticFitStrategy** is supported. \n",
" - Adjust thresholds and parameters as needed\n",
"\n",
"3. **Run and Visualize**:\n",
" - Execute cells step by step to see the analysis unfold\n",
" - Rich matplotlib visualizations show relationships and signals\n",
" - Comprehensive summary at the end\n",
"\n",
"The notebook provides exactly what you requested - a way to visualize the relationship between two instruments and their scaled dis-equilibrium, with all the stages of your pairs trading strategy clearly displayed and analyzed.\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Setup and Imports"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Setup complete!\n"
]
}
],
"source": [
"import sys\n",
"import os\n",
"sys.path.append('..')\n",
"\n",
"import pandas as pd\n",
"import numpy as np\n",
"import matplotlib.pyplot as plt\n",
"import seaborn as sns\n",
"from typing import Dict, List, Optional\n",
"\n",
"# Import our modules\n",
"from pt_trading.fit_methods import StaticFit, SlidingFit\n",
"from tools.data_loader import load_market_data\n",
"from pt_trading.trading_pair import TradingPair\n",
"from pt_trading.results import BacktestResult\n",
"\n",
"# Set plotting style\n",
"plt.style.use('seaborn-v0_8')\n",
"sns.set_palette(\"husl\")\n",
"plt.rcParams['figure.figsize'] = (12, 8)\n",
"\n",
"print(\"Setup complete!\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Configuration"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Using EQUITY configuration\n",
"Available instruments: ['COIN', 'GBTC', 'HOOD', 'MSTR', 'PYPL']\n"
]
}
],
"source": [
"# Configuration - Choose between CRYPTO_CONFIG or EQT_CONFIG\n",
"\n",
"CRYPTO_CONFIG = {\n",
" \"security_type\": \"CRYPTO\",\n",
" \"data_directory\": \"../../data/crypto\",\n",
" \"datafiles\": [\n",
" \"20250519.mktdata.ohlcv.db\",\n",
" ],\n",
" \"db_table_name\": \"bnbspot_ohlcv_1min\",\n",
" \"exchange_id\": \"BNBSPOT\",\n",
" \"instrument_id_pfx\": \"PAIR-\",\n",
" \"instruments\": [\n",
" \"BTC-USDT\",\n",
" \"BCH-USDT\",\n",
" \"ETH-USDT\",\n",
" \"LTC-USDT\",\n",
" \"XRP-USDT\",\n",
" \"ADA-USDT\",\n",
" \"SOL-USDT\",\n",
" \"DOT-USDT\",\n",
" ],\n",
" \"trading_hours\": {\n",
" \"begin_session\": \"00:00:00\",\n",
" \"end_session\": \"23:59:00\",\n",
" \"timezone\": \"UTC\",\n",
" },\n",
" \"price_column\": \"close\",\n",
" \"min_required_points\": 30,\n",
" \"zero_threshold\": 1e-10,\n",
" \"dis-equilibrium_open_trshld\": 2.0,\n",
" \"dis-equilibrium_close_trshld\": 0.5,\n",
" \"training_minutes\": 120,\n",
" \"funding_per_pair\": 2000.0,\n",
"}\n",
"\n",
"EQT_CONFIG = {\n",
" \"security_type\": \"EQUITY\",\n",
" \"data_directory\": \"../../data/equity\",\n",
" \"datafiles\": {\n",
" \"0508\": \"20250508.alpaca_sim_md.db\",\n",
" \"0509\": \"20250509.alpaca_sim_md.db\",\n",
" \"0510\": \"20250510.alpaca_sim_md.db\",\n",
" \"0511\": \"20250511.alpaca_sim_md.db\",\n",
" \"0512\": \"20250512.alpaca_sim_md.db\",\n",
" \"0513\": \"20250513.alpaca_sim_md.db\",\n",
" \"0514\": \"20250514.alpaca_sim_md.db\",\n",
" \"0515\": \"20250515.alpaca_sim_md.db\",\n",
" \"0516\": \"20250516.alpaca_sim_md.db\",\n",
" \"0517\": \"20250517.alpaca_sim_md.db\",\n",
" \"0518\": \"20250518.alpaca_sim_md.db\",\n",
" \"0519\": \"20250519.alpaca_sim_md.db\",\n",
" \"0520\": \"20250520.alpaca_sim_md.db\",\n",
" \"0521\": \"20250521.alpaca_sim_md.db\",\n",
" \"0522\": \"20250522.alpaca_sim_md.db\",\n",
" },\n",
" \"db_table_name\": \"md_1min_bars\",\n",
" \"exchange_id\": \"ALPACA\",\n",
" \"instrument_id_pfx\": \"STOCK-\",\n",
" \"instruments\": [\n",
" \"COIN\",\n",
" \"GBTC\",\n",
" \"HOOD\",\n",
" \"MSTR\",\n",
" \"PYPL\",\n",
" ],\n",
" \"trading_hours\": {\n",
" \"begin_session\": \"9:30:00\",\n",
" \"end_session\": \"16:00:00\",\n",
" \"timezone\": \"America/New_York\",\n",
" },\n",
" \"price_column\": \"close\",\n",
" \"min_required_points\": 30,\n",
" \"zero_threshold\": 1e-10,\n",
" \"dis-equilibrium_open_trshld\": 2.0,\n",
" \"dis-equilibrium_close_trshld\": 1.0, #0.5,\n",
" \"training_minutes\": 120,\n",
" \"funding_per_pair\": 2000.0,\n",
"}\n",
"\n",
"# Choose your configuration\n",
"CONFIG = EQT_CONFIG # Change to CRYPTO_CONFIG if you want to use crypto data\n",
"\n",
"print(f\"Using {CONFIG['security_type']} configuration\")\n",
"print(f\"Available instruments: {CONFIG['instruments']}\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Select Trading Pair and Data File"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Selected pair: COIN & GBTC\n",
"Data file: 20250509.alpaca_sim_md.db\n",
"Strategy: StaticFitStrategy\n"
]
}
],
"source": [
"# Select your trading pair and strategy\n",
"SYMBOL_A = \"COIN\" # Change these to your desired symbols\n",
"SYMBOL_B = \"GBTC\"\n",
"DATA_FILE = CONFIG[\"datafiles\"][\"0509\"]\n",
"\n",
"# Choose strategy\n",
"FIT_METHOD = StaticFit()\n",
"\n",
"print(f\"Selected pair: {SYMBOL_A} & {SYMBOL_B}\")\n",
"print(f\"Data file: {DATA_FILE}\")\n",
"print(f\"Strategy: {type(FIT_METHOD).__name__}\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Load Market Data"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Current working directory: /home/oleg/devel/pairs_trading/src/notebooks\n",
"Loading data from: ../../data/equity/20250509.alpaca_sim_md.db\n",
"Error: Execution failed on sql 'select tstamp, tstamp_ns as time_ns, substr(instrument_id, 7) as symbol, open, high, low, close, volume, num_trades, vwap from md_1min_bars where exchange_id ='ALPACA' and instrument_id in (\"STOCK-COIN\",\"STOCK-GBTC\",\"STOCK-HOOD\",\"STOCK-MSTR\",\"STOCK-PYPL\")': no such table: md_1min_bars\n"
]
},
{
"ename": "Exception",
"evalue": "",
"output_type": "error",
"traceback": [
"\u001b[31m---------------------------------------------------------------------------\u001b[39m",
"\u001b[31mOperationalError\u001b[39m Traceback (most recent call last)",
"\u001b[36mFile \u001b[39m\u001b[32m~/.pyenv/python3.12-venv/lib/python3.12/site-packages/pandas/io/sql.py:2664\u001b[39m, in \u001b[36mSQLiteDatabase.execute\u001b[39m\u001b[34m(self, sql, params)\u001b[39m\n\u001b[32m 2663\u001b[39m \u001b[38;5;28;01mtry\u001b[39;00m:\n\u001b[32m-> \u001b[39m\u001b[32m2664\u001b[39m \u001b[43mcur\u001b[49m\u001b[43m.\u001b[49m\u001b[43mexecute\u001b[49m\u001b[43m(\u001b[49m\u001b[43msql\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43m*\u001b[49m\u001b[43margs\u001b[49m\u001b[43m)\u001b[49m\n\u001b[32m 2665\u001b[39m \u001b[38;5;28;01mreturn\u001b[39;00m cur\n",
"\u001b[31mOperationalError\u001b[39m: no such table: md_1min_bars",
"\nThe above exception was the direct cause of the following exception:\n",
"\u001b[31mDatabaseError\u001b[39m Traceback (most recent call last)",
"\u001b[36mFile \u001b[39m\u001b[32m~/devel/pairs_trading/src/notebooks/../tools/data_loader.py:11\u001b[39m, in \u001b[36mload_sqlite_to_dataframe\u001b[39m\u001b[34m(db_path, query)\u001b[39m\n\u001b[32m 9\u001b[39m conn = sqlite3.connect(db_path)\n\u001b[32m---> \u001b[39m\u001b[32m11\u001b[39m df = \u001b[43mpd\u001b[49m\u001b[43m.\u001b[49m\u001b[43mread_sql_query\u001b[49m\u001b[43m(\u001b[49m\u001b[43mquery\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mconn\u001b[49m\u001b[43m)\u001b[49m\n\u001b[32m 12\u001b[39m \u001b[38;5;28;01mreturn\u001b[39;00m df\n",
"\u001b[36mFile \u001b[39m\u001b[32m~/.pyenv/python3.12-venv/lib/python3.12/site-packages/pandas/io/sql.py:528\u001b[39m, in \u001b[36mread_sql_query\u001b[39m\u001b[34m(sql, con, index_col, coerce_float, params, parse_dates, chunksize, dtype, dtype_backend)\u001b[39m\n\u001b[32m 527\u001b[39m \u001b[38;5;28;01mwith\u001b[39;00m pandasSQL_builder(con) \u001b[38;5;28;01mas\u001b[39;00m pandas_sql:\n\u001b[32m--> \u001b[39m\u001b[32m528\u001b[39m \u001b[38;5;28;01mreturn\u001b[39;00m \u001b[43mpandas_sql\u001b[49m\u001b[43m.\u001b[49m\u001b[43mread_query\u001b[49m\u001b[43m(\u001b[49m\n\u001b[32m 529\u001b[39m \u001b[43m \u001b[49m\u001b[43msql\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 530\u001b[39m \u001b[43m \u001b[49m\u001b[43mindex_col\u001b[49m\u001b[43m=\u001b[49m\u001b[43mindex_col\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 531\u001b[39m \u001b[43m \u001b[49m\u001b[43mparams\u001b[49m\u001b[43m=\u001b[49m\u001b[43mparams\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 532\u001b[39m \u001b[43m \u001b[49m\u001b[43mcoerce_float\u001b[49m\u001b[43m=\u001b[49m\u001b[43mcoerce_float\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 533\u001b[39m \u001b[43m \u001b[49m\u001b[43mparse_dates\u001b[49m\u001b[43m=\u001b[49m\u001b[43mparse_dates\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 534\u001b[39m \u001b[43m \u001b[49m\u001b[43mchunksize\u001b[49m\u001b[43m=\u001b[49m\u001b[43mchunksize\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 535\u001b[39m \u001b[43m \u001b[49m\u001b[43mdtype\u001b[49m\u001b[43m=\u001b[49m\u001b[43mdtype\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 536\u001b[39m \u001b[43m \u001b[49m\u001b[43mdtype_backend\u001b[49m\u001b[43m=\u001b[49m\u001b[43mdtype_backend\u001b[49m\u001b[43m,\u001b[49m\n\u001b[32m 537\u001b[39m \u001b[43m \u001b[49m\u001b[43m)\u001b[49m\n",
"\u001b[36mFile \u001b[39m\u001b[32m~/.pyenv/python3.12-venv/lib/python3.12/site-packages/pandas/io/sql.py:2728\u001b[39m, in \u001b[36mSQLiteDatabase.read_query\u001b[39m\u001b[34m(self, sql, index_col, coerce_float, parse_dates, params, chunksize, dtype, dtype_backend)\u001b[39m\n\u001b[32m 2717\u001b[39m \u001b[38;5;28;01mdef\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34mread_query\u001b[39m(\n\u001b[32m 2718\u001b[39m \u001b[38;5;28mself\u001b[39m,\n\u001b[32m 2719\u001b[39m sql,\n\u001b[32m (...)\u001b[39m\u001b[32m 2726\u001b[39m dtype_backend: DtypeBackend | Literal[\u001b[33m\"\u001b[39m\u001b[33mnumpy\u001b[39m\u001b[33m\"\u001b[39m] = \u001b[33m\"\u001b[39m\u001b[33mnumpy\u001b[39m\u001b[33m\"\u001b[39m,\n\u001b[32m 2727\u001b[39m ) -> DataFrame | Iterator[DataFrame]:\n\u001b[32m-> \u001b[39m\u001b[32m2728\u001b[39m cursor = \u001b[38;5;28;43mself\u001b[39;49m\u001b[43m.\u001b[49m\u001b[43mexecute\u001b[49m\u001b[43m(\u001b[49m\u001b[43msql\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mparams\u001b[49m\u001b[43m)\u001b[49m\n\u001b[32m 2729\u001b[39m columns = [col_desc[\u001b[32m0\u001b[39m] \u001b[38;5;28;01mfor\u001b[39;00m col_desc \u001b[38;5;129;01min\u001b[39;00m cursor.description]\n",
"\u001b[36mFile \u001b[39m\u001b[32m~/.pyenv/python3.12-venv/lib/python3.12/site-packages/pandas/io/sql.py:2676\u001b[39m, in \u001b[36mSQLiteDatabase.execute\u001b[39m\u001b[34m(self, sql, params)\u001b[39m\n\u001b[32m 2675\u001b[39m ex = DatabaseError(\u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mExecution failed on sql \u001b[39m\u001b[33m'\u001b[39m\u001b[38;5;132;01m{\u001b[39;00msql\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m'\u001b[39m\u001b[33m: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mexc\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m)\n\u001b[32m-> \u001b[39m\u001b[32m2676\u001b[39m \u001b[38;5;28;01mraise\u001b[39;00m ex \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mexc\u001b[39;00m\n",
"\u001b[31mDatabaseError\u001b[39m: Execution failed on sql 'select tstamp, tstamp_ns as time_ns, substr(instrument_id, 7) as symbol, open, high, low, close, volume, num_trades, vwap from md_1min_bars where exchange_id ='ALPACA' and instrument_id in (\"STOCK-COIN\",\"STOCK-GBTC\",\"STOCK-HOOD\",\"STOCK-MSTR\",\"STOCK-PYPL\")': no such table: md_1min_bars",
"\nThe above exception was the direct cause of the following exception:\n",
"\u001b[31mException\u001b[39m Traceback (most recent call last)",
"\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[5]\u001b[39m\u001b[32m, line 6\u001b[39m\n\u001b[32m 3\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mCurrent working directory: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mos.getcwd()\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m)\n\u001b[32m 4\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mLoading data from: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mdatafile_path\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m)\n\u001b[32m----> \u001b[39m\u001b[32m6\u001b[39m market_data_df = \u001b[43mload_market_data\u001b[49m\u001b[43m(\u001b[49m\u001b[43mdatafile_path\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mconfig\u001b[49m\u001b[43m=\u001b[49m\u001b[43mCONFIG\u001b[49m\u001b[43m)\u001b[49m\n\u001b[32m 8\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mLoaded \u001b[39m\u001b[38;5;132;01m{\u001b[39;00m\u001b[38;5;28mlen\u001b[39m(market_data_df)\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m rows of market data\u001b[39m\u001b[33m\"\u001b[39m)\n\u001b[32m 9\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mSymbols in data: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mmarket_data_df[\u001b[33m'\u001b[39m\u001b[33msymbol\u001b[39m\u001b[33m'\u001b[39m].unique()\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m)\n",
"\u001b[36mFile \u001b[39m\u001b[32m~/devel/pairs_trading/src/notebooks/../tools/data_loader.py:69\u001b[39m, in \u001b[36mload_market_data\u001b[39m\u001b[34m(datafile, config)\u001b[39m\n\u001b[32m 66\u001b[39m query += \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33m where exchange_id =\u001b[39m\u001b[33m'\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mexchange_id\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m'\u001b[39m\u001b[33m\"\u001b[39m\n\u001b[32m 67\u001b[39m query += \u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33m and instrument_id in (\u001b[39m\u001b[38;5;132;01m{\u001b[39;00m\u001b[33m'\u001b[39m\u001b[33m,\u001b[39m\u001b[33m'\u001b[39m.join(instrument_ids)\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m)\u001b[39m\u001b[33m\"\u001b[39m\n\u001b[32m---> \u001b[39m\u001b[32m69\u001b[39m df = \u001b[43mload_sqlite_to_dataframe\u001b[49m\u001b[43m(\u001b[49m\u001b[43mdb_path\u001b[49m\u001b[43m=\u001b[49m\u001b[43mdatafile\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mquery\u001b[49m\u001b[43m=\u001b[49m\u001b[43mquery\u001b[49m\u001b[43m)\u001b[49m\n\u001b[32m 71\u001b[39m \u001b[38;5;66;03m# Trading Hours\u001b[39;00m\n\u001b[32m 72\u001b[39m date_str = df[\u001b[33m\"\u001b[39m\u001b[33mtstamp\u001b[39m\u001b[33m\"\u001b[39m][\u001b[32m0\u001b[39m][\u001b[32m0\u001b[39m:\u001b[32m10\u001b[39m]\n",
"\u001b[36mFile \u001b[39m\u001b[32m~/devel/pairs_trading/src/notebooks/../tools/data_loader.py:18\u001b[39m, in \u001b[36mload_sqlite_to_dataframe\u001b[39m\u001b[34m(db_path, query)\u001b[39m\n\u001b[32m 16\u001b[39m \u001b[38;5;28;01mexcept\u001b[39;00m \u001b[38;5;167;01mException\u001b[39;00m \u001b[38;5;28;01mas\u001b[39;00m excpt:\n\u001b[32m 17\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33mf\u001b[39m\u001b[33m\"\u001b[39m\u001b[33mError: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mexcpt\u001b[38;5;132;01m}\u001b[39;00m\u001b[33m\"\u001b[39m)\n\u001b[32m---> \u001b[39m\u001b[32m18\u001b[39m \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mException\u001b[39;00m() \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mexcpt\u001b[39;00m\n\u001b[32m 19\u001b[39m \u001b[38;5;28;01mfinally\u001b[39;00m:\n\u001b[32m 20\u001b[39m \u001b[38;5;28;01mif\u001b[39;00m \u001b[33m\"\u001b[39m\u001b[33mconn\u001b[39m\u001b[33m\"\u001b[39m \u001b[38;5;129;01min\u001b[39;00m \u001b[38;5;28mlocals\u001b[39m():\n",
"\u001b[31mException\u001b[39m: "
]
}
],
"source": [
"# Load market data\n",
"datafile_path = f\"{CONFIG['data_directory']}/{DATA_FILE}\"\n",
"print(f\"Current working directory: {os.getcwd()}\")\n",
"print(f\"Loading data from: {datafile_path}\")\n",
"\n",
"market_data_df = load_market_data(datafile_path, config=CONFIG)\n",
"\n",
"print(f\"Loaded {len(market_data_df)} rows of market data\")\n",
"print(f\"Symbols in data: {market_data_df['symbol'].unique()}\")\n",
"print(f\"Time range: {market_data_df['tstamp'].min()} to {market_data_df['tstamp'].max()}\")\n",
"\n",
"# Display first few rows\n",
"market_data_df.head()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Create Trading Pair and Analyze"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Create trading pair\n",
"pair = TradingPair(\n",
" market_data=market_data_df,\n",
" symbol_a=SYMBOL_A,\n",
" symbol_b=SYMBOL_B,\n",
" price_column=CONFIG[\"price_column\"]\n",
")\n",
"\n",
"print(f\"Created trading pair: {pair}\")\n",
"print(f\"Market data shape: {pair.market_data_.shape}\")\n",
"print(f\"Column names: {pair.colnames()}\")\n",
"\n",
"# Display first few rows of pair data\n",
"pair.market_data_.head()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Split Data into Training and Testing"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Get training and testing datasets\n",
"training_minutes = CONFIG[\"training_minutes\"]\n",
"pair.get_datasets(training_minutes=training_minutes)\n",
"\n",
"print(f\"Training data: {len(pair.training_df_)} rows\")\n",
"print(f\"Testing data: {len(pair.testing_df_)} rows\")\n",
"print(f\"Training period: {pair.training_df_['tstamp'].iloc[0]} to {pair.training_df_['tstamp'].iloc[-1]}\")\n",
"print(f\"Testing period: {pair.testing_df_['tstamp'].iloc[0]} to {pair.testing_df_['tstamp'].iloc[-1]}\")\n",
"\n",
"# Check for any missing data\n",
"print(f\"Training data null values: {pair.training_df_.isnull().sum().sum()}\")\n",
"print(f\"Testing data null values: {pair.testing_df_.isnull().sum().sum()}\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Visualize Raw Price Data"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Plot raw price data\n",
"fig, axes = plt.subplots(3, 1, figsize=(15, 12))\n",
"\n",
"# Combined price plot\n",
"colname_a, colname_b = pair.colnames()\n",
"all_data = pd.concat([pair.training_df_, pair.testing_df_]).reset_index(drop=True)\n",
"\n",
"# Plot individual prices\n",
"axes[0].plot(all_data['tstamp'], all_data[colname_a], label=f'{SYMBOL_A}', alpha=0.8)\n",
"axes[0].plot(all_data['tstamp'], all_data[colname_b], label=f'{SYMBOL_B}', alpha=0.8)\n",
"axes[0].axvline(x=pair.training_df_['tstamp'].iloc[-1], color='red', linestyle='--', alpha=0.7, label='Train/Test Split')\n",
"axes[0].set_title(f'Price Comparison: {SYMBOL_A} vs {SYMBOL_B}')\n",
"axes[0].set_ylabel('Price')\n",
"axes[0].legend()\n",
"axes[0].grid(True)\n",
"\n",
"# Normalized prices for comparison\n",
"norm_a = all_data[colname_a] / all_data[colname_a].iloc[0]\n",
"norm_b = all_data[colname_b] / all_data[colname_b].iloc[0]\n",
"\n",
"axes[1].plot(all_data['tstamp'], norm_a, label=f'{SYMBOL_A} (normalized)', alpha=0.8)\n",
"axes[1].plot(all_data['tstamp'], norm_b, label=f'{SYMBOL_B} (normalized)', alpha=0.8)\n",
"axes[1].axvline(x=pair.training_df_['tstamp'].iloc[-1], color='red', linestyle='--', alpha=0.7, label='Train/Test Split')\n",
"axes[1].set_title('Normalized Price Comparison')\n",
"axes[1].set_ylabel('Normalized Price')\n",
"axes[1].legend()\n",
"axes[1].grid(True)\n",
"\n",
"# Price ratio\n",
"price_ratio = all_data[colname_a] / all_data[colname_b]\n",
"axes[2].plot(all_data['tstamp'], price_ratio, label=f'{SYMBOL_A}/{SYMBOL_B} Ratio', color='green', alpha=0.8)\n",
"axes[2].axvline(x=pair.training_df_['tstamp'].iloc[-1], color='red', linestyle='--', alpha=0.7, label='Train/Test Split')\n",
"axes[2].set_title('Price Ratio')\n",
"axes[2].set_ylabel('Ratio')\n",
"axes[2].set_xlabel('Time')\n",
"axes[2].legend()\n",
"axes[2].grid(True)\n",
"\n",
"plt.tight_layout()\n",
"plt.show()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Train the Pair and Check Cointegration"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Train the pair and check cointegration\n",
"try:\n",
" is_cointegrated = pair.train_pair()\n",
" print(f\"Pair {pair} cointegration status: {is_cointegrated}\")\n",
"\n",
" if is_cointegrated:\n",
" print(f\"VECM Beta coefficients: {pair.vecm_fit_.beta.flatten()}\")\n",
" print(f\"Training dis-equilibrium mean: {pair.training_mu_:.6f}\")\n",
" print(f\"Training dis-equilibrium std: {pair.training_std_:.6f}\")\n",
"\n",
" # Display VECM summary\n",
" print(\"\\nVECM Model Summary:\")\n",
" print(pair.vecm_fit_.summary())\n",
" else:\n",
" print(\"Pair is not cointegrated. Cannot proceed with strategy.\")\n",
"\n",
"except Exception as e:\n",
" print(f\"Training failed: {str(e)}\")\n",
" is_cointegrated = False"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Visualize Training Period Dis-equilibrium"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"if is_cointegrated:\n",
" # fig, axes = plt.subplots(, 1, figsize=(15, 10))\n",
"\n",
" # # Raw dis-equilibrium\n",
" # axes[0].plot(pair.training_df_['tstamp'], pair.training_df_['dis-equilibrium'],\n",
" # color='blue', alpha=0.8, label='Raw Dis-equilibrium')\n",
" # axes[0].axhline(y=pair.training_mu_, color='red', linestyle='--', alpha=0.7, label='Mean')\n",
" # axes[0].axhline(y=pair.training_mu_ + pair.training_std_, color='orange', linestyle='--', alpha=0.5, label='+1 Std')\n",
" # axes[0].axhline(y=pair.training_mu_ - pair.training_std_, color='orange', linestyle='--', alpha=0.5, label='-1 Std')\n",
" # axes[0].set_title('Training Period: Raw Dis-equilibrium')\n",
" # axes[0].set_ylabel('Dis-equilibrium')\n",
" # axes[0].legend()\n",
" # axes[0].grid(True)\n",
"\n",
" # Scaled dis-equilibrium\n",
" fig, axes = plt.subplots(1, 1, figsize=(15, 5))\n",
" axes.plot(pair.training_df_['tstamp'], pair.training_df_['scaled_dis-equilibrium'],\n",
" color='green', alpha=0.8, label='Scaled Dis-equilibrium')\n",
" axes.axhline(y=0, color='red', linestyle='--', alpha=0.7, label='Mean (0)')\n",
" axes.axhline(y=1, color='orange', linestyle='--', alpha=0.5, label='+1 Std')\n",
" axes.axhline(y=-1, color='orange', linestyle='--', alpha=0.5, label='-1 Std')\n",
" axes.axhline(y=CONFIG['dis-equilibrium_open_trshld'], color='purple',\n",
" linestyle=':', alpha=0.7, label=f\"Open Threshold ({CONFIG['dis-equilibrium_open_trshld']})\")\n",
" axes.axhline(y=CONFIG['dis-equilibrium_close_trshld'], color='brown',\n",
" linestyle=':', alpha=0.7, label=f\"Close Threshold ({CONFIG['dis-equilibrium_close_trshld']})\")\n",
" axes.set_title('Training Period: Scaled Dis-equilibrium')\n",
" axes.set_ylabel('Scaled Dis-equilibrium')\n",
" axes.set_xlabel('Time')\n",
" axes.legend()\n",
" axes.grid(True)\n",
"\n",
" plt.tight_layout()\n",
" plt.show()\n",
"\n",
" # Print statistics\n",
" print(f\"Training dis-equilibrium statistics:\")\n",
" print(f\" Mean: {pair.training_df_['dis-equilibrium'].mean():.6f}\")\n",
" print(f\" Std: {pair.training_df_['dis-equilibrium'].std():.6f}\")\n",
" print(f\" Min: {pair.training_df_['dis-equilibrium'].min():.6f}\")\n",
" print(f\" Max: {pair.training_df_['dis-equilibrium'].max():.6f}\")\n",
"\n",
" print(f\"\\nScaled dis-equilibrium statistics:\")\n",
" print(f\" Mean: {pair.training_df_['scaled_dis-equilibrium'].mean():.6f}\")\n",
" print(f\" Std: {pair.training_df_['scaled_dis-equilibrium'].std():.6f}\")\n",
" print(f\" Min: {pair.training_df_['scaled_dis-equilibrium'].min():.6f}\")\n",
" print(f\" Max: {pair.training_df_['scaled_dis-equilibrium'].max():.6f}\")\n",
"else:\n",
" print(\"The pair is not cointegrated\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Generate Predictions and Run Strategy"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"if is_cointegrated:\n",
" try:\n",
" # Generate predictions\n",
" pair.predict()\n",
" print(f\"Generated predictions for {len(pair.predicted_df_)} rows\")\n",
"\n",
" # Display prediction data structure\n",
" print(f\"Prediction columns: {list(pair.predicted_df_.columns)}\")\n",
" print(f\"Prediction period: {pair.predicted_df_['tstamp'].iloc[0]} to {pair.predicted_df_['tstamp'].iloc[-1]}\")\n",
"\n",
" # Run strategy\n",
" bt_result = BacktestResult(config=CONFIG)\n",
" pair_trades = FIT_METHOD.run_pair(config=CONFIG, pair=pair, bt_result=bt_result)\n",
"\n",
" if pair_trades is not None and len(pair_trades) > 0:\n",
" print(f\"\\nGenerated {len(pair_trades)} trading signals:\")\n",
" print(pair_trades)\n",
" else:\n",
" print(\"\\nNo trading signals generated\")\n",
"\n",
" except Exception as e:\n",
" print(f\"Prediction/Strategy failed: {str(e)}\")\n",
" pair_trades = None"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Visualize Predictions and Dis-equilibrium"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"if is_cointegrated and hasattr(pair, 'predicted_df_'):\n",
" fig, axes = plt.subplots(4, 1, figsize=(16, 16))\n",
"\n",
" # Actual vs Predicted Prices\n",
" colname_a, colname_b = pair.colnames()\n",
"\n",
" axes[0].plot(pair.predicted_df_['tstamp'], pair.predicted_df_[colname_a],\n",
" label=f'{SYMBOL_A} Actual', alpha=0.8)\n",
" axes[0].plot(pair.predicted_df_['tstamp'], pair.predicted_df_[f'{colname_a}_pred'],\n",
" label=f'{SYMBOL_A} Predicted', alpha=0.8, linestyle='--')\n",
" axes[0].set_title('Actual vs Predicted Prices - Symbol A')\n",
" axes[0].set_ylabel('Price')\n",
" axes[0].legend()\n",
" axes[0].grid(True)\n",
"\n",
" axes[1].plot(pair.predicted_df_['tstamp'], pair.predicted_df_[colname_b],\n",
" label=f'{SYMBOL_B} Actual', alpha=0.8)\n",
" axes[1].plot(pair.predicted_df_['tstamp'], pair.predicted_df_[f'{colname_b}_pred'],\n",
" label=f'{SYMBOL_B} Predicted', alpha=0.8, linestyle='--')\n",
" axes[1].set_title('Actual vs Predicted Prices - Symbol B')\n",
" axes[1].set_ylabel('Price')\n",
" axes[1].legend()\n",
" axes[1].grid(True)\n",
"\n",
" # Raw dis-equilibrium\n",
" axes[2].plot(pair.predicted_df_['tstamp'], pair.predicted_df_['disequilibrium'],\n",
" color='blue', alpha=0.8, label='Dis-equilibrium')\n",
" axes[2].axhline(y=pair.training_mu_, color='red', linestyle='--', alpha=0.7, label='Training Mean')\n",
" axes[2].set_title('Testing Period: Raw Dis-equilibrium')\n",
" axes[2].set_ylabel('Dis-equilibrium')\n",
" axes[2].legend()\n",
" axes[2].grid(True)\n",
"\n",
" # Scaled dis-equilibrium with trading signals\n",
" axes[3].plot(pair.predicted_df_['tstamp'], pair.predicted_df_['scaled_disequilibrium'],\n",
" color='green', alpha=0.8, label='Scaled Dis-equilibrium')\n",
"\n",
" # Add threshold lines\n",
" axes[3].axhline(y=CONFIG['dis-equilibrium_open_trshld'], color='purple',\n",
" linestyle=':', alpha=0.7, label=f\"Open Threshold ({CONFIG['dis-equilibrium_open_trshld']})\")\n",
" axes[3].axhline(y=CONFIG['dis-equilibrium_close_trshld'], color='brown',\n",
" linestyle=':', alpha=0.7, label=f\"Close Threshold ({CONFIG['dis-equilibrium_close_trshld']})\")\n",
"\n",
" # Add trading signals if they exist\n",
" if pair_trades is not None and len(pair_trades) > 0:\n",
" for _, trade in pair_trades.iterrows():\n",
" color = 'red' if 'BUY' in trade['action'] else 'blue'\n",
" marker = '^' if 'BUY' in trade['action'] else 'v'\n",
" axes[3].scatter(trade['time'], trade['scaled_disequilibrium'],\n",
" color=color, marker=marker, s=100, alpha=0.8,\n",
" label=f\"{trade['action']} {trade['symbol']}\" if _ < 2 else \"\")\n",
"\n",
" axes[3].set_title('Testing Period: Scaled Dis-equilibrium with Trading Signals')\n",
" axes[3].set_ylabel('Scaled Dis-equilibrium')\n",
" axes[3].set_xlabel('Time')\n",
" axes[3].legend()\n",
" axes[3].grid(True)\n",
"\n",
" plt.tight_layout()\n",
" plt.show()\n",
"\n",
" # Print prediction statistics\n",
" print(f\"\\nTesting dis-equilibrium statistics:\")\n",
" print(f\" Mean: {pair.predicted_df_['disequilibrium'].mean():.6f}\")\n",
" print(f\" Std: {pair.predicted_df_['disequilibrium'].std():.6f}\")\n",
" print(f\" Min: {pair.predicted_df_['disequilibrium'].min():.6f}\")\n",
" print(f\" Max: {pair.predicted_df_['disequilibrium'].max():.6f}\")\n",
"\n",
" print(f\"\\nTesting scaled dis-equilibrium statistics:\")\n",
" print(f\" Mean: {pair.predicted_df_['scaled_disequilibrium'].mean():.6f}\")\n",
" print(f\" Std: {pair.predicted_df_['scaled_disequilibrium'].std():.6f}\")\n",
" print(f\" Min: {pair.predicted_df_['scaled_disequilibrium'].min():.6f}\")\n",
" print(f\" Max: {pair.predicted_df_['scaled_disequilibrium'].max():.6f}\")\n",
"\n",
" # Count threshold crossings\n",
" open_crossings = (pair.predicted_df_['scaled_disequilibrium'] >= CONFIG['dis-equilibrium_open_trshld']).sum()\n",
" close_crossings = (pair.predicted_df_['scaled_disequilibrium'] <= CONFIG['dis-equilibrium_close_trshld']).sum()\n",
" print(f\"\\nThreshold crossings:\")\n",
" print(f\" Open threshold ({CONFIG['dis-equilibrium_open_trshld']}): {open_crossings} times\")\n",
" print(f\" Close threshold ({CONFIG['dis-equilibrium_close_trshld']}): {close_crossings} times\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Summary and Analysis"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"print(\"=\" * 60)\n",
"print(\"PAIRS TRADING ANALYSIS SUMMARY\")\n",
"print(\"=\" * 60)\n",
"\n",
"print(f\"\\nPair: {SYMBOL_A} & {SYMBOL_B}\")\n",
"print(f\"Strategy: {type(FIT_METHOD).__name__}\")\n",
"print(f\"Data file: {DATA_FILE}\")\n",
"print(f\"Training period: {training_minutes} minutes\")\n",
"\n",
"print(f\"\\nCointegration Status: {'✓ COINTEGRATED' if is_cointegrated else '✗ NOT COINTEGRATED'}\")\n",
"\n",
"if is_cointegrated:\n",
" print(f\"\\nVECM Model:\")\n",
" print(f\" Beta coefficients: {pair.vecm_fit_.beta.flatten()}\")\n",
" print(f\" Training mean: {pair.training_mu_:.6f}\")\n",
" print(f\" Training std: {pair.training_std_:.6f}\")\n",
"\n",
" if pair_trades is not None and len(pair_trades) > 0:\n",
" print(f\"\\nTrading Signals: {len(pair_trades)} generated\")\n",
" unique_times = pair_trades['time'].unique()\n",
" print(f\" Unique trade times: {len(unique_times)}\")\n",
"\n",
" # Group by time to see paired trades\n",
" for trade_time in unique_times:\n",
" trades_at_time = pair_trades[pair_trades['time'] == trade_time]\n",
" print(f\"\\n Trade at {trade_time}:\")\n",
" for _, trade in trades_at_time.iterrows():\n",
" print(f\" {trade['action']} {trade['symbol']} @ ${trade['price']:.2f} (dis-eq: {trade['scaled_disequilibrium']:.2f})\")\n",
" else:\n",
" print(f\"\\nTrading Signals: None generated\")\n",
" print(\" Possible reasons:\")\n",
" print(\" - Dis-equilibrium never exceeded open threshold\")\n",
" print(\" - Insufficient testing data\")\n",
" print(\" - Strategy-specific conditions not met\")\n",
"\n",
"else:\n",
" print(\"\\nCannot proceed with trading strategy - pair is not cointegrated\")\n",
" print(\"Consider:\")\n",
" print(\" - Trying different symbol pairs\")\n",
" print(\" - Adjusting training period length\")\n",
" print(\" - Using different data timeframe\")\n",
"\n",
"print(\"\\n\" + \"=\" * 60)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Interactive Analysis (Optional)\n",
"\n",
"You can modify the parameters below and re-run the analysis:"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Interactive parameter adjustment\n",
"print(\"Current parameters:\")\n",
"print(f\" Open threshold: {CONFIG['dis-equilibrium_open_trshld']}\")\n",
"print(f\" Close threshold: {CONFIG['dis-equilibrium_close_trshld']}\")\n",
"print(f\" Training minutes: {CONFIG['training_minutes']}\")\n",
"\n",
"# Uncomment and modify these to experiment:\n",
"# CONFIG['dis-equilibrium_open_trshld'] = 1.5\n",
"# CONFIG['dis-equilibrium_close_trshld'] = 0.3\n",
"# CONFIG['training_minutes'] = 180\n",
"\n",
"print(\"\\nTo re-run with different parameters:\")\n",
"print(\"1. Modify the parameters above\")\n",
"print(\"2. Re-run from the 'Split Data into Training and Testing' cell\")\n",
"print(\"3. Or try different symbol pairs by changing SYMBOL_A and SYMBOL_B\")"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "python3.12-venv",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.9"
}
},
"nbformat": 4,
"nbformat_minor": 4
}
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+107 -132
View File
@@ -3,132 +3,124 @@ import glob
import importlib import importlib
import os import os
from datetime import date, datetime from datetime import date, datetime
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional, Tuple
import pandas as pd import pandas as pd
from research.research_tools import create_pairs
from tools.config import expand_filename, load_config from tools.config import expand_filename, load_config
from tools.data_loader import get_available_instruments_from_db, load_market_data
from pt_trading.results import ( from pt_trading.results import (
BacktestResult, BacktestResult,
create_result_database, create_result_database,
store_config_in_database, store_config_in_database,
store_results_in_database,
) )
from pt_trading.fit_methods import PairsTradingFitMethod from pt_trading.fit_method import PairsTradingFitMethod
from pt_trading.trading_pair import TradingPair from pt_trading.trading_pair import TradingPair
DayT = str
DataFileNameT = str
def resolve_datafiles(config: Dict, cli_datafiles: Optional[str] = None) -> List[str]: def resolve_datafiles(
""" config: Dict, date_pattern: str, instruments: List[Dict[str, str]]
Resolve the list of data files to process. ) -> List[Tuple[DayT, DataFileNameT]]:
CLI datafiles take priority over config datafiles. resolved_files: List[Tuple[DayT, DataFileNameT]] = []
Supports wildcards in config but not in CLI. for inst in instruments:
""" pattern = date_pattern
if cli_datafiles: inst_type = inst["instrument_type"]
# CLI override - comma-separated list, no wildcards data_dir = config["market_data_loading"][inst_type]["data_directory"]
datafiles = [f.strip() for f in cli_datafiles.split(",")]
# Make paths absolute relative to data directory
data_dir = config.get("data_directory", "./data")
resolved_files = []
for df in datafiles:
if not os.path.isabs(df):
df = os.path.join(data_dir, df)
resolved_files.append(df)
return resolved_files
# Use config datafiles with wildcard support
config_datafiles = config.get("datafiles", [])
data_dir = config.get("data_directory", "./data")
resolved_files = []
for pattern in config_datafiles:
if "*" in pattern or "?" in pattern: if "*" in pattern or "?" in pattern:
# Handle wildcards # Handle wildcards
if not os.path.isabs(pattern): if not os.path.isabs(pattern):
pattern = os.path.join(data_dir, pattern) pattern = os.path.join(data_dir, f"{pattern}.mktdata.ohlcv.db")
matched_files = glob.glob(pattern) matched_files = glob.glob(pattern)
resolved_files.extend(matched_files) for matched_file in matched_files:
import re
match = re.search(r"(\d{8})\.mktdata\.ohlcv\.db$", matched_file)
assert match is not None
day = match.group(1)
resolved_files.append((day, matched_file))
else: else:
# Handle explicit file path # Handle explicit file path
if not os.path.isabs(pattern): if not os.path.isabs(pattern):
pattern = os.path.join(data_dir, pattern) pattern = os.path.join(data_dir, f"{pattern}.mktdata.ohlcv.db")
resolved_files.append(pattern) resolved_files.append((date_pattern, pattern))
return sorted(list(set(resolved_files))) # Remove duplicates and sort return sorted(list(set(resolved_files))) # Remove duplicates and sort
def get_instruments(args: argparse.Namespace, config: Dict) -> List[Dict[str, str]]:
instruments = [
{
"symbol": inst.split(":")[0],
"instrument_type": inst.split(":")[1],
"exchange_id": inst.split(":")[2],
"instrument_id_pfx": config["market_data_loading"][inst.split(":")[1]][
"instrument_id_pfx"
],
"db_table_name": config["market_data_loading"][inst.split(":")[1]][
"db_table_name"
],
}
for inst in args.instruments.split(",")
]
return instruments
def run_backtest( def run_backtest(
config: Dict, config: Dict,
datafile: str, datafiles: List[str],
price_column: str,
fit_method: PairsTradingFitMethod, fit_method: PairsTradingFitMethod,
instruments: List[str], instruments: List[Dict[str, str]],
) -> BacktestResult: ) -> BacktestResult:
""" """
Run backtest for all pairs using the specified instruments. Run backtest for all pairs using the specified instruments.
""" """
bt_result: BacktestResult = BacktestResult(config=config) bt_result: BacktestResult = BacktestResult(config=config)
# if len(datafiles) < 2:
# print(f"WARNING: insufficient data files: {datafiles}")
# return bt_result
def _create_pairs(config: Dict, instruments: List[str]) -> List[TradingPair]: if not all([os.path.exists(datafile) for datafile in datafiles]):
nonlocal datafile print(f"WARNING: data file {datafiles} does not exist")
all_indexes = range(len(instruments)) return bt_result
unique_index_pairs = [(i, j) for i in all_indexes for j in all_indexes if i < j]
pairs = []
# Update config to use the specified instruments
config_copy = config.copy()
config_copy["instruments"] = instruments
market_data_df = load_market_data(datafile, config=config_copy)
for a_index, b_index in unique_index_pairs:
pair = TradingPair(
market_data=market_data_df,
symbol_a=instruments[a_index],
symbol_b=instruments[b_index],
price_column=price_column,
)
pairs.append(pair)
return pairs
pairs_trades = [] pairs_trades = []
for pair in _create_pairs(config, instruments):
single_pair_trades = fit_method.run_pair( pairs = create_pairs(
pair=pair, config=config, bt_result=bt_result datafiles=datafiles,
) fit_method=fit_method,
config=config,
instruments=instruments,
)
for pair in pairs:
single_pair_trades = fit_method.run_pair(pair=pair, bt_result=bt_result)
if single_pair_trades is not None and len(single_pair_trades) > 0: if single_pair_trades is not None and len(single_pair_trades) > 0:
pairs_trades.append(single_pair_trades) pairs_trades.append(single_pair_trades)
print(f"pairs_trades:\n{pairs_trades}")
# Check if result_list has any data before concatenating # Check if result_list has any data before concatenating
if len(pairs_trades) == 0: if len(pairs_trades) == 0:
print("No trading signals found for any pairs") print("No trading signals found for any pairs")
return bt_result return bt_result
result = pd.concat(pairs_trades, ignore_index=True) bt_result.collect_single_day_results(pairs_trades)
result["time"] = pd.to_datetime(result["time"])
result = result.set_index("time").sort_index()
bt_result.collect_single_day_results(result)
return bt_result return bt_result
def main() -> None: def main() -> None:
parser = argparse.ArgumentParser(description="Run pairs trading backtest.") parser = argparse.ArgumentParser(description="Run pairs trading backtest.")
parser.add_argument( parser.add_argument(
"--config", type=str, required=True, help="Path to the configuration file." "--config", type=str, required=True, help="Path to the configuration file."
) )
parser.add_argument( parser.add_argument(
"--datafiles", "--date_pattern",
type=str, type=str,
required=False, required=True,
help="Comma-separated list of data files (overrides config). No wildcards supported.", help="Date YYYYMMDD, allows * and ? wildcards",
) )
parser.add_argument( parser.add_argument(
"--instruments", "--instruments",
type=str, type=str,
required=False, required=True,
help="Comma-separated list of instrument symbols (e.g., COIN,GBTC). If not provided, auto-detects from database.", help="Comma-separated list of instrument symbols (e.g., COIN:EQUITY,GBTC:CRYPTO)",
) )
parser.add_argument( parser.add_argument(
"--result_db", "--result_db",
@@ -142,19 +134,13 @@ def main() -> None:
config: Dict = load_config(args.config) config: Dict = load_config(args.config)
# Dynamically instantiate fit method class # Dynamically instantiate fit method class
fit_method_class_name = config.get("fit_method_class", None) fit_method = PairsTradingFitMethod.create(config)
assert fit_method_class_name is not None
module_name, class_name = fit_method_class_name.rsplit(".", 1)
module = importlib.import_module(module_name)
fit_method = getattr(module, class_name)()
# Resolve data files (CLI takes priority over config) # Resolve data files (CLI takes priority over config)
datafiles = resolve_datafiles(config, args.datafiles) instruments = get_instruments(args, config)
datafiles = resolve_datafiles(config, args.date_pattern, instruments)
if not datafiles:
print("No data files found to process.")
return
days = list(set([day for day, _ in datafiles]))
print(f"Found {len(datafiles)} data files to process:") print(f"Found {len(datafiles)} data files to process:")
for df in datafiles: for df in datafiles:
print(f" - {df}") print(f" - {df}")
@@ -166,51 +152,26 @@ def main() -> None:
# Initialize a dictionary to store all trade results # Initialize a dictionary to store all trade results
all_results: Dict[str, Dict[str, Any]] = {} all_results: Dict[str, Dict[str, Any]] = {}
is_config_stored = False
# Store configuration in database for reference
if args.result_db.upper() != "NONE":
# Get list of all instruments for storage
all_instruments = []
for datafile in datafiles:
if args.instruments:
file_instruments = [
inst.strip() for inst in args.instruments.split(",")
]
else:
file_instruments = get_available_instruments_from_db(datafile, config)
all_instruments.extend(file_instruments)
# Remove duplicates while preserving order
unique_instruments = list(dict.fromkeys(all_instruments))
store_config_in_database(
db_path=args.result_db,
config_file_path=args.config,
config=config,
fit_method_class=fit_method_class_name,
datafiles=datafiles,
instruments=unique_instruments,
)
# Process each data file # Process each data file
price_column = config["price_column"]
for datafile in datafiles: for day in sorted(days):
print(f"\n====== Processing {os.path.basename(datafile)} ======") md_datafiles = [datafile for md_day, datafile in datafiles if md_day == day]
if not all([os.path.exists(datafile) for datafile in md_datafiles]):
# Determine instruments to use print(f"WARNING: insufficient data files: {md_datafiles}")
if args.instruments:
# Use CLI-specified instruments
instruments = [inst.strip() for inst in args.instruments.split(",")]
print(f"Using CLI-specified instruments: {instruments}")
else:
# Auto-detect instruments from database
instruments = get_available_instruments_from_db(datafile, config)
print(f"Auto-detected instruments: {instruments}")
if not instruments:
print(f"No instruments found for {datafile}, skipping...")
continue continue
print(f"\n====== Processing {day} ======")
if not is_config_stored:
store_config_in_database(
db_path=args.result_db,
config_file_path=args.config,
config=config,
fit_method_class=config["fit_method_class"],
datafiles=datafiles,
instruments=instruments,
)
is_config_stored = True
# Process data for this file # Process data for this file
try: try:
@@ -218,24 +179,38 @@ def main() -> None:
bt_results = run_backtest( bt_results = run_backtest(
config=config, config=config,
datafile=datafile, datafiles=md_datafiles,
price_column=price_column,
fit_method=fit_method, fit_method=fit_method,
instruments=instruments, instruments=instruments,
) )
# Store results with file name as key if bt_results.trades is None or len(bt_results.trades) == 0:
filename = os.path.basename(datafile) print(f"No trades found for {day}")
all_results[filename] = {"trades": bt_results.trades.copy()} continue
# Store results with day name as key
filename = os.path.basename(day)
all_results[filename] = {
"trades": bt_results.trades.copy(),
"outstanding_positions": bt_results.outstanding_positions.copy(),
}
# Store results in database # Store results in database
if args.result_db.upper() != "NONE": if args.result_db.upper() != "NONE":
store_results_in_database(args.result_db, datafile, bt_results) bt_results.calculate_returns(
{
filename: {
"trades": bt_results.trades.copy(),
"outstanding_positions": bt_results.outstanding_positions.copy(),
}
}
)
bt_results.store_results_in_database(db_path=args.result_db, day=day)
print(f"Successfully processed {filename}") print(f"Successfully processed {filename}")
except Exception as err: except Exception as err:
print(f"Error processing {datafile}: {str(err)}") print(f"Error processing {day}: {str(err)}")
import traceback import traceback
traceback.print_exc() traceback.print_exc()
+93
View File
@@ -0,0 +1,93 @@
import glob
import os
from typing import Dict, List, Optional
import pandas as pd
from pt_trading.fit_method import PairsTradingFitMethod
def resolve_datafiles(config: Dict, cli_datafiles: Optional[str] = None) -> List[str]:
"""
Resolve the list of data files to process.
CLI datafiles take priority over config datafiles.
Supports wildcards in config but not in CLI.
"""
if cli_datafiles:
# CLI override - comma-separated list, no wildcards
datafiles = [f.strip() for f in cli_datafiles.split(",")]
# Make paths absolute relative to data directory
data_dir = config.get("data_directory", "./data")
resolved_files = []
for df in datafiles:
if not os.path.isabs(df):
df = os.path.join(data_dir, df)
resolved_files.append(df)
return resolved_files
# Use config datafiles with wildcard support
config_datafiles = config.get("datafiles", [])
data_dir = config.get("data_directory", "./data")
resolved_files = []
for pattern in config_datafiles:
if "*" in pattern or "?" in pattern:
# Handle wildcards
if not os.path.isabs(pattern):
pattern = os.path.join(data_dir, pattern)
matched_files = glob.glob(pattern)
resolved_files.extend(matched_files)
else:
# Handle explicit file path
if not os.path.isabs(pattern):
pattern = os.path.join(data_dir, pattern)
resolved_files.append(pattern)
return sorted(list(set(resolved_files))) # Remove duplicates and sort
def create_pairs(
datafiles: List[str],
fit_method: PairsTradingFitMethod,
config: Dict,
instruments: List[Dict[str, str]],
) -> List:
from pt_trading.trading_pair import TradingPair
from tools.data_loader import load_market_data
all_indexes = range(len(instruments))
unique_index_pairs = [(i, j) for i in all_indexes for j in all_indexes if i < j]
pairs = []
# Update config to use the specified instruments
config_copy = config.copy()
config_copy["instruments"] = instruments
market_data_df = pd.DataFrame()
extra_minutes = 0
if "execution_price" in config_copy:
extra_minutes = config_copy["execution_price"]["shift"]
for datafile in datafiles:
md_df = load_market_data(
datafile = datafile,
instruments = instruments,
db_table_name = config_copy["market_data_loading"][instruments[0]["instrument_type"]]["db_table_name"],
trading_hours=config_copy["trading_hours"],
extra_minutes=extra_minutes,
)
market_data_df = pd.concat([market_data_df, md_df])
if len(set(market_data_df["symbol"])) != 2: # both symbols must be present for a pair
print(f"WARNING: insufficient data in files: {datafiles}")
return []
for a_index, b_index in unique_index_pairs:
symbol_a=instruments[a_index]["symbol"]
symbol_b=instruments[b_index]["symbol"]
pair = fit_method.create_trading_pair(
config=config_copy,
market_data=market_data_df,
symbol_a=symbol_a,
symbol_b=symbol_b,
)
pairs.append(pair)
return pairs
+6 -1
View File
@@ -16,7 +16,12 @@ cd $(realpath $(dirname $0))/..
mkdir -p ./data/crypto mkdir -p ./data/crypto
pushd ./data/crypto pushd ./data/crypto
Cmd="rsync -ahvv cvtt@hs01.cvtt.vpn:/works/cvtt/md_archive/crypto/sim/*.gz ./" Files=$1
if [ -z "$Files" ]; then
Files="*.gz"
fi
Cmd="rsync -ahvv cvtt@hs01.cvtt.vpn:/works/cvtt/md_archive/crypto/sim/${Files} ./"
echo $Cmd echo $Cmd
eval $Cmd eval $Cmd
# ------------------------------------- # -------------------------------------
+6 -2
View File
@@ -26,8 +26,12 @@ for srcfname in $(ls *.db.gz); do
tgtfile=${dt}.mktdata.ohlcv.db tgtfile=${dt}.mktdata.ohlcv.db
echo "${srcfname} -> ${tgtfile}" echo "${srcfname} -> ${tgtfile}"
gunzip -c $srcfname > temp.db Cmd="gunzip -c $srcfname > temp.db && rm $srcfname"
rm -f ${tgtfile} && sqlite3 temp.db ".dump md_1min_bars" | sqlite3 ${tgtfile} && rm ${srcfname} echo ${Cmd}
eval ${Cmd}
Cmd="rm -f ${tgtfile} && sqlite3 temp.db '.dump md_1min_bars' | sqlite3 ${tgtfile}"
echo ${Cmd}
eval ${Cmd}
done done
rm temp.db rm temp.db
popd popd
+9 -8
View File
@@ -20,12 +20,9 @@ from pt_trading.fit_methods import PairsTradingFitMethod
from pt_trading.trading_pair import TradingPair from pt_trading.trading_pair import TradingPair
def run_strategy( def run_strategy(
config: Dict, config: Dict,
datafile: str, datafile: str,
price_column: str,
fit_method: PairsTradingFitMethod, fit_method: PairsTradingFitMethod,
instruments: List[str], instruments: List[str],
) -> BacktestResult: ) -> BacktestResult:
@@ -44,14 +41,20 @@ def run_strategy(
config_copy = config.copy() config_copy = config.copy()
config_copy["instruments"] = instruments config_copy["instruments"] = instruments
market_data_df = load_market_data(datafile, config=config_copy) market_data_df = load_market_data(
datafile=datafile,
exchange_id=config_copy["exchange_id"],
instruments=config_copy["instruments"],
instrument_id_pfx=config_copy["instrument_id_pfx"],
db_table_name=config_copy["db_table_name"],
trading_hours=config_copy["trading_hours"],
)
for a_index, b_index in unique_index_pairs: for a_index, b_index in unique_index_pairs:
pair = TradingPair( pair = fit_method.create_trading_pair(
market_data=market_data_df, market_data=market_data_df,
symbol_a=instruments[a_index], symbol_a=instruments[a_index],
symbol_b=instruments[b_index], symbol_b=instruments[b_index],
price_column=price_column,
) )
pairs.append(pair) pairs.append(pair)
return pairs return pairs
@@ -156,7 +159,6 @@ def main() -> None:
) )
# Process each data file # Process each data file
price_column = config["price_column"]
for datafile in datafiles: for datafile in datafiles:
print(f"\n====== Processing {os.path.basename(datafile)} ======") print(f"\n====== Processing {os.path.basename(datafile)} ======")
@@ -182,7 +184,6 @@ def main() -> None:
bt_results = run_strategy( bt_results = run_strategy(
config=config, config=config,
datafile=datafile, datafile=datafile,
price_column=price_column,
fit_method=fit_method, fit_method=fit_method,
instruments=instruments, instruments=instruments,
) )