[tune] logger migration to ExperimentLogger classes (#11984)

This commit is contained in:
Kai Fricke
2020-11-16 15:08:37 -08:00
committed by GitHub
parent 3dc68533a9
commit 9f5986ee58
14 changed files with 897 additions and 223 deletions
+138 -17
View File
@@ -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")
+182 -19
View File
@@ -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
+6 -6
View File
@@ -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)