mirror of
https://github.com/wassname/ray.git
synced 2026-08-05 13:21:03 +08:00
[tune] tune.track -> tune.report (#8388)
This commit is contained in:
@@ -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"))
|
||||
|
||||
Reference in New Issue
Block a user