Skip to content

Commit 2a4018b

Browse files
author
PavelMarian
committed
removed default_backend
1 parent 8036b2a commit 2a4018b

18 files changed

Lines changed: 480 additions & 432 deletions

File tree

fedot/api/api_utils/api_composer.py

Lines changed: 220 additions & 215 deletions
Large diffs are not rendered by default.

fedot/api/api_utils/api_params_repository.py

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -25,9 +25,10 @@ class ApiParamsRepository:
2525

2626
STATIC_INDIVIDUAL_METADATA_KEYS = {'use_input_preprocessing'}
2727

28-
def __init__(self, task_type: TaskTypesEnum):
28+
def __init__(self, task_type: TaskTypesEnum, context: Optional[ExecutionContext] = None):
2929
self.task_type = task_type
3030
self.default_params = ApiParamsRepository.default_params_for_task(self.task_type)
31+
self.context = context
3132

3233
@staticmethod
3334
def default_params_for_task(task_type: TaskTypesEnum) -> dict:
@@ -75,14 +76,21 @@ def get_params_for_gp_algorithm_params(self, params: dict) -> dict:
7576
if params.get('genetic_scheme') == 'steady_state':
7677
gp_algorithm_params['genetic_scheme_type'] = GeneticSchemeTypesEnum.steady_state
7778

78-
# gp_algorithm_params['mutation_types'] = ApiParamsRepository._get_default_mutations(self.task_type, params)
79-
gp_algorithm_params['mutation_types'] = context.api_params_repository__get_default_mutations(self.task_type,
80-
params)
79+
gp_algorithm_params['mutation_types'] = ApiParamsRepository._get_default_mutations(self.task_type, params,
80+
self.context)
8181
gp_algorithm_params['seed'] = params['seed']
8282
return gp_algorithm_params
8383

84+
85+
@staticmethod
86+
def _get_default_mutations(task_type: TaskTypesEnum, params, context: Optional[ExecutionContext] = None) -> Sequence[MutationTypesEnum]:
87+
if context:
88+
return context.default_mutations.get_default_mutation(task_type, params)
89+
else:
90+
return _get_default_mutations_core(task_type, params)
91+
8492
@staticmethod
85-
def _get_default_mutations(task_type: TaskTypesEnum, params) -> Sequence[MutationTypesEnum]:
93+
def _get_default_mutations_core(task_type: TaskTypesEnum, params) -> Sequence[MutationTypesEnum]:
8694
mutations = [parameter_change_mutation,
8795
MutationTypesEnum.single_change,
8896
MutationTypesEnum.single_drop,

fedot/api/main.py

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,7 @@
5151
from fedot.utilities.define_metric_by_task import MetricByTask
5252
from fedot.utilities.memory import MemoryAnalytics
5353
from fedot.utilities.project_import_export import export_project_to_zip, import_project_from_zip
54-
from fedot.core.context.context import ExecutionContext
54+
from fedot.core.contex.context import ExecutionContext
5555

5656
NOT_FITTED_ERR_MSG = 'Model not fitted yet'
5757

@@ -107,11 +107,11 @@ def __init__(self,
107107
**composer_tuner_params
108108
):
109109

110+
self.context = ExecutionContext(extension_name=context)
111+
110112
set_random_seed(seed)
111113
self.log = self._init_logger(logging_level)
112114

113-
self.context = ExecutionContext(extension_name=context)
114-
115115
# Attributes for dealing with metrics, data sources and hyperparameters
116116
self.params = ApiParams(composer_tuner_params, problem, task_params, n_jobs, timeout, seed)
117117

@@ -303,7 +303,8 @@ def tune_tensordata(self,
303303
.with_n_jobs(common_tune_plan.n_jobs)
304304
.with_metric(common_tune_plan.metric)
305305
.with_iterations(iterations)
306-
.with_timeout(timeout))
306+
.with_timeout(timeout)
307+
.with_context(self.context))
307308
pipeline_tuner = getattr(pipeline_tuner, tune_plan.builder_method_name)(
308309
tensor_data if tune_plan.use_tensor_runtime else common_tune_plan.input_data
309310
)
@@ -377,6 +378,7 @@ def tune(self,
377378
.with_metric(tune_plan.metric)
378379
.with_iterations(iterations)
379380
.with_timeout(timeout)
381+
.with_context(self.context)
380382
.build(tune_input_data))
381383

