mirror of
https://github.com/wassname/ray.git
synced 2026-09-13 13:02:57 +08:00
[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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user