Format code

This commit is contained in:
github-actions[bot]
2023-09-07 12:46:17 +00:00
parent eb83b8b58b
commit c62c25cb4e

View File

@@ -49,7 +49,8 @@ global_step = 0
def run(): def run():
dist.init_process_group( dist.init_process_group(
backend="gloo", init_method="env://" # Due to some training problem,we proposed to use gloo instead of nccl. backend="gloo",
init_method="env://", # Due to some training problem,we proposed to use gloo instead of nccl.
) # Use torchrun instead of mp.spawn ) # Use torchrun instead of mp.spawn
rank = dist.get_rank() rank = dist.get_rank()
n_gpus = dist.get_world_size() n_gpus = dist.get_world_size()
@@ -72,7 +73,7 @@ def run():
rank=rank, rank=rank,
shuffle=True, shuffle=True,
persistent_workers=True, persistent_workers=True,
prefetch_factor=4 prefetch_factor=4,
) )
collate_fn = TextAudioSpeakerCollate() collate_fn = TextAudioSpeakerCollate()
train_loader = DataLoader( train_loader = DataLoader(
@@ -82,7 +83,6 @@ def run():
pin_memory=True, pin_memory=True,
collate_fn=collate_fn, collate_fn=collate_fn,
batch_sampler=train_sampler, batch_sampler=train_sampler,
) # DataLoader config could be adjusted. ) # DataLoader config could be adjusted.
if rank == 0: if rank == 0:
eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data) eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data)