mirror of
https://github.com/wassname/ray.git
synced 2026-09-11 12:43:20 +08:00
[Ray SGD] use_local flag + Worker group abstraction (#10539)
Co-authored-by: Richard Liaw <rliaw@berkeley.edu>
This commit is contained in:
co-authored by
Richard Liaw
parent
0865d68466
commit
d5a7c53908
@@ -1,29 +1,25 @@
|
||||
from datetime import timedelta
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import logging
|
||||
import os
|
||||
import numbers
|
||||
import tempfile
|
||||
import time
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
import ray
|
||||
from ray.exceptions import RayActorError
|
||||
from ray.tune import Trainable
|
||||
from ray.tune.resources import Resources
|
||||
from ray.tune.utils.util import merge_dicts
|
||||
from ray.util import log_once
|
||||
from ray.util.sgd.torch.distributed_torch_runner import (
|
||||
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.worker_group import LocalWorkerGroup, \
|
||||
RemoteWorkerGroup, DeactivatedWorkerGroup
|
||||
from ray.util.sgd.utils import NUM_SAMPLES, BATCH_SIZE
|
||||
from ray.util.sgd.torch.constants import VALID_SCHEDULER_STEP, NCCL_TIMEOUT_S
|
||||
from ray.util.sgd.torch.utils import setup_address
|
||||
from ray.util.sgd.data import Dataset
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
RESIZE_COOLDOWN_S = 10
|
||||
|
||||
|
||||
def _validate_scheduler_step_freq(scheduler_step_freq):
|
||||
@@ -86,11 +82,17 @@ class TorchTrainer:
|
||||
that subclasses the TrainingOperator class. This class
|
||||
will be copied onto all remote workers and used to specify
|
||||
training components and custom training and validation operations.
|
||||
initialization_hook (function): A function to call on all training
|
||||
workers when they are first initialized. This could be useful to
|
||||
set environment variables for all the worker processes.
|
||||
config (dict): Custom configuration value to be passed to
|
||||
all operator constructors.
|
||||
num_workers (int): the number of workers used in distributed
|
||||
training. If 1, the worker will not be wrapped with
|
||||
DistributedDataParallel.
|
||||
DistributedDataParallel. TorchTrainer will scale down the number
|
||||
of workers if enough resources are not available, and will scale
|
||||
back up once they are. The total number of
|
||||
workers will never exceed `num_workers` amount.
|
||||
num_cpus_per_worker (int): Sets the cpu requirement for each worker.
|
||||
use_gpu (bool): Sets resource allocation for workers to 1 GPU
|
||||
if true, and automatically moves both the model and optimizer
|
||||
@@ -119,9 +121,13 @@ class TorchTrainer:
|
||||
``step`` will be called after every optimizer step. If "epoch",
|
||||
``step`` will be called after one pass of the DataLoader. If
|
||||
"manual", the scheduler will not be incremented automatically -
|
||||
you are expected to call ``trainer.update_schedulers`` manually.
|
||||
you are expected to call ``trainer.update_scheduler`` manually.
|
||||
If a scheduler is passed in, this value is expected to not be None.
|
||||
|
||||
use_local (bool): If True, 1 worker will be a local worker running
|
||||
on the driver process, and all other workers will be remote. If
|
||||
False, all workers will be remote. Set this to True for easy
|
||||
debugging of worker on driver process, but could also
|
||||
lead to issues with Cuda devices. Defaults to False.
|
||||
"""
|
||||
|
||||
# TODO: Implement autoscaling. If num_workers=-1, the trainer will use as
|
||||
@@ -146,6 +152,7 @@ class TorchTrainer:
|
||||
apex_args=None,
|
||||
add_dist_sampler=True,
|
||||
scheduler_step_freq=None,
|
||||
use_local=False,
|
||||
# Deprecated Args.
|
||||
num_replicas=None,
|
||||
batch_size=None,
|
||||
@@ -168,6 +175,13 @@ class TorchTrainer:
|
||||
"model_creator, ...) and pass in CustomOperator into "
|
||||
"TorchTrainer.")
|
||||
|
||||
if use_local and log_once("use_local"):
|
||||
logger.warning("use_local is set to True. This could lead to "
|
||||
"issues with Cuda devices. If you are seeing this "
|
||||
"issue, try setting use_local to False. For more "
|
||||
"information, see "
|
||||
"https://github.com/ray-project/ray/issues/9202.")
|
||||
|
||||
if num_workers > 1 and not dist.is_available():
|
||||
raise ValueError(
|
||||
("Distributed PyTorch is not supported on macOS. "
|
||||
@@ -225,6 +239,7 @@ class TorchTrainer:
|
||||
self.use_fp16 = use_fp16
|
||||
self.use_tqdm = use_tqdm
|
||||
self.add_dist_sampler = add_dist_sampler
|
||||
self.use_local = use_local
|
||||
|
||||
if apex_args and not isinstance(apex_args, dict):
|
||||
raise ValueError("apex_args needs to be a dict object.")
|
||||
@@ -234,9 +249,6 @@ class TorchTrainer:
|
||||
self._num_failures = 0
|
||||
self._last_resize = float("-inf")
|
||||
|
||||
self.local_worker = DeactivatedRunner()
|
||||
self.remote_workers = []
|
||||
|
||||
if scheduler_step_freq:
|
||||
_validate_scheduler_step_freq(scheduler_step_freq)
|
||||
|
||||
@@ -270,12 +282,10 @@ class TorchTrainer:
|
||||
return batch_size_per_worker
|
||||
|
||||
def _start_workers(self, num_workers):
|
||||
logger.debug(f"start_workers: Setting %d workers." % num_workers)
|
||||
worker_config = self.config.copy()
|
||||
batch_size_per_worker = self._configure_and_split_batch(num_workers)
|
||||
if batch_size_per_worker:
|
||||
worker_config[BATCH_SIZE] = batch_size_per_worker
|
||||
|
||||
params = dict(
|
||||
training_operator_cls=self.training_operator_cls,
|
||||
config=worker_config,
|
||||
@@ -286,57 +296,68 @@ class TorchTrainer:
|
||||
apex_args=self.apex_args,
|
||||
scheduler_step_freq=self.scheduler_step_freq)
|
||||
|
||||
if num_workers == 1:
|
||||
# Start local worker
|
||||
self.local_worker = TorchRunner(**params)
|
||||
if self.initialization_hook:
|
||||
self.apply_all_workers(self.initialization_hook)
|
||||
self.local_worker.setup_operator()
|
||||
dist_params = dict(
|
||||
backend=self.backend,
|
||||
add_dist_sampler=self.add_dist_sampler,
|
||||
wrap_ddp=self.wrap_ddp)
|
||||
|
||||
worker_args = {
|
||||
"max_workers": num_workers,
|
||||
"params": params,
|
||||
"dist_params": dist_params,
|
||||
"initialization_hook": self.initialization_hook,
|
||||
"num_cpus_per_worker": self.num_cpus_per_worker,
|
||||
"use_gpu": self.use_gpu,
|
||||
"timeout_s": self.timeout_s
|
||||
}
|
||||
|
||||
if self.use_local:
|
||||
self.worker_group = LocalWorkerGroup(**worker_args)
|
||||
else:
|
||||
params.update(
|
||||
backend=self.backend,
|
||||
add_dist_sampler=self.add_dist_sampler,
|
||||
wrap_ddp=self.wrap_ddp)
|
||||
self.worker_group = RemoteWorkerGroup(**worker_args)
|
||||
|
||||
# Start local worker
|
||||
self.local_worker = LocalDistributedRunner(
|
||||
num_cpus=self.num_cpus_per_worker,
|
||||
num_gpus=int(self.use_gpu),
|
||||
**params)
|
||||
# TODO(amogkam): If not enough resources are available to create
|
||||
# num_workers workers, this command will hang. Instead,
|
||||
# start_workers should take into account available resources when
|
||||
# determining how many workers to create.
|
||||
self.worker_group.start_workers(num_workers)
|
||||
|
||||
# Generate actor class
|
||||
RemoteRunner = ray.remote(
|
||||
num_cpus=self.num_cpus_per_worker,
|
||||
num_gpus=int(self.use_gpu))(DistributedTorchRunner)
|
||||
# Start workers
|
||||
self.remote_workers = [
|
||||
RemoteRunner.remote(**params) for i in range(num_workers - 1)
|
||||
]
|
||||
if self.initialization_hook:
|
||||
self.apply_all_workers(self.initialization_hook)
|
||||
def _resize_worker_group(self, max_retries=10):
|
||||
"""Resizes the number of remote workers based on available resources.
|
||||
Total number of workers will never exceed `num_workers` amount.
|
||||
|
||||
# Compute URL for initializing distributed PyTorch
|
||||
address = setup_address()
|
||||
Args:
|
||||
max_retries (int): How many times to attempt to resize workers
|
||||
before failing.
|
||||
"""
|
||||
state_dict = self.state_dict()
|
||||
old_workers = self.worker_group.num_workers
|
||||
self.worker_group.reset()
|
||||
|
||||
# Setup the process group among all workers.
|
||||
remote_pgroup_setups = [
|
||||
worker.setup_process_group.remote(address, i + 1, num_workers,
|
||||
timedelta(self.timeout_s))
|
||||
for i, worker in enumerate(self.remote_workers)
|
||||
]
|
||||
self.local_worker.setup_process_group(address, 0, num_workers,
|
||||
timedelta(self.timeout_s))
|
||||
# Get setup tasks in order to throw errors on failure
|
||||
ray.get(remote_pgroup_setups)
|
||||
|
||||
# Runs code that requires all creator functions to have run.
|
||||
remote_operator_setups = [
|
||||
worker.setup_operator.remote()
|
||||
for worker in self.remote_workers
|
||||
]
|
||||
self.local_worker.setup_operator()
|
||||
# Get setup tasks in order to throw errors on failure
|
||||
ray.get(remote_operator_setups)
|
||||
time.sleep(1)
|
||||
for i in range(max_retries):
|
||||
new_workers = self.worker_group.new_workers_size()
|
||||
if new_workers:
|
||||
self._last_resize = time.time()
|
||||
self._start_workers(int(new_workers))
|
||||
self.load_state_dict(state_dict, blocking=True)
|
||||
if self.use_local and new_workers == 1 and old_workers > 1:
|
||||
# Major hack. If we go from LocalDistributedRunner to a
|
||||
# standard TorchRunner we have to manually reset the
|
||||
# dummy actor handle global vars.
|
||||
# TODO(amog): Refactor LocalDistributedTorchRunner to
|
||||
# not use global variables for resource reservation.
|
||||
ray.util.sgd.torch.distributed_torch_runner\
|
||||
._dummy_cuda_actor = None
|
||||
ray.util.sgd.torch.distributed_torch_runner\
|
||||
._dummy_cpu_actor = None
|
||||
return
|
||||
else:
|
||||
delay = 2**i
|
||||
logger.warning(
|
||||
"No new workers found. Retrying in %d sec." % delay)
|
||||
time.sleep(delay)
|
||||
raise RuntimeError("Exceeded max_retries for relaunching workers.")
|
||||
|
||||
def train(self,
|
||||
num_steps=None,
|
||||
@@ -384,10 +405,10 @@ class TorchTrainer:
|
||||
assert isinstance(dataset, Dataset) is not None \
|
||||
or self.data_creator, \
|
||||
"Must specify either a data creator or a dataset"
|
||||
if self._should_resize():
|
||||
if self.worker_group.should_scale_up():
|
||||
logger.info("Resize opportunity detected. Attempting to scale up.")
|
||||
self._resize_workers()
|
||||
success, worker_stats = self._train_epoch(
|
||||
self._resize_worker_group()
|
||||
success, worker_stats = self.worker_group.train(
|
||||
num_steps=num_steps, profile=profile, info=info, dataset=dataset)
|
||||
# Fault handling
|
||||
for i in range(max_retries):
|
||||
@@ -395,10 +416,10 @@ class TorchTrainer:
|
||||
break
|
||||
else:
|
||||
self._num_failures += 1
|
||||
self._resize_workers()
|
||||
self._resize_worker_group()
|
||||
logger.info("Retrying training step with %d workers." %
|
||||
(len(self.remote_workers) + 1))
|
||||
success, worker_stats = self._train_epoch(
|
||||
self.worker_group.num_workers)
|
||||
success, worker_stats = self.worker_group.train(
|
||||
num_steps=num_steps,
|
||||
profile=profile,
|
||||
info=info,
|
||||
@@ -425,43 +446,6 @@ class TorchTrainer:
|
||||
stats[stat_key] = worker_stats[0][stat_key]
|
||||
return stats
|
||||
|
||||
def _train_epoch(self,
|
||||
num_steps=None,
|
||||
profile=False,
|
||||
info=None,
|
||||
dataset=None):
|
||||
params = dict(num_steps=num_steps, profile=profile, info=info)
|
||||
remote_worker_stats = []
|
||||
if dataset:
|
||||
dataset.set_num_shards(self.max_replicas)
|
||||
for i, w in enumerate(self.remote_workers):
|
||||
params = dict(num_steps=num_steps, profile=profile, info=info)
|
||||
if dataset:
|
||||
params["iterator"] = dataset.get_shard(i)
|
||||
stats = w.train_epoch.remote(**params)
|
||||
remote_worker_stats.append(stats)
|
||||
|
||||
try:
|
||||
if dataset:
|
||||
params["iterator"] = dataset.get_shard(
|
||||
len(self.remote_workers))
|
||||
local_worker_stats = self.local_worker.train_epoch(**params)
|
||||
except RuntimeError as err:
|
||||
if "gloo" in err.args[0] and "Timed out" in err.args[0]:
|
||||
logger.warning(err)
|
||||
return False, None
|
||||
if "NCCL" in err.args[0]: # there is no specific error message
|
||||
logger.warning(err)
|
||||
return False, None
|
||||
|
||||
raise err
|
||||
|
||||
success = check_for_failure(remote_worker_stats)
|
||||
if success:
|
||||
return success, [local_worker_stats] + ray.get(remote_worker_stats)
|
||||
|
||||
return success, None
|
||||
|
||||
def apply_all_workers(self, fn):
|
||||
"""Run a function on all operators on the workers.
|
||||
|
||||
@@ -472,9 +456,7 @@ class TorchTrainer:
|
||||
A list of objects returned by ``fn`` on each worker.
|
||||
|
||||
"""
|
||||
remote_calls = [w.apply.remote(fn) for w in self.remote_workers]
|
||||
local_call = self.local_worker.apply(fn)
|
||||
return [local_call] + ray.get(remote_calls)
|
||||
return self.worker_group.apply_all_workers(fn)
|
||||
|
||||
def apply_all_operators(self, fn):
|
||||
"""Run a function on all operators on the workers.
|
||||
@@ -487,11 +469,7 @@ class TorchTrainer:
|
||||
A list of objects returned by ``fn`` on each operator.
|
||||
|
||||
"""
|
||||
remote_calls = [
|
||||
w.apply_operator.remote(fn) for w in self.remote_workers
|
||||
]
|
||||
local_call = self.local_worker.apply_operator(fn)
|
||||
return [local_call] + ray.get(remote_calls)
|
||||
return self.worker_group.apply_all_operators(fn)
|
||||
|
||||
def validate(self,
|
||||
num_steps=None,
|
||||
@@ -517,13 +495,8 @@ class TorchTrainer:
|
||||
You can provide custom metrics by passing in a custom
|
||||
``training_operator_cls``.
|
||||
"""
|
||||
params = dict(num_steps=num_steps, profile=profile, info=info)
|
||||
|
||||
remote_worker_stats = [
|
||||
w.validate.remote(**params) for w in self.remote_workers
|
||||
]
|
||||
local_worker_stats = self.local_worker.validate(**params)
|
||||
worker_stats = [local_worker_stats] + ray.get(remote_worker_stats)
|
||||
worker_stats = self.worker_group.validate(
|
||||
num_steps=num_steps, profile=profile, info=info)
|
||||
|
||||
if reduce_results:
|
||||
return self._process_stats(worker_stats)
|
||||
@@ -535,13 +508,14 @@ class TorchTrainer:
|
||||
|
||||
This is useful for lr_schedulers such as ``ReduceLROnPlateau``.
|
||||
"""
|
||||
self.apply_all_operators(
|
||||
self.worker_group.apply_all_operators(
|
||||
lambda op: [sched.step(metric) for sched in op._schedulers])
|
||||
|
||||
def get_model(self):
|
||||
"""Returns the learned model(s)."""
|
||||
unwrapped = []
|
||||
for model in self.local_worker.models:
|
||||
models = self.worker_group.get_model()
|
||||
for model in models:
|
||||
unwrapped += [model.module if hasattr(model, "module") else model]
|
||||
if len(unwrapped) == 1:
|
||||
return unwrapped[0]
|
||||
@@ -556,23 +530,13 @@ class TorchTrainer:
|
||||
Returns:
|
||||
TrainingOperator: The local TrainingOperator object.
|
||||
"""
|
||||
return self.local_worker.training_operator
|
||||
return self.worker_group.get_local_operator()
|
||||
|
||||
def state_dict(self):
|
||||
return self.local_worker.state_dict()
|
||||
return self.worker_group.state_dict()
|
||||
|
||||
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.load_state_stream.remote(state_id)
|
||||
for worker in self.remote_workers
|
||||
]
|
||||
if blocking:
|
||||
ray.get(remote_calls)
|
||||
self.worker_group.load_state_dict(state_dict, blocking=blocking)
|
||||
|
||||
def save(self, checkpoint):
|
||||
"""Saves the Trainer state to the provided checkpoint path.
|
||||
@@ -596,81 +560,16 @@ class TorchTrainer:
|
||||
raise DeprecationWarning("Use `TorchTrainer.load()` instead.")
|
||||
|
||||
def shutdown(self, force=False):
|
||||
"""Shuts down workers and releases resources."""
|
||||
if not force:
|
||||
cleanup = [
|
||||
worker.shutdown.remote() for worker in self.remote_workers
|
||||
]
|
||||
self.local_worker.shutdown()
|
||||
try:
|
||||
ray.get(cleanup)
|
||||
[
|
||||
worker.__ray_terminate__.remote()
|
||||
for worker in self.remote_workers
|
||||
]
|
||||
except RayActorError:
|
||||
logger.warning(
|
||||
"Failed to shutdown gracefully, forcing a shutdown.")
|
||||
"""Shuts down workers and releases resources.
|
||||
|
||||
for worker in self.remote_workers:
|
||||
logger.warning(f"Killing worker {worker}.")
|
||||
ray.kill(worker)
|
||||
else:
|
||||
self.local_worker.shutdown()
|
||||
for worker in self.remote_workers:
|
||||
logger.debug(f"Killing worker {worker}.")
|
||||
ray.kill(worker)
|
||||
Args:
|
||||
force (bool): If True, forcefully kill all workers. If False,
|
||||
attempt a graceful shutdown first, and then forcefully kill if
|
||||
unsuccessful.
|
||||
|
||||
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.debug(f"Killing worker {worker}.")
|
||||
ray.kill(worker)
|
||||
self.local_worker = DeactivatedRunner()
|
||||
self.remote_workers = []
|
||||
|
||||
def _check_potential_remote_workers_size(self):
|
||||
# ASSUME 1 GPU + 1 CPU is already reserved for the local worker
|
||||
remote_resources = ray.available_resources()
|
||||
max_remote_workers = self.max_replicas - 1
|
||||
new_remote_workers = min(
|
||||
remote_resources.get("CPU", 0), max_remote_workers)
|
||||
if self.use_gpu:
|
||||
new_remote_workers = min(
|
||||
remote_resources.get("GPU", 0), new_remote_workers)
|
||||
return new_remote_workers
|
||||
|
||||
def _resize_workers(self, max_retries=10):
|
||||
self._reset()
|
||||
|
||||
time.sleep(1)
|
||||
for i in range(max_retries):
|
||||
new_remote_workers = self._check_potential_remote_workers_size()
|
||||
if new_remote_workers:
|
||||
self._last_resize = time.time()
|
||||
self._start_workers(int(new_remote_workers) + 1)
|
||||
self.load_state_dict(self.state_dict())
|
||||
return
|
||||
else:
|
||||
delay = 2**i
|
||||
logger.warning(
|
||||
"No new workers found. Retrying in %d sec." % delay)
|
||||
time.sleep(delay)
|
||||
raise RuntimeError("Exceeded max_retries for relaunching workers.")
|
||||
|
||||
def _should_resize(self):
|
||||
"""Returns True if past cooldown and exists resources to scale up."""
|
||||
worker_gap = self.max_replicas - 1 - len(self.remote_workers)
|
||||
past_cooldown = (time.time() - self._last_resize) > RESIZE_COOLDOWN_S
|
||||
if past_cooldown and worker_gap:
|
||||
# Assume 1 resource is already reserved for local worker.
|
||||
potential_remote_size = self._check_potential_remote_workers_size()
|
||||
return potential_remote_size > 0
|
||||
return False
|
||||
"""
|
||||
self.worker_group.shutdown(force=force)
|
||||
self.worker_group = DeactivatedWorkerGroup()
|
||||
|
||||
@classmethod
|
||||
def as_trainable(cls, *args, **kwargs):
|
||||
@@ -698,16 +597,26 @@ class TorchTrainer:
|
||||
def default_resource_request(cls, config):
|
||||
num_workers = config.get("num_workers",
|
||||
kwargs.get("num_workers", 1))
|
||||
num_cpus = config.get("num_cpus_per_worker",
|
||||
kwargs.get("num_cpus_per_worker", 1))
|
||||
num_cpus_per_worker = config.get(
|
||||
"num_cpus_per_worker", kwargs.get("num_cpus_per_worker",
|
||||
1))
|
||||
use_gpu = config.get("use_gpu", kwargs.get("use_gpu"))
|
||||
use_local = config.get("use_local",
|
||||
kwargs.get("use_local", False))
|
||||
|
||||
remote_worker_count = num_workers - 1
|
||||
if use_local:
|
||||
remote_worker_count = num_workers - 1
|
||||
local_cpus = 1
|
||||
local_gpus = int(use_gpu)
|
||||
else:
|
||||
remote_worker_count = num_workers
|
||||
local_cpus = 0
|
||||
local_gpus = 0
|
||||
|
||||
return Resources(
|
||||
cpu=num_cpus,
|
||||
gpu=int(use_gpu),
|
||||
extra_cpu=int(remote_worker_count),
|
||||
cpu=int(local_cpus * num_cpus_per_worker),
|
||||
gpu=int(local_gpus),
|
||||
extra_cpu=int(remote_worker_count * num_cpus_per_worker),
|
||||
extra_gpu=int(int(use_gpu) * remote_worker_count))
|
||||
|
||||
def _create_trainer(self, tune_config):
|
||||
|
||||
Reference in New Issue
Block a user