Feat: use DistributedLengthGroupedSampler (experimental, to be checked)

This commit is contained in:
litagin02
2024-06-22 17:21:58 +09:00
parent c1381b2c49
commit e78d2a6f7e
2 changed files with 22 additions and 5 deletions

View File

@@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
from tqdm import tqdm from tqdm import tqdm
from transformers.trainer_pt_utils import DistributedLengthGroupedSampler
# logging.getLogger("numba").setLevel(logging.WARNING) # logging.getLogger("numba").setLevel(logging.WARNING)
import default_style 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.nlp.symbols import SYMBOLS
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = ( torch.backends.cudnn.allow_tf32 = (
True # If encontered training problem,please try to disable TF32. True # If encontered training problem,please try to disable TF32.
@@ -242,15 +242,23 @@ def run():
# prefetch_factor=6, # prefetch_factor=6,
) )
else: 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_loader = DataLoader(
train_dataset, train_dataset,
# メモリ消費量を減らそうとnum_workersを1にしてみる # メモリ消費量を減らそうとnum_workersを1にしてみる
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
num_workers=1, num_workers=1,
shuffle=True, # shuffle=True,
pin_memory=True, pin_memory=True,
collate_fn=collate_fn, collate_fn=collate_fn,
# batch_sampler=train_sampler, batch_sampler=train_sampler,
batch_size=hps.train.batch_size, batch_size=hps.train.batch_size,
persistent_workers=True, persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト # これもメモリ消費量を減らそうとしてコメントアウト

View File

@@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
from tqdm import tqdm from tqdm import tqdm
from transformers.trainer_pt_utils import DistributedLengthGroupedSampler
# logging.getLogger("numba").setLevel(logging.WARNING) # logging.getLogger("numba").setLevel(logging.WARNING)
import default_style import default_style
@@ -243,15 +244,23 @@ def run():
# prefetch_factor=6, # prefetch_factor=6,
) )
else: 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_loader = DataLoader(
train_dataset, train_dataset,
# メモリ消費量を減らそうとnum_workersを1にしてみる # メモリ消費量を減らそうとnum_workersを1にしてみる
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
num_workers=1, num_workers=1,
shuffle=True, # shuffle=True,
pin_memory=True, pin_memory=True,
collate_fn=collate_fn, collate_fn=collate_fn,
# batch_sampler=train_sampler, batch_sampler=train_sampler,
batch_size=hps.train.batch_size, batch_size=hps.train.batch_size,
persistent_workers=True, persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト # これもメモリ消費量を減らそうとしてコメントアウト