hydra-ax-sweeper 1.4.0.dev4__tar.gz → 1.4.0.dev6__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.
Files changed (18) hide show
  1. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/PKG-INFO +1 -1
  2. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_ax_sweeper.egg-info/PKG-INFO +1 -1
  3. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_plugins/hydra_ax_sweeper/__init__.py +1 -1
  4. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_plugins/hydra_ax_sweeper/_core.py +32 -12
  5. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/tests/test_ax_sweeper_plugin.py +17 -0
  6. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/MANIFEST.in +0 -0
  7. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/README.md +0 -0
  8. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_ax_sweeper.egg-info/SOURCES.txt +0 -0
  9. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_ax_sweeper.egg-info/dependency_links.txt +0 -0
  10. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_ax_sweeper.egg-info/requires.txt +0 -0
  11. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_ax_sweeper.egg-info/top_level.txt +0 -0
  12. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_plugins/hydra_ax_sweeper/_earlystopper.py +0 -0
  13. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_plugins/hydra_ax_sweeper/ax_sweeper.py +0 -0
  14. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_plugins/hydra_ax_sweeper/config.py +0 -0
  15. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/hydra_plugins/hydra_ax_sweeper/py.typed +0 -0
  16. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/pyproject.toml +0 -0
  17. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/setup.cfg +0 -0
  18. {hydra_ax_sweeper-1.4.0.dev4 → hydra_ax_sweeper-1.4.0.dev6}/setup.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: hydra-ax-sweeper
3
- Version: 1.4.0.dev4
3
+ Version: 1.4.0.dev6
4
4
  Summary: Hydra Ax Sweeper plugin
5
5
  Home-page: https://github.com/facebookresearch/hydra/
6
6
  Author: Omry Yadan, Shagun Sodhani
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: hydra-ax-sweeper
3
- Version: 1.4.0.dev4
3
+ Version: 1.4.0.dev6
4
4
  Summary: Hydra Ax Sweeper plugin
5
5
  Home-page: https://github.com/facebookresearch/hydra/
6
6
  Author: Omry Yadan, Shagun Sodhani
@@ -1,3 +1,3 @@
1
1
  # Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved
2
2
 
3
- __version__ = "1.4.0.dev4"
3
+ __version__ = "1.4.0.dev6"
@@ -34,6 +34,8 @@ log = logging.getLogger(__name__)
34
34
  AxRangeParameterType = Literal["float", "int"]
35
35
  AxChoiceParameterType = Literal["float", "int", "str", "bool"]
36
36
  AxParameterConfig = Union[RangeParameterConfig, ChoiceParameterConfig]
37
+ AxMetricValue = Union[float, Tuple[float, float]]
38
+ AxRawData = Mapping[str, AxMetricValue]
37
39
 
38
40
 
39
41
  @dataclass
@@ -126,6 +128,30 @@ def create_ax_parameter_config(param: Dict[Any, Any]) -> AxParameterConfig:
126
128
  )
127
129
 
128
130
 
131
+ def create_ax_raw_data(value: Any, objective_name: str, is_noisy: bool) -> AxRawData:
132
+ def normalize_metric(metric_value: Any) -> AxMetricValue:
133
+ assert isinstance(metric_value, (int, float, tuple))
134
+ if isinstance(metric_value, (int, float)):
135
+ mean = float(metric_value)
136
+ return mean if is_noisy else (mean, 0.0)
137
+
138
+ assert len(metric_value) == 2
139
+ mean, sem = metric_value
140
+ assert isinstance(mean, (int, float))
141
+ assert sem is None or isinstance(sem, (int, float))
142
+ return (float(mean), float("nan") if sem is None else float(sem))
143
+
144
+ assert isinstance(value, (int, float, tuple, dict))
145
+ if isinstance(value, dict):
146
+ raw_data: Dict[str, AxMetricValue] = {}
147
+ for metric_name, metric_value in value.items():
148
+ assert isinstance(metric_name, str)
149
+ raw_data[metric_name] = normalize_metric(metric_value)
150
+ return raw_data
151
+
152
+ return {objective_name: normalize_metric(value)}
153
+
154
+
129
155
  def get_one_batch_of_trials(
130
156
  ax_client: Client,
131
157
  num_max_trials_to_do: int,
@@ -258,19 +284,13 @@ class CoreAxSweeper(Sweeper):
258
284
  # Alternatively, the task function can return a dict whose values
259
285
  # represent multiple metrics, where each key is the name of the metric
260
286
  # and the item can be an int, float or tuple.
261
- assert isinstance(val, (int, float, tuple, dict))
262
- # is_noisy specifies how Ax should behave when not given an error value.
263
- # if true (default), the error of each measurement is inferred by Ax.
264
- # if false, the error of each measurement is set to 0.
265
- if isinstance(val, (int, float)):
266
- if self.is_noisy:
267
- val = (val, None) # specify unknown noise
268
- else:
269
- val = (val, 0) # specify no noise
270
- if isinstance(val, tuple):
271
- val = {self.experiment.objective_name: val}
287
+ # is_noisy specifies how Ax should behave when not given an
288
+ # error value: true means unknown error, false means zero error.
289
+ raw_data = create_ax_raw_data(
290
+ val, self.experiment.objective_name, self.is_noisy
291
+ )
272
292
  ax_client.complete_trial(
273
- trial_index=batch[idx].trial_index, raw_data=val
293
+ trial_index=batch[idx].trial_index, raw_data=raw_data
274
294
  )
275
295
 
276
296
  def get_best_point(
@@ -316,6 +316,23 @@ def test_command_line_log_interval_configures_ax_log_range() -> None:
316
316
  assert ax_parameter.scaling == "log"
317
317
 
318
318
 
319
+ def test_create_ax_raw_data() -> None:
320
+ from hydra_plugins.hydra_ax_sweeper._core import create_ax_raw_data
321
+
322
+ assert create_ax_raw_data(1, "loss", is_noisy=True) == {"loss": 1.0}
323
+ assert create_ax_raw_data(1, "loss", is_noisy=False) == {"loss": (1.0, 0.0)}
324
+ assert create_ax_raw_data({"loss": 1, "accuracy": (0.5, 0.1)}, "loss", True) == {
325
+ "loss": 1.0,
326
+ "accuracy": (0.5, 0.1),
327
+ }
328
+
329
+ raw_data = create_ax_raw_data((1, None), "loss", is_noisy=True)
330
+ loss = raw_data["loss"]
331
+ assert isinstance(loss, tuple)
332
+ sem = loss[1]
333
+ assert math.isnan(sem)
334
+
335
+
319
336
  def test_ax_logging_from_hydra_app(tmpdir: Path) -> None:
320
337
  cmd = [
321
338
  "tests/apps/polynomial.py",