diff --git a/gradio_tabs/dataset.py b/gradio_tabs/dataset.py index bbafdb5..35fccb3 100644 --- a/gradio_tabs/dataset.py +++ b/gradio_tabs/dataset.py @@ -86,7 +86,7 @@ Style-Bert-VITS2の学習用データセットを作成するためのツール ## 必要なもの -学習したい音声が入った音声ファイルいくつか(形式はwav以外でも通常の音声ファイル形式なら可能)。 +学習したい音声が入った音声ファイルいくつか(形式はwav以外でもmp3等通常の音声ファイル形式なら可能)。 合計時間がある程度はあったほうがいいかも、10分とかでも大丈夫だったとの報告あり。単一ファイルでも良いし複数ファイルでもよい。 ## スライス使い方 @@ -120,7 +120,7 @@ def create_dataset_app() -> gr.Blocks: input_dir = gr.Textbox( label="元音声の入っているフォルダパス", value="inputs", - info="下記フォルダにwavファイルを入れておいてください", + info="下記フォルダにwavやmp3等のファイルを入れておいてください", ) min_sec = gr.Slider( minimum=0, diff --git a/gradio_tabs/train.py b/gradio_tabs/train.py index 8237797..f83d30b 100644 --- a/gradio_tabs/train.py +++ b/gradio_tabs/train.py @@ -329,6 +329,7 @@ def train( skip_style: bool = False, use_jp_extra: bool = True, speedup: bool = False, + use_custom_batch_sampler: bool = False, ): paths = get_path(model_name) # 学習再開の場合を考えて念のためconfig.ymlの名前等を更新 @@ -351,6 +352,8 @@ def train( cmd.append("--skip_default_style") if 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) if not success: logger.error("Train failed.") @@ -694,6 +697,11 @@ def create_train_app(): label="JP-Extra版を使う", value=True, ) + use_custom_batch_sampler = gr.Checkbox( + label="カスタムバッチサンプラーを使う", + info="Ver 2.5以降にうまく学習できなかったりVRAMが足りない場合に試してみてください", + value=False, + ) speedup = gr.Checkbox( label="ログ等をスキップして学習を高速化する", value=False, @@ -781,7 +789,13 @@ def create_train_app(): # Train train_btn.click( 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], ) tensorboard_btn.click( diff --git a/train_ms.py b/train_ms.py index 930e8c3..31670ac 100644 --- a/train_ms.py +++ b/train_ms.py @@ -210,28 +210,45 @@ def run(): writer = SummaryWriter(log_dir=model_dir) writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval")) 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() - 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=False, - pin_memory=True, - collate_fn=collate_fn, - batch_sampler=train_sampler, - persistent_workers=True, - # これもメモリ消費量を減らそうとしてコメントアウト - # prefetch_factor=4, - ) # DataLoader config could be adjusted. + if args.use_custom_batch_sampler: + 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, + ) + 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=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_loader = None if rank == 0 and not args.speedup: @@ -760,21 +777,21 @@ def train_and_evaluate( scalar_dict.update( {f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)} ) - - image_dict = { - "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() - ), - "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() - ), - } + # 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト + # image_dict = { + # "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() + # ), + # "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() + # ), + # } utils.summarize( writer=writer, 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, ) y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length - - mel = spec_to_mel_torch( - spec, - hps.data.filter_length, - hps.data.n_mel_channels, - hps.data.sampling_rate, - hps.data.mel_fmin, - hps.data.mel_fmax, - ) - y_hat_mel = mel_spectrogram_torch( - y_hat.squeeze(1).float(), - hps.data.filter_length, - hps.data.n_mel_channels, - hps.data.sampling_rate, - hps.data.hop_length, - hps.data.win_length, - hps.data.mel_fmin, - hps.data.mel_fmax, - ) - image_dict.update( - { - f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( - y_hat_mel[0].cpu().numpy() - ) - } - ) + # 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト + # mel = spec_to_mel_torch( + # spec, + # hps.data.filter_length, + # hps.data.n_mel_channels, + # hps.data.sampling_rate, + # hps.data.mel_fmin, + # hps.data.mel_fmax, + # ) + # y_hat_mel = mel_spectrogram_torch( + # y_hat.squeeze(1).float(), + # hps.data.filter_length, + # hps.data.n_mel_channels, + # hps.data.sampling_rate, + # hps.data.hop_length, + # hps.data.win_length, + # hps.data.mel_fmin, + # hps.data.mel_fmax, + # ) + # image_dict.update( + # { + # f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_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( { 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]]}) utils.summarize( diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 07cc24b..36cbb6e 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -17,7 +17,11 @@ from tqdm import tqdm # logging.getLogger("numba").setLevel(logging.WARNING) import default_style 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 mel_processing import mel_spectrogram_torch, spec_to_mel_torch from style_bert_vits2.logging import logger @@ -94,6 +98,11 @@ def run(): help="Huggingface model repo id to backup the model.", 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() # Set log file @@ -207,29 +216,45 @@ def run(): writer = SummaryWriter(log_dir=model_dir) writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval")) 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) - 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, - ) # DataLoader config could be adjusted. + if args.use_custom_batch_sampler: + 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, + ) + 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=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_loader = None if rank == 0 and not args.speedup: @@ -900,20 +925,21 @@ def train_and_evaluate( "loss/g/lm_gen": loss_lm_gen, } ) - image_dict = { - "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() - ), - "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() - ), - } + # 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト + # image_dict = { + # "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() + # ), + # "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() + # ), + # } utils.summarize( writer=writer, 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, ) y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length - - mel = spec_to_mel_torch( - spec, - hps.data.filter_length, - hps.data.n_mel_channels, - hps.data.sampling_rate, - hps.data.mel_fmin, - hps.data.mel_fmax, - ) - y_hat_mel = mel_spectrogram_torch( - y_hat.squeeze(1).float(), - hps.data.filter_length, - hps.data.n_mel_channels, - hps.data.sampling_rate, - hps.data.hop_length, - hps.data.win_length, - hps.data.mel_fmin, - hps.data.mel_fmax, - ) - image_dict.update( - { - f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy( - y_hat_mel[0].cpu().numpy() - ) - } - ) + # 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト + # mel = spec_to_mel_torch( + # spec, + # hps.data.filter_length, + # hps.data.n_mel_channels, + # hps.data.sampling_rate, + # hps.data.mel_fmin, + # hps.data.mel_fmax, + # ) + # y_hat_mel = mel_spectrogram_torch( + # y_hat.squeeze(1).float(), + # hps.data.filter_length, + # hps.data.n_mel_channels, + # hps.data.sampling_rate, + # hps.data.hop_length, + # hps.data.win_length, + # hps.data.mel_fmin, + # hps.data.mel_fmax, + # ) + # image_dict.update( + # { + # f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_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( { 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]]}) utils.summarize(