Batch sampler for backward compatibility, reduce tb log
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
155
train_ms.py
155
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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user