This commit is contained in:
Oleg Sheynin
2025-08-02 00:12:31 +00:00
parent 80c3e8d54b
commit 73f36ddcea
21 changed files with 812 additions and 321 deletions
+52 -30
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import copy
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import Any, Dict, cast
from typing import Any, Dict, Optional, cast
import numpy as np
import pandas as pd
@@ -19,17 +19,26 @@ class ModelDataPolicy(ABC):
config_: Dict[str, Any]
current_data_params_: DataWindowParams
count_: int
is_real_time_: bool
def __init__(self, config: Dict[str, Any]):
def __init__(self, config: Dict[str, Any], *args: Any, **kwargs: Any):
self.config_ = config
training_size = config.get("training_size", 120)
training_start_index = 0
if kwargs.get("is_real_time", False):
training_size = 120
training_start_index = 0
else:
training_size = config.get("training_size", 120)
self.current_data_params_ = DataWindowParams(
training_size=config.get("training_size", 120),
training_start_index=0,
)
self.count_ = 0
self.is_real_time_ = kwargs.get("is_real_time", False)
@abstractmethod
def advance(self) -> DataWindowParams:
def advance(self, mkt_data_df: Optional[pd.DataFrame] = None) -> DataWindowParams:
self.count_ += 1
print(self.count_, end="\r")
return self.current_data_params_
@@ -50,22 +59,15 @@ class ModelDataPolicy(ABC):
class RollingWindowDataPolicy(ModelDataPolicy):
def __init__(self, config: Dict[str, Any], *args: Any, **kwargs: Any):
super().__init__(config)
super().__init__(config, *args, **kwargs)
self.count_ = 1
def advance(self) -> DataWindowParams:
super().advance()
self.current_data_params_.training_start_index += 1
return self.current_data_params_
class ExpandingWindowDataPolicy(ModelDataPolicy):
def __init__(self, config: Dict[str, Any], *args: Any, **kwargs: Any):
super().__init__(config)
def advance(self) -> DataWindowParams:
super().advance()
self.current_data_params_.training_size += 1
def advance(self, mkt_data_df: Optional[pd.DataFrame] = None) -> DataWindowParams:
super().advance(mkt_data_df)
if self.is_real_time_:
self.current_data_params_.training_start_index = -self.current_data_params_.training_size
else:
self.current_data_params_.training_start_index += 1
return self.current_data_params_
@@ -79,34 +81,47 @@ class OptimizedWndDataPolicy(ModelDataPolicy, ABC):
prices_b_: np.ndarray
def __init__(self, config: Dict[str, Any], *args: Any, **kwargs: Any):
super().__init__(config)
super().__init__(config, *args, **kwargs)
assert (
kwargs.get("mkt_data") is not None and kwargs.get("pair") is not None
), "mkt_data and/or pair must be provided"
kwargs.get("pair") is not None
), "pair must be provided"
assert (
"min_training_size" in config and "max_training_size" in config
), "min_training_size and max_training_size must be provided"
self.min_training_size_ = cast(int, config.get("min_training_size"))
self.max_training_size_ = cast(int, config.get("max_training_size"))
assert self.min_training_size_ < self.max_training_size_
from pt_strategy.trading_pair import TradingPair
self.pair_ = cast(TradingPair, kwargs.get("pair"))
if "mkt_data" in kwargs:
self.mkt_data_df_ = cast(pd.DataFrame, kwargs.get("mkt_data"))
col_a, col_b = self.pair_.colnames()
self.prices_a_ = np.array(self.mkt_data_df_[col_a])
self.prices_b_ = np.array(self.mkt_data_df_[col_b])
assert self.min_training_size_ < self.max_training_size_
self.mkt_data_df_ = cast(pd.DataFrame, kwargs.get("mkt_data"))
self.pair_ = cast(TradingPair, kwargs.get("pair"))
self.end_index_ = (
self.current_data_params_.training_start_index + self.max_training_size_
)
def advance(self, mkt_data_df: Optional[pd.DataFrame] = None) -> DataWindowParams:
super().advance(mkt_data_df)
if mkt_data_df is not None:
self.mkt_data_df_ = mkt_data_df
if self.is_real_time_:
self.end_index_ = len(self.mkt_data_df_) - 1
else:
self.end_index_ = self.current_data_params_.training_start_index + self.max_training_size_
if self.end_index_ > len(self.mkt_data_df_) - 1:
self.end_index_ = len(self.mkt_data_df_) - 1
self.current_data_params_.training_start_index = self.end_index_ - self.max_training_size_
if self.current_data_params_.training_start_index < 0:
self.current_data_params_.training_start_index = 0
col_a, col_b = self.pair_.colnames()
self.prices_a_ = np.array(self.mkt_data_df_[col_a])
self.prices_b_ = np.array(self.mkt_data_df_[col_b])
def advance(self) -> DataWindowParams:
super().advance()
self.current_data_params_ = self.optimize_window_size()
self.end_index_ += 1
return self.current_data_params_
@abstractmethod
@@ -126,6 +141,9 @@ class EGOptimizedWndDataPolicy(OptimizedWndDataPolicy):
last_pvalue = 1.0
result = copy.copy(self.current_data_params_)
for trn_size in range(self.min_training_size_, self.max_training_size_):
if self.end_index_ - trn_size < 0:
break
from statsmodels.tsa.stattools import coint # type: ignore
start_index = self.end_index_ - trn_size
@@ -155,6 +173,8 @@ class ADFOptimizedWndDataPolicy(OptimizedWndDataPolicy):
last_pvalue = 1.0
result = copy.copy(self.current_data_params_)
for trn_size in range(self.min_training_size_, self.max_training_size_):
if self.end_index_ - trn_size < 0:
break
start_index = self.end_index_ - trn_size
y = self.prices_a_[start_index : self.end_index_]
x = self.prices_b_[start_index : self.end_index_]
@@ -201,6 +221,8 @@ class JohansenOptdWndDataPolicy(OptimizedWndDataPolicy):
result = copy.copy(self.current_data_params_)
for trn_size in range(self.min_training_size_, self.max_training_size_):
if self.end_index_ - trn_size < 0:
break
start_index = self.end_index_ - trn_size
series_a = self.prices_a_[start_index:self.end_index_]
series_b = self.prices_b_[start_index:self.end_index_]