from dataclasses import dataclass from typing import Any from model_manager.sientia.models import validate_frontend_date_format # Model name constants MODEL_LINEAR_REGRESSION = 'Linear Regression' MODEL_POLYNOMIAL_REGRESSION = 'Polynomial Regression' @dataclass class TrainModelParams: """ Parameters for machine learning model training. This class encapsulates all configuration parameters required for the training pipeline, including data processing settings, model configuration, and experiment tracking information. All parameters are validated upon initialization to ensure data integrity and prevent runtime errors. Use the `from_dict()` class method to create instances from dictionaries with automatic validation of all fields. Attributes: variable_columns (list[str]): List of variable column names to use as features. lag_train (dict[str, int]): Dictionary of lags per variable for training phase. lag_val (dict[str, int]): Dictionary of lags per variable for validation phase. target_variable (str): Name of the target variable to predict. rem_static_win (bool): Whether to remove static windows from data. low_lim (dict[str, float]): Dictionary of lower limits for each variable. upp_lim (dict[str, float]): Dictionary of upper limits for each variable. window (int): Window size for rolling operations. use_scaler (bool): Whether to use a scaler for data normalization. include_ar (bool): Whether to include autoregressive variables. bucket_name (str): Name of the MinIO bucket containing training data. file_name (str): Name of the training file in the MinIO bucket. validation_file_name (str | None): Optional name of the validation file in the same MinIO bucket as the training file. line_separator (str): Line separator used in the CSV file. decimal_separator (str): Decimal separator used in the CSV file. date_column (str | None): Name of the date/time column. If set with date_format, the column is parsed as datetime. date_format (str | None): Format of the date column (e.g. dd/MM/yyyy HH:mm:ss). Used when date_column is set. train_size (int): Percentage of data to use for training (0-100). shuffle (bool): Whether to shuffle the data during train/test split. experiment_run_id (int): Unique identifier for the experiment run. experiment_name (str): Name of the experiment for tracking. removed_intervals (list): List of time intervals to remove from the data. model_name (str): Name of the model type ('Linear Regression' or 'Polynomial Regression'). degree (int): Degree of polynomial features (1 for linear, >1 for polynomial). interaction_only (bool): If True, only interaction features are produced for polynomial. nan_treatment (str): Treatment for NaN values ('drop' or 'linear interpolation'). start_date (str | None): Start date for filtering data. end_date (str | None): End date for filtering data. scaler_name (str): Name of the scaler to use ('Standard Scaler' or 'None'). support_filters (dict): Custom support filters per variable. static_threshold (int | None): Threshold for static window removal (1-1000). Only used when rem_static_win is True. """ variable_columns: list[str] lag_train: dict[str, int] lag_val: dict[str, int] target_variable: str rem_static_win: bool low_lim: dict[str, float] upp_lim: dict[str, float] window: int use_scaler: bool include_ar: bool bucket_name: str file_name: str validation_file_name: str | None line_separator: str decimal_separator: str date_column: str | None date_format: str | None train_size: int shuffle: bool experiment_run_id: int experiment_name: str removed_intervals: list model_name: str degree: int interaction_only: bool nan_treatment: str start_date: str | None end_date: str | None scaler_name: str support_filters: dict static_threshold: int | None @classmethod def from_dict(cls, data: dict[str, Any]) -> 'TrainModelParams': """ Create TrainModelParams from dictionary with validation. This factory method creates a TrainModelParams instance from a dictionary, applying validation to ensure all required fields are present and have the correct types. This is the recommended way to create instances from workflow input data. Args: data: Dictionary containing training parameters with keys matching the attribute names (variable_columns, lag_train, etc.) Returns: TrainModelParams: Validated instance with all fields populated Raises: ValueError: If any required field is missing or None TypeError: If any field has an incorrect type KeyError: If any required key is missing from the dictionary """ return cls( variable_columns=cls._check_none( data.get('variable_columns'), list, 'variable_columns' ), lag_train=cls._check_none(data.get('lag_train'), dict, 'lag_train'), lag_val=cls._check_none(data.get('lag_val'), dict, 'lag_val'), target_variable=cls._check_none(data.get('target_variable'), str, 'target_variable'), rem_static_win=cls._check_none(data.get('rem_static_win'), bool, 'rem_static_win'), low_lim=cls._check_none(data.get('low_lim'), dict, 'low_lim'), upp_lim=cls._check_none(data.get('upp_lim'), dict, 'upp_lim'), window=cls._check_none(data.get('window'), int, 'window'), use_scaler=cls._check_none(data.get('use_scaler'), bool, 'use_scaler'), include_ar=cls._check_none(data.get('include_ar'), bool, 'include_ar'), bucket_name=cls._check_none(data.get('bucket_name'), str, 'bucket_name'), file_name=cls._check_none(data.get('file_name'), str, 'file_name'), validation_file_name=cls._check_type( data.get('validation_file_name'), str, 'validation_file_name' ), line_separator=cls._check_none(data.get('line_separator'), str, 'line_separator'), decimal_separator=cls._check_none( data.get('decimal_separator'), str, 'decimal_separator' ), date_column=data.get('date_column'), date_format=data.get('date_format'), train_size=cls._check_none(data.get('train_size'), int, 'train_size'), shuffle=cls._check_none(data.get('shuffle'), bool, 'shuffle'), experiment_run_id=cls._check_none( data.get('experiment_run_id'), int, 'experiment_run_id' ), experiment_name=cls._check_none(data.get('experiment_name'), str, 'experiment_name'), removed_intervals=cls._check_type( data.get('removed_intervals'), list, 'removed_intervals' ), model_name=cls._check_none(data.get('model_name'), str, 'model_name'), degree=cls._check_none(data.get('degree'), int, 'degree'), interaction_only=cls._check_none( data.get('interaction_only'), bool, 'interaction_only' ), nan_treatment=cls._check_none(data.get('nan_treatment'), str, 'nan_treatment'), start_date=cls._check_type(data.get('start_date'), str, 'start_date'), end_date=cls._check_type(data.get('end_date'), str, 'end_date'), scaler_name=cls._check_none(data.get('scaler_name'), str, 'scaler_name'), support_filters=cls._check_type(data.get('support_filters'), dict, 'support_filters') or {}, static_threshold=cls._check_type(data.get('static_threshold'), int, 'static_threshold'), ) @staticmethod def _check_none(value: Any | None, expected_type: type, field_name: str) -> Any: """ Validate that a value is not None and check its type. This method ensures that required parameters are provided and have the correct type, raising descriptive errors if validation fails. Args: value (Any | None): The value to validate. expected_type (type): The expected type of the value. field_name (str): The name of the field being validated (for error messages). Returns: Any: The validated value if it is not None and matches the expected type. Raises: ValueError: If the value is None. TypeError: If the value is not of the expected type. """ if value is None: error = f'{field_name} is required and cannot be None.' raise ValueError(error) return TrainModelParams._check_type(value, expected_type, field_name) @staticmethod def _check_type(value: Any | None, expected_type: type, field_name: str) -> Any: """ Validate that a value matches the expected type. This method checks type compatibility and raises a descriptive error if the value does not match the expected type. Args: value (Any | None): The value to validate. expected_type (type): The expected type of the value. field_name (str): The name of the field being validated (for error messages). Returns: Any: The validated value if it matches the expected type. Raises: TypeError: If the value is not of the expected type. """ if value is not None and not isinstance(value, expected_type): error = f'{field_name} must be of type {expected_type.__name__}, but got {type(value).__name__}.' raise TypeError(error) return value def validate_business_rules(self) -> None: """ Validate business rules and constraints for training parameters. This method performs additional validation beyond type checking to ensure that parameter values are within acceptable ranges and logically consistent. It implements defense-in-depth validation to catch configuration errors early in the workflow. Raises: ValueError: If any business rule is violated """ self._validate_numeric_ranges() self._validate_model_params() self._validate_intervals_and_dates() self._validate_limits() self._validate_required_strings() self._validate_date_format() self._validate_validation_file_name() def _validate_numeric_ranges(self) -> None: """Validate numeric parameters are within acceptable ranges.""" if not 10 <= self.train_size <= 100: raise ValueError(f'train_size must be between 10 and 100, got {self.train_size}') if not self.variable_columns: raise ValueError('variable_columns cannot be empty') for var, lag in self.lag_train.items(): if lag < 0: raise ValueError(f'lag_train for {var} must be non-negative, got {lag}') for var, lag in self.lag_val.items(): if lag < 0: raise ValueError(f'lag_val for {var} must be non-negative, got {lag}') if self.window < 0: raise ValueError(f'window must be non-negative, got {self.window}') # Validate static_threshold only when rem_static_win is True and value is provided if self.rem_static_win and self.static_threshold is not None: if not 1 <= self.static_threshold <= 1000: raise ValueError( f'static_threshold must be between 1 and 1000, got {self.static_threshold}' ) def _validate_model_params(self) -> None: """Validate model-related parameters.""" if self.degree < 1: raise ValueError(f'degree must be at least 1, got {self.degree}') valid_nan_treatments = ['drop', 'linear interpolation', 'fill linear'] if self.nan_treatment not in valid_nan_treatments: raise ValueError( f'nan_treatment must be one of {valid_nan_treatments}, got {self.nan_treatment}' ) valid_scalers = ['Standard Scaler', 'None'] if self.scaler_name not in valid_scalers: raise ValueError(f'scaler_name must be one of {valid_scalers}, got {self.scaler_name}') valid_models = [MODEL_LINEAR_REGRESSION, MODEL_POLYNOMIAL_REGRESSION] if self.model_name not in valid_models: raise ValueError(f'model_name must be one of {valid_models}, got {self.model_name}') if self.model_name == MODEL_POLYNOMIAL_REGRESSION and self.degree < 2: raise ValueError( f'degree must be at least 2 for {MODEL_POLYNOMIAL_REGRESSION}, got {self.degree}' ) if self.model_name == MODEL_POLYNOMIAL_REGRESSION and self.scaler_name == 'None': raise ValueError( f'scaler_name must be set (e.g., "Standard Scaler") for {MODEL_POLYNOMIAL_REGRESSION} ' 'to avoid numerical overflow with large feature values' ) if self.model_name == MODEL_LINEAR_REGRESSION and self.degree != 1: raise ValueError(f'degree must be 1 for {MODEL_LINEAR_REGRESSION}, got {self.degree}') def _validate_intervals_and_dates(self) -> None: """Validate removed_intervals format and date parameters.""" if self.removed_intervals: for i, interval in enumerate(self.removed_intervals): if not isinstance(interval, (list, tuple)): raise ValueError( f'removed_intervals[{i}] must be a list or tuple, ' f'got {type(interval).__name__}' ) if len(interval) < 2: raise ValueError( f'removed_intervals[{i}] must have at least 2 elements (start, end), ' f'got {len(interval)}' ) if self.start_date is not None and not isinstance(self.start_date, str): raise TypeError(f'start_date must be a string, got {type(self.start_date).__name__}') if self.end_date is not None and not isinstance(self.end_date, str): raise TypeError(f'end_date must be a string, got {type(self.end_date).__name__}') def _validate_limits(self) -> None: """Validate low_lim and upp_lim consistency.""" if set(self.low_lim.keys()) != set(self.upp_lim.keys()): raise ValueError( f'low_lim and upp_lim must have the same keys. ' f'low_lim keys: {set(self.low_lim.keys())}, ' f'upp_lim keys: {set(self.upp_lim.keys())}' ) for var in self.low_lim: if self.low_lim[var] >= self.upp_lim[var]: raise ValueError( f'low_lim must be less than upp_lim for variable "{var}". ' f'Got low_lim={self.low_lim[var]}, upp_lim={self.upp_lim[var]}' ) def _validate_required_strings(self) -> None: """Validate required string fields are not empty.""" if not self.target_variable.strip(): raise ValueError('target_variable cannot be empty or whitespace') if not self.bucket_name.strip(): raise ValueError('bucket_name cannot be empty or whitespace') if not self.file_name.strip(): raise ValueError('file_name cannot be empty or whitespace') if not self.experiment_name.strip(): raise ValueError('experiment_name cannot be empty or whitespace') def _validate_date_format(self) -> None: """Validate date_format is one of the allowed frontend formats when set.""" if self.date_format: validate_frontend_date_format(self.date_format) def _validate_validation_file_name(self) -> None: """ Validate that validation_file_name, when provided, is not empty or whitespace. This field is optional; when present it must point to a valid object key in the same MinIO bucket specified by bucket_name. """ if self.validation_file_name is None: return if not isinstance(self.validation_file_name, str): raise TypeError( f'validation_file_name must be a string, got {type(self.validation_file_name).__name__}' ) if not self.validation_file_name.strip(): raise ValueError('validation_file_name cannot be empty or whitespace')