mlrl-testbed-sklearn 0.12.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (60) hide show
  1. mlrl/__init__.py +0 -0
  2. mlrl/testbed_sklearn/__init__.py +0 -0
  3. mlrl/testbed_sklearn/experiments/__init__.py +8 -0
  4. mlrl/testbed_sklearn/experiments/dataset.py +141 -0
  5. mlrl/testbed_sklearn/experiments/experiment.py +164 -0
  6. mlrl/testbed_sklearn/experiments/input/__init__.py +0 -0
  7. mlrl/testbed_sklearn/experiments/input/dataset/__init__.py +0 -0
  8. mlrl/testbed_sklearn/experiments/input/dataset/preprocessors/__init__.py +0 -0
  9. mlrl/testbed_sklearn/experiments/input/dataset/preprocessors/extension.py +47 -0
  10. mlrl/testbed_sklearn/experiments/input/dataset/preprocessors/one_hot_encoder.py +63 -0
  11. mlrl/testbed_sklearn/experiments/input/dataset/splitters/__init__.py +0 -0
  12. mlrl/testbed_sklearn/experiments/input/dataset/splitters/extension.py +115 -0
  13. mlrl/testbed_sklearn/experiments/input/dataset/splitters/splitter_bipartition.py +136 -0
  14. mlrl/testbed_sklearn/experiments/input/dataset/splitters/splitter_cross_validation.py +208 -0
  15. mlrl/testbed_sklearn/experiments/output/__init__.py +0 -0
  16. mlrl/testbed_sklearn/experiments/output/characteristics/__init__.py +0 -0
  17. mlrl/testbed_sklearn/experiments/output/characteristics/data/__init__.py +10 -0
  18. mlrl/testbed_sklearn/experiments/output/characteristics/data/characteristics.py +135 -0
  19. mlrl/testbed_sklearn/experiments/output/characteristics/data/characteristics_data.py +124 -0
  20. mlrl/testbed_sklearn/experiments/output/characteristics/data/characteristics_prediction.py +26 -0
  21. mlrl/testbed_sklearn/experiments/output/characteristics/data/extension.py +93 -0
  22. mlrl/testbed_sklearn/experiments/output/characteristics/data/extension_prediction.py +94 -0
  23. mlrl/testbed_sklearn/experiments/output/characteristics/data/matrix_feature.py +71 -0
  24. mlrl/testbed_sklearn/experiments/output/characteristics/data/matrix_label.py +67 -0
  25. mlrl/testbed_sklearn/experiments/output/characteristics/data/matrix_output.py +58 -0
  26. mlrl/testbed_sklearn/experiments/output/characteristics/data/writer_data.py +39 -0
  27. mlrl/testbed_sklearn/experiments/output/characteristics/data/writer_prediction.py +46 -0
  28. mlrl/testbed_sklearn/experiments/output/dataset/__init__.py +7 -0
  29. mlrl/testbed_sklearn/experiments/output/dataset/dataset.py +47 -0
  30. mlrl/testbed_sklearn/experiments/output/dataset/dataset_ground_truth.py +22 -0
  31. mlrl/testbed_sklearn/experiments/output/dataset/dataset_prediction.py +21 -0
  32. mlrl/testbed_sklearn/experiments/output/dataset/extension_ground_truth.py +74 -0
  33. mlrl/testbed_sklearn/experiments/output/dataset/extension_prediction.py +74 -0
  34. mlrl/testbed_sklearn/experiments/output/dataset/writer_ground_truth.py +39 -0
  35. mlrl/testbed_sklearn/experiments/output/dataset/writer_prediction.py +62 -0
  36. mlrl/testbed_sklearn/experiments/output/evaluation/__init__.py +11 -0
  37. mlrl/testbed_sklearn/experiments/output/evaluation/evaluation_result.py +146 -0
  38. mlrl/testbed_sklearn/experiments/output/evaluation/extension.py +131 -0
  39. mlrl/testbed_sklearn/experiments/output/evaluation/extractor_classification.py +43 -0
  40. mlrl/testbed_sklearn/experiments/output/evaluation/extractor_ranking.py +45 -0
  41. mlrl/testbed_sklearn/experiments/output/evaluation/extractor_regression.py +33 -0
  42. mlrl/testbed_sklearn/experiments/output/evaluation/measures_classification.py +155 -0
  43. mlrl/testbed_sklearn/experiments/output/evaluation/measures_ranking.py +45 -0
  44. mlrl/testbed_sklearn/experiments/output/evaluation/measures_regression.py +40 -0
  45. mlrl/testbed_sklearn/experiments/output/evaluation/writer.py +89 -0
  46. mlrl/testbed_sklearn/experiments/output/label_vectors/__init__.py +7 -0
  47. mlrl/testbed_sklearn/experiments/output/label_vectors/extension.py +76 -0
  48. mlrl/testbed_sklearn/experiments/output/label_vectors/label_vector_histogram.py +68 -0
  49. mlrl/testbed_sklearn/experiments/output/label_vectors/label_vectors.py +59 -0
  50. mlrl/testbed_sklearn/experiments/output/label_vectors/writer.py +40 -0
  51. mlrl/testbed_sklearn/experiments/prediction/__init__.py +7 -0
  52. mlrl/testbed_sklearn/experiments/prediction/extension.py +42 -0
  53. mlrl/testbed_sklearn/experiments/prediction/predictor.py +93 -0
  54. mlrl/testbed_sklearn/experiments/prediction/predictor_global.py +67 -0
  55. mlrl/testbed_sklearn/experiments/problem_domain.py +105 -0
  56. mlrl/testbed_sklearn/runnables.py +184 -0
  57. mlrl_testbed_sklearn-0.12.0.dist-info/METADATA +38 -0
  58. mlrl_testbed_sklearn-0.12.0.dist-info/RECORD +60 -0
  59. mlrl_testbed_sklearn-0.12.0.dist-info/WHEEL +5 -0
  60. mlrl_testbed_sklearn-0.12.0.dist-info/top_level.txt +1 -0
