mirror of
https://github.com/wassname/ray.git
synced 2026-08-10 12:30:14 +08:00
[autoscaler/tune] Optional YAML Fields + Fix Pretty Printing for Tune (#1541)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user