mirror of
https://github.com/wassname/ray.git
synced 2026-08-12 12:20:11 +08:00
[RaySGD] Simplify Builder Process (#10321)
Co-authored-by: Richard Liaw <rliaw@berkeley.edu>
This commit is contained in:
co-authored by
Richard Liaw
parent
69c1a9dd08
commit
415be78cc0
@@ -13,6 +13,7 @@ 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
|
||||
@@ -49,80 +50,44 @@ class TorchTrainer:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def model_creator(config):
|
||||
return nn.Linear(1, 1)
|
||||
class MyTrainingOperator(TrainingOperator):
|
||||
|
||||
def setup(self, config):
|
||||
model = nn.Linear(1, 1)
|
||||
optimizer = torch.optim.SGD(
|
||||
model.parameters(), lr=config.get("lr", 1e-4))
|
||||
loss = torch.nn.MSELoss()
|
||||
|
||||
def optimizer_creator(model, config):
|
||||
return torch.optim.SGD(
|
||||
model.parameters(), lr=config.get("lr", 1e-4))
|
||||
batch_size = config["batch_size"]
|
||||
train_data, val_data = LinearDataset(2, 5), LinearDataset(2, 5)
|
||||
train_loader = DataLoader(train_data, batch_size=batch_size)
|
||||
val_loader = DataLoader(val_data, batch_size=batch_size)
|
||||
|
||||
self.model, self.optimizer = self.register(
|
||||
models=model,
|
||||
optimizers=optimizer,
|
||||
criterion=loss)
|
||||
|
||||
def data_creator(config):
|
||||
batch_size = config["batch_size"]
|
||||
train_data, val_data = LinearDataset(2, 5), LinearDataset(2, 5)
|
||||
train_loader = DataLoader(train_data, batch_size=batch_size)
|
||||
val_loader = DataLoader(val_data, batch_size=batch_size)
|
||||
return train_loader, val_loader
|
||||
self.register_data(
|
||||
train_loader=train_loader,
|
||||
validation_loader=val_loader)
|
||||
|
||||
|
||||
trainer = TorchTrainer(
|
||||
model_creator=model_creator,
|
||||
data_creator=data_creator,
|
||||
optimizer_creator=optimizer_creator,
|
||||
loss_creator=nn.MSELoss,
|
||||
training_operator_cls=MyTrainingOperator,
|
||||
config={"batch_size": 32},
|
||||
use_gpu=True
|
||||
)
|
||||
for i in range(4):
|
||||
trainer.train()
|
||||
|
||||
The creator functions will execute before distributed coordination and
|
||||
training is setup. This is so that creator functions that download
|
||||
large datasets will not trigger any timeouts.
|
||||
|
||||
The order of operations for creator functions are:
|
||||
|
||||
``data_creator`` -> ``model_creator`` -> ``optimizer_creator`` ->
|
||||
``scheduler_creator`` -> ``loss_creator``.
|
||||
|
||||
Args:
|
||||
model_creator (dict -> Model(s)): Constructor function that takes in
|
||||
config and returns the model(s) to be optimized. These must be
|
||||
``torch.nn.Module`` objects. If multiple models are returned,
|
||||
a ``training_operator_cls`` must be specified. You do not need to
|
||||
handle GPU/devices in this function; RaySGD will do that under
|
||||
the hood.
|
||||
data_creator (dict -> Iterable(s)): Constructor function
|
||||
that takes in the passed config and returns one or
|
||||
two Iterable objects. Note that even though two Iterable objects
|
||||
can be returned, only one will be used for training, and the
|
||||
other will be used for validation. If not provided, you must
|
||||
provide a custom TrainingOperator.
|
||||
optimizer_creator ((models, dict) -> optimizers): Constructor
|
||||
function that takes in the return values from
|
||||
``model_creator`` and the passed config and returns One or
|
||||
more Torch optimizer objects. You do not need to handle
|
||||
GPU/devices in this function; ``RaySGD`` will do that for you.
|
||||
loss_creator (torch.nn.*Loss class | dict -> loss): A constructor
|
||||
function for the training loss. This can be either a function that
|
||||
takes in the provided config for customization or a subclass
|
||||
of ``torch.nn.modules.loss._Loss``, which is most Pytorch
|
||||
loss classes. For example, ``loss_creator=torch.nn.BCELoss``.
|
||||
If not provided, you must provide a custom TrainingOperator.
|
||||
scheduler_creator ((optimizers, dict) -> scheduler):
|
||||
A constructor function for the torch scheduler. This is
|
||||
a function that takes in the generated optimizers (from
|
||||
``optimizer_creator``) provided config for customization.
|
||||
Be sure to set ``scheduler_step_freq`` to increment the
|
||||
scheduler correctly.
|
||||
training_operator_cls (type): Custom training operator class
|
||||
that subclasses the TrainingOperator class. This class
|
||||
will be copied onto all remote workers and used to specify
|
||||
custom training and validation operations. Defaults to
|
||||
TrainingOperator.
|
||||
training components and custom training and validation operations.
|
||||
config (dict): Custom configuration value to be passed to
|
||||
all creator and operator constructors.
|
||||
all operator constructors.
|
||||
num_workers (int): the number of workers used in distributed
|
||||
training. If 1, the worker will not be wrapped with
|
||||
DistributedDataParallel.
|
||||
@@ -134,10 +99,6 @@ class TorchTrainer:
|
||||
support "nccl", "gloo", and "auto". If "auto", RaySGD will
|
||||
automatically use "nccl" if `use_gpu` is True, and "gloo"
|
||||
otherwise.
|
||||
serialize_data_creation (bool): A filelock will be used
|
||||
to ensure no race conditions in data downloading among
|
||||
different workers on the same node (using the local file system).
|
||||
Defaults to True.
|
||||
wrap_ddp (bool): Whether to automatically wrap DistributedDataParallel
|
||||
over each model. If False, you are expected to call it yourself.
|
||||
timeout_s (float): Seconds before the torch process group
|
||||
@@ -171,12 +132,7 @@ class TorchTrainer:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model_creator,
|
||||
data_creator,
|
||||
optimizer_creator,
|
||||
loss_creator=None,
|
||||
scheduler_creator=None,
|
||||
training_operator_cls=None,
|
||||
training_operator_cls,
|
||||
initialization_hook=None,
|
||||
config=None,
|
||||
num_workers=1,
|
||||
@@ -185,16 +141,33 @@ class TorchTrainer:
|
||||
backend="auto",
|
||||
wrap_ddp=True,
|
||||
timeout_s=NCCL_TIMEOUT_S,
|
||||
serialize_data_creation=True,
|
||||
use_fp16=False,
|
||||
use_tqdm=False,
|
||||
apex_args=None,
|
||||
add_dist_sampler=True,
|
||||
scheduler_step_freq=None,
|
||||
# Deprecated Args.
|
||||
num_replicas=None,
|
||||
batch_size=None,
|
||||
model_creator=None,
|
||||
data_creator=None,
|
||||
optimizer_creator=None,
|
||||
scheduler_creator=None,
|
||||
loss_creator=None,
|
||||
serialize_data_creation=None,
|
||||
data_loader_args=None,
|
||||
):
|
||||
if (model_creator or data_creator or optimizer_creator
|
||||
or scheduler_creator or loss_creator):
|
||||
raise DeprecationWarning(
|
||||
"Creator functions are deprecated. You should create a "
|
||||
"custom TrainingOperator, override setup, and register all "
|
||||
"training state there. See TrainingOperator for more info. "
|
||||
"If you would still like to use creator functions, you can "
|
||||
"do CustomOperator = TrainingOperator.from_creators("
|
||||
"model_creator, ...) and pass in CustomOperator into "
|
||||
"TorchTrainer.")
|
||||
|
||||
if num_workers > 1 and not dist.is_available():
|
||||
raise ValueError(
|
||||
("Distributed PyTorch is not supported on macOS. "
|
||||
@@ -202,10 +175,6 @@ class TorchTrainer:
|
||||
"For more information, see "
|
||||
"https://github.com/pytorch/examples/issues/467."))
|
||||
|
||||
if not (callable(model_creator) and callable(optimizer_creator)):
|
||||
raise ValueError(
|
||||
"Must provide a callable model_creator and optimizer_creator.")
|
||||
|
||||
if num_replicas is not None:
|
||||
raise DeprecationWarning(
|
||||
"num_replicas is deprecated. Use num_workers instead.")
|
||||
@@ -217,24 +186,23 @@ class TorchTrainer:
|
||||
"config={ray.util.sgd.utils.BATCH_SIZE: N} to specify a "
|
||||
"batch size to be used across all workers.")
|
||||
|
||||
if serialize_data_creation is True:
|
||||
if log_once("serialize_data_creation"):
|
||||
logging.warning(
|
||||
"serialize_data_creation is deprecated and will be "
|
||||
"ignored. If you require serialized data loading you "
|
||||
"should implement this in TrainingOperator.setup. "
|
||||
"You may find FileLock useful here.")
|
||||
|
||||
if data_loader_args:
|
||||
raise ValueError(
|
||||
raise DeprecationWarning(
|
||||
"data_loader_args is deprecated. You can return a "
|
||||
"torch.utils.data.DataLoader in data_creator. Ray will "
|
||||
"automatically set a DistributedSampler if a DataLoader is "
|
||||
"returned and num_workers > 1.")
|
||||
|
||||
self.model_creator = model_creator
|
||||
self.optimizer_creator = optimizer_creator
|
||||
self.loss_creator = loss_creator
|
||||
self.data_creator = data_creator
|
||||
self.scheduler_creator = scheduler_creator
|
||||
self.training_operator_cls = training_operator_cls
|
||||
|
||||
if not training_operator_cls and not loss_creator:
|
||||
raise ValueError("If a loss_creator is not provided, you must "
|
||||
"provide a custom training operator.")
|
||||
|
||||
self.initialization_hook = initialization_hook
|
||||
self.config = {} if config is None else config
|
||||
if use_gpu == "auto":
|
||||
@@ -269,7 +237,7 @@ class TorchTrainer:
|
||||
self.local_worker = DeactivatedRunner()
|
||||
self.remote_workers = []
|
||||
|
||||
if scheduler_creator:
|
||||
if scheduler_step_freq:
|
||||
_validate_scheduler_step_freq(scheduler_step_freq)
|
||||
|
||||
self.scheduler_step_freq = scheduler_step_freq
|
||||
@@ -309,11 +277,6 @@ class TorchTrainer:
|
||||
worker_config[BATCH_SIZE] = batch_size_per_worker
|
||||
|
||||
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,
|
||||
training_operator_cls=self.training_operator_cls,
|
||||
config=worker_config,
|
||||
serialize_data_creation=self.serialize_data_creation,
|
||||
@@ -328,7 +291,7 @@ class TorchTrainer:
|
||||
self.local_worker = TorchRunner(**params)
|
||||
if self.initialization_hook:
|
||||
self.apply_all_workers(self.initialization_hook)
|
||||
self.local_worker.setup()
|
||||
self.local_worker.setup_operator()
|
||||
else:
|
||||
params.update(
|
||||
backend=self.backend,
|
||||
@@ -355,15 +318,6 @@ class TorchTrainer:
|
||||
# Compute URL for initializing distributed PyTorch
|
||||
address = setup_address()
|
||||
|
||||
# Runs the creator functions.
|
||||
remote_component_setup = [
|
||||
worker.setup_components.remote()
|
||||
for i, worker in enumerate(self.remote_workers)
|
||||
]
|
||||
self.local_worker.setup_components()
|
||||
# Get setup tasks in order to throw errors on failure
|
||||
ray.get(remote_component_setup)
|
||||
|
||||
# Setup the process group among all workers.
|
||||
remote_pgroup_setups = [
|
||||
worker.setup_process_group.remote(address, i + 1, num_workers,
|
||||
@@ -377,10 +331,10 @@ class TorchTrainer:
|
||||
|
||||
# Runs code that requires all creator functions to have run.
|
||||
remote_operator_setups = [
|
||||
worker.setup_ddp_and_operator.remote()
|
||||
worker.setup_operator.remote()
|
||||
for worker in self.remote_workers
|
||||
]
|
||||
self.local_worker.setup_ddp_and_operator()
|
||||
self.local_worker.setup_operator()
|
||||
# Get setup tasks in order to throw errors on failure
|
||||
ray.get(remote_operator_setups)
|
||||
|
||||
@@ -421,10 +375,10 @@ class TorchTrainer:
|
||||
|
||||
Returns:
|
||||
(dict | list) A dictionary of metrics for training.
|
||||
You can provide custom metrics by passing in a custom
|
||||
``training_operator_cls``. If ``reduce_results=False``,
|
||||
this will return a list of metric dictionaries whose
|
||||
length will be equal to ``num_workers``.
|
||||
You can provide custom metrics by implementing a custom
|
||||
training loop. If ``reduce_results=False``, this will return a
|
||||
list of metric dictionaries whose length will be equal to
|
||||
``num_workers``.
|
||||
"""
|
||||
assert max_retries >= 0, "`max_retries` must be non-negative."
|
||||
assert isinstance(dataset, Dataset) is not None \
|
||||
@@ -577,12 +531,12 @@ class TorchTrainer:
|
||||
return worker_stats
|
||||
|
||||
def update_scheduler(self, metric):
|
||||
"""Calls ``scheduler.step(metric)`` on all schedulers.
|
||||
"""Calls ``scheduler.step(metric)`` on all registered schedulers.
|
||||
|
||||
This is useful for lr_schedulers such as ``ReduceLROnPlateau``.
|
||||
"""
|
||||
self.apply_all_operators(
|
||||
lambda op: [sched.step(metric) for sched in op.schedulers])
|
||||
lambda op: [sched.step(metric) for sched in op._schedulers])
|
||||
|
||||
def get_model(self):
|
||||
"""Returns the learned model(s)."""
|
||||
@@ -729,10 +683,7 @@ class TorchTrainer:
|
||||
.. code-block:: python
|
||||
|
||||
TorchTrainable = TorchTrainer.as_trainable(
|
||||
model_creator=ResNet18,
|
||||
data_creator=cifar_creator,
|
||||
optimizer_creator=optimizer_creator,
|
||||
loss_creator=nn.CrossEntropyLoss,
|
||||
training_operator_cls=MyTrainingOperator,
|
||||
num_gpus=2
|
||||
)
|
||||
analysis = tune.run(
|
||||
@@ -781,10 +732,7 @@ class BaseTorchTrainable(Trainable):
|
||||
.. code-block:: python
|
||||
|
||||
TorchTrainable = TorchTrainer.as_trainable(
|
||||
model_creator=ResNet18,
|
||||
data_creator=cifar_creator,
|
||||
optimizer_creator=optimizer_creator,
|
||||
loss_creator=nn.CrossEntropyLoss,
|
||||
training_operator_cls=MyTrainingOperator,
|
||||
num_gpus=2
|
||||
)
|
||||
# TorchTrainable is subclass of BaseTorchTrainable.
|
||||
|
||||
Reference in New Issue
Block a user