mirror of
https://github.com/wassname/ray.git
synced 2026-08-13 12:30:18 +08:00
[tune] Async restores and S3/GCP-capable trial FT (#6376)
* Initial commit for asynchronous save/restore * Set stage for cloud checkpointable trainable. * Refactor log_sync and sync_client. * Add durable trainable impl. * Support delete in cmd based client * Fix some tests and such * Cleanup, comments. * Use upload_dir instead. * Revert files belonging to other PR in split. * Pass upload_dir into trainable init. * Pickle checkpoint at driver, more robust checkpoint_dir discovery. * Cleanup trainable helper functions, fix tests. * Addressed comments. * Fix bugs from cluster testing, add parameterized cluster tests. * Add trainable util test * package_ref * pbt_address * Fix bug after running pbt example (_save returning dir). * get cluster tests running, other bug fixes. * raise_errors * Fix deleter bug, add durable trainable example. * Fix cluster test bugs. * filelock * save/restore bug fixes * . * Working cluster tests. * Lint, revert to tracking memory checkpoints. * Documentation, cleanup * fixinitialsync * fix_one_test * Fix cluster test bug * nit * lint * Revert tune md change * Fix basename bug for directories. * lint * fix_tests * nit_fix * Add __init__ file. * Move to utils package Co-authored-by: Richard Liaw <rliaw@berkeley.edu>
This commit is contained in:
committed by
Richard Liaw
co-authored by
Richard Liaw
parent
57061a15cf
commit
ca651af1d7
@@ -13,10 +13,12 @@ import ray
|
||||
from ray.exceptions import RayTimeoutError
|
||||
from ray import ray_constants
|
||||
from ray.resource_spec import ResourceSpec
|
||||
from ray.tune.durable_trainable import DurableTrainable
|
||||
from ray.tune.error import AbortTrialExecution
|
||||
from ray.tune.logger import NoopLogger
|
||||
from ray.tune.trial import Trial, Checkpoint, Location
|
||||
from ray.tune.resources import Resources
|
||||
from ray.tune.trainable import TrainableUtil
|
||||
from ray.tune.trial import Trial, Checkpoint, Location
|
||||
from ray.tune.trial_executor import TrialExecutor
|
||||
from ray.tune.util import warn_if_slow
|
||||
from ray.tune.error import TuneError
|
||||
@@ -27,7 +29,6 @@ RESOURCE_REFRESH_PERIOD = 0.5 # Refresh resources every 500 ms
|
||||
BOTTLENECK_WARN_PERIOD_S = 60
|
||||
NONTRIVIAL_WAIT_TIME_THRESHOLD_S = 1e-3
|
||||
DEFAULT_GET_TIMEOUT = 30.0 # seconds
|
||||
TRIAL_START_ATTEMPTS = 3
|
||||
|
||||
|
||||
class _LocalWrapper:
|
||||
@@ -86,7 +87,7 @@ class RayTrialExecutor(TrialExecutor):
|
||||
self._cached_actor)
|
||||
existing_runner = self._cached_actor
|
||||
self._cached_actor = None
|
||||
trial.runner = existing_runner
|
||||
trial.set_runner(existing_runner)
|
||||
if not self.reset_trial(trial, trial.config, trial.experiment_tag):
|
||||
raise AbortTrialExecution(
|
||||
"Trainable runner reuse requires reset_config() to be "
|
||||
@@ -122,11 +123,16 @@ class RayTrialExecutor(TrialExecutor):
|
||||
logger.debug("Trial %s: Setting up new remote runner.", trial)
|
||||
# Logging for trials is handled centrally by TrialRunner, so
|
||||
# configure the remote runner to use a noop-logger.
|
||||
return cls.remote(config=trial.config, logger_creator=logger_creator)
|
||||
kwargs = {
|
||||
"config": trial.config,
|
||||
"logger_creator": logger_creator,
|
||||
}
|
||||
if issubclass(trial.get_trainable_cls(), DurableTrainable):
|
||||
kwargs["remote_checkpoint_dir"] = trial.remote_checkpoint_dir
|
||||
return cls.remote(**kwargs)
|
||||
|
||||
def _train(self, trial):
|
||||
"""Start one iteration of training and save remote id."""
|
||||
|
||||
if self._find_item(self._paused, trial):
|
||||
raise TuneError(
|
||||
"Should not call `train` on PAUSED trial {}. "
|
||||
@@ -168,9 +174,11 @@ class RayTrialExecutor(TrialExecutor):
|
||||
"""
|
||||
prior_status = trial.status
|
||||
self.set_status(trial, Trial.RUNNING)
|
||||
trial.runner = runner or self._setup_remote_runner(
|
||||
trial,
|
||||
reuse_allowed=checkpoint is not None or trial.has_checkpoint())
|
||||
trial.set_runner(
|
||||
runner or self._setup_remote_runner(
|
||||
trial,
|
||||
reuse_allowed=checkpoint is not None
|
||||
or trial.has_checkpoint()))
|
||||
self.restore(trial, checkpoint)
|
||||
|
||||
previous_run = self._find_item(self._paused, trial)
|
||||
@@ -178,7 +186,7 @@ class RayTrialExecutor(TrialExecutor):
|
||||
# If Trial was in flight when paused, self._paused stores result.
|
||||
self._paused.pop(previous_run[0])
|
||||
self._running[previous_run[0]] = trial
|
||||
else:
|
||||
elif not trial.is_restoring:
|
||||
self._train(trial)
|
||||
|
||||
def _stop_trial(self, trial, error=False, error_msg=None,
|
||||
@@ -194,7 +202,6 @@ class RayTrialExecutor(TrialExecutor):
|
||||
error_msg (str): Optional error message.
|
||||
stop_logger (bool): Whether to shut down the trial logger.
|
||||
"""
|
||||
|
||||
if stop_logger:
|
||||
trial.close_logger()
|
||||
|
||||
@@ -206,7 +213,7 @@ class RayTrialExecutor(TrialExecutor):
|
||||
if hasattr(trial, "runner") and trial.runner:
|
||||
if (not error and self._reuse_actors
|
||||
and self._cached_actor is None):
|
||||
logger.debug("Reusing actor for {}".format(trial.runner))
|
||||
logger.debug("Reusing actor for %s", trial.runner)
|
||||
self._cached_actor = trial.runner
|
||||
else:
|
||||
logger.debug("Trial %s: Destroying actor.", trial)
|
||||
@@ -216,7 +223,7 @@ class RayTrialExecutor(TrialExecutor):
|
||||
logger.exception("Trial %s: Error stopping runner.", trial)
|
||||
self.set_status(trial, Trial.ERROR)
|
||||
finally:
|
||||
trial.runner = None
|
||||
trial.set_runner(None)
|
||||
|
||||
def start_trial(self, trial, checkpoint=None):
|
||||
"""Starts the trial.
|
||||
@@ -229,49 +236,22 @@ class RayTrialExecutor(TrialExecutor):
|
||||
of trial.
|
||||
"""
|
||||
self._commit_resources(trial.resources)
|
||||
remote_runner = None
|
||||
attempts = 0
|
||||
while attempts < TRIAL_START_ATTEMPTS:
|
||||
attempts += 1
|
||||
if attempts > 1:
|
||||
logger.warning("Trial %s: Start attempt #%s...", trial,
|
||||
attempts)
|
||||
try:
|
||||
self._start_trial(trial, checkpoint, remote_runner)
|
||||
break
|
||||
except AbortTrialExecution:
|
||||
logger.exception("Trial %s: Error starting runner, aborting!",
|
||||
trial)
|
||||
time.sleep(2)
|
||||
error_msg = traceback.format_exc()
|
||||
self._stop_trial(trial, error=True, error_msg=error_msg)
|
||||
break # don't retry fatal Tune errors
|
||||
except RayTimeoutError:
|
||||
# Reuse the existing runner on retries.
|
||||
remote_runner = trial.runner
|
||||
warning = ("Runner task timed out. This could be due to "
|
||||
"slow worker startup.")
|
||||
if attempts == TRIAL_START_ATTEMPTS:
|
||||
error_msg = traceback.format_exc()
|
||||
self._stop_trial(trial, error=True, error_msg=error_msg)
|
||||
else:
|
||||
warning += " Reusing the same runner."
|
||||
logger.warning("Trial %s: %s", trial, warning)
|
||||
except Exception:
|
||||
logger.exception("Trial %s: Error starting runner.", trial)
|
||||
time.sleep(2)
|
||||
error_msg = traceback.format_exc()
|
||||
self._stop_trial(trial, error=True, error_msg=error_msg)
|
||||
remote_runner = None
|
||||
# This forces the trial to not start from checkpoint.
|
||||
checkpoint = None
|
||||
trial.clear_checkpoint()
|
||||
# Note that we don't return the resources, since they may
|
||||
# have been lost. TODO(ujvl): is this the right thing to do?
|
||||
else:
|
||||
logger.exception(
|
||||
"Trial %s: Aborting trial after %s start "
|
||||
"attempts!", trial, TRIAL_START_ATTEMPTS)
|
||||
try:
|
||||
self._start_trial(trial, checkpoint)
|
||||
except AbortTrialExecution:
|
||||
logger.exception("Trial %s: Error starting runner, aborting!",
|
||||
trial)
|
||||
time.sleep(2)
|
||||
error_msg = traceback.format_exc()
|
||||
self._stop_trial(trial, error=True, error_msg=error_msg)
|
||||
except Exception:
|
||||
logger.exception("Trial %s: Unexpected error starting runner.",
|
||||
trial)
|
||||
time.sleep(2)
|
||||
error_msg = traceback.format_exc()
|
||||
self._stop_trial(trial, error=True, error_msg=error_msg)
|
||||
# Note that we don't return the resources, since they may
|
||||
# have been lost. TODO(ujvl): is this the right thing to do?
|
||||
|
||||
def _find_item(self, dictionary, item):
|
||||
out = [rid for rid, t in dictionary.items() if t is item]
|
||||
@@ -332,7 +312,6 @@ class RayTrialExecutor(TrialExecutor):
|
||||
|
||||
def get_running_trials(self):
|
||||
"""Returns the running trials."""
|
||||
|
||||
return list(self._running.values())
|
||||
|
||||
def get_alive_node_ips(self):
|
||||
@@ -387,7 +366,8 @@ class RayTrialExecutor(TrialExecutor):
|
||||
"""Fetches one result of the running trials.
|
||||
|
||||
Returns:
|
||||
Result of the most recent trial training run."""
|
||||
Result of the most recent trial training run.
|
||||
"""
|
||||
trial_future = self._find_item(self._running, trial)
|
||||
if not trial_future:
|
||||
raise ValueError("Trial was not running.")
|
||||
@@ -437,6 +417,7 @@ class RayTrialExecutor(TrialExecutor):
|
||||
"Resource invalid: {}".format(resources))
|
||||
|
||||
def _update_avail_resources(self, num_retries=5):
|
||||
resources = None
|
||||
for i in range(num_retries):
|
||||
try:
|
||||
resources = ray.cluster_resources()
|
||||
@@ -520,7 +501,6 @@ class RayTrialExecutor(TrialExecutor):
|
||||
|
||||
def debug_string(self):
|
||||
"""Returns a human readable message for printing to the console."""
|
||||
|
||||
if self._resources_initialized:
|
||||
status = ("Resources requested: {}/{} CPUs, {}/{} GPUs, "
|
||||
"{}/{} GiB heap, {}/{} GiB objects".format(
|
||||
@@ -548,7 +528,6 @@ class RayTrialExecutor(TrialExecutor):
|
||||
|
||||
def resource_string(self):
|
||||
"""Returns a string describing the total resources available."""
|
||||
|
||||
if self._resources_initialized:
|
||||
res_str = ("{} CPUs, {} GPUs, "
|
||||
"{} GiB heap, {} GiB objects".format(
|
||||
@@ -570,18 +549,28 @@ class RayTrialExecutor(TrialExecutor):
|
||||
"""Before step() called, update the available resources."""
|
||||
self._update_avail_resources()
|
||||
|
||||
def save(self, trial, storage=Checkpoint.DISK, result=None):
|
||||
"""Saves the trial's state to a checkpoint."""
|
||||
result = result or trial.last_result
|
||||
def save(self, trial, storage=Checkpoint.PERSISTENT, result=None):
|
||||
"""Saves the trial's state to a checkpoint.
|
||||
|
||||
Args:
|
||||
trial (Trial): The state of this trial to be saved.
|
||||
storage (str): Where to store the checkpoint. Defaults to
|
||||
PERSISTENT.
|
||||
result (dict): The state of this trial as a dictionary to be saved.
|
||||
If result is None, the trial's last result will be used.
|
||||
|
||||
Returns:
|
||||
Checkpoint future, or None if an Exception occurs.
|
||||
"""
|
||||
result = result or trial.last_result
|
||||
if storage == Checkpoint.MEMORY:
|
||||
value = trial.runner.save_to_object.remote()
|
||||
checkpoint = Checkpoint(storage, value, result)
|
||||
else:
|
||||
with warn_if_slow("save_checkpoint_to_disk"):
|
||||
with warn_if_slow("save_checkpoint_to_storage"):
|
||||
# TODO(ujvl): Make this asynchronous.
|
||||
value = ray.get(trial.runner.save.remote())
|
||||
checkpoint = Checkpoint(storage, value, result)
|
||||
|
||||
with warn_if_slow("on_checkpoint", DEFAULT_GET_TIMEOUT) as profile:
|
||||
try:
|
||||
trial.on_checkpoint(checkpoint)
|
||||
@@ -600,13 +589,10 @@ class RayTrialExecutor(TrialExecutor):
|
||||
def restore(self, trial, checkpoint=None):
|
||||
"""Restores training state from a given model checkpoint.
|
||||
|
||||
This will also sync the trial results to a new location
|
||||
if restoring on a different node.
|
||||
|
||||
Raises:
|
||||
RuntimeError: This error is raised if no runner is found.
|
||||
RayTimeoutError: This error is raised if a remote call to the
|
||||
runner times out.
|
||||
AbortTrialExecution: This error is raised if the trial is
|
||||
ineligible for restoration, given the Tune input arguments.
|
||||
"""
|
||||
if checkpoint is None or checkpoint.value is None:
|
||||
checkpoint = trial.checkpoint
|
||||
@@ -617,19 +603,29 @@ class RayTrialExecutor(TrialExecutor):
|
||||
"Trial {}: Unable to restore - no runner found.".format(trial))
|
||||
value = checkpoint.value
|
||||
if checkpoint.storage == Checkpoint.MEMORY:
|
||||
assert not isinstance(value, Checkpoint), type(value)
|
||||
logger.debug("Trial %s: Attempting restore from object", trial)
|
||||
# Note that we don't store the remote since in-memory checkpoints
|
||||
# don't guarantee fault tolerance and don't need to be waited on.
|
||||
trial.runner.restore_from_object.remote(value)
|
||||
else:
|
||||
logger.info("Trial %s: Attempting restore from %s", trial, value)
|
||||
with warn_if_slow("get_current_ip"):
|
||||
worker_ip = ray.get(trial.runner.current_ip.remote(),
|
||||
DEFAULT_GET_TIMEOUT)
|
||||
with warn_if_slow("sync_to_new_location"):
|
||||
trial.sync_logger_to_new_location(worker_ip)
|
||||
with warn_if_slow("restore_from_disk"):
|
||||
# TODO(ujvl): Take blocking restores out of the control loop.
|
||||
ray.get(trial.runner.restore.remote(value))
|
||||
trial.last_result = checkpoint.result
|
||||
logger.debug("Trial %s: Attempting restore from %s", trial, value)
|
||||
if issubclass(trial.get_trainable_cls(), DurableTrainable):
|
||||
remote = trial.runner.restore.remote(value)
|
||||
elif trial.sync_on_checkpoint:
|
||||
# This provides FT backwards compatibility in the
|
||||
# case where a DurableTrainable is not provided.
|
||||
logger.warning("Trial %s: Reading checkpoint into memory.",
|
||||
trial)
|
||||
data_dict = TrainableUtil.pickle_checkpoint(value)
|
||||
remote = trial.runner.restore_from_object.remote(data_dict)
|
||||
else:
|
||||
raise AbortTrialExecution(
|
||||
"Pass in `sync_on_checkpoint=True` for driver-based trial"
|
||||
"restoration. Pass in an `upload_dir` and a Trainable "
|
||||
"extending `DurableTrainable` for remote storage-based "
|
||||
"restoration")
|
||||
self._running[remote] = trial
|
||||
trial.restoring_from = checkpoint
|
||||
|
||||
def export_trial_if_needed(self, trial):
|
||||
"""Exports model of this trial based on trial.export_formats.
|
||||
|
||||
Reference in New Issue
Block a user