Skip to content

Commit 58e5bce

Browse files
committed
revert change
1 parent 7e8bd78 commit 58e5bce

3 files changed

Lines changed: 91 additions & 88 deletions

File tree

src/orion/client/experiment.py

Lines changed: 39 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -26,9 +26,9 @@
2626
from orion.core.worker.trial import Trial, TrialCM
2727
from orion.core.worker.trial_pacemaker import TrialPacemaker
2828
from orion.executor.base import Executor
29+
from orion.ext.extensions import OrionExtensionManager
2930
from orion.plotting.base import PlotAccessor
3031
from orion.storage.base import FailedUpdate
31-
from orion.ext.extensions import OrionExtensionManager
3232

3333
log = logging.getLogger(__name__)
3434

@@ -772,37 +772,6 @@ def workon(
772772

773773
return sum(trials)
774774

775-
def _optimize_trial(self, fct, trial, trial_arg, kwargs, worker_broken_trials, max_broken, on_error):
776-
kwargs.update(flatten(trial.params))
777-
778-
if trial_arg:
779-
kwargs[trial_arg] = trial
780-
781-
try:
782-
with self.extensions.trial(trial):
783-
results = self.executor.wait(
784-
[self.executor.submit(fct, **unflatten(kwargs))]
785-
)[0]
786-
self.observe(trial, results=results)
787-
except (KeyboardInterrupt, InvalidResult):
788-
raise
789-
except BaseException as e:
790-
if on_error is None or on_error(self, trial, e, worker_broken_trials):
791-
log.error(traceback.format_exc())
792-
worker_broken_trials += 1
793-
else:
794-
log.error(str(e))
795-
log.debug(traceback.format_exc())
796-
797-
if worker_broken_trials >= max_broken:
798-
raise BrokenExperiment(
799-
"Worker has reached broken trials threshold"
800-
)
801-
else:
802-
self.release(trial, status="broken")
803-
804-
return worker_broken_trials
805-
806775
def _optimize(
807776
self, fct, pool_size, max_trials, max_broken, trial_arg, on_error, **kwargs
808777
):
@@ -812,24 +781,44 @@ def _optimize(
812781
max_trials = min(max_trials, self.max_trials)
813782

814783
while not self.is_done and trials - worker_broken_trials < max_trials:
815-
try:
816-
with self.suggest(pool_size=pool_size) as trial:
817-
818-
worker_broken_trials = self._optimize_trial(
819-
fct,
820-
trial,
821-
trial_arg,
822-
kwargs,
823-
worker_broken_trials,
824-
max_broken,
825-
on_error
826-
)
827-
828-
except CompletedExperiment as e:
829-
log.warning(e)
830-
break
831-
832-
trials += 1
784+
try:
785+
with self.suggest(pool_size=pool_size) as trial:
786+
787+
kwargs.update(flatten(trial.params))
788+
789+
if trial_arg:
790+
kwargs[trial_arg] = trial
791+
792+
try:
793+
with self.extensions.trial(trial):
794+
results = self.executor.wait(
795+
[self.executor.submit(fct, **unflatten(kwargs))]
796+
)[0]
797+
self.observe(trial, results=results)
798+
except (KeyboardInterrupt, InvalidResult):
799+
raise
800+
except BaseException as e:
801+
if on_error is None or on_error(
802+
self, trial, e, worker_broken_trials
803+
):
804+
log.error(traceback.format_exc())
805+
worker_broken_trials += 1
806+
else:
807+
log.error(str(e))
808+
log.debug(traceback.format_exc())
809+
810+
if worker_broken_trials >= max_broken:
811+
raise BrokenExperiment(
812+
"Worker has reached broken trials threshold"
813+
)
814+
else:
815+
self.release(trial, status="broken")
816+
817+
except CompletedExperiment as e:
818+
log.warning(e)
819+
break
820+
821+
trials += 1
833822

834823
return trials
835824

src/orion/ext/extensions.py

Lines changed: 25 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ class EventDelegate:
1414
if false events are triggered as soon as broadcast is called
1515
if true the events will need to be triggered manually
1616
"""
17+
1718
def __init__(self, name, deferred=False) -> None:
1819
self.handlers = []
1920
self.deferred_calls = []
@@ -48,7 +49,9 @@ def _execute(self, args, kwargs):
4849
fun(*args, **kwargs)
4950
except Exception as err:
5051
if self.manager:
51-
self.manager.on_extension_error.broadcast(self.name, fun, err, args=(args, kwargs))
52+
self.manager.on_extension_error.broadcast(
53+
self.name, fun, err, args=(args, kwargs)
54+
)
5255

5356
def execute(self):
5457
"""Execute all our deferred handlers if any"""
@@ -86,41 +89,40 @@ class OrionExtensionManager:
8689

8790
def __init__(self):
8891
self._events = {}
89-
self._get_event('on_extension_error')
92+
self._get_event("on_extension_error")
9093

9194
# -- Trials
92-
self._get_event('new_trial')
93-
self._get_event('on_trial_error')
94-
self._get_event('end_trial')
95+
self._get_event("new_trial")
96+
self._get_event("on_trial_error")
97+
self._get_event("end_trial")
9598

9699
# -- Experiments
97-
self._get_event('start_experiment')
98-
self._get_event('on_experiment_error')
99-
self._get_event('end_experiment')
100+
self._get_event("start_experiment")
101+
self._get_event("on_experiment_error")
102+
self._get_event("end_experiment")
100103

101104
def experiment(self, *args, **kwargs):
102105
"""Initialize a context manager that will call start/error/end events automatically"""
103106
return _DelegateStartEnd(
104-
self.start_experiment,
105-
self.on_experiment_error,
106-
self.end_experiment,
107+
self._get_event("start_experiment"),
108+
self._get_event("on_experiment_error"),
109+
self._get_event("end_experiment"),
107110
*args,
108111
**kwargs
109112
)
110113

111114
def trial(self, *args, **kwargs):
112115
"""Initialize a context manager that will call start/error/end events automatically"""
113116
return _DelegateStartEnd(
114-
self.new_trial,
115-
self.on_trial_error,
116-
self.end_trial,
117+
self._get_event("new_trial"),
118+
self._get_event("on_trial_error"),
119+
self._get_event("end_trial"),
117120
*args,
118121
**kwargs
119122
)
120123

121-
def __getattr__(self, name):
122-
if name in self._events:
123-
return self._get_event(name)
124+
def broadcast(self, name, *args, **kwargs):
125+
return self._get_event(name).broadcast(*args, **kwargs)
124126

125127
def _get_event(self, key):
126128
"""Retrieve or generate a new event delegate"""
@@ -183,7 +185,9 @@ def on_extension_error(self, name, fun, exception, args):
183185
"""
184186
return
185187

186-
def on_trial_error(self, trial, exception_type, exception_value, exception_traceback):
188+
def on_trial_error(
189+
self, trial, exception_type, exception_value, exception_traceback
190+
):
187191
"""Called when a error occur during the optimization process"""
188192
return
189193

@@ -195,7 +199,9 @@ def end_trial(self, trial):
195199
"""Called when the trial finished"""
196200
return
197201

198-
def on_experiment_error(self, experiment, exception_type, exception_value, exception_traceback):
202+
def on_experiment_error(
203+
self, experiment, exception_type, exception_value, exception_traceback
204+
):
199205
"""Called when a error occur during the optimization process"""
200206
return
201207

@@ -206,4 +212,3 @@ def start_experiment(self, experiment):
206212
def end_experiment(self, experiment):
207213
"""Called at the end of the optimization process after the worker exits"""
208214
return
209-

tests/unittests/ext/test_extension.py

Lines changed: 27 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -42,28 +42,30 @@
4242
"params": [],
4343
}
4444

45+
4546
class OrionExtensionTest:
4647
"""Base orion extension interface you need to implement"""
48+
4749
def __init__(self) -> None:
4850
self.calls = defaultdict(int)
4951

5052
def on_experiment_error(self, *args, **kwargs):
51-
self.calls['on_experiment_error'] += 1
53+
self.calls["on_experiment_error"] += 1
5254

5355
def on_trial_error(self, *args, **kwargs):
54-
self.calls['on_trial_error'] += 1
56+
self.calls["on_trial_error"] += 1
5557

5658
def start_experiment(self, *args, **kwargs):
57-
self.calls['start_experiment'] += 1
59+
self.calls["start_experiment"] += 1
5860

5961
def new_trial(self, *args, **kwargs):
60-
self.calls['new_trial'] += 1
62+
self.calls["new_trial"] += 1
6163

6264
def end_trial(self, *args, **kwargs):
63-
self.calls['end_trial'] += 1
65+
self.calls["end_trial"] += 1
6466

6567
def end_experiment(self, *args, **kwargs):
66-
self.calls['end_experiment'] += 1
68+
self.calls["end_experiment"] += 1
6769

6870

6971
def test_client_extension():
@@ -88,31 +90,38 @@ def foo(x):
8890
n_broken = len(experiment.fetch_trials_by_status("broken"))
8991
n_reserved = len(experiment.fetch_trials_by_status("reserved"))
9092

91-
assert ext.calls['new_trial'] == n_trials + n_broken - n_reserved, 'all trials should have triggered callbacks'
92-
assert ext.calls['end_trial'] == n_trials + n_broken - n_reserved, 'all trials should have triggered callbacks'
93-
assert ext.calls['on_trial_error'] == n_broken, 'failed trial should be reported '
93+
assert (
94+
ext.calls["new_trial"] == n_trials + n_broken - n_reserved
95+
), "all trials should have triggered callbacks"
96+
assert (
97+
ext.calls["end_trial"] == n_trials + n_broken - n_reserved
98+
), "all trials should have triggered callbacks"
99+
assert (
100+
ext.calls["on_trial_error"] == n_broken
101+
), "failed trial should be reported "
94102

95-
assert ext.calls['start_experiment'] == 1, 'experiment should have started'
96-
assert ext.calls['end_experiment'] == 1, 'experiment should have ended'
97-
assert ext.calls['on_experiment_error'] == 1, 'failed experiment '
103+
assert ext.calls["start_experiment"] == 1, "experiment should have started"
104+
assert ext.calls["end_experiment"] == 1, "experiment should have ended"
105+
assert ext.calls["on_experiment_error"] == 1, "failed experiment "
98106

99107
unregistered_callback = client.extensions.unregister(ext)
100108
assert unregistered_callback == 6, "All ext callbacks got unregistered"
101109

102110

103111
class BadOrionExtensionTest:
104112
"""Base orion extension interface you need to implement"""
113+
105114
def __init__(self) -> None:
106115
self.calls = defaultdict(int)
107116

108117
def on_extension_error(self, name, fun, exception, args):
109-
self.calls['on_extension_error'] += 1
118+
self.calls["on_extension_error"] += 1
110119

111120
def on_experiment_error(self, *args, **kwargs):
112-
self.calls['on_experiment_error'] += 1
121+
self.calls["on_experiment_error"] += 1
113122

114123
def on_trial_error(self, *args, **kwargs):
115-
self.calls['on_trial_error'] += 1
124+
self.calls["on_trial_error"] += 1
116125

117126
def new_trial(self, *args, **kwargs):
118127
raise RuntimeError()
@@ -132,9 +141,9 @@ def foo(x):
132141
assert client.max_trials == MAX_TRIALS
133142
client.workon(foo, max_trials=MAX_TRIALS, max_broken=MAX_BROKEN)
134143

135-
assert ext.calls['on_trial_error'] == 0, 'Orion worked as expected'
136-
assert ext.calls['on_experiment_error'] == 0, 'Orion worked as expected'
137-
assert ext.calls['on_extension_error'] == 9, 'Extension error got reported'
144+
assert ext.calls["on_trial_error"] == 0, "Orion worked as expected"
145+
assert ext.calls["on_experiment_error"] == 0, "Orion worked as expected"
146+
assert ext.calls["on_extension_error"] == 9, "Extension error got reported"
138147

139148
unregistered_callback = client.extensions.unregister(ext)
140149
assert unregistered_callback == 4, "All ext callbacks got unregistered"

0 commit comments

Comments
 (0)