From aa70540aefe4aa8d602e21a794c0821350116234 Mon Sep 17 00:00:00 2001 From: tuna2134 Date: Sun, 26 Jul 2026 22:26:11 +0900 Subject: [PATCH] fix --- train_ms_mel.py | 63 ++++++++++++++++++++++++++++++++----------------- 1 file changed, 42 insertions(+), 21 deletions(-) diff --git a/train_ms_mel.py b/train_ms_mel.py index c472c4e..a78e5bf 100644 --- a/train_ms_mel.py +++ b/train_ms_mel.py @@ -134,21 +134,47 @@ def unpack_batch( hps: HyperParameters, device: torch.device, ) -> tuple[dict[str, Any], torch.Tensor]: - ( - x, - x_lengths, - spec, - spec_lengths, - waveform, - _, - speakers, - tone, - language, - bert, - ja_bert, - en_bert, - style_vec, - ) = (item.to(device, non_blocking=True) for item in batch) + batch = tuple(item.to(device, non_blocking=True) for item in batch) + if hps.data.use_jp_extra: + ( + x, + x_lengths, + spec, + spec_lengths, + waveform, + _, + speakers, + tone, + language, + bert, + style_vec, + ) = batch + encoder_args = (tone, language, bert, style_vec) + else: + ( + x, + x_lengths, + spec, + spec_lengths, + waveform, + _, + speakers, + tone, + language, + bert, + ja_bert, + en_bert, + style_vec, + ) = batch + encoder_args = ( + tone, + language, + bert, + ja_bert, + en_bert, + style_vec, + speakers, + ) mel = ( spec if hps.model.use_mel_posterior_encoder @@ -161,11 +187,6 @@ def unpack_batch( hps.data.mel_fmax, ) ) - encoder_args = ( - (tone, language, bert, style_vec) - if hps.data.use_jp_extra - else (tone, language, bert, ja_bert, en_bert, style_vec, speakers) - ) return { "x": x, "x_lengths": x_lengths, @@ -260,7 +281,7 @@ def run() -> None: shuffle=True, num_workers=args.num_workers, pin_memory=device.type == "cuda", - collate_fn=TextAudioSpeakerCollate(), + collate_fn=TextAudioSpeakerCollate(use_jp_extra=hps.data.use_jp_extra), drop_last=True, ) output_dir = Path(args.output_dir or Path(args.config).parent / "models")