mirror of
https://github.com/wassname/ray.git
synced 2026-09-10 12:38:43 +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:
@@ -8,7 +8,7 @@ import tempfile
|
||||
import torch
|
||||
|
||||
import ray
|
||||
from ray.util.sgd.torch.constants import USE_FP16, SCHEDULER_STEP
|
||||
from ray.util.sgd.torch.constants import USE_FP16, SCHEDULER_STEP, NUM_STEPS
|
||||
from ray.util.sgd.torch.training_operator import TrainingOperator
|
||||
from ray.util.sgd import utils
|
||||
|
||||
@@ -49,6 +49,7 @@ class TorchRunner:
|
||||
training_operator_cls=None,
|
||||
config=None,
|
||||
use_fp16=False,
|
||||
use_tqdm=False,
|
||||
apex_args=None,
|
||||
scheduler_step_freq="batch"):
|
||||
self.model_creator = model_creator
|
||||
@@ -68,6 +69,7 @@ class TorchRunner:
|
||||
self.train_loader = None
|
||||
self.validation_loader = None
|
||||
self.use_fp16 = use_fp16
|
||||
self.use_tqdm = use_tqdm
|
||||
self.apex_args = apex_args or {}
|
||||
if use_fp16 and not amp:
|
||||
raise ImportError(
|
||||
@@ -133,9 +135,6 @@ class TorchRunner:
|
||||
self.models, self.optimizers = amp.initialize(
|
||||
self.models, self.optimizers, **self.apex_args)
|
||||
|
||||
def set_reporters(self, reporters):
|
||||
return self.training_operator.set_reporters(reporters)
|
||||
|
||||
def setup(self):
|
||||
"""Initializes the model."""
|
||||
logger.debug("Creating model")
|
||||
@@ -163,7 +162,8 @@ class TorchRunner:
|
||||
validation_loader=self.validation_loader,
|
||||
world_rank=0,
|
||||
schedulers=self.schedulers,
|
||||
use_fp16=self.use_fp16)
|
||||
use_fp16=self.use_fp16,
|
||||
use_tqdm=self.use_tqdm)
|
||||
|
||||
def get_node_ip(self):
|
||||
"""Returns the IP address of the current node."""
|
||||
@@ -180,6 +180,7 @@ class TorchRunner:
|
||||
self._toggle_profiling(profile=profile)
|
||||
|
||||
info.update({
|
||||
NUM_STEPS: num_steps,
|
||||
USE_FP16: self.use_fp16,
|
||||
SCHEDULER_STEP: self.scheduler_step_freq
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user