use torchrun

This commit is contained in:
Stardust·减
2023-09-05 17:26:39 +08:00
committed by GitHub
parent 042fb96fd9
commit e86c8a8991

View File

@@ -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()