Files

87 lines
2.7 KiB
Python
Raw Permalink Normal View History

from pydantic import BaseModel, validator
from typing import Literal, Optional
from labeling.types import EventFilter
from models.base import Model
from data_loader.types import DataCollection, DataSource
from feature_extractors.types import FeatureExtractor
from labeling.types import EventFilter, EventLabeller
from sklearn.base import BaseEstimator
from dataclasses import dataclass
from transformations.base import Transformation
# RawConfig is needed to ensure we can declare config presets here with static typing, we then convert it to Config
class RawConfig(BaseModel):
start_date: Optional[str]
dimensionality_reduction_ratio: float
n_features_to_select: int
initial_window_size: int
retrain_every: int
scaler: Literal["normalize", "minmax", "standardize", "robust"]
assets: list[str]
target_asset: str
other_assets: list[str]
exogenous_data: list[str]
load_non_target_asset: bool
own_features: list[str]
other_features: list[str]
exogenous_features: list[str]
event_filter: Literal["none", "cusum_vol", "cusum_fixed"]
2022-03-15 14:43:16 +01:00
event_filter_multiplier: float
remove_overlapping_events: bool
labeling: Literal["two_class", "three_class_balanced", "three_class_imbalanced"]
forecasting_horizon: int
transaction_costs: float
2022-02-19 15:06:05 +01:00
save_models: bool
ensembling_method: Literal["voting_soft", "stacking"]
directional_models: list[str]
meta_models: list[str]
@dataclass
class Config:
start_date: Optional[str]
initial_window_size: int
retrain_every: int
assets: DataCollection
target_asset: DataSource
other_assets: DataCollection
exogenous_data: DataCollection
load_non_target_asset: bool
own_features: list[tuple[str, FeatureExtractor, list[int]]]
other_features: list[tuple[str, FeatureExtractor, list[int]]]
exogenous_features: list[tuple[str, FeatureExtractor, list[int]]]
event_filter: EventFilter
remove_overlapping_events: bool
labeling: EventLabeller
forecasting_horizon: int
transaction_costs: float
no_of_classes: Literal["two", "three-balanced", "three-imbalanced"]
2022-02-19 15:06:05 +01:00
save_models: bool
mode: Literal["training", "inference"]
directional_model: Model
meta_model: Model
transformations: list[Transformation]
@validator("directional_model", "meta_model")
def check_model(cls, v):
assert isinstance(v, BaseEstimator)
return v
@validator("event_filter")
def check_event_filter(cls, v):
assert isinstance(v, EventFilter)
return v
@validator("labeling")
def check_labeling(cls, v):
assert isinstance(v, EventLabeller)
return v
# class Config:
# arbitrary_types_allowed = True