mirror of
https://github.com/wassname/ray.git
synced 2026-07-27 11:26:41 +08:00
[tune] stdout/stderr logging redirection (#9817)
* Add `log_to_file` parameter, pass to Trainable config, redirect stdout/stderr. * Add logging handler to root ray logger * Added test for `log_to_file` parameter * Added logs, reuse test * Revert debug change * Update logdir on reset, flush streams after each train() step * Remove magic keys from visible config Co-authored-by: Kai Fricke <kai@anyscale.com>
This commit is contained in:
@@ -1,7 +1,10 @@
|
||||
import os
|
||||
import unittest
|
||||
import sys
|
||||
|
||||
import ray
|
||||
from ray.tune import Trainable, run_experiments
|
||||
from ray import tune, logger
|
||||
from ray.tune import Trainable, run_experiments, register_trainable
|
||||
from ray.tune.error import TuneError
|
||||
from ray.tune.schedulers.trial_scheduler import FIFOScheduler, TrialScheduler
|
||||
|
||||
@@ -17,9 +20,15 @@ def create_resettable_class():
|
||||
self.config = config
|
||||
self.num_resets = 0
|
||||
self.iter = 0
|
||||
self.msg = config.get("message", "No message")
|
||||
|
||||
def step(self):
|
||||
self.iter += 1
|
||||
|
||||
print("PRINT_STDOUT: {}".format(self.msg))
|
||||
print("PRINT_STDERR: {}".format(self.msg), file=sys.stderr)
|
||||
logger.info("LOG_STDERR: {}".format(self.msg))
|
||||
|
||||
return {"num_resets": self.num_resets, "done": self.iter > 1}
|
||||
|
||||
def save_checkpoint(self, chkpt_dir):
|
||||
@@ -32,6 +41,7 @@ def create_resettable_class():
|
||||
if "fake_reset_not_supported" in self.config:
|
||||
return False
|
||||
self.num_resets += 1
|
||||
self.msg = new_config.get("message", "No message")
|
||||
return True
|
||||
|
||||
return MyResettableClass
|
||||
@@ -90,8 +100,50 @@ class ActorReuseTest(unittest.TestCase):
|
||||
|
||||
self.assertRaises(TuneError, lambda: run())
|
||||
|
||||
def testTrialReuseLogToFile(self):
|
||||
register_trainable("foo2", create_resettable_class())
|
||||
|
||||
# Log to default files
|
||||
[trial1, trial2] = tune.run(
|
||||
"foo2",
|
||||
config={
|
||||
"message": tune.grid_search(["First", "Second"])
|
||||
},
|
||||
log_to_file=True,
|
||||
scheduler=FrequentPausesScheduler(),
|
||||
reuse_actors=True).trials
|
||||
|
||||
# Check trial 1
|
||||
self.assertEqual(trial1.last_result["num_resets"], 1)
|
||||
self.assertTrue(os.path.exists(os.path.join(trial1.logdir, "stdout")))
|
||||
self.assertTrue(os.path.exists(os.path.join(trial1.logdir, "stderr")))
|
||||
with open(os.path.join(trial1.logdir, "stdout"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDOUT: First", content)
|
||||
self.assertNotIn("PRINT_STDOUT: Second", content)
|
||||
with open(os.path.join(trial1.logdir, "stderr"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDERR: First", content)
|
||||
self.assertIn("LOG_STDERR: First", content)
|
||||
self.assertNotIn("PRINT_STDERR: Second", content)
|
||||
self.assertNotIn("LOG_STDERR: Second", content)
|
||||
|
||||
# Check trial 2
|
||||
self.assertEqual(trial2.last_result["num_resets"], 2)
|
||||
self.assertTrue(os.path.exists(os.path.join(trial2.logdir, "stdout")))
|
||||
self.assertTrue(os.path.exists(os.path.join(trial2.logdir, "stderr")))
|
||||
with open(os.path.join(trial2.logdir, "stdout"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDOUT: Second", content)
|
||||
self.assertNotIn("PRINT_STDOUT: First", content)
|
||||
with open(os.path.join(trial2.logdir, "stderr"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDERR: Second", content)
|
||||
self.assertIn("LOG_STDERR: Second", content)
|
||||
self.assertNotIn("PRINT_STDERR: First", content)
|
||||
self.assertNotIn("LOG_STDERR: First", content)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pytest
|
||||
import sys
|
||||
sys.exit(pytest.main(["-v", __file__]))
|
||||
|
||||
@@ -973,6 +973,64 @@ class TrainableFunctionApiTest(unittest.TestCase):
|
||||
self.assertEqual(trial.status, Trial.TERMINATED)
|
||||
self.assertTrue(trial.has_checkpoint())
|
||||
|
||||
def testLogToFile(self):
|
||||
def train(config, reporter):
|
||||
import sys
|
||||
from ray import logger
|
||||
for i in range(10):
|
||||
reporter(timesteps_total=i)
|
||||
print("PRINT_STDOUT")
|
||||
print("PRINT_STDERR", file=sys.stderr)
|
||||
logger.info("LOG_STDERR")
|
||||
|
||||
register_trainable("f1", train)
|
||||
|
||||
# Do not log to file
|
||||
[trial] = tune.run("f1", log_to_file=False).trials
|
||||
self.assertFalse(os.path.exists(os.path.join(trial.logdir, "stdout")))
|
||||
self.assertFalse(os.path.exists(os.path.join(trial.logdir, "stderr")))
|
||||
|
||||
# Log to default files
|
||||
[trial] = tune.run("f1", log_to_file=True).trials
|
||||
self.assertTrue(os.path.exists(os.path.join(trial.logdir, "stdout")))
|
||||
self.assertTrue(os.path.exists(os.path.join(trial.logdir, "stderr")))
|
||||
with open(os.path.join(trial.logdir, "stdout"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDOUT", content)
|
||||
with open(os.path.join(trial.logdir, "stderr"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDERR", content)
|
||||
self.assertIn("LOG_STDERR", content)
|
||||
|
||||
# Log to one file
|
||||
[trial] = tune.run("f1", log_to_file="combined").trials
|
||||
self.assertFalse(os.path.exists(os.path.join(trial.logdir, "stdout")))
|
||||
self.assertFalse(os.path.exists(os.path.join(trial.logdir, "stderr")))
|
||||
self.assertTrue(os.path.exists(os.path.join(trial.logdir, "combined")))
|
||||
with open(os.path.join(trial.logdir, "combined"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDOUT", content)
|
||||
self.assertIn("PRINT_STDERR", content)
|
||||
self.assertIn("LOG_STDERR", content)
|
||||
|
||||
# Log to two files
|
||||
[trial] = tune.run(
|
||||
"f1", log_to_file=("alt.stdout", "alt.stderr")).trials
|
||||
self.assertFalse(os.path.exists(os.path.join(trial.logdir, "stdout")))
|
||||
self.assertFalse(os.path.exists(os.path.join(trial.logdir, "stderr")))
|
||||
self.assertTrue(
|
||||
os.path.exists(os.path.join(trial.logdir, "alt.stdout")))
|
||||
self.assertTrue(
|
||||
os.path.exists(os.path.join(trial.logdir, "alt.stderr")))
|
||||
|
||||
with open(os.path.join(trial.logdir, "alt.stdout"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDOUT", content)
|
||||
with open(os.path.join(trial.logdir, "alt.stderr"), "rt") as fp:
|
||||
content = fp.read()
|
||||
self.assertIn("PRINT_STDERR", content)
|
||||
self.assertIn("LOG_STDERR", content)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pytest
|
||||
|
||||
Reference in New Issue
Block a user