HoldoutEvaluationStrategy
Evaluation strategy implementing holdout (train/validation/test split) validation.
This strategy divides the dataset into three mutually exclusive partitions: training, validation, and test. The training set is used for model training, the validation set for HPO, and the test set for final evaluation.
The strategy handles metric aggregation at multiple levels:
- TRIAL level: Metrics during HPO trials on validation set
- LAST level: Final metrics computed on all three partitions after training
Methods
evaluate(self, model, input_dataset, output_dataset, metric)
HoldoutEvaluationStrategyEvaluate model on validation set during HPO trials.
Parameters
- model : BaseModel
- The model instance to evaluate with specific hyperparameters.
- input_dataset : DatasetDict
- DatasetDict with data partitions {"train": X_train, "validation": X_val, "test": X_test}.
- output_dataset : DatasetDict
- DatasetDict with label partitions {"train": y_train, "validation": y_val, "test": y_test}.
- metric : Metric
- The metric function to compute on predictions.
Returns
- float
- The metric score value for this hyperparameter combination.
execute(self, x, y, run: DashAI.back.dependencies.database.models.Run, db)
HoldoutEvaluationStrategyExecute holdout validation: train on training set, optimize with validation, evaluate on test.
Parameters
- x : DatasetDict
- DatasetDict with data partitions: {"train": X_train, "validation": X_val, "test": X_test}
- y : DatasetDict
- DatasetDict with label partitions: {"train": y_train, "validation": y_val, "test": y_test}
- run : Run
- Database run instance for storing results and configuration.
- db : Session
- SQLAlchemy database session for persisting metrics.
Returns
- tuple
- (trained_model, plot_paths) where: - trained_model : BaseModel - The trained model - plot_paths : list[str] - Paths to HPO visualization plot files
set_progress_reporter(self, progress_reporter: Optional[Callable[[Optional[float], Optional[str]], NoneType]]) -> None
BaseEvaluationStrategyRegister a callback that will receive progress updates.