Compare commits
62 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 809f46fe36 | |||
| 413abafe0f | |||
| 5d46c1e32c | |||
| 889f7ba1c3 | |||
| 1515b2d077 | |||
| b4ae3e715d | |||
| 6f845d32c6 | |||
| a04e8878fb | |||
| 71822c64b0 | |||
| c2f701e3a2 | |||
| 21a473a4c2 | |||
| 98a15d301a | |||
| bcf4447cb6 | |||
| 1af35000ab | |||
| 2c08b6f1a9 | |||
| 24f1f82d1f | |||
| af0a6f62a9 | |||
| a7b4777f76 | |||
| e30b0df4db | |||
| 577fb5c109 | |||
| e0138907be | |||
| b7292c11f3 | |||
| aac8b9dc50 | |||
| 9bb36dddd7 | |||
| 31eb9f800c | |||
| 0e83142d0a | |||
| b87b40a6ed | |||
| 28386cdf12 | |||
| fb3dc68a1d | |||
| c776c95d69 | |||
| ca9fff8d88 | |||
| 705330a9f7 | |||
| 2272a31765 | |||
| facf7fb0c6 | |||
| 9c34d935bd | |||
| 20f150a6b7 | |||
| d46bcb64d6 | |||
| 26659ede12 | |||
| e9995312a0 | |||
| a46c8a7576 | |||
| fe2ebbb27f | |||
| ddd9f4adb9 | |||
| 4bc947cf07 | |||
| 51944b3a2f | |||
| bff1c54b48 | |||
| 9c91f37bcc | |||
| 76547e1176 | |||
| 80cf1b60ef | |||
| 94ffb32f50 | |||
| 967c01c367 | |||
| 747ca05b16 | |||
| 30ae95a808 | |||
| bcba183768 | |||
| cc0072dcc8 | |||
| 35a1cd748e | |||
| 3b003c7811 | |||
| b24285802a | |||
| 48f18f7b4f | |||
| 85c9d2ab93 | |||
| 46072e03a2 | |||
| 352f7df269 | |||
| 191feb341d |
+1
-2
@@ -1,11 +1,10 @@
|
||||
# SpecStory explanation file
|
||||
__pycache__/
|
||||
__OLD__/
|
||||
.specstory/
|
||||
.history/
|
||||
.cursorindexingignore
|
||||
data
|
||||
.vscode/
|
||||
####.vscode/
|
||||
cvttpy
|
||||
# SpecStory explanation file
|
||||
.specstory/.what-is-this.md
|
||||
|
||||
@@ -11,6 +11,7 @@ The enhanced `pt_backtest.py` script now supports multi-day and multi-instrument
|
||||
- Support for wildcard patterns in configuration files
|
||||
- CLI override for data file specification
|
||||
|
||||
|
||||
### 2. Dynamic Instrument Selection
|
||||
- Auto-detection of instruments from database
|
||||
- CLI override for instrument specification
|
||||
|
||||
@@ -38,15 +38,12 @@ CONFIG = EQT_CONFIG # For equity data
|
||||
```
|
||||
|
||||
Each configuration dictionary specifies:
|
||||
- `security_type`: "CRYPTO" or "EQUITY".
|
||||
- `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.
|
||||
- `db_table_name`: The name of the table within the SQLite database.
|
||||
- `instruments`: A list of symbols to consider for forming trading pairs.
|
||||
- `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").
|
||||
- `min_required_points`: Minimum data points needed for statistical calculations.
|
||||
- `zero_threshold`: A small value to handle potential division by zero.
|
||||
- `stat_model_price`: The column in the data to be used as the price (e.g., "close").
|
||||
- `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.
|
||||
- `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).
|
||||
|
||||
@@ -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,
|
||||
"strategy_class": "strategies.StaticFitStrategy"
|
||||
}
|
||||
@@ -2,7 +2,7 @@
|
||||
"security_type": "EQUITY",
|
||||
"data_directory": "./data/equity",
|
||||
"datafiles": [
|
||||
"202506*.mktdata.ohlcv.db",
|
||||
"20250618.mktdata.ohlcv.db",
|
||||
],
|
||||
"db_table_name": "md_1min_bars",
|
||||
"exchange_id": "ALPACA",
|
||||
@@ -19,8 +19,9 @@
|
||||
"dis-equilibrium_close_trshld": 1.0,
|
||||
"training_minutes": 120,
|
||||
"funding_per_pair": 2000.0,
|
||||
"strategy_class": "strategies.StaticFitStrategy"
|
||||
# "strategy_class": "strategies.SlidingFitStrategy"
|
||||
"exclude_instruments": ["CAN"]
|
||||
# "fit_method_class": "pt_trading.sliding_fit.SlidingFit",
|
||||
"fit_method_class": "pt_trading.static_fit.StaticFit",
|
||||
"exclude_instruments": ["CAN"],
|
||||
"close_outstanding_positions": false
|
||||
|
||||
}
|
||||
@@ -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",
|
||||
@@ -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",
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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())
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,188 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import argparse
|
||||
from ast import Sub
|
||||
import asyncio
|
||||
from functools import partial
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Coroutine, Dict, List, Optional
|
||||
|
||||
from numpy.strings import str_len
|
||||
import websockets
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
MessageTypeT = str
|
||||
SubscriptionIdT = str
|
||||
MessageT = Dict
|
||||
UrlT = str
|
||||
CallbackT = Callable[[MessageTypeT, SubscriptionIdT, MessageT], Coroutine[None, str, None]]
|
||||
|
||||
@dataclass
|
||||
class CvttPricesSubscription:
|
||||
id_: str
|
||||
exchange_config_name_: str
|
||||
instrument_id_: str
|
||||
interval_sec_: int
|
||||
history_depth_sec_: int
|
||||
is_subscribed_: bool
|
||||
is_historical_: bool
|
||||
callback_: CallbackT
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
exchange_config_name: str,
|
||||
instrument_id: str,
|
||||
interval_sec: int,
|
||||
history_depth_sec: int,
|
||||
callback: CallbackT,
|
||||
):
|
||||
self.exchange_config_name_ = exchange_config_name
|
||||
self.instrument_id_ = instrument_id
|
||||
self.interval_sec_ = interval_sec
|
||||
self.history_depth_sec_ = history_depth_sec
|
||||
self.callback_ = callback
|
||||
self.id_ = str(uuid.uuid4())
|
||||
self.is_subscribed_ = False
|
||||
self.is_historical_ = history_depth_sec > 0
|
||||
|
||||
|
||||
class CvttPricerWebSockClient:
|
||||
# Class members with type hints
|
||||
ws_url_: UrlT
|
||||
websocket_: Optional[ClientConnection]
|
||||
subscriptions_: Dict[SubscriptionIdT, CvttPricesSubscription]
|
||||
is_connected_: bool
|
||||
logger_: logging.Logger
|
||||
|
||||
def __init__(self, url: str):
|
||||
self.ws_url_ = url
|
||||
self.websocket_ = None
|
||||
self.is_connected_ = False
|
||||
self.subscriptions_ = {}
|
||||
self.logger_ = logging.getLogger(__name__)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
|
||||
async def subscribe(
|
||||
self, subscription: CvttPricesSubscription
|
||||
) -> str: # returns subscription id
|
||||
|
||||
if not self.is_connected_:
|
||||
try:
|
||||
self.logger_.info(f"Connecting to {self.ws_url_}")
|
||||
self.websocket_ = await websockets.connect(self.ws_url_)
|
||||
self.is_connected_ = True
|
||||
except Exception as e:
|
||||
self.logger_.error(f"Unable to connect to {self.ws_url_}: {str(e)}")
|
||||
raise e
|
||||
|
||||
subscr_msg = {
|
||||
"type": "subscr",
|
||||
"id": subscription.id_,
|
||||
"subscr_type": "MD_AGGREGATE",
|
||||
"exchange_config_name": subscription.exchange_config_name_,
|
||||
"instrument_id": subscription.instrument_id_,
|
||||
"interval_sec": subscription.interval_sec_,
|
||||
}
|
||||
if subscription.is_historical_:
|
||||
subscr_msg["history_depth_sec"] = subscription.history_depth_sec_
|
||||
|
||||
assert self.websocket_ is not None
|
||||
await self.websocket_.send(json.dumps(subscr_msg))
|
||||
|
||||
response = await self.websocket_.recv()
|
||||
response_data = json.loads(response)
|
||||
if not await self.handle_subscription_response(subscription, response_data):
|
||||
await self.websocket_.close()
|
||||
self.is_connected_ = False
|
||||
raise Exception(f"Subscription failed: {str(response)}")
|
||||
|
||||
self.subscriptions_[subscription.id_] = subscription
|
||||
return subscription.id_
|
||||
|
||||
async def handle_subscription_response(
|
||||
self, subscription: CvttPricesSubscription, response: dict
|
||||
) -> bool:
|
||||
if response.get("type") != "subscr" or response.get("id") != subscription.id_:
|
||||
return False
|
||||
|
||||
if response.get("status") == "success":
|
||||
self.logger_.info(f"Subscription successful: {json.dumps(response)}")
|
||||
return True
|
||||
elif response.get("status") == "error":
|
||||
self.logger_.error(f"Subscription failed: {response.get('reason')}")
|
||||
return False
|
||||
return False
|
||||
|
||||
async def run(self) -> None:
|
||||
assert self.websocket_
|
||||
try:
|
||||
while self.is_connected_:
|
||||
try:
|
||||
message = await self.websocket_.recv()
|
||||
message_str = (
|
||||
message.decode("utf-8")
|
||||
if isinstance(message, bytes)
|
||||
else message
|
||||
)
|
||||
await self.process_message(json.loads(message_str))
|
||||
except websockets.ConnectionClosed:
|
||||
self.logger_.warning("Connection closed")
|
||||
self.is_connected_ = False
|
||||
break
|
||||
except Exception as e:
|
||||
self.logger_.error(f"Error occurred: {str(e)}")
|
||||
self.is_connected_ = False
|
||||
await asyncio.sleep(5) # Wait before reconnecting
|
||||
|
||||
async def process_message(self, message: Dict) -> None:
|
||||
message_type = message.get("type")
|
||||
if message_type in ["md_aggregate", "historical_md_aggregate"]:
|
||||
subscription_id = message.get("subscr_id")
|
||||
if subscription_id not in self.subscriptions_:
|
||||
self.logger_.warning(f"Unknown subscription id: {subscription_id}")
|
||||
return
|
||||
|
||||
subscription = self.subscriptions_[subscription_id]
|
||||
await subscription.callback_(message_type, subscription_id, message)
|
||||
else:
|
||||
self.logger_.warning(f"Unknown message type: {message.get('type')}")
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
async def on_message(message_type: MessageTypeT, subscr_id: SubscriptionIdT, message: Dict, instrument_id: str) -> None:
|
||||
print(f"{message_type=} {subscr_id=} {instrument_id}")
|
||||
if message_type == "md_aggregate":
|
||||
aggr = message.get("md_aggregate", [])
|
||||
print(f"[{aggr['tstmp'][:19]}] *** RLTM *** {message}")
|
||||
elif message_type == "historical_md_aggregate":
|
||||
for aggr in message.get("historical_data", []):
|
||||
print(f"[{aggr['tstmp'][:19]}] *** HIST *** {aggr}")
|
||||
else:
|
||||
print(f"Unknown message type: {message_type}")
|
||||
|
||||
pricer_client = CvttPricerWebSockClient(
|
||||
"ws://localhost:12346/ws"
|
||||
)
|
||||
await pricer_client.subscribe(CvttPricesSubscription(
|
||||
exchange_config_name="COINBASE_AT",
|
||||
instrument_id="PAIR-BTC-USD",
|
||||
interval_sec=60,
|
||||
history_depth_sec=60*60*24,
|
||||
callback=partial(on_message, instrument_id="PAIR-BTC-USD")
|
||||
))
|
||||
await pricer_client.subscribe(CvttPricesSubscription(
|
||||
exchange_config_name="COINBASE_AT",
|
||||
instrument_id="PAIR-ETH-USD",
|
||||
interval_sec=60,
|
||||
history_depth_sec=60*60*24,
|
||||
callback=partial(on_message, instrument_id="PAIR-ETH-USD")
|
||||
))
|
||||
|
||||
await pricer_client.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -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: ...
|
||||
|
||||
@@ -0,0 +1,743 @@
|
||||
import os
|
||||
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+
|
||||
# From: https://docs.python.org/3/library/sqlite3.html#sqlite3-adapter-converter-recipes
|
||||
def adapt_date_iso(val: date) -> str:
|
||||
"""Adapt datetime.date to ISO 8601 date."""
|
||||
return val.isoformat()
|
||||
|
||||
|
||||
def adapt_datetime_iso(val: datetime) -> str:
|
||||
"""Adapt datetime.datetime to timezone-naive ISO 8601 date."""
|
||||
return val.isoformat()
|
||||
|
||||
|
||||
def convert_date(val: bytes) -> date:
|
||||
"""Convert ISO 8601 date to datetime.date object."""
|
||||
return datetime.fromisoformat(val.decode()).date()
|
||||
|
||||
|
||||
def convert_datetime(val: bytes) -> datetime:
|
||||
"""Convert ISO 8601 datetime to datetime.datetime object."""
|
||||
return datetime.fromisoformat(val.decode())
|
||||
|
||||
|
||||
# Register the adapters and converters
|
||||
sqlite3.register_adapter(date, adapt_date_iso)
|
||||
sqlite3.register_adapter(datetime, adapt_datetime_iso)
|
||||
sqlite3.register_converter("date", convert_date)
|
||||
sqlite3.register_converter("datetime", convert_datetime)
|
||||
|
||||
|
||||
def create_result_database(db_path: str) -> None:
|
||||
"""
|
||||
Create the SQLite database and required tables if they don't exist.
|
||||
"""
|
||||
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)
|
||||
cursor = conn.cursor()
|
||||
|
||||
# Create the pt_bt_results table for completed trades
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS pt_bt_results (
|
||||
date DATE,
|
||||
pair TEXT,
|
||||
symbol TEXT,
|
||||
open_time DATETIME,
|
||||
open_side TEXT,
|
||||
open_price REAL,
|
||||
open_quantity INTEGER,
|
||||
open_disequilibrium REAL,
|
||||
close_time DATETIME,
|
||||
close_side TEXT,
|
||||
close_price REAL,
|
||||
close_quantity INTEGER,
|
||||
close_disequilibrium REAL,
|
||||
symbol_return REAL,
|
||||
pair_return REAL,
|
||||
close_condition TEXT
|
||||
)
|
||||
"""
|
||||
)
|
||||
cursor.execute("DELETE FROM pt_bt_results;")
|
||||
|
||||
# Create the outstanding_positions table for open positions
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS outstanding_positions (
|
||||
date DATE,
|
||||
pair TEXT,
|
||||
symbol TEXT,
|
||||
position_quantity REAL,
|
||||
last_price REAL,
|
||||
unrealized_return REAL,
|
||||
open_price REAL,
|
||||
open_side TEXT
|
||||
)
|
||||
"""
|
||||
)
|
||||
cursor.execute("DELETE FROM outstanding_positions;")
|
||||
|
||||
# Create the config table for storing configuration JSON for reference
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS config (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
run_timestamp DATETIME,
|
||||
config_file_path TEXT,
|
||||
config_json TEXT,
|
||||
fit_method_class TEXT,
|
||||
datafiles TEXT,
|
||||
instruments TEXT
|
||||
)
|
||||
"""
|
||||
)
|
||||
cursor.execute("DELETE FROM config;")
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error creating result database: {str(e)}")
|
||||
raise
|
||||
|
||||
|
||||
def store_config_in_database(
|
||||
db_path: str,
|
||||
config_file_path: str,
|
||||
config: Dict,
|
||||
fit_method_class: str,
|
||||
datafiles: List[Tuple[str, str]],
|
||||
instruments: List[Dict[str, str]],
|
||||
) -> None:
|
||||
"""
|
||||
Store configuration information in the database for reference.
|
||||
"""
|
||||
import json
|
||||
|
||||
if db_path.upper() == "NONE":
|
||||
return
|
||||
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.cursor()
|
||||
|
||||
# Convert config to JSON string
|
||||
config_json = json.dumps(config, indent=2, default=str)
|
||||
|
||||
# Convert lists to comma-separated strings for storage
|
||||
datafiles_str = ", ".join([f"{datafile}" for _, datafile in datafiles])
|
||||
instruments_str = ", ".join(
|
||||
[
|
||||
f"{inst['symbol']}:{inst['instrument_type']}:{inst['exchange_id']}"
|
||||
for inst in instruments
|
||||
]
|
||||
)
|
||||
|
||||
# Insert configuration record
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO config (
|
||||
run_timestamp, config_file_path, config_json, fit_method_class, datafiles, instruments
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
datetime.now(),
|
||||
config_file_path,
|
||||
config_json,
|
||||
fit_method_class,
|
||||
datafiles_str,
|
||||
instruments_str,
|
||||
),
|
||||
)
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
print(f"Configuration stored in database")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error storing configuration in database: {str(e)}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
|
||||
def convert_timestamp(timestamp: Any) -> Optional[datetime]:
|
||||
"""Convert pandas Timestamp to Python datetime object for SQLite compatibility."""
|
||||
if timestamp is None:
|
||||
return None
|
||||
if isinstance(timestamp, pd.Timestamp):
|
||||
return timestamp.to_pydatetime()
|
||||
elif isinstance(timestamp, datetime):
|
||||
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)}")
|
||||
|
||||
|
||||
|
||||
class BacktestResult:
|
||||
"""
|
||||
Class to handle backtest results, trades tracking, PnL calculations, and reporting.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
self.config = config
|
||||
self.trades: Dict[str, Dict[str, Any]] = {}
|
||||
self.total_realized_pnl = 0.0
|
||||
self.outstanding_positions: List[Dict[str, Any]] = []
|
||||
self.pairs_trades_: Dict[str, List[Dict[str, Any]]] = {}
|
||||
|
||||
def add_trade(
|
||||
self,
|
||||
pair_nm: str,
|
||||
symbol: str,
|
||||
side: str,
|
||||
action: str,
|
||||
price: Any,
|
||||
disequilibrium: Optional[float] = None,
|
||||
scaled_disequilibrium: Optional[float] = None,
|
||||
timestamp: Optional[datetime] = None,
|
||||
status: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Add a trade to the results tracking."""
|
||||
pair_nm = str(pair_nm)
|
||||
|
||||
if pair_nm not in self.trades:
|
||||
self.trades[pair_nm] = {symbol: []}
|
||||
if symbol not in self.trades[pair_nm]:
|
||||
self.trades[pair_nm][symbol] = []
|
||||
self.trades[pair_nm][symbol].append(
|
||||
{
|
||||
"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]) -> None:
|
||||
"""Add an outstanding position to tracking."""
|
||||
self.outstanding_positions.append(position)
|
||||
|
||||
def add_realized_pnl(self, realized_pnl: float) -> None:
|
||||
"""Add realized PnL to the total."""
|
||||
self.total_realized_pnl += realized_pnl
|
||||
|
||||
def get_total_realized_pnl(self) -> float:
|
||||
"""Get total realized PnL."""
|
||||
return self.total_realized_pnl
|
||||
|
||||
def get_outstanding_positions(self) -> List[Dict[str, Any]]:
|
||||
"""Get all outstanding positions."""
|
||||
return self.outstanding_positions
|
||||
|
||||
def get_trades(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""Get all trades."""
|
||||
return self.trades
|
||||
|
||||
def clear_trades(self) -> None:
|
||||
"""Clear all trades (used when processing new files)."""
|
||||
self.trades.clear()
|
||||
|
||||
def collect_single_day_results(self, pairs_trades: List[pd.DataFrame]) -> None:
|
||||
"""Collect and process single day trading results."""
|
||||
result = pd.concat(pairs_trades, ignore_index=True)
|
||||
result["time"] = pd.to_datetime(result["time"])
|
||||
result = result.set_index("time").sort_index()
|
||||
|
||||
print("\n -------------- Suggested Trades ")
|
||||
print(result)
|
||||
|
||||
for row in result.itertuples():
|
||||
side = row.side
|
||||
action = row.action
|
||||
symbol = row.symbol
|
||||
price = row.price
|
||||
disequilibrium = getattr(row, "disequilibrium", None)
|
||||
scaled_disequilibrium = getattr(row, "scaled_disequilibrium", None)
|
||||
if hasattr(row, "time"):
|
||||
timestamp = getattr(row, "time")
|
||||
else:
|
||||
timestamp = convert_timestamp(row.Index)
|
||||
status = row.status
|
||||
self.add_trade(
|
||||
pair_nm=str(row.pair),
|
||||
symbol=str(symbol),
|
||||
side=str(side),
|
||||
action=str(action),
|
||||
price=float(str(price)),
|
||||
disequilibrium=disequilibrium,
|
||||
scaled_disequilibrium=scaled_disequilibrium,
|
||||
timestamp=timestamp,
|
||||
status=str(status) if status is not None else "?",
|
||||
)
|
||||
|
||||
def print_single_day_results(self) -> None:
|
||||
"""Print single day results summary."""
|
||||
for pair, symbols in self.trades.items():
|
||||
print(f"\n--- {pair} ---")
|
||||
for symbol, trades in symbols.items():
|
||||
for trade_data in trades:
|
||||
if len(trade_data) >= 2:
|
||||
side, price = trade_data[:2]
|
||||
print(f"{symbol} {side} at ${price}")
|
||||
|
||||
def print_results_summary(self, all_results: Dict[str, Dict[str, Any]]) -> None:
|
||||
"""Print summary of all processed files."""
|
||||
print("\n====== Summary of All Processed Files ======")
|
||||
for filename, data in all_results.items():
|
||||
trade_count = sum(
|
||||
len(trades)
|
||||
for symbol_trades in data["trades"].values()
|
||||
for trades in symbol_trades.values()
|
||||
)
|
||||
print(f"{filename}: {trade_count} trades")
|
||||
|
||||
def calculate_returns(self, all_results: Dict[str, Dict[str, Any]]) -> None:
|
||||
"""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 ======")
|
||||
|
||||
trades = []
|
||||
for filename, data in all_results.items():
|
||||
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} ---")
|
||||
|
||||
self.outstanding_positions = data["outstanding_positions"]
|
||||
|
||||
day_return = 0.0
|
||||
for idx in range(0, len(trades), 4):
|
||||
symbol_a = trades[idx]["symbol"]
|
||||
trade_a_1 = trades[idx]
|
||||
trade_a_2 = trades[idx + 2]
|
||||
|
||||
symbol_b = trades[idx + 1]["symbol"]
|
||||
trade_b_1 = trades[idx + 1]
|
||||
trade_b_2 = trades[idx + 3]
|
||||
|
||||
symbol_return = 0
|
||||
assert (
|
||||
trade_a_1["timestamp"] < trade_a_2["timestamp"]
|
||||
), f"Trade 1: {trade_a_1['timestamp']} is not less than Trade 2: {trade_a_2['timestamp']}"
|
||||
assert (
|
||||
trade_a_1["action"] == "OPEN" and trade_a_2["action"] == "CLOSE"
|
||||
), f"Trade 1: {trade_a_1['action']} and Trade 2: {trade_a_2['action']} are the same"
|
||||
|
||||
# Calculate return based on action combination
|
||||
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}'
|
||||
)
|
||||
|
||||
print(
|
||||
f' {trd["open_time"].time()}-{trd["close_time"].time()} {trd["symbol"]}: '
|
||||
f' {trd["open_side"]} @ ${trd["open_price"]:.2f},'
|
||||
f' {trd["close_side"]} @ ${trd["close_price"]:.2f},'
|
||||
f' Return: {trd["symbol_return"]:.2f}%{disequil_info}'
|
||||
)
|
||||
pair_return += trd["symbol_return"]
|
||||
|
||||
print(f" Pair Total Return: {pair_return:.2f}%")
|
||||
day_return += pair_return
|
||||
|
||||
# Print day total return and add to global realized PnL
|
||||
if day_return != 0:
|
||||
print(f" Day Total Return: {day_return:.2f}%")
|
||||
self.add_realized_pnl(day_return)
|
||||
|
||||
def print_outstanding_positions(self) -> None:
|
||||
"""Print all outstanding positions with share quantities and current values."""
|
||||
if not self.get_outstanding_positions():
|
||||
print("\n====== NO OUTSTANDING POSITIONS ======")
|
||||
return
|
||||
|
||||
print(f"\n====== OUTSTANDING POSITIONS ======")
|
||||
print(
|
||||
f"{'Pair':<15}"
|
||||
f" {'Symbol':<10}"
|
||||
f" {'Side':<4}"
|
||||
f" {'Shares':<10}"
|
||||
f" {'Open $':<8}"
|
||||
f" {'Current $':<10}"
|
||||
f" {'Value $':<12}"
|
||||
f" {'Disequilibrium':<15}"
|
||||
)
|
||||
print("-" * 100)
|
||||
|
||||
total_value = 0.0
|
||||
|
||||
for pos in self.get_outstanding_positions():
|
||||
# Print position A
|
||||
print(
|
||||
f"{pos['pair']:<15}"
|
||||
f" {pos['symbol_a']:<10}"
|
||||
f" {pos['side_a']:<4}"
|
||||
f" {pos['shares_a']:<10.2f}"
|
||||
f" {pos['open_px_a']:<8.2f}"
|
||||
f" {pos['current_px_a']:<10.2f}"
|
||||
f" {pos['current_value_a']:<12.2f}"
|
||||
f" {'':<15}"
|
||||
)
|
||||
|
||||
# Print position B
|
||||
print(
|
||||
f"{'':<15}"
|
||||
f" {pos['symbol_b']:<10}"
|
||||
f" {pos['side_b']:<4}"
|
||||
f" {pos['shares_b']:<10.2f}"
|
||||
f" {pos['open_px_b']:<8.2f}"
|
||||
f" {pos['current_px_b']:<10.2f}"
|
||||
f" {pos['current_value_b']:<12.2f}"
|
||||
)
|
||||
|
||||
# Print pair totals with disequilibrium info
|
||||
print(
|
||||
f"{'':<15}"
|
||||
f" {'PAIR TOTAL':<10}"
|
||||
f" {'':<4}"
|
||||
f" {'':<10}"
|
||||
f" {'':<8}"
|
||||
f" {'':<10}"
|
||||
f" {pos['total_current_value']:<12.2f}"
|
||||
)
|
||||
|
||||
# Print disequilibrium details
|
||||
print(
|
||||
f"{'':<15}"
|
||||
f" {'DISEQUIL':<10}"
|
||||
f" {'':<4}"
|
||||
f" {'':<10}"
|
||||
f" {'':<8}"
|
||||
f" {'':<10}"
|
||||
f" Raw: {pos['current_disequilibrium']:<6.4f}"
|
||||
f" Scaled: {pos['current_scaled_disequilibrium']:<6.4f}"
|
||||
)
|
||||
|
||||
print("-" * 100)
|
||||
|
||||
total_value += pos["total_current_value"]
|
||||
|
||||
print(f"{'TOTAL OUTSTANDING VALUE':<80} ${total_value:<12.2f}")
|
||||
|
||||
def print_grand_totals(self) -> None:
|
||||
"""Print grand totals across all pairs."""
|
||||
print(f"\n====== GRAND TOTALS ACROSS ALL PAIRS ======")
|
||||
print(f"Total Realized PnL: {self.get_total_realized_pnl():.2f}%")
|
||||
|
||||
def handle_outstanding_position(
|
||||
self,
|
||||
pair: TradingPair,
|
||||
pair_result_df: pd.DataFrame,
|
||||
last_row_index: int,
|
||||
open_side_a: str,
|
||||
open_side_b: str,
|
||||
open_px_a: float,
|
||||
open_px_b: float,
|
||||
open_tstamp: datetime,
|
||||
) -> Tuple[float, float, float]:
|
||||
"""
|
||||
Handle calculation and tracking of outstanding positions when no close signal is found.
|
||||
|
||||
Args:
|
||||
pair: TradingPair object
|
||||
pair_result_df: DataFrame with pair results
|
||||
last_row_index: Index of the last row in the data
|
||||
open_side_a, open_side_b: Trading sides for symbols A and B
|
||||
open_px_a, open_px_b: Opening prices for symbols A and B
|
||||
open_tstamp: Opening timestamp
|
||||
"""
|
||||
if pair_result_df is None or pair_result_df.empty:
|
||||
return 0, 0, 0
|
||||
|
||||
last_row = pair_result_df.loc[last_row_index]
|
||||
last_tstamp = last_row["tstamp"]
|
||||
colname_a, colname_b = pair.exec_prices_colnames()
|
||||
last_px_a = last_row[colname_a]
|
||||
last_px_b = last_row[colname_b]
|
||||
|
||||
# Calculate share quantities based on funding per pair
|
||||
# Split funding equally between the two positions
|
||||
funding_per_position = self.config["funding_per_pair"] / 2
|
||||
shares_a = funding_per_position / open_px_a
|
||||
shares_b = funding_per_position / open_px_b
|
||||
|
||||
# Calculate current position values (shares * current price)
|
||||
current_value_a = shares_a * last_px_a * (-1 if open_side_a == "SELL" else 1)
|
||||
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
|
||||
|
||||
# Get disequilibrium information
|
||||
current_disequilibrium = last_row["disequilibrium"]
|
||||
current_scaled_disequilibrium = last_row["scaled_disequilibrium"]
|
||||
|
||||
# Store outstanding positions
|
||||
self.add_outstanding_position(
|
||||
{
|
||||
"pair": str(pair),
|
||||
"symbol_a": pair.symbol_a_,
|
||||
"symbol_b": pair.symbol_b_,
|
||||
"side_a": open_side_a,
|
||||
"side_b": open_side_b,
|
||||
"shares_a": shares_a,
|
||||
"shares_b": shares_b,
|
||||
"open_px_a": open_px_a,
|
||||
"open_px_b": open_px_b,
|
||||
"current_px_a": last_px_a,
|
||||
"current_px_b": last_px_b,
|
||||
"current_value_a": current_value_a,
|
||||
"current_value_b": current_value_b,
|
||||
"total_current_value": total_current_value,
|
||||
"open_time": open_tstamp,
|
||||
"last_time": last_tstamp,
|
||||
"current_abs_term": current_scaled_disequilibrium,
|
||||
"current_disequilibrium": current_disequilibrium,
|
||||
"current_scaled_disequilibrium": current_scaled_disequilibrium,
|
||||
}
|
||||
)
|
||||
|
||||
# Print position details
|
||||
print(f"{pair}: NO CLOSE SIGNAL FOUND - Position held until end of session")
|
||||
print(f" Open: {open_tstamp} | Last: {last_tstamp}")
|
||||
print(
|
||||
f" {pair.symbol_a_}: {open_side_a} {shares_a:.2f} shares @ ${open_px_a:.2f} -> ${last_px_a:.2f} | Value: ${current_value_a:.2f}"
|
||||
)
|
||||
print(
|
||||
f" {pair.symbol_b_}: {open_side_b} {shares_b:.2f} shares @ ${open_px_b:.2f} -> ${last_px_b:.2f} | Value: ${current_value_b:.2f}"
|
||||
)
|
||||
print(f" Total Value: ${total_current_value:.2f}")
|
||||
print(
|
||||
f" Disequilibrium: {current_disequilibrium:.4f} | Scaled: {current_scaled_disequilibrium:.4f}"
|
||||
)
|
||||
|
||||
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()
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,380 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import pandas as pd # type:ignore
|
||||
|
||||
|
||||
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
|
||||
symbol_a_: str
|
||||
symbol_b_: str
|
||||
stat_model_price_: str
|
||||
|
||||
training_mu_: float
|
||||
training_std_: float
|
||||
|
||||
training_df_: pd.DataFrame
|
||||
testing_df_: pd.DataFrame
|
||||
|
||||
user_data_: Dict[str, Any]
|
||||
|
||||
# predicted_df_: Optional[pd.DataFrame]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Dict[str, Any],
|
||||
market_data: pd.DataFrame,
|
||||
symbol_a: str,
|
||||
symbol_b: str,
|
||||
):
|
||||
self.symbol_a_ = symbol_a
|
||||
self.symbol_b_ = symbol_b
|
||||
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._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
|
||||
|
||||
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:
|
||||
# Select only the columns we need
|
||||
df_selected: pd.DataFrame = pd.DataFrame(
|
||||
df[["tstamp", "symbol", self.stat_model_price_]]
|
||||
)
|
||||
|
||||
# Start with unique timestamps
|
||||
result_df: pd.DataFrame = (
|
||||
pd.DataFrame(df_selected["tstamp"]).drop_duplicates().reset_index(drop=True)
|
||||
)
|
||||
|
||||
# For each unique symbol, add a corresponding close price column
|
||||
|
||||
symbols = df_selected["symbol"].unique()
|
||||
for symbol in symbols:
|
||||
# Filter rows for this symbol
|
||||
df_symbol = df_selected[df_selected["symbol"] == symbol].reset_index(
|
||||
drop=True
|
||||
)
|
||||
|
||||
# Create column name like "close-COIN"
|
||||
new_price_column = f"{self.stat_model_price_}_{symbol}"
|
||||
|
||||
# Create temporary dataframe with timestamp and price
|
||||
temp_df = pd.DataFrame(
|
||||
{
|
||||
"tstamp": df_symbol["tstamp"],
|
||||
new_price_column: df_symbol[self.stat_model_price_],
|
||||
}
|
||||
)
|
||||
|
||||
# Join with our result dataframe
|
||||
result_df = pd.merge(result_df, temp_df, on="tstamp", how="left")
|
||||
result_df = result_df.reset_index(
|
||||
drop=True
|
||||
) # do not dropna() since irrelevant symbol would affect dataset
|
||||
|
||||
return result_df.dropna()
|
||||
|
||||
def get_datasets(
|
||||
self,
|
||||
training_minutes: int,
|
||||
training_start_index: int = 0,
|
||||
testing_size: Optional[int] = None,
|
||||
) -> None:
|
||||
|
||||
testing_start_index = training_start_index + training_minutes
|
||||
self.training_df_ = self.market_data_.iloc[
|
||||
training_start_index:testing_start_index, :training_minutes
|
||||
].copy()
|
||||
assert self.training_df_ is not None
|
||||
self.training_df_ = self.training_df_.dropna().reset_index(drop=True)
|
||||
|
||||
testing_start_index = training_start_index + training_minutes
|
||||
if testing_size is None:
|
||||
self.testing_df_ = self.market_data_.iloc[testing_start_index:, :].copy()
|
||||
else:
|
||||
self.testing_df_ = self.market_data_.iloc[
|
||||
testing_start_index : testing_start_index + testing_size, :
|
||||
].copy()
|
||||
assert self.testing_df_ is not None
|
||||
self.testing_df_ = self.testing_df_.dropna().reset_index(drop=True)
|
||||
|
||||
def colnames(self) -> List[str]:
|
||||
return [
|
||||
f"{self.stat_model_price_}_{self.symbol_a_}",
|
||||
f"{self.stat_model_price_}_{self.symbol_b_}",
|
||||
]
|
||||
|
||||
def exec_prices_colnames(self) -> List[str]:
|
||||
return [
|
||||
f"exec_price_{self.symbol_a_}",
|
||||
f"exec_price_{self.symbol_b_}",
|
||||
]
|
||||
|
||||
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"]
|
||||
|
||||
# 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
|
||||
|
||||
# Concatenate with explicit dtypes to avoid warnings
|
||||
self.user_data_["trades"] = pd.concat(
|
||||
[existing_trades, trades], ignore_index=True, copy=False
|
||||
)
|
||||
|
||||
def get_trades(self) -> pd.DataFrame:
|
||||
return (
|
||||
self.user_data_["trades"] if "trades" in self.user_data_ else pd.DataFrame()
|
||||
)
|
||||
|
||||
def cointegration_check(self) -> Optional[pd.DataFrame]:
|
||||
print(f"***{self}*** STARTING....")
|
||||
config = self.config_
|
||||
|
||||
curr_training_start_idx = 0
|
||||
|
||||
COINTEGRATION_DATA_COLUMNS = {
|
||||
"tstamp": "datetime64[ns]",
|
||||
"pair": "string",
|
||||
"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)
|
||||
|
||||
training_minutes = config["training_minutes"]
|
||||
while True:
|
||||
print(curr_training_start_idx, end="\r")
|
||||
self.get_datasets(
|
||||
training_minutes=training_minutes,
|
||||
training_start_index=curr_training_start_idx,
|
||||
testing_size=1,
|
||||
)
|
||||
|
||||
if len(self.training_df_) < training_minutes:
|
||||
print(
|
||||
f"{self}: current offset={curr_training_start_idx}"
|
||||
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
|
||||
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
|
||||
|
||||
def on_open_trades(self, trades: pd.DataFrame) -> None:
|
||||
if "close_trades" in self.user_data_:
|
||||
del self.user_data_["close_trades"]
|
||||
self.user_data_["open_trades"] = trades
|
||||
|
||||
def on_close_trades(self, trades: pd.DataFrame) -> None:
|
||||
del self.user_data_["open_trades"]
|
||||
self.user_data_["close_trades"] = trades
|
||||
|
||||
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 _single_instrument_return(symbol: str) -> float:
|
||||
instrument_open_trades = open_trades[open_trades["symbol"] == symbol]
|
||||
instrument_open_price = instrument_open_trades["price"].iloc[0]
|
||||
|
||||
sign = -1 if instrument_open_trades["side"].iloc[0] == "SELL" else 1
|
||||
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
|
||||
|
||||
instrument_a_return = _single_instrument_return(self.symbol_a_)
|
||||
instrument_b_return = _single_instrument_return(self.symbol_b_)
|
||||
return instrument_a_return + instrument_b_return
|
||||
return 0.0
|
||||
|
||||
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_}"
|
||||
|
||||
@abstractmethod
|
||||
def predict(self) -> pd.DataFrame: ...
|
||||
|
||||
# @abstractmethod
|
||||
# def predicted_df(self) -> Optional[pd.DataFrame]: ...
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
)
|
||||
@@ -0,0 +1,17 @@
|
||||
import hjson
|
||||
from typing import Dict
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
def load_config(config_path: str) -> Dict:
|
||||
with open(config_path, "r") as f:
|
||||
config = hjson.load(f)
|
||||
return dict(config)
|
||||
|
||||
|
||||
def expand_filename(filename: str) -> str:
|
||||
# expand %T
|
||||
res = filename.replace("%T", datetime.now().strftime("%Y%m%d_%H%M%S"))
|
||||
# expand %D
|
||||
return res.replace("%D", datetime.now().strftime("%Y%m%d"))
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
from typing import Dict, List, cast
|
||||
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
|
||||
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
|
||||
df = pd.read_sql_query(query, conn)
|
||||
return df
|
||||
except sqlite3.Error as excpt:
|
||||
print(f"SQLite error: {excpt}")
|
||||
raise
|
||||
except Exception as excpt:
|
||||
print(f"Error: {excpt}")
|
||||
raise Exception() from excpt
|
||||
finally:
|
||||
if "conn" in locals():
|
||||
conn.close()
|
||||
|
||||
|
||||
def convert_time_to_UTC(value: str, timezone: str, extra_minutes: int = 0) -> str:
|
||||
|
||||
from zoneinfo import ZoneInfo
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
# Parse it to naive datetime object
|
||||
local_dt = datetime.strptime(value, "%Y-%m-%d %H:%M:%S")
|
||||
local_dt = local_dt + timedelta(minutes=extra_minutes)
|
||||
|
||||
zinfo = ZoneInfo(timezone)
|
||||
result: datetime = local_dt.replace(tzinfo=zinfo).astimezone(ZoneInfo("UTC"))
|
||||
|
||||
return result.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
def load_market_data(
|
||||
datafile: str,
|
||||
instruments: List[Dict[str, str]],
|
||||
db_table_name: str,
|
||||
trading_hours: Dict = {},
|
||||
extra_minutes: int = 0,
|
||||
) -> pd.DataFrame:
|
||||
|
||||
insts = [
|
||||
'"' + instrument["instrument_id_pfx"] + instrument["symbol"] + '"'
|
||||
for instrument in instruments
|
||||
]
|
||||
instrument_ids = list(set(insts))
|
||||
exchange_ids = list(
|
||||
set(['"' + instrument["exchange_id"] + '"' for instrument in instruments])
|
||||
)
|
||||
|
||||
query = "select"
|
||||
query += " tstamp"
|
||||
query += ", tstamp_ns as time_ns"
|
||||
|
||||
query += f", substr(instrument_id, instr(instrument_id, '-') + 1) as symbol"
|
||||
query += ", open"
|
||||
query += ", high"
|
||||
query += ", low"
|
||||
query += ", close"
|
||||
query += ", volume"
|
||||
query += ", num_trades"
|
||||
query += ", vwap"
|
||||
|
||||
query += f" from {db_table_name}"
|
||||
query += f" where exchange_id in ({','.join(exchange_ids)})"
|
||||
query += f" and instrument_id in ({','.join(instrument_ids)})"
|
||||
|
||||
df = load_sqlite_to_dataframe(db_path=datafile, query=query)
|
||||
|
||||
# Trading Hours
|
||||
if len(df) > 0 and len(trading_hours) > 0:
|
||||
date_str = df["tstamp"][0][0:10]
|
||||
|
||||
start_time = convert_time_to_UTC(
|
||||
f"{date_str} {trading_hours['begin_session']}", trading_hours["timezone"]
|
||||
)
|
||||
end_time = convert_time_to_UTC(
|
||||
f"{date_str} {trading_hours['end_session']}", trading_hours["timezone"], extra_minutes=extra_minutes # to get execution price
|
||||
)
|
||||
|
||||
# Perform boolean selection
|
||||
df = df[(df["tstamp"] >= start_time) & (df["tstamp"] <= end_time)]
|
||||
df["tstamp"] = pd.to_datetime(df["tstamp"])
|
||||
|
||||
return cast(pd.DataFrame, df)
|
||||
|
||||
|
||||
# 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.
|
||||
# Returns instruments without the configured prefix.
|
||||
# """
|
||||
# try:
|
||||
# conn = sqlite3.connect(datafile)
|
||||
|
||||
# # Build exclusion list with full instrument_ids
|
||||
# exclude_instruments = config.get("exclude_instruments", [])
|
||||
# prefix = config.get("instrument_id_pfx", "")
|
||||
# exclude_instrument_ids = [f"{prefix}{inst}" for inst in exclude_instruments]
|
||||
|
||||
# # Query to get distinct instrument_ids
|
||||
# query = f"""
|
||||
# SELECT DISTINCT instrument_id
|
||||
# FROM {config['db_table_name']}
|
||||
# WHERE exchange_id = ?
|
||||
# """
|
||||
|
||||
# # Add exclusion clause if there are instruments to exclude
|
||||
# if exclude_instrument_ids:
|
||||
# placeholders = ",".join(["?" for _ in exclude_instrument_ids])
|
||||
# query += f" AND instrument_id NOT IN ({placeholders})"
|
||||
# cursor = conn.execute(
|
||||
# query, (config["exchange_id"],) + tuple(exclude_instrument_ids)
|
||||
# )
|
||||
# else:
|
||||
# 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
|
||||
# instruments = []
|
||||
# for instrument_id in instrument_ids:
|
||||
# if instrument_id.startswith(prefix):
|
||||
# symbol = instrument_id[len(prefix) :]
|
||||
# instruments.append(symbol)
|
||||
# else:
|
||||
# instruments.append(instrument_id)
|
||||
|
||||
# return sorted(instruments)
|
||||
|
||||
# except Exception as e:
|
||||
# print(f"Error auto-detecting instruments from {datafile}: {str(e)}")
|
||||
# return []
|
||||
|
||||
|
||||
# if __name__ == "__main__":
|
||||
# df1 = load_sqlite_to_dataframe(sys.argv[1], table_name="md_1min_bars")
|
||||
|
||||
# print(df1)
|
||||
@@ -25,7 +25,7 @@ def list_tables(db_path: str) -> List[str]:
|
||||
conn.close()
|
||||
return tables
|
||||
|
||||
def view_table_schema(db_path: str, table_name: str):
|
||||
def view_table_schema(db_path: str, table_name: str) -> None:
|
||||
"""View the schema of a specific table."""
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.cursor()
|
||||
@@ -44,13 +44,13 @@ def view_table_schema(db_path: str, table_name: str):
|
||||
|
||||
conn.close()
|
||||
|
||||
def view_config_table(db_path: str, limit: int = 10):
|
||||
def view_config_table(db_path: str, limit: int = 10) -> None:
|
||||
"""View entries from the config table."""
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.cursor()
|
||||
|
||||
cursor.execute(f"""
|
||||
SELECT id, run_timestamp, config_file_path, strategy_class,
|
||||
SELECT id, run_timestamp, config_file_path, fit_method_class,
|
||||
datafiles, instruments, config_json
|
||||
FROM config
|
||||
ORDER BY run_timestamp DESC
|
||||
@@ -67,17 +67,17 @@ def view_config_table(db_path: str, limit: int = 10):
|
||||
print("=" * 80)
|
||||
|
||||
for row in rows:
|
||||
id, run_timestamp, config_file_path, strategy_class, datafiles, instruments, config_json = row
|
||||
id, run_timestamp, config_file_path, fit_method_class, datafiles, instruments, config_json = row
|
||||
|
||||
print(f"ID: {id} | {run_timestamp}")
|
||||
print(f"Config: {config_file_path} | Strategy: {strategy_class}")
|
||||
print(f"Config: {config_file_path} | Strategy: {fit_method_class}")
|
||||
print(f"Files: {datafiles}")
|
||||
print(f"Instruments: {instruments}")
|
||||
print("-" * 80)
|
||||
|
||||
conn.close()
|
||||
|
||||
def view_results_summary(db_path: str):
|
||||
def view_results_summary(db_path: str) -> None:
|
||||
"""View summary of trading results."""
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.cursor()
|
||||
@@ -119,7 +119,7 @@ def view_results_summary(db_path: str):
|
||||
|
||||
conn.close()
|
||||
|
||||
def main():
|
||||
def main() -> None:
|
||||
if len(sys.argv) < 2:
|
||||
print("Usage: python db_inspector.py <database_path> [command]")
|
||||
print("Commands:")
|
||||
@@ -0,0 +1,66 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=45", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "pairs-trading"
|
||||
version = "0.1.0"
|
||||
description = "Pairs Trading Backtesting Framework"
|
||||
requires-python = ">=3.8"
|
||||
|
||||
[tool.black]
|
||||
line-length = 88
|
||||
target-version = ['py38']
|
||||
include = '\.pyi?$'
|
||||
extend-exclude = '''
|
||||
/(
|
||||
# directories
|
||||
\.eggs
|
||||
| \.git
|
||||
| \.hg
|
||||
| \.mypy_cache
|
||||
| \.tox
|
||||
| \.venv
|
||||
| build
|
||||
| dist
|
||||
)/
|
||||
'''
|
||||
|
||||
[tool.flake8]
|
||||
max-line-length = 88
|
||||
extend-ignore = ["E203", "W503"]
|
||||
exclude = [
|
||||
".git",
|
||||
"__pycache__",
|
||||
"build",
|
||||
"dist",
|
||||
".venv",
|
||||
".mypy_cache",
|
||||
".tox"
|
||||
]
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.8"
|
||||
warn_return_any = true
|
||||
warn_unused_configs = true
|
||||
disallow_untyped_defs = true
|
||||
disallow_incomplete_defs = true
|
||||
check_untyped_defs = true
|
||||
disallow_untyped_decorators = true
|
||||
no_implicit_optional = true
|
||||
warn_redundant_casts = true
|
||||
warn_unused_ignores = true
|
||||
warn_no_return = true
|
||||
warn_unreachable = true
|
||||
strict_equality = true
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
module = [
|
||||
"numpy.*",
|
||||
"pandas.*",
|
||||
"matplotlib.*",
|
||||
"seaborn.*",
|
||||
"scipy.*",
|
||||
"sklearn.*"
|
||||
]
|
||||
ignore_missing_imports = true
|
||||
@@ -0,0 +1,24 @@
|
||||
{
|
||||
"include": [
|
||||
"lib"
|
||||
],
|
||||
"exclude": [
|
||||
"**/node_modules",
|
||||
"**/__pycache__",
|
||||
"**/.*",
|
||||
"results",
|
||||
"data"
|
||||
],
|
||||
"ignore": [],
|
||||
"defineConstant": {},
|
||||
"typeCheckingMode": "basic",
|
||||
"useLibraryCodeForTypes": true,
|
||||
"autoImportCompletions": true,
|
||||
"autoSearchPaths": true,
|
||||
"extraPaths": [
|
||||
"lib"
|
||||
],
|
||||
"stubPath": "./typings",
|
||||
"venvPath": ".",
|
||||
"venv": "python3.12-venv"
|
||||
}
|
||||
+115
-106
@@ -4,6 +4,7 @@ async-timeout>=4.0.2
|
||||
attrs>=21.2.0
|
||||
beautifulsoup4>=4.10.0
|
||||
black>=23.3.0
|
||||
flake8>=6.0.0
|
||||
certifi>=2020.6.20
|
||||
chardet>=4.0.0
|
||||
charset-normalizer>=3.1.0
|
||||
@@ -23,11 +24,14 @@ hjson>=3.0.2
|
||||
html5lib>=1.1
|
||||
httplib2>=0.20.2
|
||||
idna>=3.3
|
||||
ipython>=8.18.1
|
||||
ipywidgets>=8.1.1
|
||||
ifaddr>=0.1.7
|
||||
IMDbPY>=2021.4.18
|
||||
ipykernel>=6.29.5
|
||||
jeepney>=0.7.1
|
||||
jsonschema>=3.2.0
|
||||
jupyter>=1.0.0
|
||||
keyring>=23.5.0
|
||||
launchpadlib>=1.10.16
|
||||
lazr.restfulclient>=0.14.4
|
||||
@@ -41,19 +45,23 @@ more-itertools>=8.10.0
|
||||
multidict>=6.0.4
|
||||
mypy>=0.942
|
||||
mypy-extensions>=0.4.3
|
||||
nbformat>=5.10.2
|
||||
netaddr>=0.8.0
|
||||
######### netifaces>=0.11.0
|
||||
numpy>=1.26.4,<2.3.0
|
||||
oauthlib>=3.2.0
|
||||
packaging>=23.1
|
||||
pandas>=2.2.3
|
||||
pathspec>=0.11.1
|
||||
pexpect>=4.8.0
|
||||
Pillow>=9.0.1
|
||||
platformdirs>=3.2.0
|
||||
plotly>=5.19.0
|
||||
protobuf>=3.12.4
|
||||
psutil>=5.9.0
|
||||
ptyprocess>=0.7.0
|
||||
pycurl>=7.44.1
|
||||
pyelftools>=0.27
|
||||
# pyelftools>=0.27
|
||||
Pygments>=2.11.2
|
||||
pyparsing>=2.4.7
|
||||
pyrsistent>=0.18.1
|
||||
@@ -61,11 +69,12 @@ python-debian>=0.1.43 #+ubuntu1.1
|
||||
python-dotenv>=0.19.2
|
||||
python-magic>=0.4.24
|
||||
python-xlib>=0.29
|
||||
pyxdg>=0.27
|
||||
# pyxdg>=0.27
|
||||
PyYAML>=6.0
|
||||
reportlab>=3.6.8
|
||||
requests>=2.25.1
|
||||
requests-file>=1.5.1
|
||||
scipy<1.13.0
|
||||
seaborn>=0.13.2
|
||||
SecretStorage>=3.3.1
|
||||
setproctitle>=1.2.2
|
||||
@@ -73,113 +82,113 @@ six>=1.16.0
|
||||
soupsieve>=2.3.1
|
||||
ssh-import-id>=5.11
|
||||
statsmodels>=0.14.4
|
||||
texttable>=1.6.4
|
||||
# texttable>=1.6.4
|
||||
tldextract>=3.1.2
|
||||
tomli>=1.2.2
|
||||
######## typed-ast>=1.4.3
|
||||
types-aiofiles>=0.1
|
||||
types-annoy>=1.17
|
||||
types-appdirs>=1.4
|
||||
types-atomicwrites>=1.4
|
||||
types-aws-xray-sdk>=2.8
|
||||
types-babel>=2.9
|
||||
types-backports-abc>=0.5
|
||||
types-backports.ssl-match-hostname>=3.7
|
||||
types-beautifulsoup4>=4.10
|
||||
types-bleach>=4.1
|
||||
types-boto>=2.49
|
||||
types-braintree>=4.11
|
||||
types-cachetools>=4.2
|
||||
types-caldav>=0.8
|
||||
types-certifi>=2020.4
|
||||
types-characteristic>=14.3
|
||||
types-chardet>=4.0
|
||||
types-click>=7.1
|
||||
types-click-spinner>=0.1
|
||||
types-colorama>=0.4
|
||||
types-commonmark>=0.9
|
||||
types-contextvars>=0.1
|
||||
types-croniter>=1.0
|
||||
types-cryptography>=3.3
|
||||
types-dataclasses>=0.1
|
||||
types-dateparser>=1.0
|
||||
types-DateTimeRange>=0.1
|
||||
types-decorator>=0.1
|
||||
types-Deprecated>=1.2
|
||||
types-docopt>=0.6
|
||||
types-docutils>=0.17
|
||||
types-editdistance>=0.5
|
||||
types-emoji>=1.2
|
||||
types-entrypoints>=0.3
|
||||
types-enum34>=1.1
|
||||
types-filelock>=3.2
|
||||
types-first>=2.0
|
||||
types-Flask>=1.1
|
||||
types-freezegun>=1.1
|
||||
types-frozendict>=0.1
|
||||
types-futures>=3.3
|
||||
types-html5lib>=1.1
|
||||
types-httplib2>=0.19
|
||||
types-humanfriendly>=9.2
|
||||
types-ipaddress>=1.0
|
||||
types-itsdangerous>=1.1
|
||||
types-JACK-Client>=0.1
|
||||
types-Jinja2>=2.11
|
||||
types-jmespath>=0.10
|
||||
types-jsonschema>=3.2
|
||||
types-Markdown>=3.3
|
||||
types-MarkupSafe>=1.1
|
||||
types-mock>=4.0
|
||||
types-mypy-extensions>=0.4
|
||||
types-mysqlclient>=2.0
|
||||
types-oauthlib>=3.1
|
||||
types-orjson>=3.6
|
||||
types-paramiko>=2.7
|
||||
types-Pillow>=8.3
|
||||
types-polib>=1.1
|
||||
types-prettytable>=2.1
|
||||
types-protobuf>=3.17
|
||||
types-psutil>=5.8
|
||||
types-psycopg2>=2.9
|
||||
types-pyaudio>=0.2
|
||||
types-pycurl>=0.1
|
||||
types-pyfarmhash>=0.2
|
||||
types-Pygments>=2.9
|
||||
types-PyMySQL>=1.0
|
||||
types-pyOpenSSL>=20.0
|
||||
types-pyRFC3339>=0.1
|
||||
types-pysftp>=0.2
|
||||
types-pytest-lazy-fixture>=0.6
|
||||
types-python-dateutil>=2.8
|
||||
types-python-gflags>=3.1
|
||||
types-python-nmap>=0.6
|
||||
types-python-slugify>=5.0
|
||||
types-pytz>=2021.1
|
||||
types-pyvmomi>=7.0
|
||||
types-PyYAML>=5.4
|
||||
types-redis>=3.5
|
||||
types-requests>=2.25
|
||||
types-retry>=0.9
|
||||
types-selenium>=3.141
|
||||
types-Send2Trash>=1.8
|
||||
types-setuptools>=57.4
|
||||
types-simplejson>=3.17
|
||||
types-singledispatch>=3.7
|
||||
types-six>=1.16
|
||||
types-slumber>=0.7
|
||||
types-stripe>=2.59
|
||||
types-tabulate>=0.8
|
||||
types-termcolor>=1.1
|
||||
types-toml>=0.10
|
||||
types-toposort>=1.6
|
||||
types-ttkthemes>=3.2
|
||||
types-typed-ast>=1.4
|
||||
types-tzlocal>=0.1
|
||||
types-ujson>=0.1
|
||||
types-vobject>=0.9
|
||||
types-waitress>=0.1
|
||||
types-Werkzeug>=1.0
|
||||
types-xxhash>=2.0
|
||||
# types-aiofiles>=0.1
|
||||
# types-annoy>=1.17
|
||||
# types-appdirs>=1.4
|
||||
# types-atomicwrites>=1.4
|
||||
# types-aws-xray-sdk>=2.8
|
||||
# types-babel>=2.9
|
||||
# types-backports-abc>=0.5
|
||||
# types-backports.ssl-match-hostname>=3.7
|
||||
# types-beautifulsoup4>=4.10
|
||||
# types-bleach>=4.1
|
||||
# types-boto>=2.49
|
||||
# types-braintree>=4.11
|
||||
# types-cachetools>=4.2
|
||||
# types-caldav>=0.8
|
||||
# types-certifi>=2020.4
|
||||
# types-characteristic>=14.3
|
||||
# types-chardet>=4.0
|
||||
# types-click>=7.1
|
||||
# types-click-spinner>=0.1
|
||||
# types-colorama>=0.4
|
||||
# types-commonmark>=0.9
|
||||
# types-contextvars>=0.1
|
||||
# types-croniter>=1.0
|
||||
# types-cryptography>=3.3
|
||||
# types-dataclasses>=0.1
|
||||
# types-dateparser>=1.0
|
||||
# types-DateTimeRange>=0.1
|
||||
# types-decorator>=0.1
|
||||
# types-Deprecated>=1.2
|
||||
# types-docopt>=0.6
|
||||
# types-docutils>=0.17
|
||||
# types-editdistance>=0.5
|
||||
# types-emoji>=1.2
|
||||
# types-entrypoints>=0.3
|
||||
# types-enum34>=1.1
|
||||
# types-filelock>=3.2
|
||||
# types-first>=2.0
|
||||
# types-Flask>=1.1
|
||||
# types-freezegun>=1.1
|
||||
# types-frozendict>=0.1
|
||||
# types-futures>=3.3
|
||||
# types-html5lib>=1.1
|
||||
# types-httplib2>=0.19
|
||||
# types-humanfriendly>=9.2
|
||||
# types-ipaddress>=1.0
|
||||
# types-itsdangerous>=1.1
|
||||
# types-JACK-Client>=0.1
|
||||
# types-Jinja2>=2.11
|
||||
# types-jmespath>=0.10
|
||||
# types-jsonschema>=3.2
|
||||
# types-Markdown>=3.3
|
||||
# types-MarkupSafe>=1.1
|
||||
# types-mock>=4.0
|
||||
# types-mypy-extensions>=0.4
|
||||
# types-mysqlclient>=2.0
|
||||
# types-oauthlib>=3.1
|
||||
# types-orjson>=3.6
|
||||
# types-paramiko>=2.7
|
||||
# types-Pillow>=8.3
|
||||
# types-polib>=1.1
|
||||
# types-prettytable>=2.1
|
||||
# types-protobuf>=3.17
|
||||
# types-psutil>=5.8
|
||||
# types-psycopg2>=2.9
|
||||
# types-pyaudio>=0.2
|
||||
# types-pycurl>=0.1
|
||||
# types-pyfarmhash>=0.2
|
||||
# types-Pygments>=2.9
|
||||
# types-PyMySQL>=1.0
|
||||
# types-pyOpenSSL>=20.0
|
||||
# types-pyRFC3339>=0.1
|
||||
# types-pysftp>=0.2
|
||||
# types-pytest-lazy-fixture>=0.6
|
||||
# types-python-dateutil>=2.8
|
||||
# types-python-gflags>=3.1
|
||||
# types-python-nmap>=0.6
|
||||
# types-python-slugify>=5.0
|
||||
# types-pytz>=2021.1
|
||||
# types-pyvmomi>=7.0
|
||||
# types-PyYAML>=5.4
|
||||
# types-redis>=3.5
|
||||
# types-requests>=2.25
|
||||
# types-retry>=0.9
|
||||
# types-selenium>=3.141
|
||||
# types-Send2Trash>=1.8
|
||||
# types-setuptools>=57.4
|
||||
# types-simplejson>=3.17
|
||||
# types-singledispatch>=3.7
|
||||
# types-six>=1.16
|
||||
# types-slumber>=0.7
|
||||
# types-stripe>=2.59
|
||||
# types-tabulate>=0.8
|
||||
# types-termcolor>=1.1
|
||||
# types-toml>=0.10
|
||||
# types-toposort>=1.6
|
||||
# types-ttkthemes>=3.2
|
||||
# types-typed-ast>=1.4
|
||||
# types-tzlocal>=0.1
|
||||
# types-ujson>=0.1
|
||||
# types-vobject>=0.9
|
||||
# types-waitress>=0.1
|
||||
#types-Werkzeug>=1.0
|
||||
#types-xxhash>=2.0
|
||||
typing-extensions>=3.10.0.2
|
||||
Unidecode>=1.3.3
|
||||
urllib3>=1.26.5
|
||||
|
||||
@@ -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()
|
||||
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"cells": [],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python",
|
||||
"version": "3.12.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1,232 @@
|
||||
import argparse
|
||||
import glob
|
||||
import importlib
|
||||
import os
|
||||
from datetime import date, datetime
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from research.research_tools import create_pairs
|
||||
from tools.config import expand_filename, load_config
|
||||
from pt_trading.results import (
|
||||
BacktestResult,
|
||||
create_result_database,
|
||||
store_config_in_database,
|
||||
)
|
||||
from pt_trading.fit_method import PairsTradingFitMethod
|
||||
from pt_trading.trading_pair import TradingPair
|
||||
|
||||
DayT = str
|
||||
DataFileNameT = str
|
||||
|
||||
def resolve_datafiles(
|
||||
config: Dict, date_pattern: str, instruments: List[Dict[str, str]]
|
||||
) -> List[Tuple[DayT, DataFileNameT]]:
|
||||
resolved_files: List[Tuple[DayT, DataFileNameT]] = []
|
||||
for inst in instruments:
|
||||
pattern = date_pattern
|
||||
inst_type = inst["instrument_type"]
|
||||
data_dir = config["market_data_loading"][inst_type]["data_directory"]
|
||||
if "*" in pattern or "?" in pattern:
|
||||
# Handle wildcards
|
||||
if not os.path.isabs(pattern):
|
||||
pattern = os.path.join(data_dir, f"{pattern}.mktdata.ohlcv.db")
|
||||
matched_files = glob.glob(pattern)
|
||||
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:
|
||||
# Handle explicit file path
|
||||
if not os.path.isabs(pattern):
|
||||
pattern = os.path.join(data_dir, f"{pattern}.mktdata.ohlcv.db")
|
||||
resolved_files.append((date_pattern, pattern))
|
||||
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(
|
||||
config: Dict,
|
||||
datafiles: List[str],
|
||||
fit_method: PairsTradingFitMethod,
|
||||
instruments: List[Dict[str, str]],
|
||||
) -> BacktestResult:
|
||||
"""
|
||||
Run backtest for all pairs using the specified instruments.
|
||||
"""
|
||||
bt_result: BacktestResult = BacktestResult(config=config)
|
||||
# if len(datafiles) < 2:
|
||||
# print(f"WARNING: insufficient data files: {datafiles}")
|
||||
# return bt_result
|
||||
|
||||
if not all([os.path.exists(datafile) for datafile in datafiles]):
|
||||
print(f"WARNING: data file {datafiles} does not exist")
|
||||
return bt_result
|
||||
|
||||
pairs_trades = []
|
||||
|
||||
pairs = create_pairs(
|
||||
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:
|
||||
pairs_trades.append(single_pair_trades)
|
||||
print(f"pairs_trades:\n{pairs_trades}")
|
||||
# Check if result_list has any data before concatenating
|
||||
if len(pairs_trades) == 0:
|
||||
print("No trading signals found for any pairs")
|
||||
return bt_result
|
||||
|
||||
bt_result.collect_single_day_results(pairs_trades)
|
||||
return bt_result
|
||||
|
||||
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(
|
||||
"--date_pattern",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Date YYYYMMDD, allows * and ? wildcards",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--instruments",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Comma-separated list of instrument symbols (e.g., COIN:EQUITY,GBTC:CRYPTO)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--result_db",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to SQLite database for storing results. Use 'NONE' to disable database output.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
config: Dict = load_config(args.config)
|
||||
|
||||
# Dynamically instantiate fit method class
|
||||
fit_method = PairsTradingFitMethod.create(config)
|
||||
|
||||
# Resolve data files (CLI takes priority over config)
|
||||
instruments = get_instruments(args, config)
|
||||
datafiles = resolve_datafiles(config, args.date_pattern, instruments)
|
||||
|
||||
days = list(set([day for day, _ in datafiles]))
|
||||
print(f"Found {len(datafiles)} data files to process:")
|
||||
for df in datafiles:
|
||||
print(f" - {df}")
|
||||
|
||||
# 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]] = {}
|
||||
is_config_stored = False
|
||||
# Process each data file
|
||||
|
||||
for day in sorted(days):
|
||||
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]):
|
||||
print(f"WARNING: insufficient data files: {md_datafiles}")
|
||||
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
|
||||
try:
|
||||
fit_method.reset()
|
||||
|
||||
bt_results = run_backtest(
|
||||
config=config,
|
||||
datafiles=md_datafiles,
|
||||
fit_method=fit_method,
|
||||
instruments=instruments,
|
||||
)
|
||||
|
||||
if bt_results.trades is None or len(bt_results.trades) == 0:
|
||||
print(f"No trades found for {day}")
|
||||
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
|
||||
if args.result_db.upper() != "NONE":
|
||||
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}")
|
||||
|
||||
except Exception as err:
|
||||
print(f"Error processing {day}: {str(err)}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Calculate and print results using a new BacktestResult instance for aggregation
|
||||
if all_results:
|
||||
aggregate_bt_results = BacktestResult(config=config)
|
||||
aggregate_bt_results.calculate_returns(all_results)
|
||||
aggregate_bt_results.print_grand_totals()
|
||||
aggregate_bt_results.print_outstanding_positions()
|
||||
|
||||
if args.result_db.upper() != "NONE":
|
||||
print(f"\nResults stored in database: {args.result_db}")
|
||||
else:
|
||||
print("No results to display.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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
|
||||
@@ -12,11 +12,18 @@
|
||||
# -------------------------------------
|
||||
# --- Current month - all files
|
||||
# -------------------------------------
|
||||
cd $(realpath $(dirname $0))
|
||||
cd $(realpath $(dirname $0))/..
|
||||
mkdir -p ./data/crypto
|
||||
pushd ./data/crypto
|
||||
|
||||
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
|
||||
eval $Cmd
|
||||
# -------------------------------------
|
||||
|
||||
for srcfname in $(ls *.db.gz); do
|
||||
@@ -24,8 +31,12 @@ for srcfname in $(ls *.db.gz); do
|
||||
tgtfile=${dt}.mktdata.ohlcv.db
|
||||
echo "${srcfname} -> ${tgtfile}"
|
||||
|
||||
gunzip -c $srcfname > temp.db
|
||||
rm -f ${tgtfile} && sqlite3 temp.db ".dump md_1min_bars" | sqlite3 ${tgtfile} && rm ${srcfname}
|
||||
Cmd="gunzip -c $srcfname > temp.db"
|
||||
echo $Cmd
|
||||
eval $Cmd
|
||||
Cmd="rm -f ${tgtfile} && sqlite3 temp.db \".dump md_1min_bars\" | sqlite3 ${tgtfile} && rm ${srcfname}"
|
||||
echo $Cmd
|
||||
eval $Cmd
|
||||
done
|
||||
rm temp.db
|
||||
popd
|
||||
|
||||
@@ -26,8 +26,12 @@ for srcfname in $(ls *.db.gz); do
|
||||
tgtfile=${dt}.mktdata.ohlcv.db
|
||||
echo "${srcfname} -> ${tgtfile}"
|
||||
|
||||
gunzip -c $srcfname > temp.db
|
||||
rm -f ${tgtfile} && sqlite3 temp.db ".dump md_1min_bars" | sqlite3 ${tgtfile} && rm ${srcfname}
|
||||
Cmd="gunzip -c $srcfname > temp.db && rm $srcfname"
|
||||
echo ${Cmd}
|
||||
eval ${Cmd}
|
||||
Cmd="rm -f ${tgtfile} && sqlite3 temp.db '.dump md_1min_bars' | sqlite3 ${tgtfile}"
|
||||
echo ${Cmd}
|
||||
eval ${Cmd}
|
||||
done
|
||||
rm temp.db
|
||||
popd
|
||||
|
||||
@@ -1,858 +0,0 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Sliding Fit Strategy Visualization Notebook\n",
|
||||
"\n",
|
||||
"This notebook is specifically designed for the SlidingFitStrategy, which uses a sliding window approach.\n",
|
||||
"It re-trains the model every minute and shows how cointegration, model parameters, and trading signals evolve over time.\n",
|
||||
"You can visualize the dynamic nature of the sliding window and how the relationship between instruments changes."
|
||||
]
|
||||
},
|
||||
{
|
||||
"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",
|
||||
" - Choose your strategy (StaticFitStrategy or SlidingFitStrategy)\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": 1,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Trading Parameters Configuration\n",
|
||||
"# Specify your configuration file, trading symbols and date here\n",
|
||||
"\n",
|
||||
"# Configuration file selection\n",
|
||||
"CONFIG_FILE = \"equity\" # Options: \"equity\", \"crypto\", or custom filename (without .cfg extension)\n",
|
||||
"\n",
|
||||
"# Trading pair symbols\n",
|
||||
"SYMBOL_A = \"COIN\" # Change this to your desired symbol A\n",
|
||||
"SYMBOL_B = \"MSTR\" # Change this to your desired symbol B\n",
|
||||
"\n",
|
||||
"# Date for data file selection (format: YYYYMMDD)\n",
|
||||
"TRADING_DATE = \"20250605\" # Change this to your desired date\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"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",
|
||||
"from IPython.display import clear_output\n",
|
||||
"\n",
|
||||
"# Import our modules\n",
|
||||
"from strategies import SlidingFitStrategy, PairState\n",
|
||||
"from tools.data_loader import load_market_data\n",
|
||||
"from tools.trading_pair import TradingPair\n",
|
||||
"from 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'] = (15, 10)\n",
|
||||
"\n",
|
||||
"print(\"Setup complete!\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Configuration"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Load Configuration from Configuration Files using HJSON\n",
|
||||
"import hjson\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"def load_config_from_file(config_type=\"equity\"):\n",
|
||||
" \"\"\"Load configuration from configuration files using HJSON\"\"\"\n",
|
||||
" config_file = f\"../../configuration/{config_type}.cfg\"\n",
|
||||
" \n",
|
||||
" try:\n",
|
||||
" with open(config_file, 'r') as f:\n",
|
||||
" # HJSON handles comments, trailing commas, and other human-friendly features\n",
|
||||
" config = hjson.load(f)\n",
|
||||
" \n",
|
||||
" # Convert relative paths to absolute paths from notebook perspective\n",
|
||||
" if 'data_directory' in config:\n",
|
||||
" data_dir = config['data_directory']\n",
|
||||
" if data_dir.startswith('./'):\n",
|
||||
" # Convert relative path to absolute path from notebook's perspective\n",
|
||||
" config['data_directory'] = os.path.abspath(f\"../../{data_dir[2:]}\")\n",
|
||||
" \n",
|
||||
" return config\n",
|
||||
" \n",
|
||||
" except FileNotFoundError:\n",
|
||||
" print(f\"Configuration file not found: {config_file}\")\n",
|
||||
" return None\n",
|
||||
" except hjson.HjsonDecodeError as e:\n",
|
||||
" print(f\"HJSON parsing error in {config_file}: {e}\")\n",
|
||||
" return None\n",
|
||||
" except Exception as e:\n",
|
||||
" print(f\"Unexpected error loading config from {config_file}: {e}\")\n",
|
||||
" return None\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(f\"Trading Parameters:\")\n",
|
||||
"print(f\" Configuration: {CONFIG_FILE}\")\n",
|
||||
"print(f\" Symbol A: {SYMBOL_A}\")\n",
|
||||
"print(f\" Symbol B: {SYMBOL_B}\")\n",
|
||||
"print(f\" Trading Date: {TRADING_DATE}\")\n",
|
||||
"\n",
|
||||
"# Load the specified configuration\n",
|
||||
"print(f\"\\nLoading {CONFIG_FILE} configuration using HJSON...\")\n",
|
||||
"test_config = load_config_from_file(CONFIG_FILE)\n",
|
||||
"assert test_config is not None\n",
|
||||
"BT_TEST_CONFIG = test_config\n",
|
||||
"\n",
|
||||
"if BT_TEST_CONFIG:\n",
|
||||
" print(f\"✓ Successfully loaded {BT_TEST_CONFIG['security_type']} configuration\")\n",
|
||||
" print(f\" Data directory: {BT_TEST_CONFIG['data_directory']}\")\n",
|
||||
" print(f\" Database table: {BT_TEST_CONFIG['db_table_name']}\")\n",
|
||||
" print(f\" Exchange: {BT_TEST_CONFIG['exchange_id']}\")\n",
|
||||
" print(f\" Training window: {BT_TEST_CONFIG['training_minutes']} minutes\")\n",
|
||||
" print(f\" Open threshold: {BT_TEST_CONFIG['dis-equilibrium_open_trshld']}\")\n",
|
||||
" print(f\" Close threshold: {BT_TEST_CONFIG['dis-equilibrium_close_trshld']}\")\n",
|
||||
" \n",
|
||||
" # Automatically construct data file name based on date and config type\n",
|
||||
" # if CONFIG['security_type'] == \"CRYPTO\":\n",
|
||||
" DATA_FILE = f\"{TRADING_DATE}.mktdata.ohlcv.db\"\n",
|
||||
" # elif CONFIG['security_type'] == \"EQUITY\":\n",
|
||||
" # DATA_FILE = f\"{TRADING_DATE}.alpaca_sim_md.db\"\n",
|
||||
" # else:\n",
|
||||
" # DATA_FILE = f\"{TRADING_DATE}.mktdata.db\" # Default fallback\n",
|
||||
"\n",
|
||||
" # Update CONFIG with the specific data file and instruments\n",
|
||||
" BT_TEST_CONFIG[\"datafiles\"] = [DATA_FILE]\n",
|
||||
" BT_TEST_CONFIG[\"instruments\"] = [SYMBOL_A, SYMBOL_B]\n",
|
||||
" \n",
|
||||
" print(f\"\\nData Configuration:\")\n",
|
||||
" print(f\" Data File: {DATA_FILE}\")\n",
|
||||
" print(f\" Security Type: {BT_TEST_CONFIG['security_type']}\")\n",
|
||||
" \n",
|
||||
" # Verify data file exists\n",
|
||||
" import os\n",
|
||||
" data_file_path = f\"{BT_TEST_CONFIG['data_directory']}/{DATA_FILE}\"\n",
|
||||
" if os.path.exists(data_file_path):\n",
|
||||
" print(f\" ✓ Data file found: {data_file_path}\")\n",
|
||||
" else:\n",
|
||||
" print(f\" ⚠ Data file not found: {data_file_path}\")\n",
|
||||
" print(f\" Please check if the date and file exist in the data directory\")\n",
|
||||
" \n",
|
||||
" # List available files in the data directory\n",
|
||||
" try:\n",
|
||||
" data_dir = BT_TEST_CONFIG['data_directory']\n",
|
||||
" if os.path.exists(data_dir):\n",
|
||||
" available_files = [f for f in os.listdir(data_dir) if f.endswith('.db')]\n",
|
||||
" print(f\" Available files in {data_dir}:\")\n",
|
||||
" for file in sorted(available_files)[:5]: # Show first 5 files\n",
|
||||
" print(f\" - {file}\")\n",
|
||||
" if len(available_files) > 5:\n",
|
||||
" print(f\" ... and {len(available_files)-5} more files\")\n",
|
||||
" except Exception as e:\n",
|
||||
" print(f\" Could not list files in data directory: {e}\")\n",
|
||||
"else:\n",
|
||||
" print(\"⚠ Failed to load configuration. Please check the configuration file.\")\n",
|
||||
" print(\"Available configuration files:\")\n",
|
||||
" config_dir = \"../../configuration\"\n",
|
||||
" if os.path.exists(config_dir):\n",
|
||||
" config_files = [f for f in os.listdir(config_dir) if f.endswith('.cfg')]\n",
|
||||
" for file in config_files:\n",
|
||||
" print(f\" - {file}\")\n",
|
||||
" else:\n",
|
||||
" print(f\" Configuration directory not found: {config_dir}\")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Select Trading Pair and Initialize Strategy"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Initialize Strategy\n",
|
||||
"# Trading pair and data file are now defined in the previous cell\n",
|
||||
"\n",
|
||||
"# Initialize SlidingFitStrategy\n",
|
||||
"STRATEGY = SlidingFitStrategy()\n",
|
||||
"\n",
|
||||
"print(f\"Strategy Initialization:\")\n",
|
||||
"print(f\" Selected pair: {SYMBOL_A} & {SYMBOL_B}\")\n",
|
||||
"print(f\" Data file: {DATA_FILE}\")\n",
|
||||
"print(f\" Strategy: {type(STRATEGY).__name__}\")\n",
|
||||
"print(f\"\\nStrategy characteristics:\")\n",
|
||||
"print(f\" - Sliding window training every minute\")\n",
|
||||
"print(f\" - Dynamic cointegration testing\")\n",
|
||||
"print(f\" - State-based position management\")\n",
|
||||
"print(f\" - Continuous model re-training\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Load and Prepare Market Data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Load market data\n",
|
||||
"datafile_path = f\"{BT_TEST_CONFIG['data_directory']}/{DATA_FILE}\"\n",
|
||||
"print(f\"Loading data from: {datafile_path}\")\n",
|
||||
"\n",
|
||||
"market_data_df = load_market_data(datafile_path, config=BT_TEST_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",
|
||||
"# 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=BT_TEST_CONFIG[\"price_column\"]\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(f\"\\nCreated trading pair: {pair}\")\n",
|
||||
"print(f\"Market data shape: {pair.market_data_.shape}\")\n",
|
||||
"print(f\"Column names: {pair.colnames()}\")\n",
|
||||
"\n",
|
||||
"# Calculate maximum possible iterations for sliding window\n",
|
||||
"training_minutes = BT_TEST_CONFIG[\"training_minutes\"]\n",
|
||||
"max_iterations = len(pair.market_data_) - training_minutes\n",
|
||||
"print(f\"\\nSliding window analysis:\")\n",
|
||||
"print(f\" Training window size: {training_minutes} minutes\")\n",
|
||||
"print(f\" Maximum iterations: {max_iterations}\")\n",
|
||||
"print(f\" Total analysis time: ~{max_iterations} minutes\")\n",
|
||||
"\n",
|
||||
"# Display sample data\n",
|
||||
"print(f\"\\nSample data:\")\n",
|
||||
"display(pair.market_data_.head())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Run SlidingFitStrategy with Real-Time Visualization"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Run the sliding strategy with detailed tracking\n",
|
||||
"print(f\"Running SlidingFitStrategy on {pair}...\")\n",
|
||||
"print(f\"This will process {max_iterations} minutes of data with sliding training windows.\\n\")\n",
|
||||
"\n",
|
||||
"# Initialize tracking variables\n",
|
||||
"iteration_data = []\n",
|
||||
"cointegration_history = []\n",
|
||||
"beta_history = []\n",
|
||||
"alpha_history = []\n",
|
||||
"state_history = []\n",
|
||||
"disequilibrium_history = []\n",
|
||||
"scaled_disequilibrium_history = []\n",
|
||||
"timestamp_history = []\n",
|
||||
"training_mu_history = []\n",
|
||||
"training_std_history = []\n",
|
||||
"\n",
|
||||
"# Initialize the strategy state\n",
|
||||
"pair.user_data_['state'] = PairState.INITIAL\n",
|
||||
"pair.user_data_[\"trades\"] = pd.DataFrame(columns=pd.Index(STRATEGY.TRADES_COLUMNS, dtype=str))\n",
|
||||
"pair.user_data_[\"is_cointegrated\"] = False\n",
|
||||
"\n",
|
||||
"bt_result = BacktestResult(config=BT_TEST_CONFIG)\n",
|
||||
"training_minutes = BT_TEST_CONFIG[\"training_minutes\"]\n",
|
||||
"open_threshold = BT_TEST_CONFIG[\"dis-equilibrium_open_trshld\"]\n",
|
||||
"close_threshold = BT_TEST_CONFIG[\"dis-equilibrium_close_trshld\"]\n",
|
||||
"\n",
|
||||
"# Limit iterations for demonstration (change this to max_iterations for full run)\n",
|
||||
"max_demo_iterations = min(200, max_iterations) # Process first 200 minutes\n",
|
||||
"print(f\"Processing first {max_demo_iterations} iterations for demonstration...\\n\")\n",
|
||||
"\n",
|
||||
"for curr_training_start_idx in range(max_demo_iterations):\n",
|
||||
" if curr_training_start_idx % 20 == 0:\n",
|
||||
" print(f\"Processing iteration {curr_training_start_idx}/{max_demo_iterations}...\")\n",
|
||||
"\n",
|
||||
" # Get datasets for this iteration\n",
|
||||
" pair.get_datasets(\n",
|
||||
" training_minutes=training_minutes,\n",
|
||||
" training_start_index=curr_training_start_idx,\n",
|
||||
" testing_size=1\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" if len(pair.training_df_) < training_minutes:\n",
|
||||
" print(f\"Iteration {curr_training_start_idx}: Not enough training data. Stopping.\")\n",
|
||||
" break\n",
|
||||
"\n",
|
||||
" # Record timestamp for this iteration\n",
|
||||
" current_timestamp = pair.testing_df_['tstamp'].iloc[0] if len(pair.testing_df_) > 0 else None\n",
|
||||
" timestamp_history.append(current_timestamp)\n",
|
||||
"\n",
|
||||
" # Train and test cointegration\n",
|
||||
" try:\n",
|
||||
" is_cointegrated = pair.train_pair()\n",
|
||||
" cointegration_history.append(is_cointegrated)\n",
|
||||
"\n",
|
||||
" if is_cointegrated:\n",
|
||||
" # Record model parameters\n",
|
||||
" beta_history.append(pair.vecm_fit_.beta.flatten())\n",
|
||||
" alpha_history.append(pair.vecm_fit_.alpha.flatten())\n",
|
||||
" training_mu_history.append(pair.training_mu_)\n",
|
||||
" training_std_history.append(pair.training_std_)\n",
|
||||
"\n",
|
||||
" # Generate prediction for current minute\n",
|
||||
" pair.predict()\n",
|
||||
"\n",
|
||||
" if len(pair.predicted_df_) > 0:\n",
|
||||
" current_disequilibrium = pair.predicted_df_['disequilibrium'].iloc[0]\n",
|
||||
" current_scaled_disequilibrium = pair.predicted_df_['scaled_disequilibrium'].iloc[0]\n",
|
||||
" disequilibrium_history.append(current_disequilibrium)\n",
|
||||
" scaled_disequilibrium_history.append(current_scaled_disequilibrium)\n",
|
||||
" else:\n",
|
||||
" disequilibrium_history.append(np.nan)\n",
|
||||
" scaled_disequilibrium_history.append(np.nan)\n",
|
||||
" else:\n",
|
||||
" # No cointegration\n",
|
||||
" beta_history.append(None)\n",
|
||||
" alpha_history.append(None)\n",
|
||||
" training_mu_history.append(np.nan)\n",
|
||||
" training_std_history.append(np.nan)\n",
|
||||
" disequilibrium_history.append(np.nan)\n",
|
||||
" scaled_disequilibrium_history.append(np.nan)\n",
|
||||
"\n",
|
||||
" except Exception as e:\n",
|
||||
" print(f\"Iteration {curr_training_start_idx}: Training failed: {str(e)}\")\n",
|
||||
" cointegration_history.append(False)\n",
|
||||
" beta_history.append(None)\n",
|
||||
" alpha_history.append(None)\n",
|
||||
" training_mu_history.append(np.nan)\n",
|
||||
" training_std_history.append(np.nan)\n",
|
||||
" disequilibrium_history.append(np.nan)\n",
|
||||
" scaled_disequilibrium_history.append(np.nan)\n",
|
||||
"\n",
|
||||
" # Record current state\n",
|
||||
" current_state = pair.user_data_.get('state', PairState.INITIAL)\n",
|
||||
" state_history.append(current_state)\n",
|
||||
"\n",
|
||||
"print(f\"\\nCompleted {len(cointegration_history)} iterations\")\n",
|
||||
"print(f\"Cointegration rate: {sum(cointegration_history)/len(cointegration_history)*100:.1f}%\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Visualize Sliding Window Results"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Create comprehensive visualization of sliding window results\n",
|
||||
"fig, axes = plt.subplots(6, 1, figsize=(18, 24))\n",
|
||||
"\n",
|
||||
"# Filter valid timestamps\n",
|
||||
"valid_timestamps = [ts for ts in timestamp_history if ts is not None]\n",
|
||||
"n_points = len(valid_timestamps)\n",
|
||||
"\n",
|
||||
"if n_points == 0:\n",
|
||||
" print(\"No valid data points to visualize\")\n",
|
||||
"else:\n",
|
||||
" # 1. Cointegration Status Over Time\n",
|
||||
" cointegration_values = [1 if coint else 0 for coint in cointegration_history[:n_points]]\n",
|
||||
" axes[0].plot(valid_timestamps, cointegration_values, 'o-', alpha=0.7, markersize=3)\n",
|
||||
" axes[0].fill_between(valid_timestamps, cointegration_values, alpha=0.3)\n",
|
||||
" axes[0].set_title('Cointegration Status Over Time (1=Cointegrated, 0=Not Cointegrated)')\n",
|
||||
" axes[0].set_ylabel('Cointegrated')\n",
|
||||
" axes[0].set_ylim(-0.1, 1.1)\n",
|
||||
" axes[0].grid(True)\n",
|
||||
"\n",
|
||||
" # 2. Beta Coefficients Evolution\n",
|
||||
" valid_betas = []\n",
|
||||
" beta_timestamps = []\n",
|
||||
" for i, beta in enumerate(beta_history[:n_points]):\n",
|
||||
" if beta is not None and i < len(valid_timestamps):\n",
|
||||
" valid_betas.append(beta)\n",
|
||||
" beta_timestamps.append(valid_timestamps[i])\n",
|
||||
"\n",
|
||||
" if valid_betas:\n",
|
||||
" beta_array = np.array(valid_betas)\n",
|
||||
" axes[1].plot(beta_timestamps, beta_array[:, 1], 'o-', alpha=0.7, markersize=2,\n",
|
||||
" label='Beta[1]', color='red')\n",
|
||||
" axes[1].set_title('VECM Beta[1] Coefficient Evolution (Beta[0] = 1.0 by normalization)')\n",
|
||||
" axes[1].set_ylabel('Beta[1] Value')\n",
|
||||
" axes[1].legend()\n",
|
||||
" axes[1].grid(True)\n",
|
||||
"\n",
|
||||
" # 3. Training Mean and Std Evolution\n",
|
||||
" valid_mu = [mu for mu in training_mu_history[:n_points] if not np.isnan(mu)]\n",
|
||||
" valid_std = [std for std in training_std_history[:n_points] if not np.isnan(std)]\n",
|
||||
" mu_timestamps = [valid_timestamps[i] for i, mu in enumerate(training_mu_history[:n_points]) if not np.isnan(mu)]\n",
|
||||
"\n",
|
||||
" if valid_mu:\n",
|
||||
" axes[2].plot(mu_timestamps, valid_mu, 'b-', alpha=0.7, label='Training Mean', linewidth=1)\n",
|
||||
" ax2_twin = axes[2].twinx()\n",
|
||||
" ax2_twin.plot(mu_timestamps, valid_std, 'r-', alpha=0.7, label='Training Std', linewidth=1)\n",
|
||||
" axes[2].set_title('Training Dis-equilibrium Statistics Evolution')\n",
|
||||
" axes[2].set_ylabel('Mean', color='b')\n",
|
||||
" ax2_twin.set_ylabel('Std', color='r')\n",
|
||||
" axes[2].grid(True)\n",
|
||||
" axes[2].legend(loc='upper left')\n",
|
||||
" ax2_twin.legend(loc='upper right')\n",
|
||||
"\n",
|
||||
" # 4. Raw Dis-equilibrium Over Time\n",
|
||||
" valid_diseq = [diseq for diseq in disequilibrium_history[:n_points] if not np.isnan(diseq)]\n",
|
||||
" diseq_timestamps = [valid_timestamps[i] for i, diseq in enumerate(disequilibrium_history[:n_points]) if not np.isnan(diseq)]\n",
|
||||
"\n",
|
||||
" if valid_diseq:\n",
|
||||
" axes[3].plot(diseq_timestamps, valid_diseq, 'g-', alpha=0.7, linewidth=1)\n",
|
||||
" # Add rolling mean\n",
|
||||
" if len(valid_diseq) > 10:\n",
|
||||
" rolling_mean = pd.Series(valid_diseq).rolling(window=10, min_periods=1).mean()\n",
|
||||
" axes[3].plot(diseq_timestamps, rolling_mean, 'r-', alpha=0.8, linewidth=2, label='10-period MA')\n",
|
||||
" axes[3].legend()\n",
|
||||
" axes[3].set_title('Raw Dis-equilibrium Over Time')\n",
|
||||
" axes[3].set_ylabel('Dis-equilibrium')\n",
|
||||
" axes[3].grid(True)\n",
|
||||
"\n",
|
||||
" # 5. Scaled Dis-equilibrium with Thresholds\n",
|
||||
" valid_scaled_diseq = [diseq for diseq in scaled_disequilibrium_history[:n_points] if not np.isnan(diseq)]\n",
|
||||
" scaled_diseq_timestamps = [valid_timestamps[i] for i, diseq in enumerate(scaled_disequilibrium_history[:n_points]) if not np.isnan(diseq)]\n",
|
||||
"\n",
|
||||
" if valid_scaled_diseq:\n",
|
||||
" axes[4].plot(scaled_diseq_timestamps, valid_scaled_diseq, 'purple', alpha=0.7, linewidth=1)\n",
|
||||
" axes[4].axhline(y=open_threshold, color='red', linestyle='--', alpha=0.8,\n",
|
||||
" label=f'Open Threshold ({open_threshold})')\n",
|
||||
" axes[4].axhline(y=close_threshold, color='blue', linestyle='--', alpha=0.8,\n",
|
||||
" label=f'Close Threshold ({close_threshold})')\n",
|
||||
" axes[4].axhline(y=0, color='black', linestyle='-', alpha=0.5, linewidth=0.5)\n",
|
||||
" axes[4].set_title('Scaled Dis-equilibrium with Trading Thresholds')\n",
|
||||
" axes[4].set_ylabel('Scaled Dis-equilibrium')\n",
|
||||
" axes[4].legend()\n",
|
||||
" axes[4].grid(True)\n",
|
||||
"\n",
|
||||
" # 6. Price Data with Training Windows\n",
|
||||
" # Show original price data with indication of training windows\n",
|
||||
" colname_a, colname_b = pair.colnames()\n",
|
||||
" price_data = pair.market_data_[:n_points + training_minutes].copy()\n",
|
||||
"\n",
|
||||
" axes[5].plot(price_data['tstamp'], price_data[colname_a], alpha=0.7, label=f'{SYMBOL_A}', linewidth=1)\n",
|
||||
" axes[5].plot(price_data['tstamp'], price_data[colname_b], alpha=0.7, label=f'{SYMBOL_B}', linewidth=1)\n",
|
||||
"\n",
|
||||
" # Highlight training windows\n",
|
||||
" for i in range(0, min(n_points, 10), max(1, n_points//20)): # Show every 20th window\n",
|
||||
" start_idx = i\n",
|
||||
" end_idx = i + training_minutes\n",
|
||||
" if end_idx < len(price_data):\n",
|
||||
" window_data = price_data.iloc[start_idx:end_idx]\n",
|
||||
" axes[5].axvspan(window_data['tstamp'].iloc[0], window_data['tstamp'].iloc[-1],\n",
|
||||
" alpha=0.1, color='gray')\n",
|
||||
"\n",
|
||||
" axes[5].set_title(f'Price Data with Training Windows (Gray bands show some training windows)')\n",
|
||||
" axes[5].set_ylabel('Price')\n",
|
||||
" axes[5].set_xlabel('Time')\n",
|
||||
" axes[5].legend()\n",
|
||||
" axes[5].grid(True)\n",
|
||||
"\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"# Print summary statistics\n",
|
||||
"print(f\"\\n\" + \"=\"*80)\n",
|
||||
"print(f\"SLIDING WINDOW ANALYSIS SUMMARY\")\n",
|
||||
"print(f\"=\"*80)\n",
|
||||
"print(f\"Total iterations processed: {n_points}\")\n",
|
||||
"print(f\"Cointegration episodes: {sum(cointegration_history[:n_points])}\")\n",
|
||||
"print(f\"Cointegration rate: {sum(cointegration_history[:n_points])/n_points*100:.1f}%\")\n",
|
||||
"if valid_betas:\n",
|
||||
" print(f\"Beta coefficient stability: Std = {np.std(beta_array, axis=0)}\")\n",
|
||||
"if valid_scaled_diseq:\n",
|
||||
" threshold_breaches = sum(1 for x in valid_scaled_diseq if abs(x) > open_threshold)\n",
|
||||
" print(f\"Open threshold breaches: {threshold_breaches} ({threshold_breaches/len(valid_scaled_diseq)*100:.1f}%)\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Analyze Training Window Evolution"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Detailed analysis of how training windows evolve\n",
|
||||
"print(\"TRAINING WINDOW EVOLUTION ANALYSIS\")\n",
|
||||
"print(\"=\" * 50)\n",
|
||||
"\n",
|
||||
"# Analyze cointegration stability\n",
|
||||
"if len(cointegration_history) > 1:\n",
|
||||
" # Find cointegration change points\n",
|
||||
" change_points = []\n",
|
||||
" for i in range(1, len(cointegration_history)):\n",
|
||||
" if cointegration_history[i] != cointegration_history[i - 1]:\n",
|
||||
" change_points.append(\n",
|
||||
" (\n",
|
||||
" i,\n",
|
||||
" cointegration_history[i],\n",
|
||||
" valid_timestamps[i] if i < len(valid_timestamps) else None,\n",
|
||||
" )\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" print(f\"\\nCointegration Change Points:\")\n",
|
||||
" if change_points:\n",
|
||||
" for idx, status, timestamp in change_points[:10]: # Show first 10\n",
|
||||
" status_str = \"GAINED\" if status else \"LOST\"\n",
|
||||
" print(f\" Iteration {idx}: {status_str} cointegration at {timestamp}\")\n",
|
||||
" if len(change_points) > 10:\n",
|
||||
" print(f\" ... and {len(change_points)-10} more changes\")\n",
|
||||
" else:\n",
|
||||
" print(f\" No cointegration changes detected\")\n",
|
||||
"\n",
|
||||
"# Analyze beta stability when cointegrated\n",
|
||||
"if valid_betas and len(valid_betas) > 10:\n",
|
||||
" beta_df = pd.DataFrame(\n",
|
||||
" valid_betas,\n",
|
||||
" columns=pd.Index([f\"Beta_{i}\" for i in range(len(valid_betas[0]))], dtype=str),\n",
|
||||
" )\n",
|
||||
" beta_df[\"timestamp\"] = beta_timestamps\n",
|
||||
"\n",
|
||||
" print(f\"\\nBeta Coefficient Analysis:\")\n",
|
||||
" print(f\" Number of valid beta estimates: {len(valid_betas)}\")\n",
|
||||
" print(f\" Beta statistics:\")\n",
|
||||
" for col in beta_df.columns[:-1]: # Exclude timestamp\n",
|
||||
" print(\n",
|
||||
" f\" {col}: Mean={beta_df[col].mean():.4f}, Std={beta_df[col].std():.4f}\"\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" # Check for beta regime changes\n",
|
||||
" beta_changes = []\n",
|
||||
" threshold = 0.1 # 10% change threshold\n",
|
||||
" for i in range(1, len(valid_betas)):\n",
|
||||
" if np.any(\n",
|
||||
" np.abs(np.array(valid_betas[i]) - np.array(valid_betas[i - 1])) > threshold\n",
|
||||
" ):\n",
|
||||
" beta_changes.append(i)\n",
|
||||
"\n",
|
||||
" print(f\" Significant beta changes (>{threshold*100}%): {len(beta_changes)}\")\n",
|
||||
" if beta_changes:\n",
|
||||
" print(\n",
|
||||
" f\" Change frequency: {len(beta_changes)/len(valid_betas)*100:.1f}% of cointegrated periods\"\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
"# Analyze dis-equilibrium characteristics\n",
|
||||
"if valid_scaled_diseq:\n",
|
||||
" scaled_diseq_series = pd.Series(valid_scaled_diseq)\n",
|
||||
"\n",
|
||||
" print(f\"\\nDis-equilibrium Analysis:\")\n",
|
||||
" print(f\" Mean: {scaled_diseq_series.mean():.4f}\")\n",
|
||||
" print(f\" Std: {scaled_diseq_series.std():.4f}\")\n",
|
||||
" print(f\" Min: {scaled_diseq_series.min():.4f}\")\n",
|
||||
" print(f\" Max: {scaled_diseq_series.max():.4f}\")\n",
|
||||
"\n",
|
||||
" # Threshold analysis\n",
|
||||
" open_breaches = sum(1 for x in valid_scaled_diseq if abs(x) >= open_threshold)\n",
|
||||
" close_opportunities = sum(\n",
|
||||
" 1 for x in valid_scaled_diseq if abs(x) <= close_threshold\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" print(\n",
|
||||
" f\" Open threshold breaches: {open_breaches} ({open_breaches/len(valid_scaled_diseq)*100:.1f}%)\"\n",
|
||||
" )\n",
|
||||
" print(\n",
|
||||
" f\" Close opportunities: {close_opportunities} ({close_opportunities/len(valid_scaled_diseq)*100:.1f}%)\"\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" # Mean reversion analysis\n",
|
||||
" zero_crossings = 0\n",
|
||||
" for i in range(1, len(valid_scaled_diseq)):\n",
|
||||
" if (valid_scaled_diseq[i - 1] * valid_scaled_diseq[i]) < 0: # Sign change\n",
|
||||
" zero_crossings += 1\n",
|
||||
"\n",
|
||||
" print(f\" Zero crossings (mean reversion events): {zero_crossings}\")\n",
|
||||
" if zero_crossings > 0:\n",
|
||||
" print(\n",
|
||||
" f\" Average time between mean reversions: {len(valid_scaled_diseq)/zero_crossings:.1f} minutes\"\n",
|
||||
" )"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Run Complete Strategy (Optional)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Optional: Run the complete strategy to generate actual trades\n",
|
||||
"# Warning: This may take several minutes depending on data size\n",
|
||||
"\n",
|
||||
"RUN_COMPLETE_STRATEGY = False # Set to True to run full strategy\n",
|
||||
"\n",
|
||||
"if RUN_COMPLETE_STRATEGY:\n",
|
||||
" print(\"Running complete SlidingFitStrategy...\")\n",
|
||||
" print(\"This may take several minutes...\")\n",
|
||||
"\n",
|
||||
" # Reset strategy state\n",
|
||||
" STRATEGY.curr_training_start_idx_ = 0\n",
|
||||
"\n",
|
||||
" # Create new pair and result objects\n",
|
||||
" pair_full = TradingPair(\n",
|
||||
" market_data=market_data_df,\n",
|
||||
" symbol_a=SYMBOL_A,\n",
|
||||
" symbol_b=SYMBOL_B,\n",
|
||||
" price_column=BT_TEST_CONFIG[\"price_column\"]\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" bt_result_full = BacktestResult(config=BT_TEST_CONFIG)\n",
|
||||
"\n",
|
||||
" # Run strategy\n",
|
||||
" pair_trades = STRATEGY.run_pair(config=BT_TEST_CONFIG, pair=pair_full, bt_result=bt_result_full)\n",
|
||||
"\n",
|
||||
" if pair_trades is not None and len(pair_trades) > 0:\n",
|
||||
" print(f\"\\nGenerated {len(pair_trades)} trading signals:\")\n",
|
||||
" display(pair_trades)\n",
|
||||
"\n",
|
||||
" # Analyze trades\n",
|
||||
" trade_times = pair_trades['time'].unique()\n",
|
||||
" print(f\"\\nTrade Analysis:\")\n",
|
||||
" print(f\" Unique trade times: {len(trade_times)}\")\n",
|
||||
" print(f\" Trade frequency: {len(trade_times)/max_iterations*100:.2f}% of total periods\")\n",
|
||||
"\n",
|
||||
" # Group trades by time\n",
|
||||
" for trade_time in trade_times[:5]: # Show first 5 trade 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} \"\n",
|
||||
" f\"(dis-eq: {trade['scaled_disequilibrium']:.2f})\")\n",
|
||||
" else:\n",
|
||||
" print(\"\\nNo trading signals generated\")\n",
|
||||
" print(\"Possible reasons:\")\n",
|
||||
" print(\" - Insufficient cointegration periods\")\n",
|
||||
" print(\" - Dis-equilibrium never exceeded thresholds\")\n",
|
||||
" print(\" - Strategy-specific conditions not met\")\n",
|
||||
"else:\n",
|
||||
" print(\"Complete strategy execution is disabled.\")\n",
|
||||
" print(\"Set RUN_COMPLETE_STRATEGY = True to run the full strategy.\")\n",
|
||||
" print(\"Note: This may take several minutes depending on your data size.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Interactive Parameter Analysis"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Interactive analysis for parameter optimization\n",
|
||||
"print(\"PARAMETER SENSITIVITY ANALYSIS\")\n",
|
||||
"print(\"=\"*40)\n",
|
||||
"\n",
|
||||
"print(f\"Current parameters:\")\n",
|
||||
"print(f\" Training window: {BT_TEST_CONFIG['training_minutes']} minutes\")\n",
|
||||
"print(f\" Open threshold: {BT_TEST_CONFIG['dis-equilibrium_open_trshld']}\")\n",
|
||||
"print(f\" Close threshold: {BT_TEST_CONFIG['dis-equilibrium_close_trshld']}\")\n",
|
||||
"\n",
|
||||
"# Recommendations based on observed data\n",
|
||||
"if valid_scaled_diseq:\n",
|
||||
" diseq_stats = pd.Series(valid_scaled_diseq).describe()\n",
|
||||
" print(f\"\\nObserved scaled dis-equilibrium statistics:\")\n",
|
||||
" print(f\" 75th percentile: {diseq_stats['75%']:.2f}\")\n",
|
||||
" print(f\" 95th percentile: {np.percentile(valid_scaled_diseq, 95):.2f}\")\n",
|
||||
" print(f\" 99th percentile: {np.percentile(valid_scaled_diseq, 99):.2f}\")\n",
|
||||
"\n",
|
||||
" # Suggest optimal thresholds\n",
|
||||
" suggested_open = np.percentile(np.abs(valid_scaled_diseq), 85)\n",
|
||||
" suggested_close = np.percentile(np.abs(valid_scaled_diseq), 30)\n",
|
||||
"\n",
|
||||
" print(f\"\\nSuggested threshold optimization:\")\n",
|
||||
" print(f\" Suggested open threshold: {suggested_open:.2f} (85th percentile)\")\n",
|
||||
" print(f\" Suggested close threshold: {suggested_close:.2f} (30th percentile)\")\n",
|
||||
"\n",
|
||||
" if suggested_open != open_threshold or suggested_close != close_threshold:\n",
|
||||
" print(f\"\\nTo test these parameters, modify the CONFIG dictionary:\")\n",
|
||||
" print(f\" CONFIG['dis-equilibrium_open_trshld'] = {suggested_open:.2f}\")\n",
|
||||
" print(f\" CONFIG['dis-equilibrium_close_trshld'] = {suggested_close:.2f}\")\n",
|
||||
"\n",
|
||||
"# Training window recommendations\n",
|
||||
"if len(cointegration_history) > 0:\n",
|
||||
" cointegration_rate = sum(cointegration_history)/len(cointegration_history)\n",
|
||||
" print(f\"\\nTraining window analysis:\")\n",
|
||||
" print(f\" Current cointegration rate: {cointegration_rate*100:.1f}%\")\n",
|
||||
"\n",
|
||||
" if cointegration_rate < 0.3:\n",
|
||||
" print(f\" Recommendation: Consider increasing training window (current: {training_minutes})\")\n",
|
||||
" print(f\" Suggested: {int(training_minutes * 1.5)} minutes\")\n",
|
||||
" elif cointegration_rate > 0.8:\n",
|
||||
" print(f\" Recommendation: Consider decreasing training window for more responsive model\")\n",
|
||||
" print(f\" Suggested: {int(training_minutes * 0.75)} minutes\")\n",
|
||||
" else:\n",
|
||||
" print(f\" Current training window appears appropriate\")\n",
|
||||
"\n",
|
||||
"print(f\"\\nTo re-run analysis with different parameters:\")\n",
|
||||
"print(f\"1. Modify the CONFIG dictionary above\")\n",
|
||||
"print(f\"2. Re-run from the 'Run SlidingFitStrategy' cell\")\n",
|
||||
"print(f\"3. Compare results with current analysis\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Summary and Conclusions\n",
|
||||
"\n",
|
||||
"This notebook demonstrates the SlidingFitStrategy's dynamic approach to pairs trading.\n",
|
||||
"Key insights from the sliding window analysis:\n",
|
||||
"\n",
|
||||
"1. **Cointegration Stability**: How often the pair maintains cointegration\n",
|
||||
"2. **Model Parameter Evolution**: How VECM coefficients change over time\n",
|
||||
"3. **Threshold Effectiveness**: How well current thresholds capture trading opportunities\n",
|
||||
"4. **Mean Reversion Patterns**: Frequency and timing of dis-equilibrium corrections\n",
|
||||
"\n",
|
||||
"The sliding approach allows for:\n",
|
||||
"- **Adaptive modeling**: Responds to changing market conditions\n",
|
||||
"- **Dynamic thresholding**: Can be optimized based on observed patterns\n",
|
||||
"- **Real-time monitoring**: Provides continuous assessment of pair relationships\n",
|
||||
"- **Risk management**: Early detection of cointegration breakdown"
|
||||
]
|
||||
}
|
||||
],
|
||||
"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
|
||||
}
|
||||
@@ -1,710 +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": [
|
||||
"### \ud83c\udfaf 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",
|
||||
"### \ud83d\ude80 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": [],
|
||||
"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 strategies import StaticFitStrategy, SlidingFitStrategy\n",
|
||||
"from tools.data_loader import load_market_data\n",
|
||||
"from tools.trading_pair import TradingPair\n",
|
||||
"from 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": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"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": [],
|
||||
"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",
|
||||
"STRATEGY = StaticFitStrategy()\n",
|
||||
"\n",
|
||||
"print(f\"Selected pair: {SYMBOL_A} & {SYMBOL_B}\")\n",
|
||||
"print(f\"Data file: {DATA_FILE}\")\n",
|
||||
"print(f\"Strategy: {type(STRATEGY).__name__}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Load Market Data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"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 = STRATEGY.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(STRATEGY).__name__}\")\n",
|
||||
"print(f\"Data file: {DATA_FILE}\")\n",
|
||||
"print(f\"Training period: {training_minutes} minutes\")\n",
|
||||
"\n",
|
||||
"print(f\"\\nCointegration Status: {'\u2713 COINTEGRATED' if is_cointegrated else '\u2717 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.3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 4
|
||||
}
|
||||
-656
@@ -1,656 +0,0 @@
|
||||
from typing import Any, Dict, List
|
||||
import pandas as pd
|
||||
import sqlite3
|
||||
import os
|
||||
from datetime import datetime, date
|
||||
|
||||
|
||||
# Recommended replacement adapters and converters for Python 3.12+
|
||||
# From: https://docs.python.org/3/library/sqlite3.html#sqlite3-adapter-converter-recipes
|
||||
def adapt_date_iso(val):
|
||||
"""Adapt datetime.date to ISO 8601 date."""
|
||||
return val.isoformat()
|
||||
|
||||
def adapt_datetime_iso(val):
|
||||
"""Adapt datetime.datetime to timezone-naive ISO 8601 date."""
|
||||
return val.isoformat()
|
||||
|
||||
def convert_date(val):
|
||||
"""Convert ISO 8601 date to datetime.date object."""
|
||||
return datetime.fromisoformat(val.decode()).date()
|
||||
|
||||
def convert_datetime(val):
|
||||
"""Convert ISO 8601 datetime to datetime.datetime object."""
|
||||
return datetime.fromisoformat(val.decode())
|
||||
|
||||
# Register the adapters and converters
|
||||
sqlite3.register_adapter(date, adapt_date_iso)
|
||||
sqlite3.register_adapter(datetime, adapt_datetime_iso)
|
||||
sqlite3.register_converter("date", convert_date)
|
||||
sqlite3.register_converter("datetime", convert_datetime)
|
||||
|
||||
|
||||
def create_result_database(db_path: str) -> None:
|
||||
"""
|
||||
Create the SQLite database and required tables if they don't exist.
|
||||
"""
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.cursor()
|
||||
|
||||
# Create the pt_bt_results table for completed trades
|
||||
cursor.execute('''
|
||||
CREATE TABLE IF NOT EXISTS pt_bt_results (
|
||||
date DATE,
|
||||
pair TEXT,
|
||||
symbol TEXT,
|
||||
open_time DATETIME,
|
||||
open_side TEXT,
|
||||
open_price REAL,
|
||||
open_quantity INTEGER,
|
||||
open_disequilibrium REAL,
|
||||
close_time DATETIME,
|
||||
close_side TEXT,
|
||||
close_price REAL,
|
||||
close_quantity INTEGER,
|
||||
close_disequilibrium REAL,
|
||||
symbol_return REAL,
|
||||
pair_return REAL
|
||||
)
|
||||
''')
|
||||
cursor.execute("DELETE FROM pt_bt_results;")
|
||||
|
||||
# Create the outstanding_positions table for open positions
|
||||
cursor.execute('''
|
||||
CREATE TABLE IF NOT EXISTS outstanding_positions (
|
||||
date DATE,
|
||||
pair TEXT,
|
||||
symbol TEXT,
|
||||
position_quantity REAL,
|
||||
last_price REAL,
|
||||
unrealized_return REAL,
|
||||
open_price REAL,
|
||||
open_side TEXT
|
||||
)
|
||||
''')
|
||||
cursor.execute("DELETE FROM outstanding_positions;")
|
||||
|
||||
# Create the config table for storing configuration JSON for reference
|
||||
cursor.execute('''
|
||||
CREATE TABLE IF NOT EXISTS config (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
run_timestamp DATETIME,
|
||||
config_file_path TEXT,
|
||||
config_json TEXT,
|
||||
strategy_class TEXT,
|
||||
datafiles TEXT,
|
||||
instruments TEXT
|
||||
)
|
||||
''')
|
||||
cursor.execute("DELETE FROM config;")
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error creating result database: {str(e)}")
|
||||
raise
|
||||
|
||||
|
||||
def store_config_in_database(db_path: str, config_file_path: str, config: Dict, strategy_class: str, datafiles: List[str], instruments: List[str]) -> None:
|
||||
"""
|
||||
Store configuration information in the database for reference.
|
||||
"""
|
||||
import json
|
||||
|
||||
if db_path.upper() == "NONE":
|
||||
return
|
||||
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.cursor()
|
||||
|
||||
# Convert config to JSON string
|
||||
config_json = json.dumps(config, indent=2, default=str)
|
||||
|
||||
# Convert lists to comma-separated strings for storage
|
||||
datafiles_str = ', '.join(datafiles)
|
||||
instruments_str = ', '.join(instruments)
|
||||
|
||||
# Insert configuration record
|
||||
cursor.execute('''
|
||||
INSERT INTO config (
|
||||
run_timestamp, config_file_path, config_json, strategy_class, datafiles, instruments
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
''', (
|
||||
datetime.now(),
|
||||
config_file_path,
|
||||
config_json,
|
||||
strategy_class,
|
||||
datafiles_str,
|
||||
instruments_str
|
||||
))
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
print(f"Configuration stored in database")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error storing configuration in database: {str(e)}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
|
||||
def store_results_in_database(db_path: str, datafile: str, bt_result: 'BacktestResult') -> None:
|
||||
"""
|
||||
Store backtest results in the SQLite database.
|
||||
"""
|
||||
if db_path.upper() == "NONE":
|
||||
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
|
||||
|
||||
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 to handle backtest results, trades tracking, PnL calculations, and reporting.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
self.config = config
|
||||
self.trades: Dict[str, Dict[str, Any]] = {}
|
||||
self.total_realized_pnl = 0.0
|
||||
self.outstanding_positions: List[Dict[str, Any]] = []
|
||||
|
||||
def add_trade(self, pair_nm, symbol, action, price, disequilibrium=None, scaled_disequilibrium=None, timestamp=None):
|
||||
"""Add a trade to the results tracking."""
|
||||
pair_nm = str(pair_nm)
|
||||
|
||||
if pair_nm not in self.trades:
|
||||
self.trades[pair_nm] = {symbol: []}
|
||||
if symbol not in self.trades[pair_nm]:
|
||||
self.trades[pair_nm][symbol] = []
|
||||
self.trades[pair_nm][symbol].append((action, price, disequilibrium, scaled_disequilibrium, timestamp))
|
||||
|
||||
def add_outstanding_position(self, position: Dict[str, Any]):
|
||||
"""Add an outstanding position to tracking."""
|
||||
self.outstanding_positions.append(position)
|
||||
|
||||
def add_realized_pnl(self, realized_pnl: float):
|
||||
"""Add realized PnL to the total."""
|
||||
self.total_realized_pnl += realized_pnl
|
||||
|
||||
def get_total_realized_pnl(self) -> float:
|
||||
"""Get total realized PnL."""
|
||||
return self.total_realized_pnl
|
||||
|
||||
def get_outstanding_positions(self) -> List[Dict[str, Any]]:
|
||||
"""Get all outstanding positions."""
|
||||
return self.outstanding_positions
|
||||
|
||||
def get_trades(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""Get all trades."""
|
||||
return self.trades
|
||||
|
||||
def clear_trades(self):
|
||||
"""Clear all trades (used when processing new files)."""
|
||||
self.trades.clear()
|
||||
|
||||
def collect_single_day_results(self, result):
|
||||
"""Collect and process single day trading results."""
|
||||
if result is None:
|
||||
return
|
||||
|
||||
print("\n -------------- Suggested Trades ")
|
||||
print(result)
|
||||
|
||||
for row in result.itertuples():
|
||||
action = row.action
|
||||
symbol = row.symbol
|
||||
price = row.price
|
||||
disequilibrium = getattr(row, 'disequilibrium', None)
|
||||
scaled_disequilibrium = getattr(row, 'scaled_disequilibrium', None)
|
||||
timestamp = getattr(row, 'time', None)
|
||||
self.add_trade(
|
||||
pair_nm=row.pair, action=action, symbol=symbol, price=price,
|
||||
disequilibrium=disequilibrium, scaled_disequilibrium=scaled_disequilibrium,
|
||||
timestamp=timestamp
|
||||
)
|
||||
|
||||
def print_single_day_results(self):
|
||||
"""Print single day results summary."""
|
||||
for pair, symbols in self.trades.items():
|
||||
print(f"\n--- {pair} ---")
|
||||
for symbol, trades in symbols.items():
|
||||
for trade_data in trades:
|
||||
if len(trade_data) >= 2:
|
||||
side, price = trade_data[:2]
|
||||
print(f"{symbol} {side} at ${price}")
|
||||
|
||||
def print_results_summary(self, all_results):
|
||||
"""Print summary of all processed files."""
|
||||
print("\n====== Summary of All Processed Files ======")
|
||||
for filename, data in all_results.items():
|
||||
trade_count = sum(
|
||||
len(trades)
|
||||
for symbol_trades in data["trades"].values()
|
||||
for trades in symbol_trades.values()
|
||||
)
|
||||
print(f"{filename}: {trade_count} trades")
|
||||
|
||||
def calculate_returns(self, all_results: Dict):
|
||||
"""Calculate and print returns by day and pair."""
|
||||
print("\n====== Returns By Day and Pair ======")
|
||||
|
||||
for filename, data in all_results.items():
|
||||
day_return = 0
|
||||
print(f"\n--- {filename} ---")
|
||||
|
||||
# Process each pair
|
||||
for pair, symbols in data["trades"].items():
|
||||
pair_return = 0
|
||||
pair_trades = []
|
||||
|
||||
# Calculate individual symbol returns in the pair
|
||||
for symbol, trades in symbols.items():
|
||||
if len(trades) >= 2: # Need at least entry and exit
|
||||
# Get entry and exit trades - handle both old and new tuple formats
|
||||
if len(trades[0]) == 2: # Old format: (action, price)
|
||||
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_return = 0
|
||||
if entry_action == "BUY" and exit_action == "SELL":
|
||||
# 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,
|
||||
entry_action,
|
||||
entry_price,
|
||||
exit_action,
|
||||
exit_price,
|
||||
symbol_return,
|
||||
open_scaled_disequilibrium,
|
||||
close_scaled_disequilibrium,
|
||||
)
|
||||
)
|
||||
pair_return += symbol_return
|
||||
|
||||
# Print pair returns with disequilibrium information
|
||||
if pair_trades:
|
||||
print(f" {pair}:")
|
||||
for (
|
||||
symbol,
|
||||
entry_action,
|
||||
entry_price,
|
||||
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(
|
||||
f" {symbol}: {entry_action} @ ${entry_price:.2f}, {exit_action} @ ${exit_price:.2f}, Return: {symbol_return:.2f}%{disequil_info}"
|
||||
)
|
||||
print(f" Pair Total Return: {pair_return:.2f}%")
|
||||
day_return += pair_return
|
||||
|
||||
# Print day total return and add to global realized PnL
|
||||
if day_return != 0:
|
||||
print(f" Day Total Return: {day_return:.2f}%")
|
||||
self.add_realized_pnl(day_return)
|
||||
|
||||
def print_outstanding_positions(self):
|
||||
"""Print all outstanding positions with share quantities and current values."""
|
||||
if not self.get_outstanding_positions():
|
||||
print("\n====== NO OUTSTANDING POSITIONS ======")
|
||||
return
|
||||
|
||||
print(f"\n====== OUTSTANDING POSITIONS ======")
|
||||
print(
|
||||
f"{'Pair':<15}"
|
||||
f" {'Symbol':<10}"
|
||||
f" {'Side':<4}"
|
||||
f" {'Shares':<10}"
|
||||
f" {'Open $':<8}"
|
||||
f" {'Current $':<10}"
|
||||
f" {'Value $':<12}"
|
||||
f" {'Disequilibrium':<15}"
|
||||
)
|
||||
print("-" * 100)
|
||||
|
||||
total_value = 0.0
|
||||
|
||||
for pos in self.get_outstanding_positions():
|
||||
# Print position A
|
||||
print(
|
||||
f"{pos['pair']:<15}"
|
||||
f" {pos['symbol_a']:<10}"
|
||||
f" {pos['side_a']:<4}"
|
||||
f" {pos['shares_a']:<10.2f}"
|
||||
f" {pos['open_px_a']:<8.2f}"
|
||||
f" {pos['current_px_a']:<10.2f}"
|
||||
f" {pos['current_value_a']:<12.2f}"
|
||||
f" {'':<15}"
|
||||
)
|
||||
|
||||
# Print position B
|
||||
print(
|
||||
f"{'':<15}"
|
||||
f" {pos['symbol_b']:<10}"
|
||||
f" {pos['side_b']:<4}"
|
||||
f" {pos['shares_b']:<10.2f}"
|
||||
f" {pos['open_px_b']:<8.2f}"
|
||||
f" {pos['current_px_b']:<10.2f}"
|
||||
f" {pos['current_value_b']:<12.2f}"
|
||||
)
|
||||
|
||||
# Print pair totals with disequilibrium info
|
||||
print(
|
||||
f"{'':<15}"
|
||||
f" {'PAIR TOTAL':<10}"
|
||||
f" {'':<4}"
|
||||
f" {'':<10}"
|
||||
f" {'':<8}"
|
||||
f" {'':<10}"
|
||||
f" {pos['total_current_value']:<12.2f}"
|
||||
)
|
||||
|
||||
# Print disequilibrium details
|
||||
print(
|
||||
f"{'':<15}"
|
||||
f" {'DISEQUIL':<10}"
|
||||
f" {'':<4}"
|
||||
f" {'':<10}"
|
||||
f" {'':<8}"
|
||||
f" {'':<10}"
|
||||
f" Raw: {pos['current_disequilibrium']:<6.4f}"
|
||||
f" Scaled: {pos['current_scaled_disequilibrium']:<6.4f}"
|
||||
)
|
||||
|
||||
print("-" * 100)
|
||||
|
||||
total_value += pos["total_current_value"]
|
||||
|
||||
print(f"{'TOTAL OUTSTANDING VALUE':<80} ${total_value:<12.2f}")
|
||||
|
||||
def print_grand_totals(self):
|
||||
"""Print grand totals across all pairs."""
|
||||
print(f"\n====== GRAND TOTALS ACROSS ALL PAIRS ======")
|
||||
print(f"Total Realized PnL: {self.get_total_realized_pnl():.2f}%")
|
||||
|
||||
def handle_outstanding_position(self, pair, pair_result_df, last_row_index,
|
||||
open_side_a, open_side_b, open_px_a, open_px_b,
|
||||
open_tstamp):
|
||||
"""
|
||||
Handle calculation and tracking of outstanding positions when no close signal is found.
|
||||
|
||||
Args:
|
||||
pair: TradingPair object
|
||||
pair_result_df: DataFrame with pair results
|
||||
last_row_index: Index of the last row in the data
|
||||
open_side_a, open_side_b: Trading sides for symbols A and B
|
||||
open_px_a, open_px_b: Opening prices for symbols A and B
|
||||
open_tstamp: Opening timestamp
|
||||
"""
|
||||
if pair_result_df is None or pair_result_df.empty:
|
||||
return 0, 0, 0
|
||||
|
||||
last_row = pair_result_df.loc[last_row_index]
|
||||
last_tstamp = last_row["tstamp"]
|
||||
colname_a, colname_b = pair.colnames()
|
||||
last_px_a = last_row[colname_a]
|
||||
last_px_b = last_row[colname_b]
|
||||
|
||||
# Calculate share quantities based on funding per pair
|
||||
# Split funding equally between the two positions
|
||||
funding_per_position = self.config["funding_per_pair"] / 2
|
||||
shares_a = funding_per_position / open_px_a
|
||||
shares_b = funding_per_position / open_px_b
|
||||
|
||||
# Calculate current position values (shares * current price)
|
||||
current_value_a = shares_a * last_px_a
|
||||
current_value_b = shares_b * last_px_b
|
||||
total_current_value = current_value_a + current_value_b
|
||||
|
||||
# Get disequilibrium information
|
||||
current_disequilibrium = last_row["disequilibrium"]
|
||||
current_scaled_disequilibrium = last_row["scaled_disequilibrium"]
|
||||
|
||||
# Store outstanding positions
|
||||
self.add_outstanding_position(
|
||||
{
|
||||
"pair": str(pair),
|
||||
"symbol_a": pair.symbol_a_,
|
||||
"symbol_b": pair.symbol_b_,
|
||||
"side_a": open_side_a,
|
||||
"side_b": open_side_b,
|
||||
"shares_a": shares_a,
|
||||
"shares_b": shares_b,
|
||||
"open_px_a": open_px_a,
|
||||
"open_px_b": open_px_b,
|
||||
"current_px_a": last_px_a,
|
||||
"current_px_b": last_px_b,
|
||||
"current_value_a": current_value_a,
|
||||
"current_value_b": current_value_b,
|
||||
"total_current_value": total_current_value,
|
||||
"open_time": open_tstamp,
|
||||
"last_time": last_tstamp,
|
||||
"current_abs_term": current_scaled_disequilibrium,
|
||||
"current_disequilibrium": current_disequilibrium,
|
||||
"current_scaled_disequilibrium": current_scaled_disequilibrium,
|
||||
}
|
||||
)
|
||||
|
||||
# Print position details
|
||||
print(f"{pair}: NO CLOSE SIGNAL FOUND - Position held until end of session")
|
||||
print(f" Open: {open_tstamp} | Last: {last_tstamp}")
|
||||
print(f" {pair.symbol_a_}: {open_side_a} {shares_a:.2f} shares @ ${open_px_a:.2f} -> ${last_px_a:.2f} | Value: ${current_value_a:.2f}")
|
||||
print(f" {pair.symbol_b_}: {open_side_b} {shares_b:.2f} shares @ ${open_px_b:.2f} -> ${last_px_b:.2f} | Value: ${current_value_b:.2f}")
|
||||
print(f" Total Value: ${total_current_value:.2f}")
|
||||
print(f" Disequilibrium: {current_disequilibrium:.4f} | Scaled: {current_scaled_disequilibrium:.4f}")
|
||||
|
||||
return current_value_a, current_value_b, total_current_value
|
||||
@@ -1,420 +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 tools.trading_pair import TradingPair
|
||||
from results import BacktestResult
|
||||
|
||||
NanoPerMin = 1e9
|
||||
|
||||
class PairsTradingStrategy(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 StaticFitStrategy(PairsTradingStrategy):
|
||||
|
||||
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):
|
||||
pass
|
||||
|
||||
class PairState(Enum):
|
||||
INITIAL = 1
|
||||
OPEN = 2
|
||||
CLOSED = 3
|
||||
|
||||
class SlidingFitStrategy(PairsTradingStrategy):
|
||||
def __init__(self):
|
||||
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: # type: ignore
|
||||
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
|
||||
|
||||
|
||||
|
||||
@@ -1,138 +0,0 @@
|
||||
import sqlite3
|
||||
from typing import Dict, List, cast
|
||||
import pandas as pd
|
||||
|
||||
|
||||
|
||||
def load_sqlite_to_dataframe(db_path, query):
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
|
||||
df = pd.read_sql_query(query, conn)
|
||||
return df
|
||||
except sqlite3.Error as excpt:
|
||||
print(f"SQLite error: {excpt}")
|
||||
raise
|
||||
except Exception as excpt:
|
||||
print(f"Error: {excpt}")
|
||||
raise Exception() from excpt
|
||||
finally:
|
||||
if "conn" in locals():
|
||||
conn.close()
|
||||
|
||||
|
||||
def convert_time_to_UTC(value: str, timezone: str) -> str:
|
||||
|
||||
from zoneinfo import ZoneInfo
|
||||
from datetime import datetime
|
||||
|
||||
# Parse it to naive datetime object
|
||||
local_dt = datetime.strptime(value, "%Y-%m-%d %H:%M:%S")
|
||||
|
||||
zinfo = ZoneInfo(timezone)
|
||||
result: datetime = local_dt.replace(tzinfo=zinfo).astimezone(ZoneInfo("UTC"))
|
||||
|
||||
return result.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
def load_market_data(datafile: str, config: Dict) -> pd.DataFrame:
|
||||
from tools.data_loader import load_sqlite_to_dataframe
|
||||
|
||||
instrument_ids = [
|
||||
'"' + config["instrument_id_pfx"] + instrument + '"'
|
||||
for instrument in config["instruments"]
|
||||
]
|
||||
security_type = config["security_type"]
|
||||
exchange_id = config["exchange_id"]
|
||||
|
||||
query = "select"
|
||||
if security_type == "CRYPTO":
|
||||
query += " strftime('%Y-%m-%d %H:%M:%S', tstamp_ns/1000000000, 'unixepoch') as tstamp"
|
||||
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 += ", open"
|
||||
query += ", high"
|
||||
query += ", low"
|
||||
query += ", close"
|
||||
query += ", volume"
|
||||
query += ", num_trades"
|
||||
query += ", vwap"
|
||||
|
||||
query += f" from {config['db_table_name']}"
|
||||
query += f" where exchange_id ='{exchange_id}'"
|
||||
query += f" and instrument_id in ({','.join(instrument_ids)})"
|
||||
|
||||
df = load_sqlite_to_dataframe(db_path=datafile, query=query)
|
||||
|
||||
# Trading Hours
|
||||
date_str = df["tstamp"][0][0:10]
|
||||
trading_hours = config["trading_hours"]
|
||||
|
||||
start_time = convert_time_to_UTC(
|
||||
f"{date_str} {trading_hours['begin_session']}", trading_hours["timezone"]
|
||||
)
|
||||
end_time = convert_time_to_UTC(
|
||||
f"{date_str} {trading_hours['end_session']}", trading_hours["timezone"]
|
||||
)
|
||||
|
||||
# Perform boolean selection
|
||||
df = df[(df["tstamp"] >= start_time) & (df["tstamp"] <= end_time)]
|
||||
df["tstamp"] = pd.to_datetime(df["tstamp"])
|
||||
|
||||
return cast(pd.DataFrame, df)
|
||||
|
||||
|
||||
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.
|
||||
Returns instruments without the configured prefix.
|
||||
"""
|
||||
try:
|
||||
conn = sqlite3.connect(datafile)
|
||||
|
||||
# Build exclusion list with full instrument_ids
|
||||
exclude_instruments = config.get("exclude_instruments", [])
|
||||
prefix = config.get("instrument_id_pfx", "")
|
||||
exclude_instrument_ids = [f"{prefix}{inst}" for inst in exclude_instruments]
|
||||
|
||||
# Query to get distinct instrument_ids
|
||||
query = f"""
|
||||
SELECT DISTINCT instrument_id
|
||||
FROM {config['db_table_name']}
|
||||
WHERE exchange_id = ?
|
||||
"""
|
||||
|
||||
# Add exclusion clause if there are instruments to exclude
|
||||
if exclude_instrument_ids:
|
||||
placeholders = ','.join(['?' for _ in exclude_instrument_ids])
|
||||
query += f" AND instrument_id NOT IN ({placeholders})"
|
||||
cursor = conn.execute(query, (config["exchange_id"],) + tuple(exclude_instrument_ids))
|
||||
else:
|
||||
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
|
||||
instruments = []
|
||||
for instrument_id in instrument_ids:
|
||||
if instrument_id.startswith(prefix):
|
||||
symbol = instrument_id[len(prefix) :]
|
||||
instruments.append(symbol)
|
||||
else:
|
||||
instruments.append(instrument_id)
|
||||
|
||||
return sorted(instruments)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error auto-detecting instruments from {datafile}: {str(e)}")
|
||||
return []
|
||||
|
||||
|
||||
# if __name__ == "__main__":
|
||||
# df1 = load_sqlite_to_dataframe(sys.argv[1], table_name="md_1min_bars")
|
||||
|
||||
# print(df1)
|
||||
@@ -1,208 +0,0 @@
|
||||
from typing import Any, Dict, List, Optional
|
||||
import pandas as pd # type:ignore
|
||||
from statsmodels.tsa.vector_ar.vecm import VECM, VECMResults # type:ignore
|
||||
|
||||
|
||||
class TradingPair:
|
||||
market_data_: pd.DataFrame
|
||||
symbol_a_: str
|
||||
symbol_b_: str
|
||||
price_column_: str
|
||||
|
||||
training_mu_: float
|
||||
training_std_: float
|
||||
|
||||
training_df_: pd.DataFrame
|
||||
testing_df_: pd.DataFrame
|
||||
|
||||
vecm_fit_: VECMResults
|
||||
|
||||
user_data_: Dict[str, Any]
|
||||
|
||||
def __init__(
|
||||
self, market_data: pd.DataFrame, symbol_a: str, symbol_b: str, price_column: str
|
||||
):
|
||||
self.symbol_a_ = symbol_a
|
||||
self.symbol_b_ = symbol_b
|
||||
self.price_column_ = price_column
|
||||
self.market_data_ = pd.DataFrame(
|
||||
self._transform_dataframe(market_data)[["tstamp"] + self.colnames()]
|
||||
)
|
||||
|
||||
|
||||
self.user_data_ = {}
|
||||
|
||||
def _transform_dataframe(self, df: pd.DataFrame) -> pd.DataFrame:
|
||||
# Select only the columns we need
|
||||
df_selected: pd.DataFrame = pd.DataFrame(
|
||||
df[["tstamp", "symbol", self.price_column_]]
|
||||
)
|
||||
|
||||
# Start with unique timestamps
|
||||
result_df: pd.DataFrame = (
|
||||
pd.DataFrame(df_selected["tstamp"]).drop_duplicates().reset_index(drop=True)
|
||||
)
|
||||
|
||||
# For each unique symbol, add a corresponding close price column
|
||||
|
||||
symbols = df_selected["symbol"].unique()
|
||||
for symbol in symbols:
|
||||
# Filter rows for this symbol
|
||||
df_symbol = df_selected[df_selected["symbol"] == symbol].reset_index(
|
||||
drop=True
|
||||
)
|
||||
|
||||
# Create column name like "close-COIN"
|
||||
new_price_column = f"{self.price_column_}_{symbol}"
|
||||
|
||||
# Create temporary dataframe with timestamp and price
|
||||
temp_df = pd.DataFrame(
|
||||
{
|
||||
"tstamp": df_symbol["tstamp"],
|
||||
new_price_column: df_symbol[self.price_column_],
|
||||
}
|
||||
)
|
||||
|
||||
# Join with our result dataframe
|
||||
result_df = pd.merge(result_df, temp_df, on="tstamp", how="left")
|
||||
result_df = result_df.reset_index(
|
||||
drop=True
|
||||
) # do not dropna() since irrelevant symbol would affect dataset
|
||||
|
||||
return result_df
|
||||
|
||||
def get_datasets(
|
||||
self,
|
||||
training_minutes: int,
|
||||
training_start_index: int = 0,
|
||||
testing_size: Optional[int] = None,
|
||||
) -> None:
|
||||
|
||||
testing_start_index = training_start_index + training_minutes
|
||||
self.training_df_ = self.market_data_.iloc[
|
||||
training_start_index:testing_start_index, :
|
||||
].copy()
|
||||
assert self.training_df_ is not None
|
||||
self.training_df_ = self.training_df_.dropna().reset_index(drop=True)
|
||||
|
||||
testing_start_index = training_start_index + training_minutes
|
||||
if testing_size is None:
|
||||
self.testing_df_ = self.market_data_.iloc[testing_start_index:, :].copy()
|
||||
else:
|
||||
self.testing_df_ = self.market_data_.iloc[
|
||||
testing_start_index : testing_start_index + testing_size, :
|
||||
].copy()
|
||||
assert self.testing_df_ is not None
|
||||
self.testing_df_ = self.testing_df_.dropna().reset_index(drop=True)
|
||||
|
||||
def colnames(self) -> List[str]:
|
||||
return [
|
||||
f"{self.price_column_}_{self.symbol_a_}",
|
||||
f"{self.price_column_}_{self.symbol_b_}",
|
||||
]
|
||||
|
||||
def fit_VECM(self):
|
||||
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()
|
||||
|
||||
assert vecm_fit is not None
|
||||
|
||||
# URGENT check beta and alpha
|
||||
|
||||
# Check if the model converged properly
|
||||
if not hasattr(vecm_fit, "beta") or vecm_fit.beta is None:
|
||||
print(f"{self}: VECM model failed to converge properly")
|
||||
|
||||
self.vecm_fit_ = vecm_fit
|
||||
# print(f"{self}: beta={self.vecm_fit_.beta} alpha={self.vecm_fit_.alpha}" )
|
||||
# print(f"{self}: {self.vecm_fit_.summary()}")
|
||||
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 check_cointegration_engle_granger(self):
|
||||
from statsmodels.tsa.stattools import coint
|
||||
|
||||
col1, col2 = self.colnames()
|
||||
assert self.training_df_ is not None
|
||||
series1 = self.training_df_[col1].reset_index(drop=True)
|
||||
series2 = self.training_df_[col2].reset_index(drop=True)
|
||||
|
||||
# Run Engle-Granger cointegration test
|
||||
pvalue = coint(series1, series2)[1]
|
||||
# Define cointegration if p-value < 0.05 (i.e., reject null of no cointegration)
|
||||
is_cointegrated = pvalue < 0.05
|
||||
print(f"{self}: is_cointegrated={is_cointegrated} pvalue={pvalue}")
|
||||
return is_cointegrated
|
||||
|
||||
def train_pair(self) -> bool:
|
||||
is_cointegrated_johansen = self.check_cointegration_johansen()
|
||||
is_cointegrated_engle_granger = self.check_cointegration_engle_granger()
|
||||
if not is_cointegrated_johansen and not is_cointegrated_engle_granger:
|
||||
return False
|
||||
pass
|
||||
|
||||
# print('*' * 80 + '\n' + f"**************** {self} IS COINTEGRATED ****************\n" + '*' * 80)
|
||||
self.fit_VECM()
|
||||
assert self.training_df_ is not None and self.vecm_fit_ is not None
|
||||
diseq_series = self.training_df_[self.colnames()] @ self.vecm_fit_.beta
|
||||
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"] = (
|
||||
self.training_df_[self.colnames()] @ self.vecm_fit_.beta
|
||||
)
|
||||
# Normalize the dis-equilibrium
|
||||
self.training_df_["scaled_dis-equilibrium"] = (
|
||||
diseq_series - self.training_mu_
|
||||
) / self.training_std_
|
||||
|
||||
return True
|
||||
|
||||
def predict(self) -> pd.DataFrame:
|
||||
assert self.testing_df_ is not None
|
||||
assert self.vecm_fit_ is not None
|
||||
predicted_prices = self.vecm_fit_.predict(steps=len(self.testing_df_))
|
||||
|
||||
# Convert prediction to a DataFrame for readability
|
||||
# predicted_df =
|
||||
|
||||
self.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()
|
||||
|
||||
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:
|
||||
return f"{self.symbol_a_} & {self.symbol_b_}"
|
||||
@@ -1,106 +1,29 @@
|
||||
import argparse
|
||||
import hjson
|
||||
import importlib
|
||||
import asyncio
|
||||
import glob
|
||||
import importlib
|
||||
import os
|
||||
import sqlite3
|
||||
from datetime import datetime, date
|
||||
|
||||
from datetime import date, datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import hjson
|
||||
import pandas as pd
|
||||
|
||||
from tools.data_loader import get_available_instruments_from_db, load_market_data
|
||||
from tools.trading_pair import TradingPair
|
||||
from results import BacktestResult, create_result_database, store_results_in_database, store_config_in_database
|
||||
from pt_trading.results import (
|
||||
BacktestResult,
|
||||
create_result_database,
|
||||
store_config_in_database,
|
||||
store_results_in_database,
|
||||
)
|
||||
from pt_trading.fit_methods import PairsTradingFitMethod
|
||||
from pt_trading.trading_pair import TradingPair
|
||||
|
||||
|
||||
def load_config(config_path: str) -> Dict:
|
||||
with open(config_path, "r") as f:
|
||||
config = hjson.load(f)
|
||||
return config
|
||||
|
||||
|
||||
# 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.
|
||||
# Returns instruments without the configured prefix.
|
||||
# """
|
||||
# try:
|
||||
# conn = sqlite3.connect(datafile)
|
||||
|
||||
# # Query to get distinct instrument_ids
|
||||
# query = f"""
|
||||
# SELECT DISTINCT instrument_id
|
||||
# FROM {config['db_table_name']}
|
||||
# WHERE exchange_id = ?
|
||||
# """
|
||||
|
||||
# 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
|
||||
# prefix = config.get("instrument_id_pfx", "")
|
||||
# instruments = []
|
||||
# for instrument_id in instrument_ids:
|
||||
# if instrument_id.startswith(prefix):
|
||||
# symbol = instrument_id[len(prefix) :]
|
||||
# instruments.append(symbol)
|
||||
# else:
|
||||
# instruments.append(instrument_id)
|
||||
|
||||
# return sorted(instruments)
|
||||
|
||||
# except Exception as e:
|
||||
# print(f"Error auto-detecting instruments from {datafile}: {str(e)}")
|
||||
# return []
|
||||
|
||||
|
||||
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 run_backtest(
|
||||
def run_strategy(
|
||||
config: Dict,
|
||||
datafile: str,
|
||||
price_column: str,
|
||||
strategy,
|
||||
fit_method: PairsTradingFitMethod,
|
||||
instruments: List[str],
|
||||
) -> BacktestResult:
|
||||
"""
|
||||
@@ -118,21 +41,27 @@ def run_backtest(
|
||||
config_copy = config.copy()
|
||||
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:
|
||||
pair = TradingPair(
|
||||
pair = fit_method.create_trading_pair(
|
||||
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 = []
|
||||
for pair in _create_pairs(config, instruments):
|
||||
single_pair_trades = strategy.run_pair(
|
||||
single_pair_trades = fit_method.run_pair(
|
||||
pair=pair, config=config, bt_result=bt_result
|
||||
)
|
||||
if single_pair_trades is not None and len(single_pair_trades) > 0:
|
||||
@@ -179,11 +108,12 @@ def main() -> None:
|
||||
|
||||
config: Dict = load_config(args.config)
|
||||
|
||||
# Dynamically instantiate strategy class
|
||||
strategy_class_name = config.get("strategy_class", "strategies.StaticFitStrategy")
|
||||
module_name, class_name = strategy_class_name.rsplit(".", 1)
|
||||
# Dynamically instantiate fit method class
|
||||
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)
|
||||
strategy = getattr(module, class_name)()
|
||||
fit_method = getattr(module, class_name)()
|
||||
|
||||
# Resolve data files (CLI takes priority over config)
|
||||
datafiles = resolve_datafiles(config, args.datafiles)
|
||||
@@ -202,32 +132,33 @@ def main() -> None:
|
||||
|
||||
# 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(",")]
|
||||
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,
|
||||
strategy_class=strategy_class_name,
|
||||
fit_method_class=fit_method_class_name,
|
||||
datafiles=datafiles,
|
||||
instruments=unique_instruments
|
||||
instruments=unique_instruments,
|
||||
)
|
||||
|
||||
# Process each data file
|
||||
price_column = config["price_column"]
|
||||
|
||||
for datafile in datafiles:
|
||||
print(f"\n====== Processing {os.path.basename(datafile)} ======")
|
||||
@@ -248,13 +179,12 @@ def main() -> None:
|
||||
|
||||
# Process data for this file
|
||||
try:
|
||||
strategy.reset()
|
||||
|
||||
bt_results = run_backtest(
|
||||
fit_method.reset()
|
||||
|
||||
bt_results = run_strategy(
|
||||
config=config,
|
||||
datafile=datafile,
|
||||
price_column=price_column,
|
||||
strategy=strategy,
|
||||
fit_method=fit_method,
|
||||
instruments=instruments,
|
||||
)
|
||||
|
||||
@@ -288,4 +218,4 @@ def main() -> None:
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
asyncio.run(main())
|
||||
Reference in New Issue
Block a user