Fix: use sampler instead of batch_sampler for non-custom

This commit is contained in:
litagin02
2024-06-28 09:34:12 +09:00
parent 8832cfb4c9
commit b6d09bad44
2 changed files with 6 additions and 2 deletions

View File

@@ -259,7 +259,7 @@ def run():
# shuffle=True, # shuffle=True,
pin_memory=True, pin_memory=True,
collate_fn=collate_fn, collate_fn=collate_fn,
batch_sampler=train_sampler, sampler=train_sampler,
batch_size=hps.train.batch_size, batch_size=hps.train.batch_size,
persistent_workers=True, persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト # これもメモリ消費量を減らそうとしてコメントアウト

View File

@@ -260,12 +260,16 @@ def run():
# shuffle=True, # shuffle=True,
pin_memory=True, pin_memory=True,
collate_fn=collate_fn, collate_fn=collate_fn,
batch_sampler=train_sampler, sampler=train_sampler,
batch_size=hps.train.batch_size, batch_size=hps.train.batch_size,
persistent_workers=True, persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト # これもメモリ消費量を減らそうとしてコメントアウト
# prefetch_factor=6, # 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_dataset = None
eval_loader = None eval_loader = None
if rank == 0 and not args.speedup: if rank == 0 and not args.speedup: