Batch sampler for backward compatibility, reduce tb log

This commit is contained in:
litagin02
2024-05-26 08:23:38 +09:00
parent a92b0cabff
commit 5fa6210176
4 changed files with 199 additions and 142 deletions

View File

@@ -86,7 +86,7 @@ Style-Bert-VITS2の学習用データセットを作成するためのツール
## 必要なもの ## 必要なもの
学習したい音声が入った音声ファイルいくつか形式はwav以外でも通常の音声ファイル形式なら可能 学習したい音声が入った音声ファイルいくつか形式はwav以外でもmp3等通常の音声ファイル形式なら可能)。
合計時間がある程度はあったほうがいいかも、10分とかでも大丈夫だったとの報告あり。単一ファイルでも良いし複数ファイルでもよい。 合計時間がある程度はあったほうがいいかも、10分とかでも大丈夫だったとの報告あり。単一ファイルでも良いし複数ファイルでもよい。
## スライス使い方 ## スライス使い方
@@ -120,7 +120,7 @@ def create_dataset_app() -> gr.Blocks:
input_dir = gr.Textbox( input_dir = gr.Textbox(
label="元音声の入っているフォルダパス", label="元音声の入っているフォルダパス",
value="inputs", value="inputs",
info="下記フォルダにwavファイルを入れておいてください", info="下記フォルダにwavやmp3等のファイルを入れておいてください",
) )
min_sec = gr.Slider( min_sec = gr.Slider(
minimum=0, minimum=0,

View File

@@ -329,6 +329,7 @@ def train(
skip_style: bool = False, skip_style: bool = False,
use_jp_extra: bool = True, use_jp_extra: bool = True,
speedup: bool = False, speedup: bool = False,
use_custom_batch_sampler: bool = False,
): ):
paths = get_path(model_name) paths = get_path(model_name)
# 学習再開の場合を考えて念のためconfig.ymlの名前等を更新 # 学習再開の場合を考えて念のためconfig.ymlの名前等を更新
@@ -351,6 +352,8 @@ def train(
cmd.append("--skip_default_style") cmd.append("--skip_default_style")
if speedup: if speedup:
cmd.append("--speedup") cmd.append("--speedup")
if use_custom_batch_sampler:
cmd.append("--use_custom_batch_sampler")
success, message = run_script_with_log(cmd, ignore_warning=True) success, message = run_script_with_log(cmd, ignore_warning=True)
if not success: if not success:
logger.error("Train failed.") logger.error("Train failed.")
@@ -694,6 +697,11 @@ def create_train_app():
label="JP-Extra版を使う", label="JP-Extra版を使う",
value=True, value=True,
) )
use_custom_batch_sampler = gr.Checkbox(
label="カスタムバッチサンプラーを使う",
info="Ver 2.5以降にうまく学習できなかったりVRAMが足りない場合に試してみてください",
value=False,
)
speedup = gr.Checkbox( speedup = gr.Checkbox(
label="ログ等をスキップして学習を高速化する", label="ログ等をスキップして学習を高速化する",
value=False, value=False,
@@ -781,7 +789,13 @@ def create_train_app():
# Train # Train
train_btn.click( train_btn.click(
second_elem_of(train), second_elem_of(train),
inputs=[model_name, skip_style, use_jp_extra_train, speedup], inputs=[
model_name,
skip_style,
use_jp_extra_train,
speedup,
use_custom_batch_sampler,
],
outputs=[info_train], outputs=[info_train],
) )
tensorboard_btn.click( tensorboard_btn.click(

View File

@@ -210,28 +210,45 @@ def run():
writer = SummaryWriter(log_dir=model_dir) writer = SummaryWriter(log_dir=model_dir)
writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval")) writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data) train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
train_sampler = DistributedBucketSampler(
train_dataset,
hps.train.batch_size,
[32, 300, 400, 500, 600, 700, 800, 900, 1000],
num_replicas=n_gpus,
rank=rank,
shuffle=True,
)
collate_fn = TextAudioSpeakerCollate() collate_fn = TextAudioSpeakerCollate()
train_loader = DataLoader( if args.use_custom_batch_sampler:
train_dataset, train_sampler = DistributedBucketSampler(
# メモリ消費量を減らそうとnum_workersを1にしてみる train_dataset,
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), hps.train.batch_size,
num_workers=1, [32, 300, 400, 500, 600, 700, 800, 900, 1000],
shuffle=False, num_replicas=n_gpus,
pin_memory=True, rank=rank,
collate_fn=collate_fn, shuffle=True,
batch_sampler=train_sampler, )
persistent_workers=True, train_loader = DataLoader(
# これもメモリ消費量を減らそうとしてコメントアウト train_dataset,
# prefetch_factor=4, # メモリ消費量を減らそうとnum_workersを1にしてみる
) # DataLoader config could be adjusted. # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
num_workers=1,
shuffle=False,
pin_memory=True,
collate_fn=collate_fn,
batch_sampler=train_sampler,
# batch_size=hps.train.batch_size,
persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト
# prefetch_factor=6,
)
else:
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,
pin_memory=True,
collate_fn=collate_fn,
# batch_sampler=train_sampler,
batch_size=hps.train.batch_size,
persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト
# prefetch_factor=6,
)
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:
@@ -760,21 +777,21 @@ def train_and_evaluate(
scalar_dict.update( scalar_dict.update(
{f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)} {f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)}
) )
# 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト
image_dict = { # image_dict = {
"slice/mel_org": utils.plot_spectrogram_to_numpy( # "slice/mel_org": utils.plot_spectrogram_to_numpy(
y_mel[0].data.cpu().numpy() # y_mel[0].data.cpu().numpy()
), # ),
"slice/mel_gen": utils.plot_spectrogram_to_numpy( # "slice/mel_gen": utils.plot_spectrogram_to_numpy(
y_hat_mel[0].data.cpu().numpy() # y_hat_mel[0].data.cpu().numpy()
), # ),
"all/mel": utils.plot_spectrogram_to_numpy( # "all/mel": utils.plot_spectrogram_to_numpy(
mel[0].data.cpu().numpy() # mel[0].data.cpu().numpy()
), # ),
"all/attn": utils.plot_alignment_to_numpy( # "all/attn": utils.plot_alignment_to_numpy(
attn[0, 0].data.cpu().numpy() # attn[0, 0].data.cpu().numpy()
), # ),
} # }
utils.summarize( utils.summarize(
writer=writer, writer=writer,
global_step=global_step, global_step=global_step,
@@ -906,32 +923,39 @@ def evaluate(hps, generator, eval_loader, writer_eval):
sdp_ratio=0.0 if not use_sdp else 1.0, sdp_ratio=0.0 if not use_sdp else 1.0,
) )
y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length
# 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト
mel = spec_to_mel_torch( # mel = spec_to_mel_torch(
spec, # spec,
hps.data.filter_length, # hps.data.filter_length,
hps.data.n_mel_channels, # hps.data.n_mel_channels,
hps.data.sampling_rate, # hps.data.sampling_rate,
hps.data.mel_fmin, # hps.data.mel_fmin,
hps.data.mel_fmax, # hps.data.mel_fmax,
) # )
y_hat_mel = mel_spectrogram_torch( # y_hat_mel = mel_spectrogram_torch(
y_hat.squeeze(1).float(), # y_hat.squeeze(1).float(),
hps.data.filter_length, # hps.data.filter_length,
hps.data.n_mel_channels, # hps.data.n_mel_channels,
hps.data.sampling_rate, # hps.data.sampling_rate,
hps.data.hop_length, # hps.data.hop_length,
hps.data.win_length, # hps.data.win_length,
hps.data.mel_fmin, # hps.data.mel_fmin,
hps.data.mel_fmax, # hps.data.mel_fmax,
) # )
image_dict.update( # image_dict.update(
{ # {
f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( # f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(
y_hat_mel[0].cpu().numpy() # y_hat_mel[0].cpu().numpy()
) # )
} # }
) # )
# image_dict.update(
# {
# f"gt/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(
# mel[0].cpu().numpy()
# )
# }
# )
audio_dict.update( audio_dict.update(
{ {
f"gen/audio_{batch_idx}_{use_sdp}": y_hat[ f"gen/audio_{batch_idx}_{use_sdp}": y_hat[
@@ -939,13 +963,6 @@ def evaluate(hps, generator, eval_loader, writer_eval):
] ]
} }
) )
image_dict.update(
{
f"gt/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(
mel[0].cpu().numpy()
)
}
)
audio_dict.update({f"gt/audio_{batch_idx}": y[0, :, : y_lengths[0]]}) audio_dict.update({f"gt/audio_{batch_idx}": y[0, :, : y_lengths[0]]})
utils.summarize( utils.summarize(

View File

@@ -17,7 +17,11 @@ from tqdm import tqdm
# logging.getLogger("numba").setLevel(logging.WARNING) # logging.getLogger("numba").setLevel(logging.WARNING)
import default_style import default_style
from config import get_config from config import get_config
from data_utils import TextAudioSpeakerCollate, TextAudioSpeakerLoader from data_utils import (
TextAudioSpeakerCollate,
TextAudioSpeakerLoader,
DistributedBucketSampler,
)
from losses import WavLMLoss, discriminator_loss, feature_loss, generator_loss, kl_loss from losses import WavLMLoss, discriminator_loss, feature_loss, generator_loss, kl_loss
from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from mel_processing import mel_spectrogram_torch, spec_to_mel_torch
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
@@ -94,6 +98,11 @@ def run():
help="Huggingface model repo id to backup the model.", help="Huggingface model repo id to backup the model.",
default=None, default=None,
) )
parser.add_argument(
"--use_custom_batch_sampler",
help="Use custom batch sampler for training, which was used in the version < 2.5",
action="store_true",
)
args = parser.parse_args() args = parser.parse_args()
# Set log file # Set log file
@@ -207,29 +216,45 @@ def run():
writer = SummaryWriter(log_dir=model_dir) writer = SummaryWriter(log_dir=model_dir)
writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval")) writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data) train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
# train_sampler = DistributedBucketSampler(
# train_dataset,
# hps.train.batch_size,
# [32, 300, 400, 500, 600, 700, 800, 900, 1000],
# num_replicas=n_gpus,
# rank=rank,
# shuffle=True,
# )
collate_fn = TextAudioSpeakerCollate(use_jp_extra=True) collate_fn = TextAudioSpeakerCollate(use_jp_extra=True)
train_loader = DataLoader( if args.use_custom_batch_sampler:
train_dataset, train_sampler = DistributedBucketSampler(
# メモリ消費量を減らそうとnum_workersを1にしてみる train_dataset,
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), hps.train.batch_size,
num_workers=1, [32, 300, 400, 500, 600, 700, 800, 900, 1000],
shuffle=True, num_replicas=n_gpus,
pin_memory=True, rank=rank,
collate_fn=collate_fn, shuffle=True,
# batch_sampler=train_sampler, )
batch_size=hps.train.batch_size, train_loader = DataLoader(
persistent_workers=True, train_dataset,
# これもメモリ消費量を減らそうとしてコメントアウト # メモリ消費量を減らそうとnum_workersを1にしてみる
# prefetch_factor=6, # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
) # DataLoader config could be adjusted. num_workers=1,
shuffle=False,
pin_memory=True,
collate_fn=collate_fn,
batch_sampler=train_sampler,
# batch_size=hps.train.batch_size,
persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト
# prefetch_factor=6,
)
else:
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,
pin_memory=True,
collate_fn=collate_fn,
# batch_sampler=train_sampler,
batch_size=hps.train.batch_size,
persistent_workers=True,
# これもメモリ消費量を減らそうとしてコメントアウト
# prefetch_factor=6,
)
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:
@@ -900,20 +925,21 @@ def train_and_evaluate(
"loss/g/lm_gen": loss_lm_gen, "loss/g/lm_gen": loss_lm_gen,
} }
) )
image_dict = { # 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト
"slice/mel_org": utils.plot_spectrogram_to_numpy( # image_dict = {
y_mel[0].data.cpu().numpy() # "slice/mel_org": utils.plot_spectrogram_to_numpy(
), # y_mel[0].data.cpu().numpy()
"slice/mel_gen": utils.plot_spectrogram_to_numpy( # ),
y_hat_mel[0].data.cpu().numpy() # "slice/mel_gen": utils.plot_spectrogram_to_numpy(
), # y_hat_mel[0].data.cpu().numpy()
"all/mel": utils.plot_spectrogram_to_numpy( # ),
mel[0].data.cpu().numpy() # "all/mel": utils.plot_spectrogram_to_numpy(
), # mel[0].data.cpu().numpy()
"all/attn": utils.plot_alignment_to_numpy( # ),
attn[0, 0].data.cpu().numpy() # "all/attn": utils.plot_alignment_to_numpy(
), # attn[0, 0].data.cpu().numpy()
} # ),
# }
utils.summarize( utils.summarize(
writer=writer, writer=writer,
global_step=global_step, global_step=global_step,
@@ -1046,32 +1072,39 @@ def evaluate(hps, generator, eval_loader, writer_eval):
sdp_ratio=0.0 if not use_sdp else 1.0, sdp_ratio=0.0 if not use_sdp else 1.0,
) )
y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length
# 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト
mel = spec_to_mel_torch( # mel = spec_to_mel_torch(
spec, # spec,
hps.data.filter_length, # hps.data.filter_length,
hps.data.n_mel_channels, # hps.data.n_mel_channels,
hps.data.sampling_rate, # hps.data.sampling_rate,
hps.data.mel_fmin, # hps.data.mel_fmin,
hps.data.mel_fmax, # hps.data.mel_fmax,
) # )
y_hat_mel = mel_spectrogram_torch( # y_hat_mel = mel_spectrogram_torch(
y_hat.squeeze(1).float(), # y_hat.squeeze(1).float(),
hps.data.filter_length, # hps.data.filter_length,
hps.data.n_mel_channels, # hps.data.n_mel_channels,
hps.data.sampling_rate, # hps.data.sampling_rate,
hps.data.hop_length, # hps.data.hop_length,
hps.data.win_length, # hps.data.win_length,
hps.data.mel_fmin, # hps.data.mel_fmin,
hps.data.mel_fmax, # hps.data.mel_fmax,
) # )
image_dict.update( # image_dict.update(
{ # {
f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( # f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(
y_hat_mel[0].cpu().numpy() # y_hat_mel[0].cpu().numpy()
) # )
} # }
) # )
# image_dict.update(
# {
# f"gt/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(
# mel[0].cpu().numpy()
# )
# }
# )
audio_dict.update( audio_dict.update(
{ {
f"gen/audio_{batch_idx}_{use_sdp}": y_hat[ f"gen/audio_{batch_idx}_{use_sdp}": y_hat[
@@ -1079,13 +1112,6 @@ def evaluate(hps, generator, eval_loader, writer_eval):
] ]
} }
) )
image_dict.update(
{
f"gt/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(
mel[0].cpu().numpy()
)
}
)
audio_dict.update({f"gt/audio_{batch_idx}": y[0, :, : y_lengths[0]]}) audio_dict.update({f"gt/audio_{batch_idx}": y[0, :, : y_lengths[0]]})
utils.summarize( utils.summarize(