mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
Replaces ddp .spawn with subprocess (#2029)
* replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * replace ddp spawn with subprocess * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix * hot fix
This commit is contained in:
@@ -957,7 +957,7 @@ class LightningModule(ABC, DeviceDtypeModuleMixin, GradInformation, ModelIO, Mod
|
||||
f"is not equal to the computed world size ({world_size}). Ignored.")
|
||||
|
||||
torch_backend = "nccl" if self.trainer.on_gpu else "gloo"
|
||||
log.info(f"initializing proc_rank {proc_rank} world {world_size}")
|
||||
log.info(f"initializing ddp: LOCAL_RANK: {proc_rank}/{world_size - 1} WORLD_SIZE:{world_size}")
|
||||
torch_distrib.init_process_group(torch_backend, rank=proc_rank, world_size=world_size)
|
||||
|
||||
def configure_apex(
|
||||
|
||||
@@ -117,6 +117,11 @@ import os
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Union
|
||||
import subprocess
|
||||
import sys
|
||||
from time import sleep
|
||||
import numpy as np
|
||||
from os.path import abspath
|
||||
|
||||
import torch
|
||||
from pytorch_lightning import _logger as log
|
||||
@@ -311,7 +316,7 @@ class TrainerDDPMixin(ABC):
|
||||
os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
|
||||
|
||||
# when slurm is managing the task it sets the visible devices
|
||||
if not is_slurm_managing_tasks:
|
||||
if not is_slurm_managing_tasks and 'CUDA_VISIBLE_DEVICES' not in os.environ:
|
||||
if isinstance(data_parallel_device_ids, int):
|
||||
id_str = ','.join(str(x) for x in list(range(data_parallel_device_ids)))
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = id_str
|
||||
@@ -322,7 +327,74 @@ class TrainerDDPMixin(ABC):
|
||||
# don't make this debug... this is good UX
|
||||
log.info(f'CUDA_VISIBLE_DEVICES: [{os.environ["CUDA_VISIBLE_DEVICES"]}]')
|
||||
|
||||
def ddp_train(self, process_idx, model):
|
||||
def __set_random_port(self):
|
||||
"""
|
||||
When running DDP NOT managed by SLURM, the ports might collide
|
||||
:return:
|
||||
"""
|
||||
try:
|
||||
default_port = os.environ['MASTER_PORT']
|
||||
except Exception:
|
||||
import random
|
||||
default_port = random.randint(10000, 19000)
|
||||
os.environ['MASTER_PORT'] = str(default_port)
|
||||
|
||||
def spawn_ddp_children(self, model):
|
||||
self.__set_random_port()
|
||||
port = os.environ['MASTER_PORT']
|
||||
|
||||
master_address = '127.0.0.1' if 'MASTER_ADDR' not in os.environ else os.environ['MASTER_ADDR']
|
||||
os.environ['MASTER_PORT'] = f'{port}'
|
||||
os.environ['MASTER_ADDR'] = f'{master_address}'
|
||||
|
||||
# allow the user to pass the node rank
|
||||
node_rank = '0'
|
||||
if 'NODE_RANK' in os.environ:
|
||||
node_rank = os.environ['NODE_RANK']
|
||||
if 'GROUP_RANK' in os.environ:
|
||||
node_rank = os.environ['GROUP_RANK']
|
||||
|
||||
os.environ['NODE_RANK'] = node_rank
|
||||
os.environ['LOCAL_RANK'] = '0'
|
||||
|
||||
# pull out the commands used to run the script and resolve the abs file path
|
||||
command = sys.argv
|
||||
full_path = abspath(command[0])
|
||||
command[0] = full_path
|
||||
command = ['python'] + command
|
||||
|
||||
# since this script sets the visible devices we replace the gpus flag with a number
|
||||
num_gpus = os.environ['CUDA_VISIBLE_DEVICES'].split(',').__len__()
|
||||
|
||||
# if script called without a flag, pass in a flag anyhow
|
||||
if '--gpus' not in command:
|
||||
arg_gpus = len(self.gpus) if isinstance(self.gpus, list) else self.gpus
|
||||
command += ['--gpus', arg_gpus]
|
||||
|
||||
gpu_flag_idx = command.index('--gpus')
|
||||
command[gpu_flag_idx + 1] = f'{num_gpus}'
|
||||
|
||||
os.environ['WORLD_SIZE'] = f'{num_gpus * self.num_nodes}'
|
||||
|
||||
self.interactive_ddp_procs = []
|
||||
for local_rank in range(1, self.num_processes):
|
||||
env_copy = os.environ.copy()
|
||||
env_copy['LOCAL_RANK'] = f'{local_rank}'
|
||||
|
||||
# import pdb; pdb.set_trace()
|
||||
# start process
|
||||
proc = subprocess.Popen(command, env=env_copy)
|
||||
self.interactive_ddp_procs.append(proc)
|
||||
|
||||
# starting all processes at once can cause issues
|
||||
# with dataloaders delay between 1-10 seconds
|
||||
delay = np.random.uniform(1, 5, 1)[0]
|
||||
sleep(delay)
|
||||
|
||||
local_rank = 0
|
||||
self.ddp_train(local_rank, model, is_master=True)
|
||||
|
||||
def ddp_train(self, process_idx, model, is_master=False):
|
||||
"""
|
||||
Entry point into a DP thread
|
||||
:param gpu_idx:
|
||||
@@ -359,7 +431,14 @@ class TrainerDDPMixin(ABC):
|
||||
# MODEL
|
||||
# copy model to each gpu
|
||||
if self.on_gpu:
|
||||
self.root_gpu = process_idx
|
||||
gpu_idx = process_idx
|
||||
if is_master:
|
||||
# source of truth is cuda for gpu idx
|
||||
gpus = os.environ['CUDA_VISIBLE_DEVICES'].split(',')
|
||||
local_rank = int(os.environ['LOCAL_RANK'])
|
||||
gpu_idx = int(gpus[local_rank])
|
||||
|
||||
self.root_gpu = gpu_idx
|
||||
torch.cuda.set_device(self.root_gpu)
|
||||
model.cuda(self.root_gpu)
|
||||
|
||||
@@ -388,9 +467,6 @@ class TrainerDDPMixin(ABC):
|
||||
# continue training routine
|
||||
self.run_pretrain_routine(model)
|
||||
|
||||
# when ddp ends, we save the model
|
||||
self.save_spawn_weights(model)
|
||||
|
||||
def save_spawn_weights(self, model):
|
||||
"""
|
||||
Dump a temporary checkpoint after ddp ends to get weights out of the process
|
||||
|
||||
@@ -685,8 +685,18 @@ def sanitize_gpu_ids(gpus):
|
||||
:return: unmodified gpus variable
|
||||
"""
|
||||
all_available_gpus = get_all_available_gpus()
|
||||
misconfig = False
|
||||
for gpu in gpus:
|
||||
if gpu not in all_available_gpus:
|
||||
misconfig = True
|
||||
|
||||
if misconfig:
|
||||
# sometimes auto ddp might have different flags
|
||||
# but this is not what the user intended
|
||||
# correct for the user
|
||||
if len(gpus) == len(all_available_gpus):
|
||||
gpus = all_available_gpus
|
||||
else:
|
||||
raise MisconfigurationException(f"""
|
||||
You requested GPUs: {gpus}
|
||||
But your machine only has: {all_available_gpus}
|
||||
|
||||
@@ -35,7 +35,6 @@ from pytorch_lightning.trainer.lr_finder import TrainerLRFinderMixin
|
||||
from pytorch_lightning.utilities.exceptions import MisconfigurationException
|
||||
from pytorch_lightning.utilities import rank_zero_warn, parsing
|
||||
|
||||
|
||||
try:
|
||||
from apex import amp
|
||||
except ImportError:
|
||||
@@ -119,7 +118,7 @@ class Trainer(
|
||||
distributed_backend: Optional[str] = None,
|
||||
precision: int = 32,
|
||||
print_nan_grads: bool = False, # backward compatible, todo: remove in v0.9.0
|
||||
weights_summary: Optional[str] = 'full',
|
||||
weights_summary: Optional[str] = 'top',
|
||||
weights_save_path: Optional[str] = None,
|
||||
num_sanity_val_steps: int = 2,
|
||||
truncated_bptt_steps: Optional[int] = None,
|
||||
@@ -494,6 +493,7 @@ class Trainer(
|
||||
# init flags for SLURM+ddp to work
|
||||
self.proc_rank = 0
|
||||
self.world_size = 1
|
||||
self.interactive_ddp_procs = []
|
||||
self.configure_slurm_ddp(self.num_nodes)
|
||||
self.node_rank = self.determine_ddp_node_rank()
|
||||
|
||||
@@ -871,16 +871,12 @@ class Trainer(
|
||||
task = int(os.environ['LOCAL_RANK'])
|
||||
self.ddp_train(task, model)
|
||||
|
||||
else:
|
||||
self.__set_random_port()
|
||||
# track for predict
|
||||
elif self.distributed_backend == 'cpu_ddp':
|
||||
self.model = model
|
||||
# train
|
||||
mp.spawn(self.ddp_train, nprocs=self.num_processes, args=(model,))
|
||||
# load weights if not interrupted
|
||||
if self.on_colab_kaggle:
|
||||
self.load_spawn_weights(model)
|
||||
self.model = model
|
||||
|
||||
elif self.distributed_backend == 'ddp':
|
||||
self.spawn_ddp_children(model)
|
||||
|
||||
# 1 gpu or dp option triggers training using DP module
|
||||
# easier to avoid NCCL issues
|
||||
@@ -928,18 +924,6 @@ class Trainer(
|
||||
# used for testing or when we need to know that training succeeded
|
||||
return 1
|
||||
|
||||
def __set_random_port(self):
|
||||
"""
|
||||
When running DDP NOT managed by SLURM, the ports might collide
|
||||
:return:
|
||||
"""
|
||||
try:
|
||||
default_port = os.environ['MASTER_PORT']
|
||||
except Exception:
|
||||
import random
|
||||
default_port = random.randint(10000, 19000)
|
||||
os.environ['MASTER_PORT'] = str(default_port)
|
||||
|
||||
def __attach_dataloaders(self, model, train_dataloader=None, val_dataloaders=None, test_dataloaders=None):
|
||||
# when dataloader is passed via fit, patch the train_dataloader
|
||||
# functions to overwrite with these implementations
|
||||
@@ -1046,7 +1030,10 @@ class Trainer(
|
||||
|
||||
# clear cache before training
|
||||
if self.on_gpu:
|
||||
torch.cuda.empty_cache()
|
||||
# use context because of:
|
||||
# https://discuss.pytorch.org/t/out-of-memory-when-i-use-torch-cuda-empty-cache/57898
|
||||
with torch.cuda.device(f'cuda:{self.root_gpu}'):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# CORE TRAINING LOOP
|
||||
self.train()
|
||||
@@ -1096,7 +1083,10 @@ class Trainer(
|
||||
if model is not None:
|
||||
self.model = model
|
||||
self.fit(model)
|
||||
elif self.use_ddp or self.use_tpu: # pragma: no-cover
|
||||
|
||||
# on tpu, .spawn means we don't have a trained model
|
||||
# TODO: remove TPU spawn
|
||||
elif self.use_tpu: # pragma: no-cover
|
||||
# attempt to load weights from a spawn
|
||||
path = os.path.join(self.default_root_dir, '__temp_weight_ddp_end.ckpt')
|
||||
test_model = self.model
|
||||
|
||||
@@ -158,6 +158,7 @@ from pytorch_lightning.loggers import LightningLoggerBase
|
||||
from pytorch_lightning.trainer.supporters import TensorRunningAccum
|
||||
from pytorch_lightning.utilities import rank_zero_warn
|
||||
from pytorch_lightning.utilities.exceptions import MisconfigurationException
|
||||
import subprocess
|
||||
|
||||
try:
|
||||
from apex import amp
|
||||
@@ -305,13 +306,13 @@ class TrainerTrainLoopMixin(ABC):
|
||||
|
||||
def train(self):
|
||||
# add signal handlers for process kills
|
||||
def _signal_kill_handler(*args):
|
||||
return TrainerTrainLoopMixin.run_training_teardown(self)
|
||||
|
||||
orig_signal_handlers = {}
|
||||
for sig_name in SIGNAL_TERMINATE:
|
||||
orig_signal_handlers[sig_name] = signal.signal(getattr(signal, sig_name),
|
||||
_signal_kill_handler)
|
||||
# def _signal_kill_handler(*args):
|
||||
# return TrainerTrainLoopMixin.run_training_teardown(self)
|
||||
#
|
||||
# orig_signal_handlers = {}
|
||||
# for sig_name in SIGNAL_TERMINATE:
|
||||
# orig_signal_handlers[sig_name] = signal.signal(getattr(signal, sig_name),
|
||||
# _signal_kill_handler)
|
||||
|
||||
# get model
|
||||
model = self.get_model()
|
||||
@@ -384,15 +385,17 @@ class TrainerTrainLoopMixin(ABC):
|
||||
|
||||
self.run_training_teardown()
|
||||
|
||||
# reset signal handlers
|
||||
for sig_name in SIGNAL_TERMINATE:
|
||||
signal.signal(getattr(signal, sig_name), orig_signal_handlers[sig_name])
|
||||
|
||||
except KeyboardInterrupt:
|
||||
if self.proc_rank == 0:
|
||||
log.info('Detected KeyboardInterrupt, attempting graceful shutdown...')
|
||||
self.interrupted = True
|
||||
self.run_training_teardown()
|
||||
rank_zero_warn('Detected KeyboardInterrupt, attempting graceful shutdown...')
|
||||
|
||||
# user could press ctrl+c many times... only shutdown once
|
||||
if not self.interrupted:
|
||||
self.interrupted = True
|
||||
|
||||
for proc in self.interactive_ddp_procs:
|
||||
subprocess.Popen.kill(proc)
|
||||
|
||||
self.run_training_teardown()
|
||||
|
||||
def run_training_epoch(self):
|
||||
|
||||
@@ -678,7 +681,7 @@ class TrainerTrainLoopMixin(ABC):
|
||||
opt_idx = np.argmax(optimizer_freq_cumsum > current_place_in_loop)
|
||||
return [(opt_idx, self.optimizers[opt_idx])]
|
||||
|
||||
@atexit.register
|
||||
# @atexit.register
|
||||
def run_training_teardown(self):
|
||||
if hasattr(self, '_teardown_already_run') and self._teardown_already_run:
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user