[Ray SGD] LightningModule integration + MNIST Example (#11042)

Co-authored-by: Richard Liaw <rliaw@berkeley.edu>
This commit is contained in:
Amog Kamsetty
2020-10-01 20:07:32 -07:00
committed by GitHub
co-authored by Richard Liaw
parent 9dc7b7b11d
commit 874da9a25f
9 changed files with 1120 additions and 118 deletions
+13 -43
View File
@@ -1,6 +1,8 @@
import logging
import io
import itertools
import ray
import torch
from ray.util.sgd.torch.constants import USE_FP16, NUM_STEPS
@@ -50,6 +52,8 @@ class TorchRunner:
self.training_operator = self.training_operator_cls(
self.config,
world_rank=0,
local_rank=0,
is_distributed=False,
use_gpu=self.use_gpu,
use_fp16=self.use_fp16,
use_tqdm=self.use_tqdm,
@@ -69,6 +73,7 @@ class TorchRunner:
info.update({
NUM_STEPS: num_steps,
USE_FP16: self.use_fp16,
"epoch_idx": self.epochs,
})
with self.timers.record("train_epoch"):
if iterator is None:
@@ -94,10 +99,6 @@ class TorchRunner:
def validate(self, num_steps=None, profile=False, info=None):
"""Evaluates the model on the validation data set."""
if self.validation_loader is None:
raise ValueError("No validation dataloader provided. Make sure"
"you pass in a validation_loader to "
"TrainingOperator.register_data.")
info = info or {}
self._toggle_profiling(profile=profile)
validation_loader = self.validation_loader
@@ -204,62 +205,31 @@ class TorchRunner:
"""Getter method. Needed for remote actor calls."""
return self.models
def get_node_ip(self):
return ray.services.get_node_ip_address()
@property
def models(self):
if not hasattr(self.training_operator, "_original_models"):
raise RuntimeError("Training Operator does not have any "
"registered models. Are you calling "
"self.register(...) inside the setup method "
"of your Training Operator?")
return self.training_operator._original_models
return self.training_operator._get_original_models()
@property
def optimizers(self):
if not hasattr(self.training_operator, "_optimizers"):
raise RuntimeError("Training Operator does not have any "
"registered optimizers. Are you calling "
"self.register(...) inside the setup method "
"of your Training Operator?")
return self.training_operator._optimizers
return self.training_operator._get_optimizers()
@property
def schedulers(self):
if not hasattr(self.training_operator, "_schedulers"):
raise RuntimeError("Training Operator does not have any "
"registered schedulers. Are you calling "
"self.register(...) inside the setup method "
"of your Training Operator?")
return self.training_operator._schedulers
return self.training_operator._get_schedulers()
@property
def train_loader(self):
if not hasattr(self.training_operator, "_train_loader"):
logger.warning("Training Operator does not have any "
"registered train loader. If this is "
"unexepected, make sure to call "
"self.register_data(...) inside the setup method "
"of your Training Operator.")
return None
return self.training_operator._train_loader
return self.training_operator._get_train_loader()
@property
def validation_loader(self):
if not hasattr(self.training_operator, "_validation_loader"):
logger.warning("Training Operator does not have any "
"registered validation loader. If this is "
"unexepected, make sure to call "
"self.register_data(...) inside the setup method "
"of your Training Operator.")
return None
return self.training_operator._validation_loader
return self.training_operator._get_validation_loader()
@property
def criterion(self):
if not hasattr(self.training_operator, "_criterion"):
raise RuntimeError("Training Operator does not have any "
"registered criterion. Are you calling "
"self.register(...) inside the setup method "
"of your Training Operator?")
return self.training_operator._criterion
@property