[tune] tune.track -> tune.report (#8388)

This commit is contained in:
Richard Liaw
2020-05-16 12:55:08 -07:00
committed by GitHub
parent c8cd716295
commit 67c01455fe
20 changed files with 228 additions and 395 deletions
+12 -11
View File
@@ -1,27 +1,26 @@
from tensorflow import keras
from ray.tune import track
class TuneReporterCallback(keras.callbacks.Callback):
"""Tune Callback for Keras."""
def __init__(self, reporter=None, freq="batch", logs={}):
def __init__(self, reporter=None, freq="batch", logs=None):
"""Initializer.
Args:
reporter (StatusReporter|tune.track.log|None): Tune object for
returning results.
freq (str): Sets the frequency of reporting intermediate results.
One of ["batch", "epoch"].
"""
self.reporter = reporter or track.log
self.iteration = 0
logs = logs or {}
if freq not in ["batch", "epoch"]:
raise ValueError("{} not supported as a frequency.".format(freq))
self.freq = freq
super(TuneReporterCallback, self).__init__()
def on_batch_end(self, batch, logs={}):
def on_batch_end(self, batch, logs=None):
from ray import tune
logs = logs or {}
if not self.freq == "batch":
return
self.iteration += 1
@@ -29,11 +28,13 @@ class TuneReporterCallback(keras.callbacks.Callback):
if "loss" in metric and "neg_" not in metric:
logs["neg_" + metric] = -logs[metric]
if "acc" in logs:
self.reporter(keras_info=logs, mean_accuracy=logs["acc"])
tune.report(keras_info=logs, mean_accuracy=logs["acc"])
else:
self.reporter(keras_info=logs, mean_accuracy=logs.get("accuracy"))
tune.report(keras_info=logs, mean_accuracy=logs.get("accuracy"))
def on_epoch_end(self, batch, logs={}):
def on_epoch_end(self, batch, logs=None):
from ray import tune
logs = logs or {}
if not self.freq == "epoch":
return
self.iteration += 1
@@ -41,6 +42,6 @@ class TuneReporterCallback(keras.callbacks.Callback):
if "loss" in metric and "neg_" not in metric:
logs["neg_" + metric] = -logs[metric]
if "acc" in logs:
self.reporter(keras_info=logs, mean_accuracy=logs["acc"])
tune.report(keras_info=logs, mean_accuracy=logs["acc"])
else:
self.reporter(keras_info=logs, mean_accuracy=logs.get("accuracy"))
tune.report(keras_info=logs, mean_accuracy=logs.get("accuracy"))