From b6d09bad445de51a517018c060b33e243f58f65f Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 28 Jun 2024 09:34:12 +0900 Subject: [PATCH] Fix: use sampler instead of batch_sampler for non-custom --- train_ms.py | 2 +- train_ms_jp_extra.py | 6 +++++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/train_ms.py b/train_ms.py index 9f8ab5f..7244bb8 100644 --- a/train_ms.py +++ b/train_ms.py @@ -259,7 +259,7 @@ def run(): # shuffle=True, pin_memory=True, collate_fn=collate_fn, - batch_sampler=train_sampler, + 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 c2539ca..8c33858 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -260,12 +260,16 @@ def run(): # shuffle=True, pin_memory=True, collate_fn=collate_fn, - batch_sampler=train_sampler, + sampler=train_sampler, batch_size=hps.train.batch_size, persistent_workers=True, # これもメモリ消費量を減らそうとしてコメントアウト # prefetch_factor=6, ) + logger.info("Using DistributedLengthGroupedSampler for training.") + logger.debug(f"len(train_dataset): {len(train_dataset)}") + logger.debug(f"len(train_loader): {len(train_loader)}") + eval_dataset = None eval_loader = None if rank == 0 and not args.speedup: