Feat: use DistributedLengthGroupedSampler (experimental, to be checked)
This commit is contained in:
14
train_ms.py
14
train_ms.py
@@ -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,
|
||||||
# これもメモリ消費量を減らそうとしてコメントアウト
|
# これもメモリ消費量を減らそうとしてコメントアウト
|
||||||
|
|||||||
@@ -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,
|
||||||
# これもメモリ消費量を減らそうとしてコメントアウト
|
# これもメモリ消費量を減らそうとしてコメントアウト
|
||||||
|
|||||||
Reference in New Issue
Block a user