This commit is contained in:
Oleg Sheynin
2026-01-12 21:26:15 +00:00
parent bd6cf1d4d0
commit c0fabcb429
15 changed files with 1205 additions and 568 deletions
+26 -23
View File
@@ -22,7 +22,7 @@ from cvttpy_trading.trading.trading_instructions import TargetPositionSignal
from pairs_trading.lib.pt_strategy.model_data_policy import ModelDataPolicy
from pairs_trading.lib.pt_strategy.pt_model import Prediction
from pairs_trading.lib.pt_strategy.trading_pair import LiveTradingPair
from pairs_trading.apps.pairs_trader import PairsTrader
from pairs_trading.apps.pair_trader import PairTrader
from pairs_trading.lib.pt_strategy.pt_market_data import LiveMarketData
@@ -37,7 +37,7 @@ class PtLiveStrategy(NamedObject):
trading_pair_: LiveTradingPair
model_data_policy_: ModelDataPolicy
pairs_trader_: PairsTrader
pairs_trader_: PairTrader
# for presentation: history of prediction values and trading signals
predictions_df_: pd.DataFrame
@@ -46,22 +46,28 @@ class PtLiveStrategy(NamedObject):
def __init__(
self,
config: Config,
pairs_trader: PairsTrader,
pairs_trader: PairTrader,
):
# import copy
# self.config_ = Config(json_src=copy.deepcopy(config.data()))
self.config_ = config
self.pairs_trader_ = pairs_trader
self.trading_pair_ = LiveTradingPair(
config=config,
instruments=self.pairs_trader_.instruments_,
)
self.model_data_policy_ = ModelDataPolicy.create(
self.config_,
is_real_time=True,
pair=self.trading_pair_,
)
assert (
self.model_data_policy_ is not None
), f"{self.fname()}: Unable to create ModelDataPolicy"
self.predictions_df_ = pd.DataFrame()
self.trading_signals_df_ = pd.DataFrame()
# self.book_ = book
import copy
# modified config must be passed to PtMarketData
self.config_ = Config(json_src=copy.deepcopy(config.data()))
self.instruments_ = self.pairs_trader_.instruments_
@@ -71,25 +77,27 @@ class PtLiveStrategy(NamedObject):
async def _on_config(self) -> None:
self.interval_sec_ = self.config_.get_value("interval_sec", 0)
assert self.interval_sec_ > 0, "interval_sec cannot be 0"
self.history_depth_sec_ = (
self.config_.get_value("history_depth_hours", 0) * SecPerHour
)
assert self.history_depth_sec_ > 0, "history_depth_hours cannot be 0"
await self.pairs_trader_.subscribe_md()
self.open_threshold_ = self.config_.get_value(
"dis-equilibrium_open_trshld", 0.0
"model/disequilibrium/open_trshld", 0.0
)
self.close_threshold_ = self.config_.get_value(
"dis-equilibrium_close_trshld", 0.0
"model/disequilibrium/close_trshld", 0.0
)
assert (
self.open_threshold_ > 0
), "dis-equilibrium_open_trshld must be greater than 0"
), "disequilibrium/open_trshld must be greater than 0"
assert (
self.close_threshold_ > 0
), "dis-equilibrium_close_trshld must be greater than 0"
), "disequilibrium/close_trshld must be greater than 0"
def __repr__(self) -> str:
return f"{self.classname()}: trading_pair={self.trading_pair_}, mdp={self.model_data_policy_.__class__.__name__}, "
@@ -106,16 +114,8 @@ class PtLiveStrategy(NamedObject):
return
self.trading_pair_.market_data_ = market_data_df
self.model_data_policy_ = ModelDataPolicy.create(
self.config_,
is_real_time=True,
pair=self.trading_pair_,
mkt_data=market_data_df,
)
assert (
self.model_data_policy_ is not None
), f"{self.fname()}: Unable to create ModelDataPolicy"
Log.info(f"{self.fname()}: Running prediction for pair: {self.trading_pair_}")
prediction = self.trading_pair_.run(
market_data_df, self.model_data_policy_.advance()
)
@@ -132,13 +132,16 @@ class PtLiveStrategy(NamedObject):
await self._send_trading_instructions(trading_instructions)
def _is_md_actual(self, hist_aggr: List[MdTradesAggregate]) -> bool:
curr_ns = current_nanoseconds()
LAG_THRESHOLD = 5 * NanoPerSec
if len(hist_aggr) == 0:
Log.warning(f"{self.fname()} list of aggregates IS EMPTY")
return False
# MAYBE check market data length
if current_nanoseconds() - hist_aggr[-1].time_ns_ > LAG_THRESHOLD:
lag_ns = curr_ns - hist_aggr[-1].time_ns_
if lag_ns > LAG_THRESHOLD:
Log.warning(f"{self.fname()} {hist_aggr[-1].exch_inst_.details_short()} Lagging {int(lag_ns/NanoPerSec)} seconds")
return False
return True
+11 -9
View File
@@ -25,7 +25,7 @@ class ModelDataPolicy(ABC):
def __init__(self, config: Config, *args: Any, **kwargs: Any):
self.config_ = config
self.current_data_params_ = DataWindowParams(
training_size_=config.get_value("training_size", 120),
training_size_=config.get_value("model/training_size", 120),
training_start_index_=0,
)
self.count_ = 0
@@ -34,14 +34,15 @@ class ModelDataPolicy(ABC):
@abstractmethod
def advance(self, mkt_data_df: Optional[pd.DataFrame] = None) -> DataWindowParams:
self.count_ += 1
print(self.count_, end="\r")
if not self.is_real_time_:
print(self.count_, end="\r")
return self.current_data_params_
@staticmethod
def create(config: Config, *args: Any, **kwargs: Any) -> ModelDataPolicy:
import importlib
model_data_policy_class_name = config.get_value("model_data_policy_class", None)
model_data_policy_class_name = config.get_value("model/model_data_policy_class", None)
assert model_data_policy_class_name is not None
module_name, class_name = model_data_policy_class_name.rsplit(".", 1)
module = importlib.import_module(module_name)
@@ -59,7 +60,9 @@ class RollingWindowDataPolicy(ModelDataPolicy):
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_
self.current_data_params_.training_start_index_ = 0
if mkt_data_df and len(mkt_data_df) > self.curren_data_params_.training_size_:
self.current_data_params_.training_start_index_ = -self.curren_data_params_.training_size_
else:
self.current_data_params_.training_start_index_ += 1
return self.current_data_params_
@@ -79,11 +82,10 @@ class OptimizedWndDataPolicy(ModelDataPolicy, ABC):
assert (
kwargs.get("pair") is not None
), "pair must be provided"
assert (
"min_training_size" in config.data() and "max_training_size" in config.data()
), "min_training_size and max_training_size must be provided"
self.min_training_size_ = cast(int, config.get_value("min_training_size"))
self.max_training_size_ = cast(int, config.get_value("max_training_size"))
assert (config.key_exists("model/max_training_size") and config.key_exists("model/min_training_size")
), "min_training_size and max_training_size must be provided"
self.min_training_size_ = cast(int, config.get_value("model/min_training_size"))
self.max_training_size_ = cast(int, config.get_value("model/max_training_size"))
from pairs_trading.lib.pt_strategy.trading_pair import TradingPair
self.pair_ = cast(TradingPair, kwargs.get("pair"))
+1 -1
View File
@@ -30,7 +30,7 @@ class PtMarketData(NamedObject, ABC):
self.config_ = config
self.origin_mkt_data_df_ = pd.DataFrame()
self.market_data_df_ = pd.DataFrame()
self.stat_model_price_ = self.config_.get_value("stat_model_price")
self.stat_model_price_ = self.config_.get_value("model/stat_model_price")
self.instruments_ = instruments
assert len(self.instruments_) > 0, "No instruments found in config"
+1 -1
View File
@@ -19,7 +19,7 @@ class PairsTradingModel(ABC):
def create(config: Config) -> PairsTradingModel:
import importlib
model_class_name = config.get_value("model_class", None)
model_class_name = config.get_value("model/model_class", None)
assert model_class_name is not None
module_name, class_name = model_class_name.rsplit(".", 1)
module = importlib.import_module(module_name)
+2 -2
View File
@@ -95,8 +95,8 @@ class PtResearchStrategy:
pair = self.trading_pair_
trades = None
open_threshold = self.config_.get_value("dis-equilibrium_open_trshld")
close_threshold = self.config_.get_value("dis-equilibrium_close_trshld")
open_threshold = self.config_.get_value("model/disequilibrium/open_trshld")
close_threshold = self.config_.get_value("model/disequilibrium/close_trshld")
scaled_disequilibrium = prediction.scaled_disequilibrium_
abs_scaled_disequilibrium = abs(scaled_disequilibrium)
+2 -3
View File
@@ -50,11 +50,11 @@ class TradingPair(NamedObject, ABC):
self.instruments_ = instruments
self.instruments_[0].user_data_["symbol"] = instruments[0].instrument_id().split("-", 1)[1]
self.instruments_[1].user_data_["symbol"] = instruments[1].instrument_id().split("-", 1)[1]
self.stat_model_price_ = config.get_value("model/stat_model_price")
def run(self, market_data: pd.DataFrame, data_params: DataWindowParams) -> Prediction: # type: ignore[assignment]
self.market_data_ = market_data[
data_params.training_start_index_ : data_params.training_start_index_
+ data_params.training_size_
data_params.training_start_index_ : data_params.training_start_index_ + data_params.training_size_
]
return self.model_.predict(pair=self)
@@ -93,7 +93,6 @@ class ResearchTradingPair(TradingPair):
assert len(instruments) == 2, "Trading pair must have exactly 2 instruments"
super().__init__(config=config, instruments=instruments)
self.stat_model_price_ = config.get_value("stat_model_price")
self.user_data_ = {
"state": PairState.INITIAL,
}