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.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,
# これもメモリ消費量を減らそうとしてコメントアウト