use torchrun
This commit is contained in:
26
train_ms.py
26
train_ms.py
@@ -46,20 +46,13 @@ torch.backends.cuda.enable_mem_efficient_sdp(True) # Not avaliable if torch ver
|
||||
torch.backends.cuda.enable_math_sdp(True)
|
||||
global_step = 0
|
||||
|
||||
|
||||
def main():
|
||||
"""Assume Single Node Multi GPUs Training Only"""
|
||||
assert torch.cuda.is_available(), "CPU training is not allowed."
|
||||
|
||||
n_gpus = torch.cuda.device_count()
|
||||
os.environ['MASTER_ADDR'] = 'localhost'
|
||||
os.environ['MASTER_PORT'] = '65280'
|
||||
|
||||
def run():
|
||||
dist.init_process_group(backend="nccl", init_method="env://") # Use torchrun instead of mp.spawn
|
||||
rank = dist.get_rank()
|
||||
n_gpus = dist.get_world_size()
|
||||
hps = utils.get_hparams()
|
||||
mp.spawn(run, nprocs=n_gpus, args=(n_gpus, hps,))
|
||||
|
||||
|
||||
def run(rank, n_gpus, hps):
|
||||
torch.manual_seed(hps.train.seed)
|
||||
torch.cuda.set_device(rank)
|
||||
global global_step
|
||||
if rank == 0:
|
||||
logger = utils.get_logger(hps.model_dir)
|
||||
@@ -67,11 +60,6 @@ def run(rank, n_gpus, hps):
|
||||
utils.check_git_hash(hps.model_dir)
|
||||
writer = SummaryWriter(log_dir=hps.model_dir)
|
||||
writer_eval = SummaryWriter(log_dir=os.path.join(hps.model_dir, "eval"))
|
||||
|
||||
dist.init_process_group(backend='nccl', init_method='env://', world_size=n_gpus, rank=rank)
|
||||
torch.manual_seed(hps.train.seed)
|
||||
torch.cuda.set_device(rank)
|
||||
|
||||
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
|
||||
train_sampler = DistributedBucketSampler(
|
||||
train_dataset,
|
||||
@@ -391,4 +379,4 @@ def evaluate(hps, generator, eval_loader, writer_eval):
|
||||
generator.train()
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
run()
|
||||
|
||||
Reference in New Issue
Block a user