mlrl/__init__.py ADDED
File without changes
File without changes
@@ -0,0 +1,8 @@
1
+ """
2
+ Author Michael Rapp (michael.rapp.ml@gmail.com)
3
+
4
+ Provides classes that allow running experiments using the scikit-learn framework.
5
+ """
6
+ from mlrl.testbed_sklearn.experiments.experiment import SkLearnExperiment
7
+ from mlrl.testbed_sklearn.experiments.problem_domain import SkLearnClassificationProblem, SkLearnProblem, \
8
+ SkLearnRegressionProblem
@@ -0,0 +1,141 @@
1
+ """
2
+ Author: Michael Rapp (michael.rapp.ml@gmail.com)
3
+
4
+ Provides classes for representing tabular datasets.
5
+ """
6
+ from dataclasses import dataclass, replace
7
+ from enum import Enum, auto
8
+ from functools import reduce
9
+ from typing import List, Optional
10
+
11
+ from scipy.sparse import lil_array
12
+
13
+ from mlrl.testbed.experiments.dataset import Dataset
14
+
15
+ from mlrl.util.arrays import is_sparse
16
+
17
+
18
+ class AttributeType(Enum):
19
+ """
20
+ All supported types of attributes.
21
+ """
22
+ NUMERICAL = auto()
23
+ ORDINAL = auto()
24
+ NOMINAL = auto()
25
+
26
+
27
+ @dataclass
28
+ class Attribute:
29
+ """
30
+ An attribute, e.g., a feature, a ground truth label, or a regression score, that is contained by a dataset.
31
+
32
+ Attributes:
33
+ name: The name of the attribute
34
+ attribute_type: The type of the attribute
35
+ nominal_values: A list that contains the possible values in case of a nominal feature
36
+ """
37
+ name: str
38
+ attribute_type: AttributeType
39
+ nominal_values: Optional[List[str]] = None
40
+
41
+
42
+ @dataclass
43
+ class TabularDataset(Dataset):
44
+ """
45
+ A tabular dataset consisting of two matrices `x` and `y`, storing the features of examples and their respective
46
+ ground truth, respectively.
47
+
48
+ Attributes:
49
+ x: A `lil_array`, shape `(num_examples, num_features)`, that stores the features of examples
50
+ y: A `lil_array`, shape `(num_examples, num_features)`, that stores the ground truth of examples
51
+ features: A list that contains all features in the dataset
52
+ outputs: A list that contains all outputs in the dataset
53
+ """
54
+ x: lil_array
55
+ y: lil_array
56
+ features: List[Attribute]
57
+ outputs: List[Attribute]
58
+
59
+ @property
60
+ def num_examples(self) -> int:
61
+ """
62
+ The number of examples in the dataset.
63
+ """
64
+ return self.x.shape[0]
65
+
66
+ @property
67
+ def num_features(self) -> int:
68
+ """
69
+ The number of features in the dataset.
70
+ """
71
+ return self.x.shape[1]
72
+
73
+ @property
74
+ def num_outputs(self) -> int:
75
+ """
76
+ The number of outputs in the dataset.
77
+ """
78
+ return self.y.shape[1]
79
+
80
+ @property
81
+ def has_sparse_features(self) -> bool:
82
+ """
83
+ True, if feature values in the dataset are sparse, False otherwise.
84
+ """
85
+ return is_sparse(self.x)
86
+
87
+ @property
88
+ def has_sparse_outputs(self) -> bool:
89
+ """
90
+ True, if the ground truth in the dataset is sparse, False otherwise.
91
+ """
92
+ return is_sparse(self.y)
93
+
94
+ def enforce_dense_features(self) -> 'TabularDataset':
95
+ """
96
+ Creates and returns a copy of this dataset, where the feature values have been converted into a dense format.
97
+
98
+ :return: The dataset that has been created
99
+ """
100
+ if self.has_sparse_features:
101
+ return replace(self, x=self.x.toarray())
102
+ return self
103
+
104
+ def enforce_dense_outputs(self) -> 'TabularDataset':
105
+ """
106
+ Creates and returns a copy of this dataset, where the ground truth has been converted into a dense format.
107
+
108
+ :return: The dataset that has been created
109
+ """
110
+ if self.has_sparse_outputs:
111
+ return replace(self, y=self.y.toarray())
112
+ return self
113
+
114
+ def get_num_features(self, *feature_types: AttributeType) -> int:
115
+ """
116
+ Returns the number of features with one out of a given set of types. If no types are given, all features are
117
+ counted.
118
+
119
+ :param feature_types: The types of the features to be counted
120
+ :return: The number of features of the given types
121
+ """
122
+ feature_types = set(feature_types)
123
+
124
+ if feature_types:
125
+ return reduce(lambda aggr, feature: aggr + (1 if feature.attribute_type in feature_types else 0),
126
+ self.features, 0)
127
+
128
+ return len(self.features)
129
+
130
+ def get_feature_indices(self, *feature_types: AttributeType) -> List[int]:
131
+ """
132
+ Returns a list that contains the indices of all features with one out of a given set of types (in ascending
133
+ order). If no types are given, all indices are returned.
134
+
135
+ :param feature_types: The types of the features whose indices should be returned
136
+ :return: A list that contains the indices of all features of the given types
137
+ """
138
+ feature_types = set(feature_types)
139
+ return [
140
+ i for i, feature in enumerate(self.features) if not feature_types or feature.attribute_type in feature_types
141
+ ]
@@ -0,0 +1,164 @@
1
+ """
2
+ Author: Michael Rapp (michael.rapp.ml@gmail.com)
3
+
4
+ Provides classes for performing experiments using the scikit-learn framework.
5
+ """
6
+ import logging as log
7
+
8
+ from dataclasses import replace
9
+ from functools import reduce
10
+ from typing import Any, Dict, Generator, Optional
11
+
12
+ from sklearn.base import BaseEstimator, clone
13
+
14
+ from mlrl.testbed_sklearn.experiments.dataset import TabularDataset
15
+ from mlrl.testbed_sklearn.experiments.output.characteristics.data.writer_data import DataCharacteristicsWriter
16
+ from mlrl.testbed_sklearn.experiments.output.characteristics.data.writer_prediction import \
17
+ PredictionCharacteristicsWriter
18
+ from mlrl.testbed_sklearn.experiments.output.dataset.writer_ground_truth import GroundTruthWriter
19
+ from mlrl.testbed_sklearn.experiments.output.dataset.writer_prediction import PredictionWriter
20
+ from mlrl.testbed_sklearn.experiments.output.evaluation.writer import EvaluationWriter
21
+ from mlrl.testbed_sklearn.experiments.output.label_vectors import LabelVectorWriter
22
+ from mlrl.testbed_sklearn.experiments.problem_domain import SkLearnProblem
23
+
24
+ from mlrl.testbed.experiments.dataset import Dataset
25
+ from mlrl.testbed.experiments.experiment import Experiment
26
+ from mlrl.testbed.experiments.input.dataset.splitters.splitter import DatasetSplitter
27
+ from mlrl.testbed.experiments.output.model.writer import ModelWriter
28
+ from mlrl.testbed.experiments.output.parameters.writer import ParameterWriter
29
+ from mlrl.testbed.experiments.problem_domain import ProblemDomain
30
+ from mlrl.testbed.experiments.state import ExperimentState, ParameterDict, PredictionState, TrainingState
31
+ from mlrl.testbed.experiments.timer import Timer
32
+
33
+
34
+ class SkLearnExperiment(Experiment):
35
+ """
36
+ An experiment that trains and evaluates a machine learning model using the scikit-learn framework.
37
+ """
38
+
39
+ class Builder(Experiment.Builder):
40
+ """
41
+ Allows to configure and create instances of the class `SkLearnExperiment`.
42
+ """
43
+
44
+ def __init__(self, problem_domain: ProblemDomain, dataset_splitter: DatasetSplitter):
45
+ """
46
+ :param problem_domain: The problem domain, the experiment should be concerned with
47
+ :param dataset_splitter: The method to be used for splitting the dataset into training and test datasets
48
+ """
49
+ super().__init__(problem_domain=problem_domain, dataset_splitter=dataset_splitter)
50
+ self.data_characteristics_writer = DataCharacteristicsWriter()
51
+ self.prediction_characteristics_writer = PredictionCharacteristicsWriter()
52
+ self.ground_truth_writer = GroundTruthWriter()
53
+ self.prediction_writer = PredictionWriter()
54
+ self.label_vector_writer = LabelVectorWriter()
55
+ self.evaluation_writer = EvaluationWriter()
56
+ self.model_writer = ModelWriter()
57
+ self.parameter_writer = ParameterWriter()
58
+ self.add_pre_training_output_writers(
59
+ self.data_characteristics_writer,
60
+ self.parameter_writer,
61
+ )
62
+ self.add_post_training_output_writers(
63
+ self.model_writer,
64
+ self.label_vector_writer,
65
+ )
66
+ self.add_prediction_output_writers(
67
+ self.prediction_characteristics_writer,
68
+ self.ground_truth_writer,
69
+ self.prediction_writer,
70
+ self.evaluation_writer,
71
+ )
72
+
73
+ def _create_experiment(self, problem_domain: ProblemDomain, dataset_splitter: DatasetSplitter) -> Experiment:
74
+ return SkLearnExperiment(problem_domain=problem_domain, dataset_splitter=dataset_splitter)
75
+
76
+ def __create_learner(self, parameters: ParameterDict) -> BaseEstimator:
77
+ learner = clone(self.problem_domain.base_learner)
78
+
79
+ if parameters:
80
+ learner.set_params(**parameters)
81
+ log.info('Successfully applied parameter setting: %s', parameters)
82
+
83
+ return learner
84
+
85
+ @staticmethod
86
+ def __check_for_parameter_changes(expected_parameters: Dict[str, Any], actual_parameters: Dict[str, Any]):
87
+ changes = []
88
+
89
+ for key, expected_value in expected_parameters.items():
90
+ expected_value = str(expected_value)
91
+ actual_value = str(actual_parameters[key])
92
+
93
+ if actual_value != expected_value:
94
+ changes.append((key, expected_value, actual_value))
95
+
96
+ if changes:
97
+ log.warning(
98
+ 'The loaded model\'s values for the following parameters differ from the expected configuration: %s',
99
+ reduce(
100
+ lambda aggr, change: aggr +
101
+ (', '
102
+ if aggr else '') + '"' + change[0] + '" is "' + change[2] + '" instead of "' + change[1] + '"',
103
+ changes, ''))
104
+
105
+ # pylint: disable=useless-parent-delegation
106
+ def __init__(self, problem_domain: SkLearnProblem, dataset_splitter: DatasetSplitter):
107
+ """
108
+ :param problem_domain: The problem domain, the experiment is concerned with
109
+ :param dataset_splitter: The method to be used for splitting the dataset into training and test datasets
110
+ """
111
+ super().__init__(problem_domain=problem_domain, dataset_splitter=dataset_splitter)
112
+
113
+ def _fit(self, estimator: BaseEstimator, dataset: TabularDataset,
114
+ fit_kwargs: Optional[Dict[str, Any]]) -> Timer.Duration:
115
+ """
116
+ May be overridden by subclasses in order to fit a scikit-learn estimator to a dataset.
117
+
118
+ :param estimator: A scikit-learn estimator
119
+ :param fit_kwargs: Optional keyword arguments to be passed to the estimator
120
+
121
+ """
122
+ fit_kwargs = fit_kwargs if fit_kwargs else {}
123
+
124
+ try:
125
+ start_time = Timer.start()
126
+ estimator.fit(dataset.x, dataset.y, **fit_kwargs)
127
+ return Timer.stop(start_time)
128
+ except ValueError as error:
129
+ if dataset.has_sparse_features:
130
+ return self._fit(estimator, dataset.enforce_dense_features(), fit_kwargs)
131
+ if dataset.has_sparse_outputs:
132
+ return self._fit(estimator, dataset.enforce_dense_outputs(), fit_kwargs)
133
+ raise error
134
+
135
+ def _train(self, learner: Optional[Any], parameters: ParameterDict, dataset: Dataset) -> TrainingState:
136
+ new_learner = self.__create_learner(parameters=parameters)
137
+
138
+ # Use existing model, if possible, otherwise train a new model...
139
+ if isinstance(learner, type(new_learner)):
140
+ self.__check_for_parameter_changes(expected_parameters=parameters, actual_parameters=learner.get_params())
141
+ return TrainingState(learner=learner)
142
+
143
+ log.info('Fitting model to %s training examples...', dataset.num_examples)
144
+ training_duration = self._fit(new_learner, dataset, fit_kwargs=self.problem_domain.fit_kwargs)
145
+ log.info('Successfully fit model in %s', training_duration)
146
+ return TrainingState(learner=new_learner, training_duration=training_duration)
147
+
148
+ def _predict(self, state: ExperimentState) -> Generator[PredictionState, None, None]:
149
+ dataset = state.dataset_as(self, TabularDataset)
150
+ learner = state.learner_as(self, BaseEstimator)
151
+
152
+ if dataset and learner:
153
+ try:
154
+ problem_domain = self.problem_domain
155
+ predict_kwargs = problem_domain.predict_kwargs
156
+ predict_kwargs = predict_kwargs if predict_kwargs else {}
157
+ predictor = problem_domain.predictor_factory.create()
158
+ dataset_type = state.dataset_type
159
+ yield from predictor.obtain_predictions(learner, dataset, dataset_type, **predict_kwargs)
160
+ except ValueError as error:
161
+ if dataset.has_sparse_features:
162
+ yield self._predict(replace(state, dataset=dataset.enforce_dense_features()))
163
+
164
+ raise error
File without changes
@@ -0,0 +1,47 @@
1
+ """
2
+ Author: Michael Rapp (michael.rapp.ml@gmail.com)
3
+
4
+ Provides classes that allow configuring the functionality to preprocess tabular datasets.
5
+ """
6
+ from argparse import Namespace
7
+ from typing import List, Set
8
+
9
+ from mlrl.testbed_sklearn.experiments.input.dataset.preprocessors.one_hot_encoder import OneHotEncoder
10
+
11
+ from mlrl.testbed.experiments.input.dataset.preprocessors.preprocessor import Preprocessor
12
+ from mlrl.testbed.extensions.extension import Extension
13
+
14
+ from mlrl.util.cli import Argument, BoolArgument
15
+
16
+
17
+ class PreprocessorExtension(Extension):
18
+ """
19
+ An extension that configures the functionality to preprocess tabular datasets.
20
+ """
21
+
22
+ ONE_HOT_ENCODING = BoolArgument(
23
+ '--one-hot-encoding',
24
+ default=False,
25
+ description='Whether one-hot-encoding should be used to encode nominal features or not.',
26
+ )
27
+
28
+ def _get_arguments(self) -> Set[Argument]:
29
+ """
30
+ See :func:`mlrl.testbed.extensions.extension.Extension._get_arguments`
31
+ """
32
+ return {self.ONE_HOT_ENCODING}
33
+
34
+ @staticmethod
35
+ def get_preprocessors(args: Namespace) -> List[Preprocessor]:
36
+ """
37
+ Returns the preprocessors to be used for preprocessing datasets according to the configuration.
38
+
39
+ :param args: The command line arguments specified by the user
40
+ :return: The preprocessors to be used
41
+ """
42
+ preprocessors = []
43
+
44
+ if PreprocessorExtension.ONE_HOT_ENCODING.get_value(args):
45
+ preprocessors.append(OneHotEncoder())
46
+
47
+ return preprocessors
@@ -0,0 +1,63 @@
1
+ """
2
+ Author: Michael Rapp (michael.rapp.ml@gmail.com)
3
+
4
+ Provides classes for preprocessing datasets.
5
+ """
6
+ import logging as log
7
+
8
+ from dataclasses import replace
9
+
10
+ from sklearn.compose import ColumnTransformer
11
+ from sklearn.preprocessing import OneHotEncoder as SkLearnOneHotEncoder
12
+
13
+ from mlrl.testbed_sklearn.experiments.dataset import AttributeType
14
+
15
+ from mlrl.testbed.experiments.dataset import Dataset
16
+ from mlrl.testbed.experiments.input.dataset.preprocessors.preprocessor import Preprocessor
17
+
18
+
19
+ class OneHotEncoder(Preprocessor):
20
+ """
21
+ Allows one-hot-encoding all nominal features contained in a tabular dataset, if any.
22
+ """
23
+
24
+ class Encoder(Preprocessor.Encoder):
25
+ """
26
+ Allows one-hot-encoding all nominal features contained in a tabular dataset.
27
+ """
28
+
29
+ def __init__(self):
30
+ self.encoder = None
31
+
32
+ def encode(self, dataset: Dataset) -> Dataset:
33
+ """
34
+ See :func:`mlrl.testbed.experiments.input.dataset.preprocessors.Preprocessor.Encoder.encode`
35
+ """
36
+ nominal_indices = dataset.get_feature_indices(AttributeType.NOMINAL)
37
+ num_nominal_features = len(nominal_indices)
38
+ log.info('Dataset contains %s nominal and %s numerical features.', num_nominal_features,
39
+ (len(dataset.features) - num_nominal_features))
40
+
41
+ if num_nominal_features > 0:
42
+ dataset = dataset.enforce_dense_features()
43
+
44
+ encoder = self.encoder
45
+
46
+ if not encoder:
47
+ log.info('Applying one-hot encoding...')
48
+ one_hot_encoder = SkLearnOneHotEncoder(handle_unknown='ignore', sparse_output=False)
49
+ transformers = [('one_hot_encoder', one_hot_encoder, nominal_indices)]
50
+ encoder = ColumnTransformer(transformers, remainder='passthrough')
51
+ encoder.fit(dataset.x, dataset.y)
52
+ self.encoder = encoder
53
+
54
+ return replace(dataset, x=encoder.transform(dataset.x), features=[])
55
+
56
+ log.debug('No need to apply one-hot encoding, as the dataset does not contain any nominal features.')
57
+ return dataset
58
+
59
+ def create_encoder(self) -> Preprocessor.Encoder:
60
+ """
61
+ See :func:`mlrl.testbed.experiments.input.dataset.preprocessors.Preprocessor.create_encoder`
62
+ """
63
+ return OneHotEncoder.Encoder()
@@ -0,0 +1,115 @@
1
+ """
2
+ Author: Michael Rapp (michael.rapp.ml@gmail.com)
3
+
4
+ Provides classes that allow configuring the functionality to split datasets into training and test datasets.
5
+ """
6
+ from argparse import Namespace
7
+ from typing import Set
8
+
9
+ from mlrl.testbed_sklearn.experiments.input.dataset.preprocessors.extension import PreprocessorExtension
10
+ from mlrl.testbed_sklearn.experiments.input.dataset.splitters.splitter_bipartition import BipartitionSplitter
11
+ from mlrl.testbed_sklearn.experiments.input.dataset.splitters.splitter_cross_validation import CrossValidationSplitter
12
+
13
+ from mlrl.testbed.experiments.input.dataset.splitters.extension import DatasetFileExtension
14
+ from mlrl.testbed.experiments.input.dataset.splitters.splitter import DatasetSplitter
15
+ from mlrl.testbed.experiments.input.dataset.splitters.splitter_no import NoSplitter
16
+ from mlrl.testbed.extensions.extension import Extension
17
+
18
+ from mlrl.util.cli import NONE, Argument, IntArgument, SetArgument
19
+ from mlrl.util.validation import assert_greater, assert_greater_or_equal, assert_less, assert_less_or_equal
20
+
21
+ VALUE_TRAIN_TEST = 'train-test'
22
+
23
+ OPTION_TEST_SIZE = 'test_size'
24
+
25
+ VALUE_CROSS_VALIDATION = 'cross-validation'
26
+
27
+ OPTION_NUM_FOLDS = 'num_folds'
28
+
29
+ OPTION_FIRST_FOLD = 'last_fold'
30
+
31
+ OPTION_LAST_FOLD = 'first_fold'
32
+
33
+
34
+ class DatasetSplitterExtension(Extension):
35
+ """
36
+ An extension that configures the functionality to split tabular datasets into training and test datasets.
37
+ """
38
+
39
+ RANDOM_STATE = IntArgument(
40
+ '--random-state',
41
+ description='The seed to be used by random number generators. Must be at least 1.',
42
+ default=1,
43
+ )
44
+
45
+ DATASET_SPLITTER = SetArgument(
46
+ '--data-split',
47
+ description='The strategy to be used for splitting the available data into training and test sets.',
48
+ default=VALUE_TRAIN_TEST,
49
+ values={
50
+ NONE: {},
51
+ VALUE_TRAIN_TEST: {OPTION_TEST_SIZE},
52
+ VALUE_CROSS_VALIDATION: {OPTION_NUM_FOLDS, OPTION_FIRST_FOLD, OPTION_LAST_FOLD}
53
+ },
54
+ )
55
+
56
+ def __init__(self, *dependencies: Extension):
57
+ """
58
+ :param dependencies: Other extensions, this extension depends on
59
+ """
60
+ super().__init__(PreprocessorExtension(), DatasetFileExtension(), *dependencies)
61
+
62
+ def _get_arguments(self) -> Set[Argument]:
63
+ """
64
+ See :func:`mlrl.testbed.extensions.extension.Extension._get_arguments`
65
+ """
66
+ return {self.RANDOM_STATE, self.DATASET_SPLITTER}
67
+
68
+ @staticmethod
69
+ def get_random_state(args: Namespace) -> int:
70
+ """
71
+ Returns the seed to be used by random number generators.
72
+
73
+ :param args: The command line arguments specified by the user
74
+ :return: The seed to be used
75
+ """
76
+ random_state = DatasetSplitterExtension.RANDOM_STATE.get_value(args)
77
+ assert_greater_or_equal(DatasetSplitterExtension.RANDOM_STATE.name, random_state, 1)
78
+ return random_state
79
+
80
+ @staticmethod
81
+ def get_dataset_splitter(args: Namespace) -> DatasetSplitter:
82
+ """
83
+ Returns the `DatasetSplitter` to be used for splitting datasets into training and test datasets according to the
84
+ configuration.
85
+
86
+ :param args: The command line arguments specified by the user
87
+ :return: The `DatasetSplitter` to be used
88
+ """
89
+ dataset_reader = DatasetFileExtension.get_dataset_reader(args)
90
+ dataset_reader.add_preprocessors(*PreprocessorExtension.get_preprocessors(args))
91
+ dataset_splitter, options = DatasetSplitterExtension.DATASET_SPLITTER.get_value(args)
92
+
93
+ if dataset_splitter == VALUE_CROSS_VALIDATION:
94
+ num_folds = options.get_int(OPTION_NUM_FOLDS, 10)
95
+ assert_greater_or_equal(OPTION_NUM_FOLDS, num_folds, 2)
96
+ first_fold = options.get_int(OPTION_FIRST_FOLD, 1)
97
+ assert_greater_or_equal(OPTION_FIRST_FOLD, first_fold, 1)
98
+ assert_less_or_equal(OPTION_FIRST_FOLD, first_fold, num_folds)
99
+ last_fold = options.get_int(OPTION_LAST_FOLD, num_folds)
100
+ assert_greater_or_equal(OPTION_LAST_FOLD, last_fold, first_fold)
101
+ assert_less_or_equal(OPTION_LAST_FOLD, last_fold, num_folds)
102
+ return CrossValidationSplitter(dataset_reader,
103
+ num_folds=num_folds,
104
+ first_fold=first_fold - 1,
105
+ last_fold=last_fold,
106
+ random_state=DatasetSplitterExtension.get_random_state(args))
107
+ if dataset_splitter == VALUE_TRAIN_TEST:
108
+ test_size = options.get_float(OPTION_TEST_SIZE, 0.33)
109
+ assert_greater(OPTION_TEST_SIZE, test_size, 0)
110
+ assert_less(OPTION_TEST_SIZE, test_size, 1)
111
+ return BipartitionSplitter(dataset_reader,
112
+ test_size=test_size,
113
+ random_state=DatasetSplitterExtension.get_random_state(args))
114
+
115
+ return NoSplitter(dataset_reader)