[RaySGD] Convert the head worker to a local model (#7746)

Why are these changes needed?

Running a worker on head (locally, not as a Ray actor) allows for easier handling of stateful stuff like logging and for easier debugging.
This commit is contained in:
Maksim Smolin
2020-03-27 20:19:15 -07:00
committed by GitHub
parent 875309fc48
commit 7b27ce2b23
10 changed files with 341 additions and 301 deletions
+163 -134
View File
@@ -4,22 +4,18 @@ import logging
import numbers
import tempfile
import time
import asyncio
import torch
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.util.sgd.torch.distributed_torch_runner import (
DistributedTorchRunner)
DistributedTorchRunner, LocalDistributedRunner)
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,
BATCH_LOGS_RATE_LIMIT)
from ray.util.sgd.torch.tqdm_handler import TqdmHandler
from ray.util.sgd.torch.constants import VALID_SCHEDULER_STEP
logger = logging.getLogger(__name__)
RESIZE_COOLDOWN_S = 10
@@ -149,7 +145,7 @@ class TorchTrainer:
use_gpu=False,
backend="auto",
use_fp16=False,
tqdm=False,
use_tqdm=False,
apex_args=None,
scheduler_step_freq="batch",
num_replicas=None,
@@ -212,6 +208,7 @@ class TorchTrainer:
self.max_replicas = num_workers
self.use_fp16 = use_fp16
self.use_tqdm = use_tqdm
if apex_args and not isinstance(apex_args, dict):
raise ValueError("apex_args needs to be a dict object.")
@@ -221,10 +218,6 @@ class TorchTrainer:
self._num_failures = 0
self._last_resize = float("-inf")
self.handlers = []
if tqdm:
self.handlers.append(TqdmHandler())
_validate_scheduler_step_freq(scheduler_step_freq)
self.scheduler_step_freq = scheduler_step_freq
@@ -256,68 +249,71 @@ class TorchTrainer:
batch_size_per_worker = self._configure_and_split_batch(num_workers)
if batch_size_per_worker:
worker_config[BATCH_SIZE] = batch_size_per_worker
self.local_worker = None
self.remote_workers = []
if num_workers == 1:
# Generate actor class
Runner = ray.remote(
num_cpus=1, num_gpus=int(self.use_gpu))(TorchRunner)
# Start workers
self.workers = [
Runner.remote(
model_creator=self.model_creator,
data_creator=self.data_creator,
optimizer_creator=self.optimizer_creator,
loss_creator=self.loss_creator,
scheduler_creator=self.scheduler_creator,
training_operator_cls=self.training_operator_cls,
config=worker_config,
use_fp16=self.use_fp16,
apex_args=self.apex_args,
scheduler_step_freq=self.scheduler_step_freq,
)
]
# Start local worker
self.local_worker = TorchRunner(
model_creator=self.model_creator,
data_creator=self.data_creator,
optimizer_creator=self.optimizer_creator,
loss_creator=self.loss_creator,
scheduler_creator=self.scheduler_creator,
training_operator_cls=self.training_operator_cls,
config=worker_config,
use_fp16=self.use_fp16,
use_tqdm=self.use_tqdm,
apex_args=self.apex_args,
scheduler_step_freq=self.scheduler_step_freq)
if self.initialization_hook:
self.apply_all_workers(self.initialization_hook)
# Get setup tasks in order to throw errors on failure
ray.get(self.workers[0].setup.remote())
ray.get(self.workers[0].set_reporters.remote(
[h.create_reporter() for h in self.handlers]))
self.local_worker.setup()
else:
params = dict(
model_creator=self.model_creator,
data_creator=self.data_creator,
optimizer_creator=self.optimizer_creator,
loss_creator=self.loss_creator,
scheduler_creator=self.scheduler_creator,
backend=self.backend,
training_operator_cls=self.training_operator_cls,
config=worker_config,
use_fp16=self.use_fp16,
use_tqdm=self.use_tqdm,
apex_args=self.apex_args,
scheduler_step_freq=self.scheduler_step_freq)
# Start local worker
self.local_worker = LocalDistributedRunner(
num_cpus=1, num_gpus=int(self.use_gpu), **params)
# Generate actor class
Runner = ray.remote(
RemoteRunner = ray.remote(
num_cpus=1, num_gpus=int(self.use_gpu))(DistributedTorchRunner)
# Start workers
self.workers = [
Runner.remote(
model_creator=self.model_creator,
data_creator=self.data_creator,
optimizer_creator=self.optimizer_creator,
loss_creator=self.loss_creator,
scheduler_creator=self.scheduler_creator,
backend=self.backend,
training_operator_cls=self.training_operator_cls,
config=worker_config,
use_fp16=self.use_fp16,
apex_args=self.apex_args,
scheduler_step_freq=self.scheduler_step_freq)
for i in range(num_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)
# Compute URL for initializing distributed PyTorch
ip = ray.get(self.workers[0].get_node_ip.remote())
port = ray.get(self.workers[0].find_free_port.remote())
ip = ray.services.get_node_ip_address()
port = self.local_worker.find_free_port()
address = "tcp://{ip}:{port}".format(ip=ip, port=port)
remote_setups = [
worker.setup.remote(address, i + 1, num_workers)
for i, worker in enumerate(self.remote_workers)
]
self.local_worker.setup(address, 0, num_workers)
# Get setup tasks in order to throw errors on failure
ray.get([
worker.setup.remote(address, i, len(self.workers))
for i, worker in enumerate(self.workers)
])
ray.get([
w.set_reporters.remote(
[h.create_reporter() for h in self.handlers])
for w in self.workers
])
ray.get(remote_setups)
def train(self,
num_steps=None,
@@ -374,9 +370,6 @@ class TorchTrainer:
logger.info("Resize opportunity detected. Attempting to scale up.")
self._resize_workers(checkpoint=checkpoint)
for h in self.handlers:
h.record_train_info(info, num_steps)
success, worker_stats = self._train_epoch(
num_steps=num_steps, profile=profile, info=info)
# Fault handling
@@ -386,14 +379,13 @@ class TorchTrainer:
else:
self._num_failures += 1
self._resize_workers(checkpoint=checkpoint)
logger.info(
"Retrying training step with %d workers." % len(self.workers))
logger.info("Retrying training step with %d workers." %
(len(self.remote_workers) + 1))
success, worker_stats = self._train_epoch(
num_steps=num_steps, profile=profile, info=info)
if not success:
raise RuntimeError("Training run failed.")
worker_stats = ray.get(worker_stats)
if reduce_results:
return self._process_stats(worker_stats)
else:
@@ -413,42 +405,30 @@ class TorchTrainer:
stats[stat_key] = worker_stats[0][stat_key]
return stats
def _train_epoch(self,
num_steps=None,
profile=False,
info=None,
batch_logs_handler=None):
worker_trains = [
w.train_epoch.remote(
num_steps=num_steps, profile=profile, info=info)
for w in self.workers
def _train_epoch(self, num_steps=None, profile=False, info=None):
params = dict(num_steps=num_steps, profile=profile, info=info)
remote_worker_stats = [
w.train_epoch.remote(**params) for w in self.remote_workers
]
if not self.handlers:
success = check_for_failure(worker_trains)
return success, worker_trains
unfinished = worker_trains
try:
while len(unfinished) > 0:
finished, unfinished = ray.wait(
unfinished, timeout=BATCH_LOGS_RATE_LIMIT)
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
# throw errors on agent failure
finished = ray.get(finished)
raise err
futures = [h.update() for h in self.handlers]
loop = asyncio.get_event_loop()
if loop.is_closed():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
loop.run_until_complete(asyncio.wait(futures))
loop.close()
success = check_for_failure(remote_worker_stats)
if success:
return success, [local_worker_stats] + ray.get(remote_worker_stats)
return True, worker_trains
except RayActorError as exc:
logger.exception(str(exc))
return False, worker_trains
return success, None
def apply_all_workers(self, fn):
"""Run a function on all operators on the workers.
@@ -460,7 +440,9 @@ class TorchTrainer:
A list of objects returned by ``fn`` on each worker.
"""
return ray.get([w.apply.remote(fn) for w in self.workers])
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)
def apply_all_operators(self, fn):
"""Run a function on all operators on the workers.
@@ -473,7 +455,11 @@ class TorchTrainer:
A list of objects returned by ``fn`` on each operator.
"""
return ray.get([w.apply_operator.remote(fn) for w in self.workers])
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)
def validate(self, num_steps=None, profile=False, info=None):
"""Evaluates the model on the validation data set.
@@ -491,12 +477,15 @@ class TorchTrainer:
You can provide custom metrics by passing in a custom
``training_operator_cls``.
"""
worker_stats = ray.get([
w.validate.remote(num_steps=num_steps, profile=profile, info=info)
for w in self.workers
])
params = dict(num_steps=num_steps, profile=profile, info=info)
return self._process_stats(worker_stats)
remote_worker_stats = [
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))
def update_scheduler(self, metric):
"""Calls ``scheduler.step(metric)`` on all schedulers.
@@ -509,7 +498,7 @@ class TorchTrainer:
def get_model(self):
"""Returns the learned model(s)."""
models = self.model_creator(self.config)
state = ray.get(self.workers[0].get_state.remote())
state = self.local_worker.get_state()
if len(state["models"]) == 1:
models.load_state_dict(state["models"][0])
else:
@@ -517,6 +506,18 @@ class TorchTrainer:
model.load_state_dict(state_dict)
return models
def state_dict(self):
return self.local_worker.get_state()
def load_state_dict(self, state):
state_id = ray.put(state)
remote_calls = [
worker.set_state.remote(state_id) for worker in self.remote_workers
]
self.local_worker.set_state(state)
ray.get(remote_calls)
def save(self, checkpoint):
"""Saves the model(s) to the provided checkpoint.
@@ -526,8 +527,7 @@ class TorchTrainer:
Returns:
checkpoint (str): Path to target checkpoint file.
"""
state = ray.get(self.workers[0].get_state.remote())
torch.save(state, checkpoint)
torch.save(self.state_dict(), checkpoint)
return checkpoint
def restore(self, checkpoint):
@@ -537,36 +537,67 @@ class TorchTrainer:
checkpoint (str): Path to target checkpoint file.
"""
state = torch.load(checkpoint)
state_id = ray.put(state)
ray.get([worker.set_state.remote(state_id) for worker in self.workers])
self.load_state_dict(state)
def shutdown(self, force=False):
"""Shuts down workers and releases resources."""
if not force:
cleanup = [worker.shutdown.remote() for worker in self.workers]
ray.get(cleanup)
[worker.__ray_terminate__.remote() for worker in self.workers]
else:
for worker in self.workers:
logger.warning("Killing worker {}.".format(worker))
worker.__ray_kill__()
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.")
self.workers = []
for worker in self.remote_workers:
logger.warning("Killing worker {}.".format(worker))
ray.kill(worker)
else:
self.local_worker.shutdown()
for worker in self.remote_workers:
logger.warning("Killing worker {}.".format(worker))
ray.kill(worker)
self.local_worker = None
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))
ray.kill(worker)
self.local_worker = None
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, checkpoint, max_retries=10):
# check available resources
self.shutdown(force=True)
self._reset()
assert checkpoint, "Cannot restore without checkpoint."
time.sleep(1)
for i in range(max_retries):
resources = ray.available_resources()
new_workers = min(resources.get("CPU", 0), self.max_replicas)
if self.use_gpu:
new_workers = min(resources.get("GPU", 0), new_workers)
if new_workers:
new_remote_workers = self._check_potential_remote_workers_size()
if new_remote_workers:
self._last_resize = time.time()
self._start_workers(int(new_workers))
self._start_workers(int(new_remote_workers) + 1)
self.restore(checkpoint)
return
else:
@@ -578,26 +609,24 @@ class TorchTrainer:
def _should_resize(self):
"""Returns True if past cooldown and exists resources to scale up."""
worker_gap = self.max_replicas - len(self.workers)
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:
resources = ray.available_resources()
potential_workers = min(resources.get("CPU", 0), self.max_replicas)
if self.use_gpu:
potential_workers = min(
resources.get("GPU", 0), potential_workers)
return potential_workers > 0
# 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
class TorchTrainable(Trainable):
@classmethod
def default_resource_request(cls, config):
remote_worker_count = config["num_workers"] - 1
return Resources(
cpu=0,
gpu=0,
extra_cpu=config["num_workers"],
extra_gpu=int(config["use_gpu"]) * config["num_workers"])
cpu=1,
gpu=int(config["use_gpu"]),
extra_cpu=int(remote_worker_count),
extra_gpu=int(int(config["use_gpu"]) * remote_worker_count))
def _setup(self, config):
self._trainer = TorchTrainer(**config)