From e78d2a6f7e21c574408653af143c60feeeae6e07 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 22 Jun 2024 17:21:58 +0900 Subject: [PATCH] Feat: use DistributedLengthGroupedSampler (experimental, to be checked) --- train_ms.py | 14 +++++++++++--- train_ms_jp_extra.py | 13 +++++++++++-- 2 files changed, 22 insertions(+), 5 deletions(-) diff --git a/train_ms.py b/train_ms.py index 1adee4e..b1cc3a4 100644 --- a/train_ms.py +++ b/train_ms.py @@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm +from transformers.trainer_pt_utils import DistributedLengthGroupedSampler # logging.getLogger("numba").setLevel(logging.WARNING) import default_style @@ -35,7 +36,6 @@ from style_bert_vits2.models.models import ( from style_bert_vits2.nlp.symbols import SYMBOLS from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT - torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( True # If encontered training problem,please try to disable TF32. @@ -242,15 +242,23 @@ def run(): # prefetch_factor=6, ) else: + train_sampler = DistributedLengthGroupedSampler( + dataset=train_dataset, + batch_size=hps.train.batch_size, + num_replicas=n_gpus, + rank=rank, + lengths=train_dataset.lengths, + drop_last=True, + ) train_loader = DataLoader( train_dataset, # メモリ消費量を減らそうとnum_workersを1にしてみる # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), num_workers=1, - shuffle=True, + # shuffle=True, pin_memory=True, collate_fn=collate_fn, - # batch_sampler=train_sampler, + batch_sampler=train_sampler, batch_size=hps.train.batch_size, persistent_workers=True, # これもメモリ消費量を減らそうとしてコメントアウト diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index e5c5bd1..c2539ca 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm +from transformers.trainer_pt_utils import DistributedLengthGroupedSampler # logging.getLogger("numba").setLevel(logging.WARNING) import default_style @@ -243,15 +244,23 @@ def run(): # prefetch_factor=6, ) else: + train_sampler = DistributedLengthGroupedSampler( + dataset=train_dataset, + batch_size=hps.train.batch_size, + num_replicas=n_gpus, + rank=rank, + lengths=train_dataset.lengths, + drop_last=True, + ) train_loader = DataLoader( train_dataset, # メモリ消費量を減らそうとnum_workersを1にしてみる # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), num_workers=1, - shuffle=True, + # shuffle=True, pin_memory=True, collate_fn=collate_fn, - # batch_sampler=train_sampler, + batch_sampler=train_sampler, batch_size=hps.train.batch_size, persistent_workers=True, # これもメモリ消費量を減らそうとしてコメントアウト