diff --git a/train_ms.py b/train_ms.py index 66f235d..2e673b1 100644 --- a/train_ms.py +++ b/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()