382384
self.current_pipeline = pipeline_tuner.tune(self.current_pipeline, show_progress=show_progress)
@@ -682,7 +684,8 @@ def get_metrics(self,
682684
data_producer=lambda: (yield self.train_data, self.test_data),
683685
validation_blocks=validation_blocks,
684686
eval_n_jobs=self.params.n_jobs,
685-
do_unfit=False)
687+
do_unfit=False,
688+
context=self.context)
686689

687690
metrics = obj_eval.evaluate(self.current_pipeline).values
688691
metrics = {metric_name: round(abs(metric), rounding_order) for (metric_name, metric) in

fedot/core/composer/composer_builder.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -13,9 +13,9 @@
1313
from fedot.core.caching.operations_cache import OperationsCache
1414
from fedot.core.caching.preprocessing_cache import PreprocessingCache
1515
from fedot.core.caching.predictions_cache import PredictionsCache
16-
from fedot.core.context import ExecutionContext
1716
from fedot.core.composer.composer import Composer
1817
from fedot.core.composer.gp_composer.gp_composer import GPComposer
18+
from fedot.core.context.context import ExecutionContext
1919
from fedot.core.optimisers.objective.metrics_objective import MetricsObjective
2020
from fedot.core.pipelines.pipeline import Pipeline
2121
from fedot.core.pipelines.pipeline_composer_requirements import PipelineComposerRequirements
@@ -61,8 +61,9 @@ def __init__(self, task: Task):
6161

6262
self.context: Optional[ExecutionContext] = None
6363

64-
def with_context(self, context: ExecutionContext):
65-
self.context = context or ExecutionContext()
64+
def with_context(self, context):
65+
if context:
66+
self.context = context
6667
return self
6768

6869
def with_composer(self, composer_cls: Optional[Type[Composer]]):
@@ -118,8 +119,7 @@ def with_cache(self,
118119
@staticmethod
119120
def _get_default_composer_params(task: Task) -> PipelineComposerRequirements:
120121
# Get all available operations for task
121-
# operations = get_operations_for_task(task=task, mode='all')
122-
operations = self.context.operation_registry.get_operation_for_task(task=task, mode='all')
122+
operations = get_operations_for_task(task=task, mode='all')
123123
return PipelineComposerRequirements(primary=operations, secondary=operations)
124124

125125
def _get_default_graph_generation_params(self) -> GraphGenerationParams:
@@ -169,7 +169,6 @@ def build(self) -> Composer:
169169
self.composer_requirements,
170170
self.operations_cache,
171171
self.preprocessing_cache,
172-
self.predictions_cache,
173-
self.context)
172+
self.predictions_cache)
174173

175174
return composer

fedot/core/composer/gp_composer/gp_composer.py

Lines changed: 9 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -11,10 +11,10 @@
1111
from fedot.core.caching.operations_cache import OperationsCache
1212
from fedot.core.caching.predictions_cache import PredictionsCache
1313
from fedot.core.caching.preprocessing_cache import PreprocessingCache
14-
from fedot.core.context.context import ExecutionContext
1514
from fedot.core.composer.composer import Composer
1615
from fedot.core.data.data import InputData
1716
from fedot.core.data.multi_modal import MultiModalData
17+
from fedot.core.context.context import ExecutionContext
1818
from fedot.core.optimisers.objective.data_objective_eval import (
1919
PipelineObjectiveEvaluate,
2020
)
@@ -25,7 +25,6 @@
2525
)
2626
from fedot.core.utils import default_fedot_data_dir
2727

