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