hydra-optuna-sweeper 1.3.0.dev0__tar.gz → 1.4.0.dev4__tar.gz
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.
- hydra_optuna_sweeper-1.4.0.dev4/MANIFEST.in +3 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/PKG-INFO +21 -7
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_optuna_sweeper.egg-info/PKG-INFO +21 -7
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_optuna_sweeper.egg-info/SOURCES.txt +3 -1
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_optuna_sweeper.egg-info/requires.txt +1 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_plugins/hydra_optuna_sweeper/__init__.py +1 -1
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_plugins/hydra_optuna_sweeper/_impl.py +4 -6
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_plugins/hydra_optuna_sweeper/config.py +0 -1
- hydra_optuna_sweeper-1.4.0.dev4/hydra_plugins/hydra_optuna_sweeper/py.typed +0 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/setup.py +7 -5
- hydra_optuna_sweeper-1.4.0.dev4/tests/test_optuna_sweeper_plugin.py +402 -0
- hydra-optuna-sweeper-1.3.0.dev0/MANIFEST.in +0 -3
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/README.md +0 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_optuna_sweeper.egg-info/dependency_links.txt +0 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_optuna_sweeper.egg-info/top_level.txt +0 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/hydra_plugins/hydra_optuna_sweeper/optuna_sweeper.py +0 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/pyproject.toml +0 -0
- {hydra-optuna-sweeper-1.3.0.dev0 → hydra_optuna_sweeper-1.4.0.dev4}/setup.cfg +0 -0
|
@@ -1,20 +1,34 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
2
|
Name: hydra-optuna-sweeper
|
|
3
|
-
Version: 1.
|
|
3
|
+
Version: 1.4.0.dev4
|
|
4
4
|
Summary: Hydra Optuna Sweeper plugin
|
|
5
5
|
Home-page: https://github.com/facebookresearch/hydra/
|
|
6
6
|
Author: Toshihiko Yanase, Hiroyuki Vincent Yamazaki
|
|
7
7
|
Author-email: toshihiko.yanase@gmail.com, hiroyuki.vincent.yamazaki@gmail.com
|
|
8
|
-
|
|
9
|
-
Classifier: Programming Language :: Python :: 3.6
|
|
10
|
-
Classifier: Programming Language :: Python :: 3.7
|
|
11
|
-
Classifier: Programming Language :: Python :: 3.8
|
|
12
|
-
Classifier: Programming Language :: Python :: 3.9
|
|
8
|
+
License: MIT
|
|
13
9
|
Classifier: Programming Language :: Python :: 3.10
|
|
10
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
11
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
12
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.14
|
|
14
14
|
Classifier: Operating System :: POSIX :: Linux
|
|
15
15
|
Classifier: Operating System :: MacOS
|
|
16
16
|
Classifier: Development Status :: 4 - Beta
|
|
17
|
+
Requires-Python: >=3.10
|
|
17
18
|
Description-Content-Type: text/markdown
|
|
19
|
+
Requires-Dist: hydra-core>=1.1.0.dev7
|
|
20
|
+
Requires-Dist: optuna<3.0.0,>=2.10.0
|
|
21
|
+
Requires-Dist: sqlalchemy~=1.3.0
|
|
22
|
+
Dynamic: author
|
|
23
|
+
Dynamic: author-email
|
|
24
|
+
Dynamic: classifier
|
|
25
|
+
Dynamic: description
|
|
26
|
+
Dynamic: description-content-type
|
|
27
|
+
Dynamic: home-page
|
|
28
|
+
Dynamic: license
|
|
29
|
+
Dynamic: requires-dist
|
|
30
|
+
Dynamic: requires-python
|
|
31
|
+
Dynamic: summary
|
|
18
32
|
|
|
19
33
|
# Hydra Optuna Sweeper
|
|
20
34
|
|
|
@@ -1,20 +1,34 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
2
|
Name: hydra-optuna-sweeper
|
|
3
|
-
Version: 1.
|
|
3
|
+
Version: 1.4.0.dev4
|
|
4
4
|
Summary: Hydra Optuna Sweeper plugin
|
|
5
5
|
Home-page: https://github.com/facebookresearch/hydra/
|
|
6
6
|
Author: Toshihiko Yanase, Hiroyuki Vincent Yamazaki
|
|
7
7
|
Author-email: toshihiko.yanase@gmail.com, hiroyuki.vincent.yamazaki@gmail.com
|
|
8
|
-
|
|
9
|
-
Classifier: Programming Language :: Python :: 3.6
|
|
10
|
-
Classifier: Programming Language :: Python :: 3.7
|
|
11
|
-
Classifier: Programming Language :: Python :: 3.8
|
|
12
|
-
Classifier: Programming Language :: Python :: 3.9
|
|
8
|
+
License: MIT
|
|
13
9
|
Classifier: Programming Language :: Python :: 3.10
|
|
10
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
11
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
12
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.14
|
|
14
14
|
Classifier: Operating System :: POSIX :: Linux
|
|
15
15
|
Classifier: Operating System :: MacOS
|
|
16
16
|
Classifier: Development Status :: 4 - Beta
|
|
17
|
+
Requires-Python: >=3.10
|
|
17
18
|
Description-Content-Type: text/markdown
|
|
19
|
+
Requires-Dist: hydra-core>=1.1.0.dev7
|
|
20
|
+
Requires-Dist: optuna<3.0.0,>=2.10.0
|
|
21
|
+
Requires-Dist: sqlalchemy~=1.3.0
|
|
22
|
+
Dynamic: author
|
|
23
|
+
Dynamic: author-email
|
|
24
|
+
Dynamic: classifier
|
|
25
|
+
Dynamic: description
|
|
26
|
+
Dynamic: description-content-type
|
|
27
|
+
Dynamic: home-page
|
|
28
|
+
Dynamic: license
|
|
29
|
+
Dynamic: requires-dist
|
|
30
|
+
Dynamic: requires-python
|
|
31
|
+
Dynamic: summary
|
|
18
32
|
|
|
19
33
|
# Hydra Optuna Sweeper
|
|
20
34
|
|
|
@@ -10,4 +10,6 @@ hydra_optuna_sweeper.egg-info/top_level.txt
|
|
|
10
10
|
hydra_plugins/hydra_optuna_sweeper/__init__.py
|
|
11
11
|
hydra_plugins/hydra_optuna_sweeper/_impl.py
|
|
12
12
|
hydra_plugins/hydra_optuna_sweeper/config.py
|
|
13
|
-
hydra_plugins/hydra_optuna_sweeper/optuna_sweeper.py
|
|
13
|
+
hydra_plugins/hydra_optuna_sweeper/optuna_sweeper.py
|
|
14
|
+
hydra_plugins/hydra_optuna_sweeper/py.typed
|
|
15
|
+
tests/test_optuna_sweeper_plugin.py
|
|
@@ -49,7 +49,7 @@ log = logging.getLogger(__name__)
|
|
|
49
49
|
|
|
50
50
|
|
|
51
51
|
def create_optuna_distribution_from_config(
|
|
52
|
-
config: MutableMapping[str, Any]
|
|
52
|
+
config: MutableMapping[str, Any],
|
|
53
53
|
) -> BaseDistribution:
|
|
54
54
|
kwargs = dict(config)
|
|
55
55
|
if isinstance(config["type"], str):
|
|
@@ -192,13 +192,11 @@ class OptunaSweeperImpl(Sweeper):
|
|
|
192
192
|
)
|
|
193
193
|
else:
|
|
194
194
|
deprecation_warning(
|
|
195
|
-
message=dedent(
|
|
196
|
-
f"""\
|
|
195
|
+
message=dedent(f"""\
|
|
197
196
|
`hydra.sweeper.search_space` is deprecated and will be removed in the next major release.
|
|
198
197
|
Please configure with `hydra.sweeper.params`.
|
|
199
198
|
{url}
|
|
200
|
-
"""
|
|
201
|
-
),
|
|
199
|
+
"""),
|
|
202
200
|
)
|
|
203
201
|
self.search_space_distributions = {
|
|
204
202
|
str(x): create_optuna_distribution_from_config(y)
|
|
@@ -290,7 +288,7 @@ class OptunaSweeperImpl(Sweeper):
|
|
|
290
288
|
|
|
291
289
|
is_grid_sampler = (
|
|
292
290
|
isinstance(self.sampler, functools.partial)
|
|
293
|
-
and self.sampler.func == optuna.samplers.GridSampler
|
|
291
|
+
and self.sampler.func == optuna.samplers.GridSampler
|
|
294
292
|
)
|
|
295
293
|
|
|
296
294
|
(
|
|
File without changes
|
|
@@ -11,24 +11,26 @@ setup(
|
|
|
11
11
|
author="Toshihiko Yanase, Hiroyuki Vincent Yamazaki",
|
|
12
12
|
author_email="toshihiko.yanase@gmail.com, hiroyuki.vincent.yamazaki@gmail.com",
|
|
13
13
|
description="Hydra Optuna Sweeper plugin",
|
|
14
|
+
license="MIT",
|
|
14
15
|
long_description=(Path(__file__).parent / "README.md").read_text(),
|
|
15
16
|
long_description_content_type="text/markdown",
|
|
16
17
|
url="https://github.com/facebookresearch/hydra/",
|
|
17
18
|
packages=find_namespace_packages(include=["hydra_plugins.*"]),
|
|
18
19
|
classifiers=[
|
|
19
|
-
"License :: OSI Approved :: MIT License",
|
|
20
|
-
"Programming Language :: Python :: 3.6",
|
|
21
|
-
"Programming Language :: Python :: 3.7",
|
|
22
|
-
"Programming Language :: Python :: 3.8",
|
|
23
|
-
"Programming Language :: Python :: 3.9",
|
|
24
20
|
"Programming Language :: Python :: 3.10",
|
|
21
|
+
"Programming Language :: Python :: 3.11",
|
|
22
|
+
"Programming Language :: Python :: 3.12",
|
|
23
|
+
"Programming Language :: Python :: 3.13",
|
|
24
|
+
"Programming Language :: Python :: 3.14",
|
|
25
25
|
"Operating System :: POSIX :: Linux",
|
|
26
26
|
"Operating System :: MacOS",
|
|
27
27
|
"Development Status :: 4 - Beta",
|
|
28
28
|
],
|
|
29
|
+
python_requires=">=3.10",
|
|
29
30
|
install_requires=[
|
|
30
31
|
"hydra-core>=1.1.0.dev7",
|
|
31
32
|
"optuna>=2.10.0,<3.0.0",
|
|
33
|
+
"sqlalchemy~=1.3.0", # TODO: Unpin when upgrading to optuna v3.0
|
|
32
34
|
],
|
|
33
35
|
include_package_data=True,
|
|
34
36
|
)
|
|
@@ -0,0 +1,402 @@
|
|
|
1
|
+
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved
|
|
2
|
+
import os
|
|
3
|
+
import sys
|
|
4
|
+
from functools import partial
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any, List, Optional
|
|
7
|
+
|
|
8
|
+
import optuna
|
|
9
|
+
from hydra.core.override_parser.overrides_parser import OverridesParser
|
|
10
|
+
from hydra.core.plugins import Plugins
|
|
11
|
+
from hydra.plugins.sweeper import Sweeper
|
|
12
|
+
from hydra.test_utils.test_utils import (
|
|
13
|
+
TSweepRunner,
|
|
14
|
+
chdir_plugin_root,
|
|
15
|
+
run_process,
|
|
16
|
+
run_python_script,
|
|
17
|
+
)
|
|
18
|
+
from omegaconf import DictConfig, OmegaConf
|
|
19
|
+
from optuna.distributions import (
|
|
20
|
+
BaseDistribution,
|
|
21
|
+
CategoricalDistribution,
|
|
22
|
+
DiscreteUniformDistribution,
|
|
23
|
+
IntLogUniformDistribution,
|
|
24
|
+
IntUniformDistribution,
|
|
25
|
+
LogUniformDistribution,
|
|
26
|
+
UniformDistribution,
|
|
27
|
+
)
|
|
28
|
+
from optuna.samplers import RandomSampler
|
|
29
|
+
from pytest import mark, warns
|
|
30
|
+
|
|
31
|
+
from hydra_plugins.hydra_optuna_sweeper import _impl
|
|
32
|
+
from hydra_plugins.hydra_optuna_sweeper._impl import OptunaSweeperImpl
|
|
33
|
+
from hydra_plugins.hydra_optuna_sweeper.config import Direction
|
|
34
|
+
from hydra_plugins.hydra_optuna_sweeper.optuna_sweeper import OptunaSweeper
|
|
35
|
+
|
|
36
|
+
chdir_plugin_root()
|
|
37
|
+
|
|
38
|
+
SQLALCHEMY_UTC_WARNING_FILTER = (
|
|
39
|
+
r"-W ignore:datetime.datetime.utcfromtimestamp() is deprecated:"
|
|
40
|
+
r"DeprecationWarning:sqlalchemy.sql.sqltypes"
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def run_optuna_script(cmd: List[str]) -> None:
|
|
45
|
+
run_python_script([SQLALCHEMY_UTC_WARNING_FILTER, *cmd])
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def test_discovery() -> None:
|
|
49
|
+
assert OptunaSweeper.__name__ in [
|
|
50
|
+
x.__name__ for x in Plugins.instance().discover(Sweeper)
|
|
51
|
+
]
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def check_distribution(expected: BaseDistribution, actual: BaseDistribution) -> None:
|
|
55
|
+
if not isinstance(expected, CategoricalDistribution):
|
|
56
|
+
assert expected == actual
|
|
57
|
+
return
|
|
58
|
+
|
|
59
|
+
assert isinstance(actual, CategoricalDistribution)
|
|
60
|
+
# shuffle() will randomize the order of items in choices.
|
|
61
|
+
assert set(expected.choices) == set(actual.choices)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
@mark.parametrize(
|
|
65
|
+
"input, expected",
|
|
66
|
+
[
|
|
67
|
+
(
|
|
68
|
+
{"type": "categorical", "choices": [1, 2, 3]},
|
|
69
|
+
CategoricalDistribution([1, 2, 3]),
|
|
70
|
+
),
|
|
71
|
+
({"type": "int", "low": 0, "high": 10}, IntUniformDistribution(0, 10)),
|
|
72
|
+
(
|
|
73
|
+
{"type": "int", "low": 0, "high": 10, "step": 2},
|
|
74
|
+
IntUniformDistribution(0, 10, step=2),
|
|
75
|
+
),
|
|
76
|
+
({"type": "int", "low": 0, "high": 5}, IntUniformDistribution(0, 5)),
|
|
77
|
+
(
|
|
78
|
+
{"type": "int", "low": 1, "high": 100, "log": True},
|
|
79
|
+
IntLogUniformDistribution(1, 100),
|
|
80
|
+
),
|
|
81
|
+
({"type": "float", "low": 0, "high": 1}, UniformDistribution(0, 1)),
|
|
82
|
+
(
|
|
83
|
+
{"type": "float", "low": 0, "high": 10, "step": 2},
|
|
84
|
+
DiscreteUniformDistribution(0, 10, 2),
|
|
85
|
+
),
|
|
86
|
+
(
|
|
87
|
+
{"type": "float", "low": 1, "high": 100, "log": True},
|
|
88
|
+
LogUniformDistribution(1, 100),
|
|
89
|
+
),
|
|
90
|
+
],
|
|
91
|
+
)
|
|
92
|
+
def test_create_optuna_distribution_from_config(input: Any, expected: Any) -> None:
|
|
93
|
+
actual = _impl.create_optuna_distribution_from_config(input)
|
|
94
|
+
check_distribution(expected, actual)
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
@mark.parametrize(
|
|
98
|
+
"input, expected",
|
|
99
|
+
[
|
|
100
|
+
("key=choice(1,2)", CategoricalDistribution([1, 2])),
|
|
101
|
+
("key=choice(true, false)", CategoricalDistribution([True, False])),
|
|
102
|
+
("key=choice('hello', 'world')", CategoricalDistribution(["hello", "world"])),
|
|
103
|
+
("key=shuffle(range(1,3))", CategoricalDistribution((1, 2))),
|
|
104
|
+
("key=range(1,3)", IntUniformDistribution(1, 3)),
|
|
105
|
+
("key=interval(1, 5)", UniformDistribution(1, 5)),
|
|
106
|
+
("key=int(interval(1, 5))", IntUniformDistribution(1, 5)),
|
|
107
|
+
("key=tag(log, interval(1, 5))", LogUniformDistribution(1, 5)),
|
|
108
|
+
("key=tag(log, int(interval(1, 5)))", IntLogUniformDistribution(1, 5)),
|
|
109
|
+
("key=range(0.5, 5.5, step=1)", DiscreteUniformDistribution(0.5, 5.5, 1)),
|
|
110
|
+
],
|
|
111
|
+
)
|
|
112
|
+
def test_create_optuna_distribution_from_override(input: Any, expected: Any) -> None:
|
|
113
|
+
parser = OverridesParser.create()
|
|
114
|
+
parsed = parser.parse_overrides([input])[0]
|
|
115
|
+
actual = _impl.create_optuna_distribution_from_override(parsed)
|
|
116
|
+
check_distribution(expected, actual)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
@mark.parametrize(
|
|
120
|
+
"input, expected",
|
|
121
|
+
[
|
|
122
|
+
(["key=choice(1,2)"], ({"key": CategoricalDistribution([1, 2])}, {})),
|
|
123
|
+
(["key=5"], ({}, {"key": "5"})),
|
|
124
|
+
(
|
|
125
|
+
["key1=choice(1,2)", "key2=5"],
|
|
126
|
+
({"key1": CategoricalDistribution([1, 2])}, {"key2": "5"}),
|
|
127
|
+
),
|
|
128
|
+
(
|
|
129
|
+
["key1=choice(1,2)", "key2=5", "key3=range(1,3)"],
|
|
130
|
+
(
|
|
131
|
+
{
|
|
132
|
+
"key1": CategoricalDistribution([1, 2]),
|
|
133
|
+
"key3": IntUniformDistribution(1, 3),
|
|
134
|
+
},
|
|
135
|
+
{"key2": "5"},
|
|
136
|
+
),
|
|
137
|
+
),
|
|
138
|
+
],
|
|
139
|
+
)
|
|
140
|
+
def test_create_params_from_overrides(input: Any, expected: Any) -> None:
|
|
141
|
+
actual = _impl.create_params_from_overrides(input)
|
|
142
|
+
assert actual == expected
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def test_launch_jobs(hydra_sweep_runner: TSweepRunner) -> None:
|
|
146
|
+
sweep = hydra_sweep_runner(
|
|
147
|
+
calling_file=None,
|
|
148
|
+
calling_module="hydra.test_utils.a_module",
|
|
149
|
+
config_path="configs",
|
|
150
|
+
config_name="compose.yaml",
|
|
151
|
+
task_function=None,
|
|
152
|
+
overrides=[
|
|
153
|
+
"hydra/sweeper=optuna",
|
|
154
|
+
"hydra/launcher=basic",
|
|
155
|
+
"hydra.sweeper.n_trials=8",
|
|
156
|
+
"hydra.sweeper.n_jobs=3",
|
|
157
|
+
],
|
|
158
|
+
)
|
|
159
|
+
with sweep:
|
|
160
|
+
assert sweep.returns is None
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
@mark.parametrize("with_commandline", (True, False))
|
|
164
|
+
def test_optuna_example(with_commandline: bool, tmpdir: Path) -> None:
|
|
165
|
+
storage = "sqlite:///" + os.path.join(str(tmpdir), "test.db")
|
|
166
|
+
study_name = "test-optuna-example"
|
|
167
|
+
cmd = [
|
|
168
|
+
"example/sphere.py",
|
|
169
|
+
"--multirun",
|
|
170
|
+
"hydra.sweep.dir=" + str(tmpdir),
|
|
171
|
+
"hydra.job.chdir=True",
|
|
172
|
+
"hydra.sweeper.n_trials=20",
|
|
173
|
+
"hydra.sweeper.n_jobs=1",
|
|
174
|
+
f"hydra.sweeper.storage={storage}",
|
|
175
|
+
f"hydra.sweeper.study_name={study_name}",
|
|
176
|
+
"hydra/sweeper/sampler=tpe",
|
|
177
|
+
"hydra.sweeper.sampler.seed=123",
|
|
178
|
+
"~z",
|
|
179
|
+
]
|
|
180
|
+
if with_commandline:
|
|
181
|
+
cmd += [
|
|
182
|
+
"x=choice(0, 1, 2)",
|
|
183
|
+
"y=0", # Fixed parameter
|
|
184
|
+
]
|
|
185
|
+
run_optuna_script(cmd)
|
|
186
|
+
returns = OmegaConf.load(f"{tmpdir}/optimization_results.yaml")
|
|
187
|
+
study = optuna.load_study(storage=storage, study_name=study_name)
|
|
188
|
+
best_trial = study.best_trial
|
|
189
|
+
assert isinstance(returns, DictConfig)
|
|
190
|
+
assert returns.name == "optuna"
|
|
191
|
+
assert returns["best_params"]["x"] == best_trial.params["x"]
|
|
192
|
+
if with_commandline:
|
|
193
|
+
assert "y" not in returns["best_params"]
|
|
194
|
+
assert "y" not in best_trial.params
|
|
195
|
+
else:
|
|
196
|
+
assert returns["best_params"]["y"] == best_trial.params["y"]
|
|
197
|
+
assert returns["best_value"] == best_trial.value
|
|
198
|
+
# Check the search performance of the TPE sampler.
|
|
199
|
+
# The threshold is the 95th percentile calculated with 1000 different seed values
|
|
200
|
+
# to make the test robust against the detailed implementation of the sampler.
|
|
201
|
+
# See https://github.com/facebookresearch/hydra/pull/1746#discussion_r681549830.
|
|
202
|
+
assert returns["best_value"] <= 2.27
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
@mark.parametrize("num_trials", (10, 1))
|
|
206
|
+
def test_example_with_grid_sampler(
|
|
207
|
+
tmpdir: Path,
|
|
208
|
+
num_trials: int,
|
|
209
|
+
) -> None:
|
|
210
|
+
storage = "sqlite:///" + os.path.join(str(tmpdir), "test.db")
|
|
211
|
+
study_name = "test-grid-sampler"
|
|
212
|
+
cmd = [
|
|
213
|
+
"example/sphere.py",
|
|
214
|
+
"--multirun",
|
|
215
|
+
"--config-dir=tests/conf",
|
|
216
|
+
"--config-name=test_grid",
|
|
217
|
+
"hydra.sweep.dir=" + str(tmpdir),
|
|
218
|
+
"hydra.job.chdir=False",
|
|
219
|
+
f"hydra.sweeper.n_trials={num_trials}",
|
|
220
|
+
"hydra.sweeper.n_jobs=1",
|
|
221
|
+
f"hydra.sweeper.storage={storage}",
|
|
222
|
+
f"hydra.sweeper.study_name={study_name}",
|
|
223
|
+
]
|
|
224
|
+
run_optuna_script(cmd)
|
|
225
|
+
returns = OmegaConf.load(f"{tmpdir}/optimization_results.yaml")
|
|
226
|
+
assert isinstance(returns, DictConfig)
|
|
227
|
+
bv, bx, by, bz = (
|
|
228
|
+
returns["best_value"],
|
|
229
|
+
returns["best_params"]["x"],
|
|
230
|
+
returns["best_params"]["y"],
|
|
231
|
+
returns["best_params"]["z"],
|
|
232
|
+
)
|
|
233
|
+
if num_trials >= 12:
|
|
234
|
+
assert bv == 1 and abs(bx) == 1 and by == 0
|
|
235
|
+
else:
|
|
236
|
+
assert bx in [-1, 1] and by in [-1, 0]
|
|
237
|
+
assert bz in ["foo", "bar"]
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
@mark.parametrize("with_commandline", (True, False))
|
|
241
|
+
def test_optuna_multi_objective_example(with_commandline: bool, tmpdir: Path) -> None:
|
|
242
|
+
cmd = [
|
|
243
|
+
"example/multi-objective.py",
|
|
244
|
+
"--multirun",
|
|
245
|
+
"hydra.sweep.dir=" + str(tmpdir),
|
|
246
|
+
"hydra.job.chdir=True",
|
|
247
|
+
"hydra.sweeper.n_trials=20",
|
|
248
|
+
"hydra.sweeper.n_jobs=1",
|
|
249
|
+
"hydra/sweeper/sampler=random",
|
|
250
|
+
"hydra.sweeper.sampler.seed=123",
|
|
251
|
+
]
|
|
252
|
+
if with_commandline:
|
|
253
|
+
cmd += [
|
|
254
|
+
"x=range(0, 5)",
|
|
255
|
+
"y=range(0, 3)",
|
|
256
|
+
]
|
|
257
|
+
run_optuna_script(cmd)
|
|
258
|
+
returns = OmegaConf.load(f"{tmpdir}/optimization_results.yaml")
|
|
259
|
+
assert isinstance(returns, DictConfig)
|
|
260
|
+
assert returns.name == "optuna"
|
|
261
|
+
if with_commandline:
|
|
262
|
+
for trial_x in returns["solutions"]:
|
|
263
|
+
assert trial_x["params"]["x"] % 1 == 0
|
|
264
|
+
assert trial_x["params"]["y"] % 1 == 0
|
|
265
|
+
# The trials must not dominate each other.
|
|
266
|
+
for trial_y in returns["solutions"]:
|
|
267
|
+
assert not _dominates(trial_x, trial_y)
|
|
268
|
+
else:
|
|
269
|
+
for trial_x in returns["solutions"]:
|
|
270
|
+
assert trial_x["params"]["x"] % 1 in {0, 0.5}
|
|
271
|
+
assert trial_x["params"]["y"] % 1 in {0, 0.5}
|
|
272
|
+
# The trials must not dominate each other.
|
|
273
|
+
for trial_y in returns["solutions"]:
|
|
274
|
+
assert not _dominates(trial_x, trial_y)
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
def _dominates(values_x: List[float], values_y: List[float]) -> bool:
|
|
278
|
+
return all(x <= y for x, y in zip(values_x, values_y)) and any(
|
|
279
|
+
x < y for x, y in zip(values_x, values_y)
|
|
280
|
+
)
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
def test_optuna_custom_search_space_example(tmpdir: Path) -> None:
|
|
284
|
+
max_z_difference_from_x = 0.3
|
|
285
|
+
cmd = [
|
|
286
|
+
"example/custom-search-space-objective.py",
|
|
287
|
+
"--multirun",
|
|
288
|
+
"hydra.sweep.dir=" + str(tmpdir),
|
|
289
|
+
"hydra.job.chdir=True",
|
|
290
|
+
"hydra.sweeper.n_trials=20",
|
|
291
|
+
"hydra.sweeper.n_jobs=1",
|
|
292
|
+
"hydra/sweeper/sampler=random",
|
|
293
|
+
"hydra.sweeper.sampler.seed=123",
|
|
294
|
+
f"max_z_difference_from_x={max_z_difference_from_x}",
|
|
295
|
+
]
|
|
296
|
+
run_optuna_script(cmd)
|
|
297
|
+
returns = OmegaConf.load(f"{tmpdir}/optimization_results.yaml")
|
|
298
|
+
assert isinstance(returns, DictConfig)
|
|
299
|
+
assert returns.name == "optuna"
|
|
300
|
+
assert (
|
|
301
|
+
abs(returns["best_params"]["x"] - returns["best_params"]["z"])
|
|
302
|
+
<= max_z_difference_from_x
|
|
303
|
+
)
|
|
304
|
+
w = returns["best_params"]["+w"]
|
|
305
|
+
assert 0 <= w <= 1
|
|
306
|
+
|
|
307
|
+
|
|
308
|
+
@mark.parametrize(
|
|
309
|
+
"search_space,params,raise_warning,msg",
|
|
310
|
+
[
|
|
311
|
+
(None, None, False, None),
|
|
312
|
+
(
|
|
313
|
+
{},
|
|
314
|
+
{},
|
|
315
|
+
True,
|
|
316
|
+
r"Both hydra.sweeper.params and hydra.sweeper.search_space are configured.*",
|
|
317
|
+
),
|
|
318
|
+
(
|
|
319
|
+
{},
|
|
320
|
+
None,
|
|
321
|
+
True,
|
|
322
|
+
r"`hydra.sweeper.search_space` is deprecated and will be removed in the next major release.*",
|
|
323
|
+
),
|
|
324
|
+
(None, {}, False, None),
|
|
325
|
+
],
|
|
326
|
+
)
|
|
327
|
+
def test_warnings(
|
|
328
|
+
tmpdir: Path,
|
|
329
|
+
search_space: Optional[DictConfig],
|
|
330
|
+
params: Optional[DictConfig],
|
|
331
|
+
raise_warning: bool,
|
|
332
|
+
msg: Optional[str],
|
|
333
|
+
) -> None:
|
|
334
|
+
partial_sweeper = partial(
|
|
335
|
+
OptunaSweeperImpl,
|
|
336
|
+
sampler=RandomSampler(),
|
|
337
|
+
direction=Direction.minimize,
|
|
338
|
+
storage=None,
|
|
339
|
+
study_name="test",
|
|
340
|
+
n_trials=1,
|
|
341
|
+
n_jobs=1,
|
|
342
|
+
max_failure_rate=0.0,
|
|
343
|
+
custom_search_space=None,
|
|
344
|
+
)
|
|
345
|
+
if search_space is not None:
|
|
346
|
+
search_space = OmegaConf.create(search_space)
|
|
347
|
+
if params is not None:
|
|
348
|
+
params = OmegaConf.create(params)
|
|
349
|
+
sweeper = partial_sweeper(search_space=search_space, params=params)
|
|
350
|
+
if raise_warning:
|
|
351
|
+
with warns(
|
|
352
|
+
UserWarning,
|
|
353
|
+
match=msg,
|
|
354
|
+
):
|
|
355
|
+
sweeper._process_searchspace_config()
|
|
356
|
+
else:
|
|
357
|
+
sweeper._process_searchspace_config()
|
|
358
|
+
|
|
359
|
+
|
|
360
|
+
@mark.parametrize("max_failure_rate", (0.5, 1.0))
|
|
361
|
+
def test_failure_rate(max_failure_rate: float, tmpdir: Path) -> None:
|
|
362
|
+
cmd = [
|
|
363
|
+
sys.executable,
|
|
364
|
+
"example/sphere.py",
|
|
365
|
+
"--multirun",
|
|
366
|
+
"hydra.sweep.dir=" + str(tmpdir),
|
|
367
|
+
"hydra.job.chdir=True",
|
|
368
|
+
"hydra.sweeper.n_trials=20",
|
|
369
|
+
"hydra.sweeper.n_jobs=2",
|
|
370
|
+
"hydra/sweeper/sampler=random",
|
|
371
|
+
"hydra.sweeper.sampler.seed=123",
|
|
372
|
+
f"hydra.sweeper.max_failure_rate={max_failure_rate}",
|
|
373
|
+
"error=true",
|
|
374
|
+
]
|
|
375
|
+
out, err = run_process(cmd, print_error=False, raise_exception=False)
|
|
376
|
+
error_string = "RuntimeError: cfg.error is True"
|
|
377
|
+
if max_failure_rate < 1.0:
|
|
378
|
+
assert error_string in err
|
|
379
|
+
else:
|
|
380
|
+
assert error_string not in err
|
|
381
|
+
|
|
382
|
+
|
|
383
|
+
def test_example_with_deprecated_search_space(
|
|
384
|
+
tmpdir: Path,
|
|
385
|
+
) -> None:
|
|
386
|
+
cmd = [
|
|
387
|
+
"-W ignore::UserWarning",
|
|
388
|
+
"example/sphere.py",
|
|
389
|
+
"--multirun",
|
|
390
|
+
"--config-dir=tests/conf",
|
|
391
|
+
"--config-name=test_deprecated_search_space",
|
|
392
|
+
"hydra.sweep.dir=" + str(tmpdir),
|
|
393
|
+
"hydra.job.chdir=True",
|
|
394
|
+
"hydra.sweeper.n_trials=20",
|
|
395
|
+
"hydra.sweeper.n_jobs=1",
|
|
396
|
+
]
|
|
397
|
+
|
|
398
|
+
run_optuna_script(cmd)
|
|
399
|
+
returns = OmegaConf.load(f"{tmpdir}/optimization_results.yaml")
|
|
400
|
+
assert isinstance(returns, DictConfig)
|
|
401
|
+
assert returns.name == "optuna"
|
|
402
|
+
assert abs(returns["best_params"]["x"]) <= 5.5
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|