28-
from functools import partial
2928

3029
class GPComposer(Composer):
3130
"""
@@ -43,23 +42,23 @@ def __init__(self, optimizer: GraphOptimizer,
4342
operations_cache: Optional[OperationsCache] = None,
4443
preprocessing_cache: Optional[PreprocessingCache] = None,
4544
predictions_cache: Optional[PredictionsCache] = None,
46-
context: Optional[ExecutionContext] = None):
45+
context: Optional[str] = None,):
4746
super().__init__(optimizer, composer_requirements)
4847
self.composer_requirements = composer_requirements
4948
self.operations_cache: Optional[OperationsCache] = operations_cache
5049
self.preprocessing_cache: Optional[PreprocessingCache] = preprocessing_cache
5150
self.predictions_cache: Optional[PredictionsCache] = predictions_cache
5251

5352
self.best_models: Collection[Pipeline] = ()
54-
55-
self.context = context or ExecutionContext()
53+
self.context = ExecutionContext(extension_name=context)
5654

5755
def compose_pipeline(self, data: Union[InputData, MultiModalData]) -> Union[Pipeline, Sequence[Pipeline]]:
5856
# Define data source
5957
data_splitter = DataSourceSplitter(self.composer_requirements.cv_folds,
6058
shuffle=True)
61-
62-
data_producer = self.context.data_source_splitter_build(data_splitter, data)
59+
if self.context:
60+
data_splitter.build = self.context.data_source_splitter.build
61+
data_producer = data_splitter.build(data)
6362

6463
parallelization_mode = self.composer_requirements.parallelization_mode
6564
if parallelization_mode == 'populational':
@@ -77,10 +76,10 @@ def compose_pipeline(self, data: Union[InputData, MultiModalData]) -> Union[Pipe
7776
preprocessing_cache=self.preprocessing_cache,
7877
predictions_cache=self.predictions_cache,
7978
validation_blocks=data_splitter.validation_blocks,
80-
eval_n_jobs=n_jobs_for_evaluation)
79+
eval_n_jobs=n_jobs_for_evaluation,
80+
context=context)
8181

82-
# objective_function = objective_evaluator.evaluate
83-
objective_function = partial(self.context.evaluator_evaluate, objective_evaluator)
82+
objective_function = objective_evaluator.evaluate
8483

8584
# Define callback for computing intermediate metrics if needed
8685
if self.composer_requirements.collect_intermediate_metric:

fedot/core/context/context.py

Lines changed: 18 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -1,69 +1,33 @@
11
from typing import Dict, Any, Optional, Callable
22
from fedot.extensions.registry import get_registered_extension, register_extension
3-
from fedot.core.context.industrial_manifest import FEDOT_INDUSTRIAL_MANIFEST
43

5-
register_extension(FEDOT_INDUSTRIAL_MANIFEST)
64

75
class ExecutionContext:
8-
def __init__(self, extension_name: str = "core", extra_params: Optional[Dict[str, Any]] = None):
6+
def __init__(self, extension_name: str = "industrial", extra_params: Optional[Dict[str, Any]] = None):
97
self.extension_name = extension_name
108
self.extra_params = extra_params or {}
119
self._instances: Dict[str, Any] = {}
1210
self._overridden: Dict[str, Any] = {}
1311

14-
self._manifest = None
15-
if extension_name != "core":
16-
from fedot.extensions.registry import _REGISTERED_EXTENSIONS
17-
manifest = _REGISTERED_EXTENSIONS.get(extension_name)
18-
if manifest is None:
19-
raise ValueError(f"Extension '{extension_name}' not registered")
20-
self._manifest = manifest
21-
22-
self._core_implementations = self._get_core_implementations()
23-
24-
self._protocol_classes = self._core_implementations.copy()
25-
if self._manifest and self._manifest.protocols:
26-
self._protocol_classes.update(self._manifest.protocols)
27-
28-
def _get_core_implementations(self) -> Dict[str, Callable]:
29-
from fedot.core.context.default_backend import (
30-
CoreSplitter, CoreDataMerger, CoreImageMerger,
31-
CoreTSMerger, CoreTextMerger, CoreTuner,
32-
CoreDataSourceSplitter, CoreOperationPredict,
33-
CoreLaggedTransformer, CoreTopologicalFeatures,
34-
CoreTsSmoothing, CoreApiComposerTune, CoreReproduction,
35-
CoreSearchSpace, CoreDefaultMutations, CoreEvaluator
36-
)
37-
return {
38-
"splitter": CoreSplitter,
39-
"data_merger": CoreDataMerger,
40-
"image_merger": CoreImageMerger,
41-
"ts_merger": CoreTSMerger,
42-
"text_merger": CoreTextMerger,
43-
"tuner_class": CoreTuner,
44-
"data_source_splitter": CoreDataSourceSplitter,
45-
"operation_predict": CoreOperationPredict,
46-
"lagged_transformer": CoreLaggedTransformer,
47-
"topological_features": CoreTopologicalFeatures,
48-
"ts_smoothing": CoreTsSmoothing,
49-
"api_composer_tune": CoreApiComposerTune,
50-
"reproduction": CoreReproduction,
51-
"search_space": CoreSearchSpace,
52-
"default_mutations": CoreDefaultMutations,
53-
"evaluator": CoreEvaluator,
54-
}
12+
manifest = get_registered_extension(extension_name)
13+
if manifest is None:
14+
raise ValueError(f"Extension '{extension_name}' not registered")
15+
self._manifest = manifest
5516

56-
def _get_protocol_class(self, protocol_name: str) -> Callable:
57-
if self._manifest and self._manifest.protocols:
58-
if protocol_name in self._manifest.protocols:
59-
return self._manifest.protocols[protocol_name]
60-
61-
if protocol_name in self._core_implementations:
62-
return self._core_implementations[protocol_name]
17+
self._protocol_classes = self._manifest.protocols or {}
6318

19+
def _get_protocol_class(self, protocol_name: str) -> Callable:
20+
if protocol_name in self._protocol_classes:
21+
return self._protocol_classes[protocol_name]
6422
raise ValueError(f"No implementation for protocol '{protocol_name}'")
6523

6624
def _get_instance(self, protocol_name: str) -> Any:
25+
if protocol_name in self._overridden:
26+
override = self._overridden[protocol_name]
27+
if isinstance(override, type):
28+
return override(**self.extra_params)
29+
return override
30+
6731
if protocol_name not in self._instances:
6832
protocol_class = self._get_protocol_class(protocol_name)
6933
self._instances[protocol_name] = protocol_class(**self.extra_params)
@@ -135,7 +99,7 @@ def evaluator(self):
13599

136100
def __setattr__(self, name: str, value: Any) -> None:
137101
if name in ('extra_params', '_instances', '_overridden', '_protocol_classes',
138-
'_manifest', '_core_implementations', 'extension_name'):
102+
'_manifest', 'extension_name'):
139103
super().__setattr__(name, value)
140104
else:
141105
self._overridden[name] = value
@@ -144,7 +108,7 @@ def __getattr__(self, name: str):
144108
if name in self._overridden:
145109
return self._overridden[name]
146110

147-
if name in ('_protocol_classes', '_instances', '_manifest', '_core_implementations'):
111+
if name in ('_protocol_classes', '_instances', '_manifest'):
148112
return super().__getattribute__(name)
149113

150-
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
114+
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")

fedot/core/context/factories.py

Lines changed: 0 additions & 60 deletions
This file was deleted.

0 commit comments

Comments
 (0)