[autoscaler/tune] Optional YAML Fields + Fix Pretty Printing for Tune (#1541)

This commit is contained in:
Richard Liaw
2018-03-04 23:35:58 -08:00
committed by GitHub
parent 061e435411
commit 162d063f0d
12 changed files with 199 additions and 70 deletions
+12
View File
@@ -6,6 +6,7 @@ import csv
import json
import numpy as np
import os
import yaml
from ray.tune.result import TrainingResult
from ray.tune.log_sync import get_syncer
@@ -176,3 +177,14 @@ class _CustomEncoder(json.JSONEncoder):
return float(value)
if np.issubdtype(value, int):
return int(value)
def pretty_print(result):
result = result._replace(config=None) # drop config from pretty print
out = {}
for k, v in result._asdict().items():
if v is not None:
out[k] = v
cleaned = json.dumps(out, cls=_CustomEncoder)
return yaml.dump(json.loads(cleaned), default_flow_style=False)
-18
View File
@@ -3,14 +3,8 @@ from __future__ import division
from __future__ import print_function
from collections import namedtuple
import json
import os
try:
import yaml
except ImportError:
print("Could not import YAML module, falling back to JSON pretty-printing")
yaml = None
"""
When using ray.tune with custom training scripts, you must periodically report
@@ -93,16 +87,4 @@ TrainingResult = namedtuple("TrainingResult", [
])
def pretty_print(result):
result = result._replace(config=None) # drop config from pretty print
out = {}
for k, v in result._asdict().items():
if v is not None:
out[k] = v
if yaml:
return yaml.safe_dump(out, default_flow_style=False)
else:
return json.dumps(out) + "\n"
TrainingResult.__new__.__defaults__ = (None,) * len(TrainingResult._fields)
+2 -2
View File
@@ -11,9 +11,9 @@ import ray
import os
from ray.tune import TuneError
from ray.tune.logger import NoopLogger, UnifiedLogger
from ray.tune.logger import NoopLogger, UnifiedLogger, pretty_print
from ray.tune.registry import _default_registry, get_registry, TRAINABLE_CLASS
from ray.tune.result import TrainingResult, DEFAULT_RESULTS_DIR, pretty_print
from ray.tune.result import TrainingResult, DEFAULT_RESULTS_DIR
from ray.utils import random_string, binary_to_hex
DEBUG_PRINT_INTERVAL = 5