mirror of
https://github.com/wassname/ray.git
synced 2026-07-20 12:40:20 +08:00
[tune] logger migration to ExperimentLogger classes (#11984)
This commit is contained in:
@@ -8,15 +8,22 @@ import numpy as np
|
||||
|
||||
from ray.tune import Trainable
|
||||
from ray.tune.function_runner import wrap_function
|
||||
from ray.tune.integration.wandb import _WandbLoggingProcess, \
|
||||
from ray.tune.integration.wandb import WandbLoggerCallback, \
|
||||
_WandbLoggingProcess, \
|
||||
_WANDB_QUEUE_END, WandbLogger, WANDB_ENV_VAR, WandbTrainableMixin, \
|
||||
wandb_mixin
|
||||
from ray.tune.result import TRIAL_INFO
|
||||
from ray.tune.trial import TrialInfo
|
||||
|
||||
Trial = namedtuple("MockTrial",
|
||||
["config", "trial_id", "trial_name", "trainable_name"])
|
||||
Trial.__str__ = lambda t: t.trial_name
|
||||
|
||||
class Trial(
|
||||
namedtuple("MockTrial",
|
||||
["config", "trial_id", "trial_name", "trainable_name"])):
|
||||
def __hash__(self):
|
||||
return hash(self.trial_id)
|
||||
|
||||
def __str__(self):
|
||||
return self.trial_name
|
||||
|
||||
|
||||
class _MockWandbLoggingProcess(_WandbLoggingProcess):
|
||||
@@ -37,9 +44,21 @@ class _MockWandbLoggingProcess(_WandbLoggingProcess):
|
||||
self.logs.put(log)
|
||||
|
||||
|
||||
class WandbTestLogger(WandbLogger):
|
||||
class WandbTestExperimentLogger(WandbLoggerCallback):
|
||||
_logger_process_cls = _MockWandbLoggingProcess
|
||||
|
||||
@property
|
||||
def trial_processes(self):
|
||||
return self._trial_processes
|
||||
|
||||
|
||||
class WandbTestLogger(WandbLogger):
|
||||
_experiment_logger_cls = WandbTestExperimentLogger
|
||||
|
||||
@property
|
||||
def trial_process(self):
|
||||
return self._trial_experiment_logger.trial_processes[self.trial]
|
||||
|
||||
|
||||
class _MockWandbAPI(object):
|
||||
def init(self, *args, **kwargs):
|
||||
@@ -63,7 +82,7 @@ class WandbIntegrationTest(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
pass
|
||||
|
||||
def testWandbLoggerConfig(self):
|
||||
def testWandbLegacyLoggerConfig(self):
|
||||
trial_config = {"par1": 4, "par2": 9.12345678}
|
||||
trial = Trial(trial_config, 0, "trial_0", "trainable")
|
||||
|
||||
@@ -115,11 +134,13 @@ class WandbIntegrationTest(unittest.TestCase):
|
||||
trial_config["wandb"] = {"project": "test_project"}
|
||||
|
||||
logger = WandbTestLogger(trial_config, "/tmp", trial)
|
||||
self.assertEqual(logger._wandb.kwargs["project"], "test_project")
|
||||
self.assertEqual(logger._wandb.kwargs["id"], trial.trial_id)
|
||||
self.assertEqual(logger._wandb.kwargs["name"], trial.trial_name)
|
||||
self.assertEqual(logger._wandb.kwargs["group"], trial.trainable_name)
|
||||
self.assertIn("config", logger._wandb._exclude)
|
||||
self.assertEqual(logger.trial_process.kwargs["project"],
|
||||
"test_project")
|
||||
self.assertEqual(logger.trial_process.kwargs["id"], trial.trial_id)
|
||||
self.assertEqual(logger.trial_process.kwargs["name"], trial.trial_name)
|
||||
self.assertEqual(logger.trial_process.kwargs["group"],
|
||||
trial.trainable_name)
|
||||
self.assertIn("config", logger.trial_process._exclude)
|
||||
|
||||
logger.close()
|
||||
|
||||
@@ -127,8 +148,8 @@ class WandbIntegrationTest(unittest.TestCase):
|
||||
trial_config["wandb"] = {"project": "test_project", "log_config": True}
|
||||
|
||||
logger = WandbTestLogger(trial_config, "/tmp", trial)
|
||||
self.assertNotIn("config", logger._wandb._exclude)
|
||||
self.assertNotIn("metric", logger._wandb._exclude)
|
||||
self.assertNotIn("config", logger.trial_process._exclude)
|
||||
self.assertNotIn("metric", logger.trial_process._exclude)
|
||||
|
||||
logger.close()
|
||||
|
||||
@@ -139,12 +160,12 @@ class WandbIntegrationTest(unittest.TestCase):
|
||||
}
|
||||
|
||||
logger = WandbTestLogger(trial_config, "/tmp", trial)
|
||||
self.assertIn("config", logger._wandb._exclude)
|
||||
self.assertIn("metric", logger._wandb._exclude)
|
||||
self.assertIn("config", logger.trial_process._exclude)
|
||||
self.assertIn("metric", logger.trial_process._exclude)
|
||||
|
||||
logger.close()
|
||||
|
||||
def testWandbLoggerReporting(self):
|
||||
def testWandbLegacyLoggerReporting(self):
|
||||
trial_config = {"par1": 4, "par2": 9.12345678}
|
||||
trial = Trial(trial_config, 0, "trial_0", "trainable")
|
||||
|
||||
@@ -166,7 +187,7 @@ class WandbIntegrationTest(unittest.TestCase):
|
||||
|
||||
logger.on_result(r1)
|
||||
|
||||
logged = logger._wandb.logs.get(timeout=10)
|
||||
logged = logger.trial_process.logs.get(timeout=10)
|
||||
self.assertIn("metric1", logged)
|
||||
self.assertNotIn("metric2", logged)
|
||||
self.assertIn("metric3", logged)
|
||||
@@ -176,6 +197,106 @@ class WandbIntegrationTest(unittest.TestCase):
|
||||
|
||||
logger.close()
|
||||
|
||||
def testWandbLoggerConfig(self):
|
||||
trial_config = {"par1": 4, "par2": 9.12345678}
|
||||
trial = Trial(trial_config, 0, "trial_0", "trainable")
|
||||
|
||||
if WANDB_ENV_VAR in os.environ:
|
||||
del os.environ[WANDB_ENV_VAR]
|
||||
|
||||
# No API key
|
||||
with self.assertRaises(ValueError):
|
||||
logger = WandbTestExperimentLogger(project="test_project")
|
||||
|
||||
# API Key in config
|
||||
logger = WandbTestExperimentLogger(
|
||||
project="test_project", api_key="1234")
|
||||
self.assertEqual(os.environ[WANDB_ENV_VAR], "1234")
|
||||
|
||||
del logger
|
||||
del os.environ[WANDB_ENV_VAR]
|
||||
|
||||
# API Key file
|
||||
with tempfile.NamedTemporaryFile("wt") as fp:
|
||||
fp.write("5678")
|
||||
fp.flush()
|
||||
|
||||
logger = WandbTestExperimentLogger(
|
||||
project="test_project", api_key_file=fp.name)
|
||||
self.assertEqual(os.environ[WANDB_ENV_VAR], "5678")
|
||||
|
||||
del logger
|
||||
del os.environ[WANDB_ENV_VAR]
|
||||
|
||||
# API Key in env
|
||||
os.environ[WANDB_ENV_VAR] = "9012"
|
||||
logger = WandbTestExperimentLogger(project="test_project")
|
||||
del logger
|
||||
|
||||
# From now on, the API key is in the env variable.
|
||||
|
||||
logger = WandbTestExperimentLogger(project="test_project")
|
||||
logger.log_trial_start(trial)
|
||||
|
||||
self.assertEqual(logger.trial_processes[trial].kwargs["project"],
|
||||
"test_project")
|
||||
self.assertEqual(logger.trial_processes[trial].kwargs["id"],
|
||||
trial.trial_id)
|
||||
self.assertEqual(logger.trial_processes[trial].kwargs["name"],
|
||||
trial.trial_name)
|
||||
self.assertEqual(logger.trial_processes[trial].kwargs["group"],
|
||||
trial.trainable_name)
|
||||
self.assertIn("config", logger.trial_processes[trial]._exclude)
|
||||
|
||||
del logger
|
||||
|
||||
# log config.
|
||||
logger = WandbTestExperimentLogger(
|
||||
project="test_project", log_config=True)
|
||||
logger.log_trial_start(trial)
|
||||
self.assertNotIn("config", logger.trial_processes[trial]._exclude)
|
||||
self.assertNotIn("metric", logger.trial_processes[trial]._exclude)
|
||||
|
||||
del logger
|
||||
|
||||
# Exclude metric.
|
||||
logger = WandbTestExperimentLogger(
|
||||
project="test_project", excludes=["metric"])
|
||||
logger.log_trial_start(trial)
|
||||
self.assertIn("config", logger.trial_processes[trial]._exclude)
|
||||
self.assertIn("metric", logger.trial_processes[trial]._exclude)
|
||||
|
||||
del logger
|
||||
|
||||
def testWandbLoggerReporting(self):
|
||||
trial_config = {"par1": 4, "par2": 9.12345678}
|
||||
trial = Trial(trial_config, 0, "trial_0", "trainable")
|
||||
|
||||
logger = WandbTestExperimentLogger(
|
||||
project="test_project", api_key="1234", excludes=["metric2"])
|
||||
logger.on_trial_start(0, [], trial)
|
||||
|
||||
r1 = {
|
||||
"metric1": 0.8,
|
||||
"metric2": 1.4,
|
||||
"metric3": np.asarray(32.0),
|
||||
"metric4": np.float32(32.0),
|
||||
"const": "text",
|
||||
"config": trial_config
|
||||
}
|
||||
|
||||
logger.on_trial_result(0, [], trial, r1)
|
||||
|
||||
logged = logger.trial_processes[trial].logs.get(timeout=10)
|
||||
self.assertIn("metric1", logged)
|
||||
self.assertNotIn("metric2", logged)
|
||||
self.assertIn("metric3", logged)
|
||||
self.assertIn("metric4", logged)
|
||||
self.assertNotIn("const", logged)
|
||||
self.assertNotIn("config", logged)
|
||||
|
||||
del logger
|
||||
|
||||
def testWandbMixinConfig(self):
|
||||
config = {"par1": 4, "par2": 9.12345678}
|
||||
trial = Trial(config, 0, "trial_0", "trainable")
|
||||
|
||||
@@ -1,12 +1,33 @@
|
||||
import csv
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
from collections import namedtuple
|
||||
import unittest
|
||||
import tempfile
|
||||
import shutil
|
||||
import numpy as np
|
||||
from ray.cloudpickle import cloudpickle
|
||||
|
||||
from ray.tune.logger import JsonLogger, CSVLogger, TBXLogger
|
||||
from ray.tune.logger import CSVLoggerCallback, JsonLoggerCallback, \
|
||||
JsonLogger, CSVLogger, \
|
||||
TBXLoggerCallback, TBXLogger
|
||||
from ray.tune.result import EXPR_PARAM_FILE, EXPR_PARAM_PICKLE_FILE, \
|
||||
EXPR_PROGRESS_FILE, \
|
||||
EXPR_RESULT_FILE
|
||||
|
||||
Trial = namedtuple("MockTrial", ["evaluated_params", "trial_id"])
|
||||
|
||||
class Trial(
|
||||
namedtuple("MockTrial", ["evaluated_params", "trial_id", "logdir"])):
|
||||
@property
|
||||
def config(self):
|
||||
return self.evaluated_params
|
||||
|
||||
def init_logdir(self):
|
||||
return
|
||||
|
||||
def __hash__(self):
|
||||
return hash(self.trial_id)
|
||||
|
||||
|
||||
def result(t, rew, **kwargs):
|
||||
@@ -28,24 +49,116 @@ class LoggerSuite(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.test_dir, ignore_errors=True)
|
||||
|
||||
def testCSV(self):
|
||||
def testLegacyCSV(self):
|
||||
config = {"a": 2, "b": 5, "c": {"c": {"D": 123}, "e": None}}
|
||||
t = Trial(evaluated_params=config, trial_id="csv")
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="csv", logdir=self.test_dir)
|
||||
logger = CSVLogger(config=config, logdir=self.test_dir, trial=t)
|
||||
logger.on_result(result(2, 4))
|
||||
logger.on_result(result(2, 4))
|
||||
logger.on_result(result(2, 4, score=[1, 2, 3], hello={"world": 1}))
|
||||
logger.on_result(result(2, 5))
|
||||
logger.on_result(result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
logger.close()
|
||||
|
||||
self._validate_csv_result()
|
||||
|
||||
def testCSV(self):
|
||||
config = {"a": 2, "b": 5, "c": {"c": {"D": 123}, "e": None}}
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="csv", logdir=self.test_dir)
|
||||
logger = CSVLoggerCallback()
|
||||
logger.on_trial_result(0, [], t, result(0, 4))
|
||||
logger.on_trial_result(1, [], t, result(1, 5))
|
||||
logger.on_trial_result(
|
||||
2, [], t, result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
|
||||
logger.on_trial_complete(3, [], t)
|
||||
self._validate_csv_result()
|
||||
|
||||
def _validate_csv_result(self):
|
||||
results = []
|
||||
result_file = os.path.join(self.test_dir, EXPR_PROGRESS_FILE)
|
||||
with open(result_file, "rt") as fp:
|
||||
reader = csv.DictReader(fp)
|
||||
for row in reader:
|
||||
results.append(row)
|
||||
|
||||
self.assertEqual(len(results), 3)
|
||||
self.assertSequenceEqual(
|
||||
[int(row["episode_reward_mean"]) for row in results], [4, 5, 6])
|
||||
|
||||
def testJSONLegacyLogger(self):
|
||||
config = {"a": 2, "b": 5, "c": {"c": {"D": 123}, "e": None}}
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="json", logdir=self.test_dir)
|
||||
logger = JsonLogger(config=config, logdir=self.test_dir, trial=t)
|
||||
logger.on_result(result(0, 4))
|
||||
logger.on_result(result(1, 5))
|
||||
logger.on_result(result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
logger.close()
|
||||
|
||||
self._validate_json_result(config)
|
||||
|
||||
def testJSON(self):
|
||||
config = {"a": 2, "b": 5, "c": {"c": {"D": 123}, "e": None}}
|
||||
t = Trial(evaluated_params=config, trial_id="json")
|
||||
logger = JsonLogger(config=config, logdir=self.test_dir, trial=t)
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="json", logdir=self.test_dir)
|
||||
logger = JsonLoggerCallback()
|
||||
logger.on_trial_result(0, [], t, result(0, 4))
|
||||
logger.on_trial_result(1, [], t, result(1, 5))
|
||||
logger.on_trial_result(
|
||||
2, [], t, result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
|
||||
logger.on_trial_complete(3, [], t)
|
||||
self._validate_json_result(config)
|
||||
|
||||
def _validate_json_result(self, config):
|
||||
# Check result logs
|
||||
results = []
|
||||
result_file = os.path.join(self.test_dir, EXPR_RESULT_FILE)
|
||||
with open(result_file, "rt") as fp:
|
||||
for row in fp.readlines():
|
||||
results.append(json.loads(row))
|
||||
|
||||
self.assertEqual(len(results), 3)
|
||||
self.assertSequenceEqual(
|
||||
[int(row["episode_reward_mean"]) for row in results], [4, 5, 6])
|
||||
|
||||
# Check json saved config file
|
||||
config_file = os.path.join(self.test_dir, EXPR_PARAM_FILE)
|
||||
with open(config_file, "rt") as fp:
|
||||
loaded_config = json.load(fp)
|
||||
|
||||
self.assertEqual(loaded_config, config)
|
||||
|
||||
# Check pickled config file
|
||||
config_file = os.path.join(self.test_dir, EXPR_PARAM_PICKLE_FILE)
|
||||
with open(config_file, "rb") as fp:
|
||||
loaded_config = cloudpickle.load(fp)
|
||||
|
||||
self.assertEqual(loaded_config, config)
|
||||
|
||||
def testLegacyTBX(self):
|
||||
config = {
|
||||
"a": 2,
|
||||
"b": [1, 2],
|
||||
"c": {
|
||||
"c": {
|
||||
"D": 123
|
||||
}
|
||||
},
|
||||
"d": np.int64(1),
|
||||
"e": np.bool8(True)
|
||||
}
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="tbx", logdir=self.test_dir)
|
||||
logger = TBXLogger(config=config, logdir=self.test_dir, trial=t)
|
||||
logger.on_result(result(0, 4))
|
||||
logger.on_result(result(1, 4))
|
||||
logger.on_result(result(2, 4, score=[1, 2, 3], hello={"world": 1}))
|
||||
logger.on_result(result(1, 5))
|
||||
logger.on_result(result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
logger.close()
|
||||
|
||||
self._validate_tbx_result()
|
||||
|
||||
def testTBX(self):
|
||||
config = {
|
||||
"a": 2,
|
||||
@@ -58,16 +171,40 @@ class LoggerSuite(unittest.TestCase):
|
||||
"d": np.int64(1),
|
||||
"e": np.bool8(True)
|
||||
}
|
||||
t = Trial(evaluated_params=config, trial_id="tbx")
|
||||
logger = TBXLogger(config=config, logdir=self.test_dir, trial=t)
|
||||
logger.on_result(result(0, 4))
|
||||
logger.on_result(result(1, 4))
|
||||
logger.on_result(result(2, 4, score=[1, 2, 3], hello={"world": 1}))
|
||||
logger.close()
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="tbx", logdir=self.test_dir)
|
||||
logger = TBXLoggerCallback()
|
||||
logger.on_trial_result(0, [], t, result(0, 4))
|
||||
logger.on_trial_result(1, [], t, result(1, 5))
|
||||
logger.on_trial_result(
|
||||
2, [], t, result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
|
||||
def testBadTBX(self):
|
||||
logger.on_trial_complete(3, [], t)
|
||||
|
||||
self._validate_tbx_result()
|
||||
|
||||
def _validate_tbx_result(self):
|
||||
try:
|
||||
from tensorflow.python.summary.summary_iterator \
|
||||
import summary_iterator
|
||||
except ImportError:
|
||||
print("Skipping rest of test as tensorflow is not installed.")
|
||||
return
|
||||
|
||||
events_file = list(glob.glob(f"{self.test_dir}/events*"))[0]
|
||||
results = []
|
||||
for event in summary_iterator(events_file):
|
||||
for v in event.summary.value:
|
||||
if v.tag == "ray/tune/episode_reward_mean":
|
||||
results.append(v.simple_value)
|
||||
|
||||
self.assertEqual(len(results), 3)
|
||||
self.assertSequenceEqual([int(res) for res in results], [4, 5, 6])
|
||||
|
||||
def testLegacyBadTBX(self):
|
||||
config = {"b": (1, 2, 3)}
|
||||
t = Trial(evaluated_params=config, trial_id="tbx")
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="tbx", logdir=self.test_dir)
|
||||
logger = TBXLogger(config=config, logdir=self.test_dir, trial=t)
|
||||
logger.on_result(result(0, 4))
|
||||
logger.on_result(result(2, 4, score=[1, 2, 3], hello={"world": 1}))
|
||||
@@ -76,7 +213,8 @@ class LoggerSuite(unittest.TestCase):
|
||||
assert "INFO" in cm.output[0]
|
||||
|
||||
config = {"None": None}
|
||||
t = Trial(evaluated_params=config, trial_id="tbx")
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="tbx", logdir=self.test_dir)
|
||||
logger = TBXLogger(config=config, logdir=self.test_dir, trial=t)
|
||||
logger.on_result(result(0, 4))
|
||||
logger.on_result(result(2, 4, score=[1, 2, 3], hello={"world": 1}))
|
||||
@@ -84,6 +222,31 @@ class LoggerSuite(unittest.TestCase):
|
||||
logger.close()
|
||||
assert "INFO" in cm.output[0]
|
||||
|
||||
def testBadTBX(self):
|
||||
config = {"b": (1, 2, 3)}
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="tbx", logdir=self.test_dir)
|
||||
logger = TBXLoggerCallback()
|
||||
logger.on_trial_result(0, [], t, result(0, 4))
|
||||
logger.on_trial_result(1, [], t, result(1, 5))
|
||||
logger.on_trial_result(
|
||||
2, [], t, result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
with self.assertLogs("ray.tune.logger", level="INFO") as cm:
|
||||
logger.on_trial_complete(3, [], t)
|
||||
assert "INFO" in cm.output[0]
|
||||
|
||||
config = {"None": None}
|
||||
t = Trial(
|
||||
evaluated_params=config, trial_id="tbx", logdir=self.test_dir)
|
||||
logger = TBXLoggerCallback()
|
||||
logger.on_trial_result(0, [], t, result(0, 4))
|
||||
logger.on_trial_result(1, [], t, result(1, 5))
|
||||
logger.on_trial_result(
|
||||
2, [], t, result(2, 6, score=[1, 2, 3], hello={"world": 1}))
|
||||
with self.assertLogs("ray.tune.logger", level="INFO") as cm:
|
||||
logger.on_trial_complete(3, [], t)
|
||||
assert "INFO" in cm.output[0]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pytest
|
||||
|
||||
@@ -7,7 +7,7 @@ from ray.rllib import _register_all
|
||||
from ray.tune.result import TIMESTEPS_TOTAL
|
||||
from ray.tune import Trainable, TuneError
|
||||
from ray.tune import register_trainable, run_experiments
|
||||
from ray.tune.logger import LegacyExperimentLogger, Logger
|
||||
from ray.tune.logger import LegacyLoggerCallback, Logger
|
||||
from ray.tune.experiment import Experiment
|
||||
from ray.tune.trial import Trial, ExportFormat
|
||||
|
||||
@@ -191,7 +191,7 @@ class RunExperimentTest(unittest.TestCase):
|
||||
}
|
||||
}
|
||||
},
|
||||
callbacks=[LegacyExperimentLogger(logger_classes=[CustomLogger])])
|
||||
callbacks=[LegacyLoggerCallback(logger_classes=[CustomLogger])])
|
||||
self.assertTrue(os.path.exists(os.path.join(trial.logdir, "test.log")))
|
||||
self.assertFalse(
|
||||
os.path.exists(os.path.join(trial.logdir, "params.json")))
|
||||
@@ -204,7 +204,7 @@ class RunExperimentTest(unittest.TestCase):
|
||||
}
|
||||
}
|
||||
})
|
||||
self.assertTrue(
|
||||
self.assertFalse(
|
||||
os.path.exists(os.path.join(trial.logdir, "params.json")))
|
||||
|
||||
[trial] = run_experiments(
|
||||
@@ -216,7 +216,7 @@ class RunExperimentTest(unittest.TestCase):
|
||||
}
|
||||
}
|
||||
},
|
||||
callbacks=[LegacyExperimentLogger(logger_classes=[])])
|
||||
callbacks=[LegacyLoggerCallback(logger_classes=[])])
|
||||
self.assertFalse(
|
||||
os.path.exists(os.path.join(trial.logdir, "params.json")))
|
||||
|
||||
@@ -239,7 +239,7 @@ class RunExperimentTest(unittest.TestCase):
|
||||
}
|
||||
}
|
||||
},
|
||||
callbacks=[LegacyExperimentLogger(logger_classes=[CustomLogger])])
|
||||
callbacks=[LegacyLoggerCallback(logger_classes=[CustomLogger])])
|
||||
self.assertTrue(os.path.exists(os.path.join(trial.logdir, "test.log")))
|
||||
self.assertTrue(
|
||||
os.path.exists(os.path.join(trial.logdir, "params.json")))
|
||||
@@ -264,7 +264,7 @@ class RunExperimentTest(unittest.TestCase):
|
||||
}
|
||||
}
|
||||
},
|
||||
callbacks=[LegacyExperimentLogger(logger_classes=[])])
|
||||
callbacks=[LegacyLoggerCallback(logger_classes=[])])
|
||||
self.assertTrue(
|
||||
os.path.exists(os.path.join(trial.logdir, "params.json")))
|
||||
|
||||
|
||||
@@ -9,8 +9,8 @@ import ray
|
||||
from ray import tune
|
||||
from ray.rllib import _register_all
|
||||
from ray.tune.checkpoint_manager import Checkpoint
|
||||
from ray.tune.logger import DEFAULT_LOGGERS, ExperimentLogger, \
|
||||
LegacyExperimentLogger
|
||||
from ray.tune.logger import DEFAULT_LOGGERS, LoggerCallback, \
|
||||
LegacyLoggerCallback
|
||||
from ray.tune.ray_trial_executor import RayTrialExecutor
|
||||
from ray.tune.result import TRAINING_ITERATION
|
||||
from ray.tune.syncer import SyncConfig, SyncerCallback
|
||||
@@ -205,14 +205,14 @@ class TrialRunnerCallbacks(unittest.TestCase):
|
||||
"delay")
|
||||
|
||||
def testCallbackReordering(self):
|
||||
"""SyncerCallback should come after ExperimentLogger callbacks"""
|
||||
"""SyncerCallback should come after LoggerCallback callbacks"""
|
||||
|
||||
def get_positions(callbacks):
|
||||
first_logger_pos = None
|
||||
last_logger_pos = None
|
||||
syncer_pos = None
|
||||
for i, callback in enumerate(callbacks):
|
||||
if isinstance(callback, ExperimentLogger):
|
||||
if isinstance(callback, LoggerCallback):
|
||||
if first_logger_pos is None:
|
||||
first_logger_pos = i
|
||||
last_logger_pos = i
|
||||
@@ -233,8 +233,8 @@ class TrialRunnerCallbacks(unittest.TestCase):
|
||||
self.assertLess(last_logger_pos, syncer_pos)
|
||||
|
||||
# Auto creation of loggers with existing logger (but no CSV/JSON)
|
||||
callbacks = create_default_callbacks([ExperimentLogger()],
|
||||
SyncConfig(), None)
|
||||
callbacks = create_default_callbacks([LoggerCallback()], SyncConfig(),
|
||||
None)
|
||||
first_logger_pos, last_logger_pos, syncer_pos = get_positions(
|
||||
callbacks)
|
||||
self.assertLess(last_logger_pos, syncer_pos)
|
||||
@@ -242,13 +242,12 @@ class TrialRunnerCallbacks(unittest.TestCase):
|
||||
# This should throw an error as the syncer comes before the logger
|
||||
with self.assertRaises(ValueError):
|
||||
callbacks = create_default_callbacks(
|
||||
[SyncerCallback(None),
|
||||
ExperimentLogger()], SyncConfig(), None)
|
||||
[SyncerCallback(None), LoggerCallback()], SyncConfig(), None)
|
||||
|
||||
# This should be reordered but preserve the regular callback order
|
||||
[mc1, mc2, mc3] = [Callback(), Callback(), Callback()]
|
||||
# Has to be legacy logger to avoid logger callback creation
|
||||
lc = LegacyExperimentLogger(logger_classes=DEFAULT_LOGGERS)
|
||||
lc = LegacyLoggerCallback(logger_classes=DEFAULT_LOGGERS)
|
||||
callbacks = create_default_callbacks([mc1, mc2, lc, mc3], SyncConfig(),
|
||||
None)
|
||||
print(callbacks)
|
||||
|
||||
Reference in New Issue
Block a user