mirror of
https://github.com/wassname/ray.git
synced 2026-08-12 12:20:11 +08:00
[tune/raysgd] Tune API for TorchTrainer + Fix State Restoration (#7547)
This commit is contained in:
@@ -1,6 +1,6 @@
|
||||
import numpy as np
|
||||
import os
|
||||
import logging
|
||||
import os
|
||||
import numbers
|
||||
import tempfile
|
||||
import time
|
||||
@@ -10,9 +10,10 @@ import torch.distributed as dist
|
||||
import ray
|
||||
from ray.exceptions import RayActorError
|
||||
from ray.tune import Trainable
|
||||
from ray.tune.trial import Resources
|
||||
from ray.tune.resources import Resources
|
||||
from ray.tune.utils.util import merge_dicts
|
||||
from ray.util.sgd.torch.distributed_torch_runner import (
|
||||
DistributedTorchRunner, LocalDistributedRunner)
|
||||
DistributedTorchRunner, LocalDistributedRunner, DeactivatedRunner)
|
||||
from ray.util.sgd.utils import check_for_failure, NUM_SAMPLES, BATCH_SIZE
|
||||
from ray.util.sgd.torch.torch_runner import TorchRunner
|
||||
from ray.util.sgd.torch.constants import VALID_SCHEDULER_STEP
|
||||
@@ -130,6 +131,11 @@ class TorchTrainer:
|
||||
|
||||
"""
|
||||
|
||||
# TODO: Implement autoscaling. If num_workers=-1, the trainer will use as
|
||||
# many resources as available. Upon each train call, TorchTrainer will
|
||||
# query the Ray global state for total available resources and resize
|
||||
# its remote workers to consume all available resources.
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
@@ -218,6 +224,9 @@ class TorchTrainer:
|
||||
self._num_failures = 0
|
||||
self._last_resize = float("-inf")
|
||||
|
||||
self.local_worker = DeactivatedRunner()
|
||||
self.remote_workers = []
|
||||
|
||||
_validate_scheduler_step_freq(scheduler_step_freq)
|
||||
self.scheduler_step_freq = scheduler_step_freq
|
||||
|
||||
@@ -250,9 +259,6 @@ class TorchTrainer:
|
||||
if batch_size_per_worker:
|
||||
worker_config[BATCH_SIZE] = batch_size_per_worker
|
||||
|
||||
self.local_worker = None
|
||||
self.remote_workers = []
|
||||
|
||||
if num_workers == 1:
|
||||
# Start local worker
|
||||
self.local_worker = TorchRunner(
|
||||
@@ -319,8 +325,7 @@ class TorchTrainer:
|
||||
num_steps=None,
|
||||
profile=False,
|
||||
reduce_results=True,
|
||||
max_retries=0,
|
||||
checkpoint="auto",
|
||||
max_retries=3,
|
||||
info=None):
|
||||
"""Runs a training epoch.
|
||||
|
||||
@@ -339,14 +344,12 @@ class TorchTrainer:
|
||||
all workers into one dict. If a metric is a non-numerical
|
||||
value (or nested dictionaries), one value will be randomly
|
||||
selected among the workers. If False, returns a list of dicts.
|
||||
max_retries (int): Must be non-negative. If set to N, will
|
||||
kill all current workers, query the Ray global state for
|
||||
total available resources, and re-launch up to the
|
||||
available resources. Behavior is not well-defined
|
||||
in case of shared cluster usage.
|
||||
checkpoint (str): Path to checkpoint to restore from if retrying.
|
||||
If max_retries is set and ``checkpoint == "auto"``,
|
||||
TorchTrainer will save a checkpoint before starting to train.
|
||||
max_retries (int): Must be non-negative. If set to N, TorchTrainer
|
||||
will detect and recover from training failure. The recovery
|
||||
process will kill all current workers, query the Ray
|
||||
global state for total available resources, and re-launch up to
|
||||
the available resources. Behavior is not well-defined
|
||||
in case of shared cluster usage. Defaults to 3.
|
||||
info (dict): Optional dictionary passed to the training
|
||||
operator for ``train_epoch`` and ``train_batch``.
|
||||
|
||||
@@ -358,18 +361,9 @@ class TorchTrainer:
|
||||
length will be equal to ``num_workers``.
|
||||
"""
|
||||
assert max_retries >= 0, "`max_retries` must be non-negative."
|
||||
if max_retries:
|
||||
if checkpoint == "auto":
|
||||
logger.debug("Retrying detected. Automatically checkpointing.")
|
||||
checkpoint = self.save(
|
||||
os.path.join(self.temp_dir, "tmp_checkpoint"))
|
||||
elif not checkpoint:
|
||||
raise ValueError("Cannot retry from empty checkpoint.")
|
||||
|
||||
if checkpoint and self._should_resize():
|
||||
if self._should_resize():
|
||||
logger.info("Resize opportunity detected. Attempting to scale up.")
|
||||
self._resize_workers(checkpoint=checkpoint)
|
||||
|
||||
self._resize_workers()
|
||||
success, worker_stats = self._train_epoch(
|
||||
num_steps=num_steps, profile=profile, info=info)
|
||||
# Fault handling
|
||||
@@ -378,7 +372,7 @@ class TorchTrainer:
|
||||
break
|
||||
else:
|
||||
self._num_failures += 1
|
||||
self._resize_workers(checkpoint=checkpoint)
|
||||
self._resize_workers()
|
||||
logger.info("Retrying training step with %d workers." %
|
||||
(len(self.remote_workers) + 1))
|
||||
success, worker_stats = self._train_epoch(
|
||||
@@ -483,7 +477,6 @@ class TorchTrainer:
|
||||
w.validate.remote(**params) for w in self.remote_workers
|
||||
]
|
||||
local_worker_stats = self.local_worker.validate(**params)
|
||||
|
||||
return self._process_stats([local_worker_stats] +
|
||||
ray.get(remote_worker_stats))
|
||||
|
||||
@@ -497,47 +490,49 @@ class TorchTrainer:
|
||||
|
||||
def get_model(self):
|
||||
"""Returns the learned model(s)."""
|
||||
models = self.model_creator(self.config)
|
||||
state = self.local_worker.get_state()
|
||||
if len(state["models"]) == 1:
|
||||
models.load_state_dict(state["models"][0])
|
||||
else:
|
||||
for model, state_dict in zip(models, state["models"]):
|
||||
model.load_state_dict(state_dict)
|
||||
return models
|
||||
unwrapped = []
|
||||
for model in self.local_worker.models:
|
||||
unwrapped += [model.module if hasattr(model, "module") else model]
|
||||
if len(unwrapped) == 1:
|
||||
return unwrapped[0]
|
||||
return unwrapped
|
||||
|
||||
def state_dict(self):
|
||||
return self.local_worker.get_state()
|
||||
return self.local_worker.state_dict()
|
||||
|
||||
def load_state_dict(self, state):
|
||||
state_id = ray.put(state)
|
||||
def load_state_dict(self, state_dict, blocking=False):
|
||||
# This is not the most efficient because you have to wait for
|
||||
# the local worker to save then dump to buffer.
|
||||
self.local_worker.load_state_dict(state_dict)
|
||||
state_id = ray.put(self.local_worker.state_stream())
|
||||
|
||||
remote_calls = [
|
||||
worker.set_state.remote(state_id) for worker in self.remote_workers
|
||||
worker.load_state_stream.remote(state_id)
|
||||
for worker in self.remote_workers
|
||||
]
|
||||
self.local_worker.set_state(state)
|
||||
ray.get(remote_calls)
|
||||
if blocking:
|
||||
ray.get(remote_calls)
|
||||
|
||||
def save(self, checkpoint):
|
||||
"""Saves the model(s) to the provided checkpoint.
|
||||
"""Saves the Trainer state to the provided checkpoint path.
|
||||
|
||||
Args:
|
||||
checkpoint (str): Path to target checkpoint file.
|
||||
|
||||
Returns:
|
||||
checkpoint (str): Path to target checkpoint file.
|
||||
"""
|
||||
torch.save(self.state_dict(), checkpoint)
|
||||
return checkpoint
|
||||
|
||||
def restore(self, checkpoint):
|
||||
"""Restores the Trainer and all workers from the provided checkpoint.
|
||||
def load(self, checkpoint):
|
||||
"""Loads the Trainer and all workers from the provided checkpoint.
|
||||
|
||||
Args:
|
||||
checkpoint (str): Path to target checkpoint file.
|
||||
"""
|
||||
state = torch.load(checkpoint)
|
||||
self.load_state_dict(state)
|
||||
state_dict = torch.load(checkpoint)
|
||||
self.load_state_dict(state_dict)
|
||||
|
||||
def restore(self, *args):
|
||||
raise DeprecationWarning("Use `TorchTrainer.load()` instead.")
|
||||
|
||||
def shutdown(self, force=False):
|
||||
"""Shuts down workers and releases resources."""
|
||||
@@ -562,19 +557,19 @@ class TorchTrainer:
|
||||
else:
|
||||
self.local_worker.shutdown()
|
||||
for worker in self.remote_workers:
|
||||
logger.warning("Killing worker {}.".format(worker))
|
||||
logger.debug("Killing worker {}.".format(worker))
|
||||
ray.kill(worker)
|
||||
|
||||
self.local_worker = None
|
||||
self.local_worker = DeactivatedRunner()
|
||||
self.remote_workers = []
|
||||
|
||||
def _reset(self):
|
||||
"""Terminates models without giving up local resource reservation."""
|
||||
self.local_worker.shutdown(cleanup=False)
|
||||
for worker in self.remote_workers:
|
||||
logger.warning("Killing worker {}.".format(worker))
|
||||
logger.debug("Killing worker {}.".format(worker))
|
||||
ray.kill(worker)
|
||||
self.local_worker = None
|
||||
self.local_worker = DeactivatedRunner()
|
||||
self.remote_workers = []
|
||||
|
||||
def _check_potential_remote_workers_size(self):
|
||||
@@ -588,9 +583,8 @@ class TorchTrainer:
|
||||
remote_resources.get("GPU", 0), new_remote_workers)
|
||||
return new_remote_workers
|
||||
|
||||
def _resize_workers(self, checkpoint, max_retries=10):
|
||||
def _resize_workers(self, max_retries=10):
|
||||
self._reset()
|
||||
assert checkpoint, "Cannot restore without checkpoint."
|
||||
|
||||
time.sleep(1)
|
||||
for i in range(max_retries):
|
||||
@@ -598,7 +592,7 @@ class TorchTrainer:
|
||||
if new_remote_workers:
|
||||
self._last_resize = time.time()
|
||||
self._start_workers(int(new_remote_workers) + 1)
|
||||
self.restore(checkpoint)
|
||||
self.load_state_dict(self.state_dict())
|
||||
return
|
||||
else:
|
||||
delay = 2**i
|
||||
@@ -617,32 +611,128 @@ class TorchTrainer:
|
||||
return potential_remote_size > 0
|
||||
return False
|
||||
|
||||
|
||||
class TorchTrainable(Trainable):
|
||||
@classmethod
|
||||
def default_resource_request(cls, config):
|
||||
remote_worker_count = config["num_workers"] - 1
|
||||
return Resources(
|
||||
cpu=1,
|
||||
gpu=int(config["use_gpu"]),
|
||||
extra_cpu=int(remote_worker_count),
|
||||
extra_gpu=int(int(config["use_gpu"]) * remote_worker_count))
|
||||
def as_trainable(cls, *args, **kwargs):
|
||||
"""Creates a BaseTorchTrainable class compatible with Tune.
|
||||
|
||||
Any configuration parameters will be overriden by the Tune
|
||||
Trial configuration. You can also subclass the provided Trainable
|
||||
to implement your own iterative optimization routine.
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
TorchTrainable = TorchTrainer.as_trainable(
|
||||
model_creator=ResNet18,
|
||||
data_creator=cifar_creator,
|
||||
optimizer_creator=optimizer_creator,
|
||||
loss_creator=nn.CrossEntropyLoss,
|
||||
num_gpus=2
|
||||
)
|
||||
analysis = tune.run(
|
||||
TorchTrainable,
|
||||
config={"lr": tune.grid_search([0.01, 0.1])}
|
||||
)
|
||||
|
||||
"""
|
||||
|
||||
class TorchTrainable(BaseTorchTrainable):
|
||||
@classmethod
|
||||
def default_resource_request(cls, config):
|
||||
num_workers = config.get("num_workers",
|
||||
kwargs.get("num_workers", 1))
|
||||
use_gpu = config.get("use_gpu", kwargs.get("use_gpu"))
|
||||
|
||||
remote_worker_count = num_workers - 1
|
||||
|
||||
return Resources(
|
||||
cpu=1,
|
||||
gpu=int(use_gpu),
|
||||
extra_cpu=int(remote_worker_count),
|
||||
extra_gpu=int(int(use_gpu) * remote_worker_count))
|
||||
|
||||
def _create_trainer(self, tune_config):
|
||||
"""Overrides the provided config with Tune config."""
|
||||
provided_config = kwargs.get("config", {}).copy()
|
||||
provided_config.update(tune_config)
|
||||
kwargs["config"] = provided_config
|
||||
trainer = TorchTrainer(*args, **kwargs)
|
||||
return trainer
|
||||
|
||||
return TorchTrainable
|
||||
|
||||
|
||||
class BaseTorchTrainable(Trainable):
|
||||
"""Base class for converting TorchTrainer to a Trainable class.
|
||||
|
||||
This class is produced when you call ``TorchTrainer.as_trainable(...)``.
|
||||
|
||||
You can override the produced Trainable to implement custom iterative
|
||||
training procedures:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
TorchTrainable = TorchTrainer.as_trainable(
|
||||
model_creator=ResNet18,
|
||||
data_creator=cifar_creator,
|
||||
optimizer_creator=optimizer_creator,
|
||||
loss_creator=nn.CrossEntropyLoss,
|
||||
num_gpus=2
|
||||
)
|
||||
# TorchTrainable is subclass of BaseTorchTrainable.
|
||||
|
||||
class CustomTrainable(TorchTrainable):
|
||||
def _train(self):
|
||||
for i in range(5):
|
||||
train_stats = self.trainer.train()
|
||||
validation_stats = self.trainer.validate()
|
||||
train_stats.update(validation_stats)
|
||||
return train_stats
|
||||
|
||||
analysis = tune.run(
|
||||
CustomTrainable,
|
||||
config={"lr": tune.grid_search([0.01, 0.1])}
|
||||
)
|
||||
|
||||
"""
|
||||
|
||||
def _setup(self, config):
|
||||
self._trainer = TorchTrainer(**config)
|
||||
"""Constructs a TorchTrainer object as `self.trainer`."""
|
||||
self._trainer = self._create_trainer(config)
|
||||
|
||||
def _train(self):
|
||||
train_stats = self._trainer.train()
|
||||
validation_stats = self._trainer.validate()
|
||||
"""Calls `self.trainer.train()` and `self.trainer.validate()` once.
|
||||
|
||||
train_stats.update(validation_stats)
|
||||
return train_stats
|
||||
You may want to override this if using a custom LR scheduler.
|
||||
"""
|
||||
train_stats = self.trainer.train(max_retries=10, profile=True)
|
||||
validation_stats = self.trainer.validate(profile=True)
|
||||
stats = merge_dicts(train_stats, validation_stats)
|
||||
return stats
|
||||
|
||||
def _save(self, checkpoint_dir):
|
||||
return self._trainer.save(os.path.join(checkpoint_dir, "model.pth"))
|
||||
"""Returns a path containing the trainer state."""
|
||||
checkpoint_path = os.path.join(checkpoint_dir, "trainer.checkpoint")
|
||||
self.trainer.save(checkpoint_path)
|
||||
return checkpoint_path
|
||||
|
||||
def _restore(self, checkpoint_path):
|
||||
return self._trainer.restore(checkpoint_path)
|
||||
"""Restores the trainer state.
|
||||
|
||||
Override this if you have state external to the Trainer object.
|
||||
"""
|
||||
return self.trainer.load(checkpoint_path)
|
||||
|
||||
def _stop(self):
|
||||
self._trainer.shutdown()
|
||||
"""Shuts down the trainer."""
|
||||
self.trainer.shutdown()
|
||||
|
||||
def _create_trainer(self, config):
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def trainer(self):
|
||||
"""An instantiated TorchTrainer object.
|
||||
|
||||
Use this when specifying custom training procedures for Tune.
|
||||
"""
|
||||
return self._trainer
|
||||
|
||||
Reference in New Issue
Block a user