progress
This commit is contained in:
@@ -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"))
|
||||
|
||||
Reference in New Issue
Block a user