From c776c082355c6e374058cd02918ea1b47a9ba1e1 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Mon, 11 Mar 2024 09:47:47 +0900 Subject: [PATCH] Apply black formatter --- .vscode/settings.json | 4 + app.py | 7 +- style_bert_vits2/constants.py | 2 + style_bert_vits2/logging.py | 6 +- style_bert_vits2/models/attentions.py | 36 +++-- style_bert_vits2/models/commons.py | 22 ++- style_bert_vits2/models/infer.py | 120 ++++++++------- style_bert_vits2/models/models.py | 30 +++- style_bert_vits2/models/models_jp_extra.py | 30 +++- style_bert_vits2/models/modules.py | 35 ++++- .../models/monotonic_alignment.py | 4 +- style_bert_vits2/models/transforms.py | 6 +- style_bert_vits2/models/utils/__init__.py | 12 +- style_bert_vits2/models/utils/checkpoints.py | 16 +- style_bert_vits2/nlp/__init__.py | 7 +- style_bert_vits2/nlp/bert_models.py | 27 +++- style_bert_vits2/nlp/chinese/g2p.py | 8 +- style_bert_vits2/nlp/english/normalizer.py | 4 +- style_bert_vits2/nlp/japanese/g2p.py | 14 +- style_bert_vits2/nlp/japanese/g2p_utils.py | 12 +- .../japanese/pyopenjtalk_worker/__init__.py | 12 +- .../pyopenjtalk_worker/worker_client.py | 15 +- .../pyopenjtalk_worker/worker_server.py | 5 +- .../nlp/japanese/user_dict/__init__.py | 14 +- style_bert_vits2/tts_model.py | 137 +++++++++--------- style_bert_vits2/utils/stdout_wrapper.py | 7 +- style_bert_vits2/utils/strenum.py | 13 +- style_bert_vits2/utils/subprocess.py | 16 +- tests/test_main.py | 26 ++-- train_ms.py | 48 +++--- train_ms_jp_extra.py | 66 +++++---- 31 files changed, 463 insertions(+), 298 deletions(-) diff --git a/.vscode/settings.json b/.vscode/settings.json index 342c024..2a583eb 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -19,4 +19,8 @@ "reportUnusedFunction": "none", "reportUnusedVariable": "information", }, + "[python]": { + "editor.defaultFormatter": "ms-python.black-formatter", + "editor.formatOnType": true, + }, } \ No newline at end of file diff --git a/app.py b/app.py index b4fbc3e..60b91a6 100644 --- a/app.py +++ b/app.py @@ -51,4 +51,9 @@ with gr.Blocks(theme=GRADIO_THEME) as app: create_merge_app(model_holder=model_holder) -app.launch(server_name=args.host, server_port=args.port, inbrowser=not args.no_autolaunch, share=args.share) +app.launch( + server_name=args.host, + server_port=args.port, + inbrowser=not args.no_autolaunch, + share=args.share, +) diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py index 735aebc..9d84690 100644 --- a/style_bert_vits2/constants.py +++ b/style_bert_vits2/constants.py @@ -9,6 +9,7 @@ VERSION = "2.4" # Style-Bert-VITS2 のベースディレクトリ BASE_DIR = Path(__file__).parent.parent + # 利用可能な言語 ## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない class Languages(StrEnum): @@ -16,6 +17,7 @@ class Languages(StrEnum): EN = "EN" ZH = "ZH" + # 言語ごとのデフォルトの BERT トークナイザーのパス DEFAULT_BERT_TOKENIZER_PATHS = { Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm", diff --git a/style_bert_vits2/logging.py b/style_bert_vits2/logging.py index eec887c..e5c216a 100644 --- a/style_bert_vits2/logging.py +++ b/style_bert_vits2/logging.py @@ -9,7 +9,7 @@ logger.remove() # Add a new handler logger.add( SAFE_STDOUT, - format = "{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}", - backtrace = True, - diagnose = True, + format="{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}", + backtrace=True, + diagnose=True, ) diff --git a/style_bert_vits2/models/attentions.py b/style_bert_vits2/models/attentions.py index 03b238d..9a10112 100644 --- a/style_bert_vits2/models/attentions.py +++ b/style_bert_vits2/models/attentions.py @@ -24,7 +24,9 @@ class LayerNorm(nn.Module): @torch.jit.script # type: ignore -def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor, n_channels: list[int]) -> torch.Tensor: +def fused_add_tanh_sigmoid_multiply( + input_a: torch.Tensor, input_b: torch.Tensor, n_channels: list[int] +) -> torch.Tensor: n_channels_int = n_channels[0] in_act = input_a + input_b t_act = torch.tanh(in_act[:, :n_channels_int, :]) @@ -44,7 +46,7 @@ class Encoder(nn.Module): p_dropout: float = 0.0, window_size: int = 4, isflow: bool = True, - **kwargs: Any + **kwargs: Any, ) -> None: super().__init__() self.hidden_channels = hidden_channels @@ -99,7 +101,9 @@ class Encoder(nn.Module): ) self.norm_layers_2.append(LayerNorm(hidden_channels)) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None + ) -> torch.Tensor: attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1) x = x * x_mask for i in range(self.n_layers): @@ -131,7 +135,7 @@ class Decoder(nn.Module): p_dropout: float = 0.0, proximal_bias: bool = False, proximal_init: bool = True, - **kwargs: Any + **kwargs: Any, ) -> None: super().__init__() self.hidden_channels = hidden_channels @@ -180,7 +184,13 @@ class Decoder(nn.Module): ) self.norm_layers_2.append(LayerNorm(hidden_channels)) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, h: torch.Tensor, h_mask: torch.Tensor) -> torch.Tensor: + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + h: torch.Tensor, + h_mask: torch.Tensor, + ) -> torch.Tensor: """ x: decoder input h: encoder output @@ -262,7 +272,9 @@ class MultiHeadAttention(nn.Module): assert self.conv_q.bias is not None self.conv_k.bias.copy_(self.conv_q.bias) - def forward(self, x: torch.Tensor, c: torch.Tensor, attn_mask: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, c: torch.Tensor, attn_mask: Optional[torch.Tensor] = None + ) -> torch.Tensor: q = self.conv_q(x) k = self.conv_k(c) v = self.conv_v(c) @@ -329,7 +341,9 @@ class MultiHeadAttention(nn.Module): ) # [b, n_h, t_t, d_k] -> [b, d, t_t] return output, p_attn - def _matmul_with_relative_values(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + def _matmul_with_relative_values( + self, x: torch.Tensor, y: torch.Tensor + ) -> torch.Tensor: """ x: [b, h, l, m] y: [h or 1, m, d] @@ -338,7 +352,9 @@ class MultiHeadAttention(nn.Module): ret = torch.matmul(x, y.unsqueeze(0)) return ret - def _matmul_with_relative_keys(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + def _matmul_with_relative_keys( + self, x: torch.Tensor, y: torch.Tensor + ) -> torch.Tensor: """ x: [b, h, l, d] y: [h or 1, m, d] @@ -347,7 +363,9 @@ class MultiHeadAttention(nn.Module): ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1)) return ret - def _get_relative_embeddings(self, relative_embeddings: torch.Tensor, length: int) -> torch.Tensor: + def _get_relative_embeddings( + self, relative_embeddings: torch.Tensor, length: int + ) -> torch.Tensor: assert self.window_size is not None 2 * self.window_size + 1 # type: ignore # Pad first before slice to avoid using cond ops. diff --git a/style_bert_vits2/models/commons.py b/style_bert_vits2/models/commons.py index 1106b21..da89930 100644 --- a/style_bert_vits2/models/commons.py +++ b/style_bert_vits2/models/commons.py @@ -67,7 +67,9 @@ def intersperse(lst: list[Any], item: Any) -> list[Any]: return result -def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4) -> torch.Tensor: +def slice_segments( + x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4 +) -> torch.Tensor: """ テンソルからセグメントをスライスする @@ -85,7 +87,9 @@ def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4 return torch.gather(x, 2, gather_indices) -def rand_slice_segments(x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4) -> tuple[torch.Tensor, torch.Tensor]: +def rand_slice_segments( + x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4 +) -> tuple[torch.Tensor, torch.Tensor]: """ ランダムなセグメントをスライスする @@ -121,7 +125,9 @@ def subsequent_mask(length: int) -> torch.Tensor: @torch.jit.script # type: ignore -def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor) -> torch.Tensor: +def fused_add_tanh_sigmoid_multiply( + input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor +) -> torch.Tensor: """ 加算、tanh、sigmoid の活性化関数を組み合わせた演算を行う @@ -141,7 +147,9 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor return acts -def sequence_mask(length: torch.Tensor, max_length: Optional[int] = None) -> torch.Tensor: +def sequence_mask( + length: torch.Tensor, max_length: Optional[int] = None +) -> torch.Tensor: """ シーケンスマスクを生成する @@ -180,7 +188,11 @@ def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: return path -def clip_grad_value_(parameters: Union[torch.Tensor, list[torch.Tensor]], clip_value: Optional[float], norm_type: float = 2.0) -> float: +def clip_grad_value_( + parameters: Union[torch.Tensor, list[torch.Tensor]], + clip_value: Optional[float], + norm_type: float = 2.0, +) -> float: """ 勾配の値をクリップする diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index d3ec963..b0ab9d2 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -9,8 +9,14 @@ from style_bert_vits2.models import commons from style_bert_vits2.models import utils from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.models.models import SynthesizerTrn -from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra -from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature +from style_bert_vits2.models.models_jp_extra import ( + SynthesizerTrn as SynthesizerTrnJPExtra, +) +from style_bert_vits2.nlp import ( + clean_text, + cleaned_text_to_sequence, + extract_bert_feature, +) from style_bert_vits2.nlp.symbols import SYMBOLS @@ -18,69 +24,71 @@ def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters): if version.endswith("JP-Extra"): logger.info("Using JP-Extra model") net_g = SynthesizerTrnJPExtra( - n_vocab = len(SYMBOLS), - spec_channels = hps.data.filter_length // 2 + 1, - segment_size = hps.train.segment_size // hps.data.hop_length, - n_speakers = hps.data.n_speakers, + n_vocab=len(SYMBOLS), + spec_channels=hps.data.filter_length // 2 + 1, + segment_size=hps.train.segment_size // hps.data.hop_length, + n_speakers=hps.data.n_speakers, # hps.model 以下のすべての値を引数に渡す - use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, - use_noise_scaled_mas = hps.model.use_noise_scaled_mas, - use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, - use_duration_discriminator = hps.model.use_duration_discriminator, - use_wavlm_discriminator = hps.model.use_wavlm_discriminator, - inter_channels = hps.model.inter_channels, - hidden_channels = hps.model.hidden_channels, - filter_channels = hps.model.filter_channels, - n_heads = hps.model.n_heads, - n_layers = hps.model.n_layers, - kernel_size = hps.model.kernel_size, - p_dropout = hps.model.p_dropout, - resblock = hps.model.resblock, - resblock_kernel_sizes = hps.model.resblock_kernel_sizes, - resblock_dilation_sizes = hps.model.resblock_dilation_sizes, - upsample_rates = hps.model.upsample_rates, - upsample_initial_channel = hps.model.upsample_initial_channel, - upsample_kernel_sizes = hps.model.upsample_kernel_sizes, - n_layers_q = hps.model.n_layers_q, - use_spectral_norm = hps.model.use_spectral_norm, - gin_channels = hps.model.gin_channels, - slm = hps.model.slm, + use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas=hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder, + use_duration_discriminator=hps.model.use_duration_discriminator, + use_wavlm_discriminator=hps.model.use_wavlm_discriminator, + inter_channels=hps.model.inter_channels, + hidden_channels=hps.model.hidden_channels, + filter_channels=hps.model.filter_channels, + n_heads=hps.model.n_heads, + n_layers=hps.model.n_layers, + kernel_size=hps.model.kernel_size, + p_dropout=hps.model.p_dropout, + resblock=hps.model.resblock, + resblock_kernel_sizes=hps.model.resblock_kernel_sizes, + resblock_dilation_sizes=hps.model.resblock_dilation_sizes, + upsample_rates=hps.model.upsample_rates, + upsample_initial_channel=hps.model.upsample_initial_channel, + upsample_kernel_sizes=hps.model.upsample_kernel_sizes, + n_layers_q=hps.model.n_layers_q, + use_spectral_norm=hps.model.use_spectral_norm, + gin_channels=hps.model.gin_channels, + slm=hps.model.slm, ).to(device) else: logger.info("Using normal model") net_g = SynthesizerTrn( - n_vocab = len(SYMBOLS), - spec_channels = hps.data.filter_length // 2 + 1, - segment_size = hps.train.segment_size // hps.data.hop_length, + n_vocab=len(SYMBOLS), + spec_channels=hps.data.filter_length // 2 + 1, + segment_size=hps.train.segment_size // hps.data.hop_length, n_speakers=hps.data.n_speakers, # hps.model 以下のすべての値を引数に渡す - use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, - use_noise_scaled_mas = hps.model.use_noise_scaled_mas, - use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, - use_duration_discriminator = hps.model.use_duration_discriminator, - use_wavlm_discriminator = hps.model.use_wavlm_discriminator, - inter_channels = hps.model.inter_channels, - hidden_channels = hps.model.hidden_channels, - filter_channels = hps.model.filter_channels, - n_heads = hps.model.n_heads, - n_layers = hps.model.n_layers, - kernel_size = hps.model.kernel_size, - p_dropout = hps.model.p_dropout, - resblock = hps.model.resblock, - resblock_kernel_sizes = hps.model.resblock_kernel_sizes, - resblock_dilation_sizes = hps.model.resblock_dilation_sizes, - upsample_rates = hps.model.upsample_rates, - upsample_initial_channel = hps.model.upsample_initial_channel, - upsample_kernel_sizes = hps.model.upsample_kernel_sizes, - n_layers_q = hps.model.n_layers_q, - use_spectral_norm = hps.model.use_spectral_norm, - gin_channels = hps.model.gin_channels, - slm = hps.model.slm, + use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas=hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder, + use_duration_discriminator=hps.model.use_duration_discriminator, + use_wavlm_discriminator=hps.model.use_wavlm_discriminator, + inter_channels=hps.model.inter_channels, + hidden_channels=hps.model.hidden_channels, + filter_channels=hps.model.filter_channels, + n_heads=hps.model.n_heads, + n_layers=hps.model.n_layers, + kernel_size=hps.model.kernel_size, + p_dropout=hps.model.p_dropout, + resblock=hps.model.resblock, + resblock_kernel_sizes=hps.model.resblock_kernel_sizes, + resblock_dilation_sizes=hps.model.resblock_dilation_sizes, + upsample_rates=hps.model.upsample_rates, + upsample_initial_channel=hps.model.upsample_initial_channel, + upsample_kernel_sizes=hps.model.upsample_kernel_sizes, + n_layers_q=hps.model.n_layers_q, + use_spectral_norm=hps.model.use_spectral_norm, + gin_channels=hps.model.gin_channels, + slm=hps.model.slm, ).to(device) net_g.state_dict() _ = net_g.eval() if model_path.endswith(".pth") or model_path.endswith(".pt"): - _ = utils.checkpoints.load_checkpoint(model_path, net_g, None, skip_optimizer=True) + _ = utils.checkpoints.load_checkpoint( + model_path, net_g, None, skip_optimizer=True + ) elif model_path.endswith(".safetensors"): _ = utils.safetensors.load_safetensors(model_path, net_g, True) else: @@ -102,8 +110,8 @@ def get_text( norm_text, phone, tone, word2ph = clean_text( text, language_str, - use_jp_extra = use_jp_extra, - raise_yomi_error = False, + use_jp_extra=use_jp_extra, + raise_yomi_error=False, ) if given_tone is not None: if len(given_tone) != len(phone): diff --git a/style_bert_vits2/models/models.py b/style_bert_vits2/models/models.py index 829ca4a..21a0be4 100644 --- a/style_bert_vits2/models/models.py +++ b/style_bert_vits2/models/models.py @@ -21,7 +21,7 @@ class DurationDiscriminator(nn.Module): # vits2 filter_channels: int, kernel_size: int, p_dropout: float, - gin_channels: int = 0 + gin_channels: int = 0, ) -> None: super().__init__() @@ -330,7 +330,9 @@ class DurationPredictor(nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, in_channels, 1) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None + ) -> torch.Tensor: x = torch.detach(x) if g is not None: g = torch.detach(g) @@ -582,7 +584,9 @@ class Generator(torch.nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) - def forward(self, x: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, g: Optional[torch.Tensor] = None + ) -> torch.Tensor: x = self.conv_pre(x) if g is not None: x = x + self.cond(g) @@ -613,7 +617,13 @@ class Generator(torch.nn.Module): class DiscriminatorP(torch.nn.Module): - def __init__(self, period: int, kernel_size: int = 5, stride: int = 3, use_spectral_norm: bool = False) -> None: + def __init__( + self, + period: int, + kernel_size: int = 5, + stride: int = 3, + use_spectral_norm: bool = False, + ) -> None: super(DiscriminatorP, self).__init__() self.period = period self.use_spectral_norm = use_spectral_norm @@ -736,7 +746,9 @@ class MultiPeriodDiscriminator(torch.nn.Module): self, y: torch.Tensor, y_hat: torch.Tensor, - ) -> tuple[list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]]: + ) -> tuple[ + list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor] + ]: y_d_rs = [] y_d_gs = [] fmap_rs = [] @@ -787,7 +799,9 @@ class ReferenceEncoder(nn.Module): ) self.proj = nn.Linear(128, gin_channels) - def forward(self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None + ) -> torch.Tensor: N = inputs.size(0) out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs] for conv in self.convs: @@ -805,7 +819,9 @@ class ReferenceEncoder(nn.Module): return self.proj(out.squeeze(0)) - def calculate_channels(self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int) -> int: + def calculate_channels( + self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int + ) -> int: for i in range(n_convs): L = (L - kernel_size + 2 * pad) // stride + 1 return L diff --git a/style_bert_vits2/models/models_jp_extra.py b/style_bert_vits2/models/models_jp_extra.py index a43a715..00cc02f 100644 --- a/style_bert_vits2/models/models_jp_extra.py +++ b/style_bert_vits2/models/models_jp_extra.py @@ -21,7 +21,7 @@ class DurationDiscriminator(nn.Module): # vits2 filter_channels: int, kernel_size: int, p_dropout: float, - gin_channels: int = 0 + gin_channels: int = 0, ) -> None: super().__init__() @@ -313,7 +313,9 @@ class DurationPredictor(nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, in_channels, 1) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None + ) -> torch.Tensor: x = torch.detach(x) if g is not None: g = torch.detach(g) @@ -587,7 +589,9 @@ class Generator(torch.nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) - def forward(self, x: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, g: Optional[torch.Tensor] = None + ) -> torch.Tensor: x = self.conv_pre(x) if g is not None: x = x + self.cond(g) @@ -618,7 +622,13 @@ class Generator(torch.nn.Module): class DiscriminatorP(torch.nn.Module): - def __init__(self, period: int, kernel_size: int = 5, stride: int = 3, use_spectral_norm: bool = False) -> None: + def __init__( + self, + period: int, + kernel_size: int = 5, + stride: int = 3, + use_spectral_norm: bool = False, + ) -> None: super(DiscriminatorP, self).__init__() self.period = period self.use_spectral_norm = use_spectral_norm @@ -741,7 +751,9 @@ class MultiPeriodDiscriminator(torch.nn.Module): self, y: torch.Tensor, y_hat: torch.Tensor, - ) -> tuple[list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]]: + ) -> tuple[ + list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor] + ]: y_d_rs = [] y_d_gs = [] fmap_rs = [] @@ -845,7 +857,9 @@ class ReferenceEncoder(nn.Module): ) self.proj = nn.Linear(128, gin_channels) - def forward(self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None + ) -> torch.Tensor: N = inputs.size(0) out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs] for conv in self.convs: @@ -863,7 +877,9 @@ class ReferenceEncoder(nn.Module): return self.proj(out.squeeze(0)) - def calculate_channels(self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int) -> int: + def calculate_channels( + self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int + ) -> int: for i in range(n_convs): L = (L - kernel_size + 2 * pad) // stride + 1 return L diff --git a/style_bert_vits2/models/modules.py b/style_bert_vits2/models/modules.py index 8eed963..ebc5273 100644 --- a/style_bert_vits2/models/modules.py +++ b/style_bert_vits2/models/modules.py @@ -88,7 +88,9 @@ class DDSConv(nn.Module): Dialted and Depth-Separable Convolution """ - def __init__(self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0) -> None: + def __init__( + self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0 + ) -> None: super().__init__() self.channels = channels self.kernel_size = kernel_size @@ -117,7 +119,9 @@ class DDSConv(nn.Module): self.norms_1.append(LayerNorm(channels)) self.norms_2.append(LayerNorm(channels)) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None + ) -> torch.Tensor: if g is not None: x = x + g for i in range(self.n_layers): @@ -184,7 +188,13 @@ class WN(torch.nn.Module): res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight") self.res_skip_layers.append(res_skip_layer) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, **kwargs: Any) -> torch.Tensor: + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + **kwargs: Any, + ) -> torch.Tensor: output = torch.zeros_like(x) n_channels_tensor = torch.IntTensor([self.hidden_channels]) @@ -221,7 +231,12 @@ class WN(torch.nn.Module): class ResBlock1(torch.nn.Module): - def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int, int] = (1, 3, 5)) -> None: + def __init__( + self, + channels: int, + kernel_size: int = 3, + dilation: tuple[int, int, int] = (1, 3, 5), + ) -> None: super(ResBlock1, self).__init__() self.convs1 = nn.ModuleList( [ @@ -295,7 +310,9 @@ class ResBlock1(torch.nn.Module): ) self.convs2.apply(commons.init_weights) - def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None + ) -> torch.Tensor: for c1, c2 in zip(self.convs1, self.convs2): xt = F.leaky_relu(x, LRELU_SLOPE) if x_mask is not None: @@ -318,7 +335,9 @@ class ResBlock1(torch.nn.Module): class ResBlock2(torch.nn.Module): - def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)) -> None: + def __init__( + self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3) + ) -> None: super(ResBlock2, self).__init__() self.convs = nn.ModuleList( [ @@ -346,7 +365,9 @@ class ResBlock2(torch.nn.Module): ) self.convs.apply(commons.init_weights) - def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor: + def forward( + self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None + ) -> torch.Tensor: for c in self.convs: xt = F.leaky_relu(x, LRELU_SLOPE) if x_mask is not None: diff --git a/style_bert_vits2/models/monotonic_alignment.py b/style_bert_vits2/models/monotonic_alignment.py index b499ad0..d33631e 100644 --- a/style_bert_vits2/models/monotonic_alignment.py +++ b/style_bert_vits2/models/monotonic_alignment.py @@ -40,8 +40,8 @@ def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: numba.int32[::1], numba.int32[::1], ), - nopython = True, - nogil = True, + nopython=True, + nogil=True, ) # type: ignore def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None: """ diff --git a/style_bert_vits2/models/transforms.py b/style_bert_vits2/models/transforms.py index 61306ad..b6f4420 100644 --- a/style_bert_vits2/models/transforms.py +++ b/style_bert_vits2/models/transforms.py @@ -39,12 +39,14 @@ def piecewise_rational_quadratic_transform( min_bin_width=min_bin_width, min_bin_height=min_bin_height, min_derivative=min_derivative, - **spline_kwargs # type: ignore + **spline_kwargs, # type: ignore ) return outputs, logabsdet -def searchsorted(bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6) -> torch.Tensor: +def searchsorted( + bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6 +) -> torch.Tensor: bin_locations[..., -1] += eps return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1 diff --git a/style_bert_vits2/models/utils/__init__.py b/style_bert_vits2/models/utils/__init__.py index e0d922f..0fd3a47 100644 --- a/style_bert_vits2/models/utils/__init__.py +++ b/style_bert_vits2/models/utils/__init__.py @@ -107,7 +107,9 @@ def plot_spectrogram_to_numpy(spectrogram: NDArray[Any]) -> NDArray[Any]: return data -def plot_alignment_to_numpy(alignment: NDArray[Any], info: Optional[str] = None) -> NDArray[Any]: +def plot_alignment_to_numpy( + alignment: NDArray[Any], info: Optional[str] = None +) -> NDArray[Any]: """ 指定されたアライメントを画像データに変換する @@ -163,7 +165,9 @@ def load_wav_to_torch(full_path: Union[str, Path]) -> tuple[torch.FloatTensor, i return torch.FloatTensor(data.astype(np.float32)), sampling_rate -def load_filepaths_and_text(filename: Union[str, Path], split: str = "|") -> list[list[str]]: +def load_filepaths_and_text( + filename: Union[str, Path], split: str = "|" +) -> list[list[str]]: """ 指定されたファイルからファイルパスとテキストを読み込む @@ -180,7 +184,9 @@ def load_filepaths_and_text(filename: Union[str, Path], split: str = "|") -> lis return filepaths_and_text -def get_logger(model_dir_path: Union[str, Path], filename: str = "train.log") -> logging.Logger: +def get_logger( + model_dir_path: Union[str, Path], filename: str = "train.log" +) -> logging.Logger: """ ロガーを取得する diff --git a/style_bert_vits2/models/utils/checkpoints.py b/style_bert_vits2/models/utils/checkpoints.py index f26f8fc..768973b 100644 --- a/style_bert_vits2/models/utils/checkpoints.py +++ b/style_bert_vits2/models/utils/checkpoints.py @@ -14,7 +14,7 @@ def load_checkpoint( model: torch.nn.Module, optimizer: Optional[torch.optim.Optimizer] = None, skip_optimizer: bool = False, - for_infer: bool = False + for_infer: bool = False, ) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]: """ 指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。 @@ -107,7 +107,9 @@ def save_checkpoint( iteration (int): イテレーション数 checkpoint_path (Union[str, Path]): 保存先のパス """ - logger.info(f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}") + logger.info( + f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}" + ) if hasattr(model, "module"): state_dict = model.module.state_dict() else: @@ -123,7 +125,11 @@ def save_checkpoint( ) -def clean_checkpoints(model_dir_path: Union[str, Path] = "logs/44k/", n_ckpts_to_keep: int = 2, sort_by_time: bool = True) -> None: +def clean_checkpoints( + model_dir_path: Union[str, Path] = "logs/44k/", + n_ckpts_to_keep: int = 2, + sort_by_time: bool = True, +) -> None: """ 指定されたディレクトリから古いチェックポイントを削除して空き容量を確保する @@ -172,7 +178,9 @@ def clean_checkpoints(model_dir_path: Union[str, Path] = "logs/44k/", n_ckpts_to [del_routine(fn) for fn in to_del] -def get_latest_checkpoint_path(model_dir_path: Union[str, Path], regex: str = "G_*.pth") -> str: +def get_latest_checkpoint_path( + model_dir_path: Union[str, Path], regex: str = "G_*.pth" +) -> str: """ 指定されたディレクトリから最新のチェックポイントのパスを取得する diff --git a/style_bert_vits2/nlp/__init__.py b/style_bert_vits2/nlp/__init__.py index 683d6d4..5f3d63f 100644 --- a/style_bert_vits2/nlp/__init__.py +++ b/style_bert_vits2/nlp/__init__.py @@ -74,16 +74,19 @@ def clean_text( if language == Languages.JP: from style_bert_vits2.nlp.japanese.g2p import g2p from style_bert_vits2.nlp.japanese.normalizer import normalize_text + norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error) elif language == Languages.EN: from style_bert_vits2.nlp.english.g2p import g2p from style_bert_vits2.nlp.english.normalizer import normalize_text + norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text) elif language == Languages.ZH: from style_bert_vits2.nlp.chinese.g2p import g2p from style_bert_vits2.nlp.chinese.normalizer import normalize_text + norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text) else: @@ -92,7 +95,9 @@ def clean_text( return norm_text, phones, tones, word2ph -def cleaned_text_to_sequence(cleaned_phones: list[str], tones: list[int], language: Languages) -> tuple[list[int], list[int], list[int]]: +def cleaned_text_to_sequence( + cleaned_phones: list[str], tones: list[int], language: Languages +) -> tuple[list[int], list[int], list[int]]: """ テキスト文字列を、テキスト内の記号に対応する一連の ID に変換する diff --git a/style_bert_vits2/nlp/bert_models.py b/style_bert_vits2/nlp/bert_models.py index 4385d64..220d584 100644 --- a/style_bert_vits2/nlp/bert_models.py +++ b/style_bert_vits2/nlp/bert_models.py @@ -30,7 +30,9 @@ from style_bert_vits2.logging import logger __loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {} # 各言語ごとのロード済みの BERT トークナイザーを格納する辞書 -__loaded_tokenizers: dict[Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]] = {} +__loaded_tokenizers: dict[ + Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer] +] = {} def load_model( @@ -63,18 +65,24 @@ def load_model( # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: - assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \ - f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model." + assert DEFAULT_BERT_TOKENIZER_PATHS[ + language + ].exists(), f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model." pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language]) # BERT モデルをロードし、辞書に格納して返す ## 英語のみ DebertaV2Model でロードする必要がある if language == Languages.EN: - model = cast(DebertaV2Model, DebertaV2Model.from_pretrained(pretrained_model_name_or_path)) + model = cast( + DebertaV2Model, + DebertaV2Model.from_pretrained(pretrained_model_name_or_path), + ) else: model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path) __loaded_models[language] = model - logger.info(f"Loaded the {language} BERT model from {pretrained_model_name_or_path}") + logger.info( + f"Loaded the {language} BERT model from {pretrained_model_name_or_path}" + ) return model @@ -109,8 +117,9 @@ def load_tokenizer( # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: - assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \ - f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." + assert DEFAULT_BERT_TOKENIZER_PATHS[ + language + ].exists(), f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language]) # BERT トークナイザーをロードし、辞書に格納して返す @@ -120,7 +129,9 @@ def load_tokenizer( else: tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path) __loaded_tokenizers[language] = tokenizer - logger.info(f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}") + logger.info( + f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}" + ) return tokenizer diff --git a/style_bert_vits2/nlp/chinese/g2p.py b/style_bert_vits2/nlp/chinese/g2p.py index 1cb3839..b5744cd 100644 --- a/style_bert_vits2/nlp/chinese/g2p.py +++ b/style_bert_vits2/nlp/chinese/g2p.py @@ -95,7 +95,11 @@ def __g2p(segments: list[str]) -> tuple[list[str], list[int], list[int]]: if pinyin[0] in single_rep_map.keys(): pinyin = single_rep_map[pinyin[0]] + pinyin[1:] - assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (pinyin, seg, raw_pinyin) + assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), ( + pinyin, + seg, + raw_pinyin, + ) phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ") word2ph.append(len(phone)) @@ -125,7 +129,7 @@ if __name__ == "__main__": text = normalize_text(text) print(text) phones, tones, word2ph = g2p(text) - bert = extract_bert_feature(text, word2ph, 'cuda') + bert = extract_bert_feature(text, word2ph, "cuda") print(phones, tones, word2ph, bert.shape) diff --git a/style_bert_vits2/nlp/english/normalizer.py b/style_bert_vits2/nlp/english/normalizer.py index 81b71d7..f6ddc90 100644 --- a/style_bert_vits2/nlp/english/normalizer.py +++ b/style_bert_vits2/nlp/english/normalizer.py @@ -121,7 +121,9 @@ def __expand_number(m: re.Match[str]) -> str: else: return __INFLECT.number_to_words( num, andword="", zero="oh", group=2 # type: ignore - ).replace(", ", " ") # type: ignore + ).replace( + ", ", " " + ) # type: ignore else: return __INFLECT.number_to_words(num, andword="") # type: ignore diff --git a/style_bert_vits2/nlp/japanese/g2p.py b/style_bert_vits2/nlp/japanese/g2p.py index 7fc97f2..1f6b450 100644 --- a/style_bert_vits2/nlp/japanese/g2p.py +++ b/style_bert_vits2/nlp/japanese/g2p.py @@ -10,9 +10,7 @@ from style_bert_vits2.nlp.symbols import PUNCTUATIONS def g2p( - norm_text: str, - use_jp_extra: bool = True, - raise_yomi_error: bool = False + norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False ) -> tuple[list[str], list[int], list[int]]: """ 他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、 @@ -93,8 +91,7 @@ def g2p( def text_to_sep_kata( - norm_text: str, - raise_yomi_error: bool = False + norm_text: str, raise_yomi_error: bool = False ) -> tuple[list[str], list[str]]: """ `normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、 @@ -212,7 +209,9 @@ def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]: return result -def __pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]: +def __pyopenjtalk_g2p_prosody( + text: str, drop_unvoiced_vowels: bool = True +) -> list[str]: """ ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。 ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py @@ -414,8 +413,7 @@ def __kata_to_phoneme_list(text: str) -> list[str]: def __align_tones( - phones_with_punct: list[str], - phone_tone_list: list[tuple[str, int]] + phones_with_punct: list[str], phone_tone_list: list[tuple[str, int]] ) -> list[tuple[str, int]]: """ 例: …私は、、そう思う。 diff --git a/style_bert_vits2/nlp/japanese/g2p_utils.py b/style_bert_vits2/nlp/japanese/g2p_utils.py index 893d3b5..511793f 100644 --- a/style_bert_vits2/nlp/japanese/g2p_utils.py +++ b/style_bert_vits2/nlp/japanese/g2p_utils.py @@ -34,11 +34,13 @@ def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, i """ # 子音の集合 - CONSONANTS = set([ - consonant - for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values() - if consonant is not None - ]) + CONSONANTS = set( + [ + consonant + for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values() + if consonant is not None + ] + ) phone_tone = phone_tone[1:] # 最初の("_", 0)を無視 phones = [phone for phone, _ in phone_tone] diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py index 212d21a..3a146b6 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py @@ -25,6 +25,7 @@ def run_frontend(text: str) -> list[dict[str, Any]]: else: # without worker import pyopenjtalk + return pyopenjtalk.run_frontend(text) @@ -36,6 +37,7 @@ def make_label(njd_features: Any) -> list[str]: else: # without worker import pyopenjtalk + return pyopenjtalk.make_label(njd_features) @@ -45,6 +47,7 @@ def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None) - else: # without worker import pyopenjtalk + pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab) @@ -54,6 +57,7 @@ def update_global_jtalk_with_user_dict(path: str) -> None: else: # without worker import pyopenjtalk + pyopenjtalk.update_global_jtalk_with_user_dict(path) @@ -63,6 +67,7 @@ def unset_user_dict() -> None: else: # without worker import pyopenjtalk + pyopenjtalk.unset_user_dict() @@ -102,7 +107,12 @@ def initialize_worker(port: int = WORKER_PORT) -> None: else: # align with Windows behavior # start_new_session is same as specifying setsid in preexec_fn - subprocess.Popen(args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, start_new_session=True) + subprocess.Popen( + args, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + start_new_session=True, + ) # wait until server listening count = 0 diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py index 425cf25..c4c5606 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py @@ -2,12 +2,15 @@ import socket from typing import Any, cast from style_bert_vits2.logging import logger -from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import RequestType, receive_data, send_data +from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import ( + RequestType, + receive_data, + send_data, +) class WorkerClient: - """ pyopenjtalk worker client """ - + """pyopenjtalk worker client""" def __init__(self, port: int) -> None: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) @@ -16,19 +19,15 @@ class WorkerClient: sock.connect((socket.gethostname(), port)) self.sock = sock - def __enter__(self) -> "WorkerClient": return self - def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None: self.close() - def close(self) -> None: self.sock.close() - def dispatch_pyopenjtalk(self, func: str, *args: Any, **kwargs: Any) -> Any: data = { "request-type": RequestType.PYOPENJTALK, @@ -43,7 +42,6 @@ class WorkerClient: logger.trace(f"client received response: {response}") return response.get("return") - def status(self) -> int: data = {"request-type": RequestType.STATUS} logger.trace(f"client sends request: {data}") @@ -53,7 +51,6 @@ class WorkerClient: logger.trace(f"client received response: {response}") return cast(int, response.get("client-count")) - def quit_server(self) -> None: data = {"request-type": RequestType.QUIT_SERVER} logger.trace(f"client sends request: {data}") diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py index 149323a..ed0d4e7 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py @@ -26,14 +26,12 @@ PYOPENJTALK_FUNC_DICT = { class WorkerServer: - """ pyopenjtalk worker server """ - + """pyopenjtalk worker server""" def __init__(self) -> None: self.client_count: int = 0 self.quit: bool = False - def handle_request(self, request: dict[str, Any]) -> dict[str, Any]: request_type = None try: @@ -70,7 +68,6 @@ class WorkerServer: return response - def start_server(self, port: int, no_client_timeout: int = 30) -> None: logger.info("start pyopenjtalk worker server") with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket: diff --git a/style_bert_vits2/nlp/japanese/user_dict/__init__.py b/style_bert_vits2/nlp/japanese/user_dict/__init__.py index a2cc43e..2a4aa2f 100644 --- a/style_bert_vits2/nlp/japanese/user_dict/__init__.py +++ b/style_bert_vits2/nlp/japanese/user_dict/__init__.py @@ -18,7 +18,11 @@ from fastapi import HTTPException from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes -from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data +from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import ( + MAX_PRIORITY, + MIN_PRIORITY, + part_of_speech_data, +) # root_dir = engine_root() # save_dir = get_save_dir() @@ -26,9 +30,13 @@ from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIO # if not save_dir.is_dir(): # save_dir.mkdir(parents=True) -default_dict_path = DEFAULT_USER_DICT_DIR / "default.csv" # VOICEVOXデフォルト辞書ファイルのパス +default_dict_path = ( + DEFAULT_USER_DICT_DIR / "default.csv" +) # VOICEVOXデフォルト辞書ファイルのパス user_dict_path = DEFAULT_USER_DICT_DIR / "user_dict.json" # ユーザー辞書ファイルのパス -compiled_dict_path = DEFAULT_USER_DICT_DIR / "user.dic" # コンパイル済み辞書ファイルのパス +compiled_dict_path = ( + DEFAULT_USER_DICT_DIR / "user.dic" +) # コンパイル済み辞書ファイルのパス # # 同時書き込みの制御 diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 527b4ab..a769ae6 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -25,7 +25,9 @@ from style_bert_vits2.constants import ( from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.models.infer import get_net_g, infer from style_bert_vits2.models.models import SynthesizerTrn -from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra +from style_bert_vits2.models.models_jp_extra import ( + SynthesizerTrn as SynthesizerTrnJPExtra, +) from style_bert_vits2.logging import logger from style_bert_vits2.voice import adjust_voice @@ -36,7 +38,6 @@ class TTSModel: モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。 """ - def __init__( self, model_path: Path, @@ -59,7 +60,9 @@ class TTSModel: self.config_path: Path = config_path self.style_vec_path: Path = style_vec_path self.device: str = device - self.hyper_parameters: HyperParameters = HyperParameters.load_from_json(self.config_path) + self.hyper_parameters: HyperParameters = HyperParameters.load_from_json( + self.config_path + ) self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()} @@ -82,19 +85,17 @@ class TTSModel: self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None - def load(self) -> None: """ 音声合成モデルをデバイスにロードする。 """ self.__net_g = get_net_g( - model_path = str(self.model_path), - version = self.hyper_parameters.version, - device = self.device, - hps = self.hyper_parameters, + model_path=str(self.model_path), + version=self.hyper_parameters.version, + device=self.device, + hps=self.hyper_parameters, ) - def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: """ スタイルベクトルを取得する。 @@ -111,8 +112,9 @@ class TTSModel: style_vec = mean + (style_vec - mean) * weight return style_vec - - def __get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]: + def __get_style_vector_from_audio( + self, audio_path: str, weight: float = 1.0 + ) -> NDArray[Any]: """ 音声からスタイルベクトルを推論する。 @@ -126,8 +128,10 @@ class TTSModel: # スタイルベクトルを取得するための推論モデルを初期化 if self.__style_vector_inference is None: self.__style_vector_inference = pyannote.audio.Inference( - model = pyannote.audio.Model.from_pretrained("pyannote/wespeaker-voxceleb-resnet34-LM"), - window = "whole", + model=pyannote.audio.Model.from_pretrained( + "pyannote/wespeaker-voxceleb-resnet34-LM" + ), + window="whole", ) self.__style_vector_inference.to(torch.device(self.device)) @@ -137,7 +141,6 @@ class TTSModel: xvec = mean + (xvec - mean) * weight return xvec - def infer( self, text: str, @@ -209,20 +212,20 @@ class TTSModel: if not line_split: with torch.no_grad(): audio = infer( - text = text, - sdp_ratio = sdp_ratio, - noise_scale = noise, - noise_scale_w = noise_w, - length_scale = length, - sid = speaker_id, - language = language, - hps = self.hyper_parameters, - net_g = self.__net_g, - device = self.device, - assist_text = assist_text, - assist_text_weight = assist_text_weight, - style_vec = style_vector, - given_tone = given_tone, + text=text, + sdp_ratio=sdp_ratio, + noise_scale=noise, + noise_scale_w=noise_w, + length_scale=length, + sid=speaker_id, + language=language, + hps=self.hyper_parameters, + net_g=self.__net_g, + device=self.device, + assist_text=assist_text, + assist_text_weight=assist_text_weight, + style_vec=style_vector, + given_tone=given_tone, ) else: texts = text.split("\n") @@ -232,19 +235,19 @@ class TTSModel: for i, t in enumerate(texts): audios.append( infer( - text = t, - sdp_ratio = sdp_ratio, - noise_scale = noise, - noise_scale_w = noise_w, - length_scale = length, - sid = speaker_id, - language = language, - hps = self.hyper_parameters, - net_g = self.__net_g, - device = self.device, - assist_text = assist_text, - assist_text_weight = assist_text_weight, - style_vec = style_vector, + text=t, + sdp_ratio=sdp_ratio, + noise_scale=noise, + noise_scale_w=noise_w, + length_scale=length, + sid=speaker_id, + language=language, + hps=self.hyper_parameters, + net_g=self.__net_g, + device=self.device, + assist_text=assist_text, + assist_text_weight=assist_text_weight, + style_vec=style_vector, ) ) if i != len(texts) - 1: @@ -253,10 +256,10 @@ class TTSModel: logger.info("Audio data generated successfully") if not (pitch_scale == 1.0 and intonation_scale == 1.0): _, audio = adjust_voice( - fs = self.hyper_parameters.data.sampling_rate, - wave = audio, - pitch_scale = pitch_scale, - intonation_scale = intonation_scale, + fs=self.hyper_parameters.data.sampling_rate, + wave=audio, + pitch_scale=pitch_scale, + intonation_scale=intonation_scale, ) with warnings.catch_warnings(): warnings.simplefilter("ignore") @@ -277,7 +280,6 @@ class TTSModelHolder: model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。 """ - def __init__(self, model_root_dir: Path, device: str) -> None: """ Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。 @@ -308,7 +310,6 @@ class TTSModelHolder: self.models_info: list[TTSModelInfo] = [] self.refresh() - def refresh(self) -> None: """ 音声合成モデルの一覧を更新する。 @@ -342,13 +343,14 @@ class TTSModelHolder: styles = list(style2id.keys()) spk2id: dict[str, int] = hyper_parameters.data.spk2id speakers = list(spk2id.keys()) - self.models_info.append(TTSModelInfo( - name = model_dir.name, - files = [str(f) for f in model_files], - styles = styles, - speakers = speakers, - )) - + self.models_info.append( + TTSModelInfo( + name=model_dir.name, + files=[str(f) for f in model_files], + styles=styles, + speakers=speakers, + ) + ) def get_model(self, model_name: str, model_path_str: str) -> TTSModel: """ @@ -370,16 +372,17 @@ class TTSModelHolder: raise ValueError(f"Model file `{model_path}` is not found") if self.current_model is None or self.current_model.model_path != model_path: self.current_model = TTSModel( - model_path = model_path, - config_path = self.root_dir / model_name / "config.json", - style_vec_path = self.root_dir / model_name / "style_vectors.npy", - device = self.device, + model_path=model_path, + config_path=self.root_dir / model_name / "config.json", + style_vec_path=self.root_dir / model_name / "style_vectors.npy", + device=self.device, ) return self.current_model - - def get_model_for_gradio(self, model_name: str, model_path_str: str) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: + def get_model_for_gradio( + self, model_name: str, model_path_str: str + ) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: model_path = Path(model_path_str) if model_name not in self.model_files_dict: raise ValueError(f"Model `{model_name}` is not found") @@ -398,10 +401,10 @@ class TTSModelHolder: gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) self.current_model = TTSModel( - model_path = model_path, - config_path = self.root_dir / model_name / "config.json", - style_vec_path = self.root_dir / model_name / "style_vectors.npy", - device = self.device, + model_path=model_path, + config_path=self.root_dir / model_name / "config.json", + style_vec_path=self.root_dir / model_name / "style_vectors.npy", + device=self.device, ) speakers = list(self.current_model.spk2id.keys()) styles = list(self.current_model.style2id.keys()) @@ -411,13 +414,13 @@ class TTSModelHolder: gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) - def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown: model_files = self.model_files_dict[model_name] return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore - - def update_model_names_for_gradio(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: + def update_model_names_for_gradio( + self, + ) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: self.refresh() initial_model_name = self.model_names[0] initial_model_files = self.model_files_dict[initial_model_name] diff --git a/style_bert_vits2/utils/stdout_wrapper.py b/style_bert_vits2/utils/stdout_wrapper.py index 09254ad..174d1fd 100644 --- a/style_bert_vits2/utils/stdout_wrapper.py +++ b/style_bert_vits2/utils/stdout_wrapper.py @@ -8,40 +8,35 @@ class StdoutWrapper(TextIO): `sys.stdout` wrapper for both Google Colab and local environment. """ - def __init__(self) -> None: self.temp_file = tempfile.NamedTemporaryFile( mode="w+", delete=False, encoding="utf-8" ) self.original_stdout = sys.stdout - def write(self, message: str) -> int: result = self.temp_file.write(message) self.temp_file.flush() print(message, end="", file=self.original_stdout) return result - def flush(self) -> None: self.temp_file.flush() - def read(self, n: int = -1) -> str: self.temp_file.seek(0) return self.temp_file.read(n) - def close(self) -> None: self.temp_file.close() - def fileno(self) -> int: return self.temp_file.fileno() try: import google.colab # type: ignore + SAFE_STDOUT = StdoutWrapper() except ImportError: SAFE_STDOUT = sys.stdout diff --git a/style_bert_vits2/utils/strenum.py b/style_bert_vits2/utils/strenum.py index 40d3b0c..b2e2e3c 100644 --- a/style_bert_vits2/utils/strenum.py +++ b/style_bert_vits2/utils/strenum.py @@ -9,27 +9,28 @@ class StrEnum(str, enum.Enum): def __new__(cls, *values: str) -> "StrEnum": "values must already be of type `str`" if len(values) > 3: - raise TypeError('too many arguments for str(): %r' % (values, )) + raise TypeError("too many arguments for str(): %r" % (values,)) if len(values) == 1: # it must be a string if not isinstance(values[0], str): # type: ignore - raise TypeError('%r is not a string' % (values[0], )) + raise TypeError("%r is not a string" % (values[0],)) if len(values) >= 2: # check that encoding argument is a string if not isinstance(values[1], str): # type: ignore - raise TypeError('encoding must be a string, not %r' % (values[1], )) + raise TypeError("encoding must be a string, not %r" % (values[1],)) if len(values) == 3: # check that errors argument is a string if not isinstance(values[2], str): # type: ignore - raise TypeError('errors must be a string, not %r' % (values[2])) + raise TypeError("errors must be a string, not %r" % (values[2])) value = str(*values) member = str.__new__(cls, value) member._value_ = value return member - @staticmethod - def _generate_next_value_(name: str, start: int, count: int, last_values: list[str]) -> str: + def _generate_next_value_( + name: str, start: int, count: int, last_values: list[str] + ) -> str: """ Return the lower-cased version of the member name. """ diff --git a/style_bert_vits2/utils/subprocess.py b/style_bert_vits2/utils/subprocess.py index 542f94b..f8e8e94 100644 --- a/style_bert_vits2/utils/subprocess.py +++ b/style_bert_vits2/utils/subprocess.py @@ -6,7 +6,9 @@ from style_bert_vits2.logging import logger from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT -def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[bool, str]: +def run_script_with_log( + cmd: list[str], ignore_warning: bool = False +) -> tuple[bool, str]: """ 指定されたコマンドを実行し、そのログを記録する。 @@ -21,10 +23,10 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b logger.info(f"Running: {' '.join(cmd)}") result = subprocess.run( [sys.executable] + cmd, - stdout = SAFE_STDOUT, - stderr = subprocess.PIPE, - text = True, - encoding = "utf-8", + stdout=SAFE_STDOUT, + stderr=subprocess.PIPE, + text=True, + encoding="utf-8", ) if result.returncode != 0: logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}") @@ -37,7 +39,9 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b return True, "" -def second_elem_of(original_function: Callable[..., tuple[Any, Any]]) -> Callable[..., Any]: +def second_elem_of( + original_function: Callable[..., tuple[Any, Any]] +) -> Callable[..., Any]: """ 与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する。 diff --git a/tests/test_main.py b/tests/test_main.py index 0c0fe77..f00d908 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -5,15 +5,15 @@ from style_bert_vits2.constants import BASE_DIR, Languages from style_bert_vits2.tts_model import TTSModelHolder -def synthesize(device: str = 'cpu'): +def synthesize(device: str = "cpu"): # 音声合成モデルが配置されていれば、音声合成を実行 - model_holder = TTSModelHolder(BASE_DIR / 'model_assets', device) + model_holder = TTSModelHolder(BASE_DIR / "model_assets", device) if len(model_holder.models_info) > 0: # jvnv-F2-jp モデルを探す for model_info in model_holder.models_info: - if model_info.name == 'jvnv-F2-jp': + if model_info.name == "jvnv-F2-jp": # すべてのスタイルに対して音声合成を実行 for style in model_info.styles: @@ -23,21 +23,21 @@ def synthesize(device: str = 'cpu'): sample_rate, audio_data = model.infer( "あらゆる現実を、すべて自分のほうへねじ曲げたのだ。", # 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ) - language = Languages.JP, + language=Languages.JP, # 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0) - speaker_id = 0, + speaker_id=0, # 感情表現の強さ (0.0 〜 1.0) - sdp_ratio = 0.4, + sdp_ratio=0.4, # スタイル (Neutral, Happy など) - style = style, + style=style, # スタイルの強さ (0.0 〜 100.0) - style_weight = 6.0, + style_weight=6.0, ) # 音声データを保存 - (BASE_DIR / 'tests/wavs').mkdir(exist_ok=True, parents=True) - wav_file_path = BASE_DIR / f'tests/wavs/{style}.wav' - with open(wav_file_path, 'wb') as f: + (BASE_DIR / "tests/wavs").mkdir(exist_ok=True, parents=True) + wav_file_path = BASE_DIR / f"tests/wavs/{style}.wav" + with open(wav_file_path, "wb") as f: wavfile.write(f, sample_rate, audio_data) # 音声データが保存されたことを確認 @@ -48,8 +48,8 @@ def synthesize(device: str = 'cpu'): def test_synthesize_cpu(): - synthesize(device='cpu') + synthesize(device="cpu") def test_synthesize_cuda(): - synthesize(device='cuda') + synthesize(device="cuda") diff --git a/train_ms.py b/train_ms.py index 3d50c03..50a2453 100644 --- a/train_ms.py +++ b/train_ms.py @@ -281,28 +281,28 @@ def run(): mas_noise_scale_initial=mas_noise_scale_initial, noise_scale_delta=noise_scale_delta, # hps.model 以下のすべての値を引数に渡す - use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, - use_noise_scaled_mas = hps.model.use_noise_scaled_mas, - use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, - use_duration_discriminator = hps.model.use_duration_discriminator, - use_wavlm_discriminator = hps.model.use_wavlm_discriminator, - inter_channels = hps.model.inter_channels, - hidden_channels = hps.model.hidden_channels, - filter_channels = hps.model.filter_channels, - n_heads = hps.model.n_heads, - n_layers = hps.model.n_layers, - kernel_size = hps.model.kernel_size, - p_dropout = hps.model.p_dropout, - resblock = hps.model.resblock, - resblock_kernel_sizes = hps.model.resblock_kernel_sizes, - resblock_dilation_sizes = hps.model.resblock_dilation_sizes, - upsample_rates = hps.model.upsample_rates, - upsample_initial_channel = hps.model.upsample_initial_channel, - upsample_kernel_sizes = hps.model.upsample_kernel_sizes, - n_layers_q = hps.model.n_layers_q, - use_spectral_norm = hps.model.use_spectral_norm, - gin_channels = hps.model.gin_channels, - slm = hps.model.slm, + use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas=hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder, + use_duration_discriminator=hps.model.use_duration_discriminator, + use_wavlm_discriminator=hps.model.use_wavlm_discriminator, + inter_channels=hps.model.inter_channels, + hidden_channels=hps.model.hidden_channels, + filter_channels=hps.model.filter_channels, + n_heads=hps.model.n_heads, + n_layers=hps.model.n_layers, + kernel_size=hps.model.kernel_size, + p_dropout=hps.model.p_dropout, + resblock=hps.model.resblock, + resblock_kernel_sizes=hps.model.resblock_kernel_sizes, + resblock_dilation_sizes=hps.model.resblock_dilation_sizes, + upsample_rates=hps.model.upsample_rates, + upsample_initial_channel=hps.model.upsample_initial_channel, + upsample_kernel_sizes=hps.model.upsample_kernel_sizes, + n_layers_q=hps.model.n_layers_q, + use_spectral_norm=hps.model.use_spectral_norm, + gin_channels=hps.model.gin_channels, + slm=hps.model.slm, ).cuda(local_rank) if getattr(hps.train, "freeze_ZH_bert", False): @@ -389,7 +389,9 @@ def run(): epoch_str = max(epoch_str, 1) # global_step = (epoch_str - 1) * len(train_loader) global_step = int( - utils.get_steps(utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth")) + utils.get_steps( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth") + ) ) logger.info( f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************" diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 5ed68a1..321a040 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -288,28 +288,28 @@ def run(): mas_noise_scale_initial=mas_noise_scale_initial, noise_scale_delta=noise_scale_delta, # hps.model 以下のすべての値を引数に渡す - use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, - use_noise_scaled_mas = hps.model.use_noise_scaled_mas, - use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, - use_duration_discriminator = hps.model.use_duration_discriminator, - use_wavlm_discriminator = hps.model.use_wavlm_discriminator, - inter_channels = hps.model.inter_channels, - hidden_channels = hps.model.hidden_channels, - filter_channels = hps.model.filter_channels, - n_heads = hps.model.n_heads, - n_layers = hps.model.n_layers, - kernel_size = hps.model.kernel_size, - p_dropout = hps.model.p_dropout, - resblock = hps.model.resblock, - resblock_kernel_sizes = hps.model.resblock_kernel_sizes, - resblock_dilation_sizes = hps.model.resblock_dilation_sizes, - upsample_rates = hps.model.upsample_rates, - upsample_initial_channel = hps.model.upsample_initial_channel, - upsample_kernel_sizes = hps.model.upsample_kernel_sizes, - n_layers_q = hps.model.n_layers_q, - use_spectral_norm = hps.model.use_spectral_norm, - gin_channels = hps.model.gin_channels, - slm = hps.model.slm, + use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas=hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder, + use_duration_discriminator=hps.model.use_duration_discriminator, + use_wavlm_discriminator=hps.model.use_wavlm_discriminator, + inter_channels=hps.model.inter_channels, + hidden_channels=hps.model.hidden_channels, + filter_channels=hps.model.filter_channels, + n_heads=hps.model.n_heads, + n_layers=hps.model.n_layers, + kernel_size=hps.model.kernel_size, + p_dropout=hps.model.p_dropout, + resblock=hps.model.resblock, + resblock_kernel_sizes=hps.model.resblock_kernel_sizes, + resblock_dilation_sizes=hps.model.resblock_dilation_sizes, + upsample_rates=hps.model.upsample_rates, + upsample_initial_channel=hps.model.upsample_initial_channel, + upsample_kernel_sizes=hps.model.upsample_kernel_sizes, + n_layers_q=hps.model.n_layers_q, + use_spectral_norm=hps.model.use_spectral_norm, + gin_channels=hps.model.gin_channels, + slm=hps.model.slm, ).cuda(local_rank) if getattr(hps.train, "freeze_JP_bert", False): logger.info("Freezing (JP) bert encoder !!!") @@ -383,7 +383,9 @@ def run(): if net_dur_disc is not None: try: _, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( - utils.checkpoints.get_latest_checkpoint_path(model_dir, "DUR_*.pth"), + utils.checkpoints.get_latest_checkpoint_path( + model_dir, "DUR_*.pth" + ), net_dur_disc, optim_dur_disc, skip_optimizer=hps.train.skip_optimizer, @@ -396,11 +398,15 @@ def run(): print("Initialize dur_disc") if net_wd is not None: try: - _, optim_wd, wd_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( - utils.checkpoints.get_latest_checkpoint_path(model_dir, "WD_*.pth"), - net_wd, - optim_wd, - skip_optimizer=hps.train.skip_optimizer, + _, optim_wd, wd_resume_lr, epoch_str = ( + utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path( + model_dir, "WD_*.pth" + ), + net_wd, + optim_wd, + skip_optimizer=hps.train.skip_optimizer, + ) ) if not optim_wd.param_groups[0].get("initial_lr"): optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr @@ -430,7 +436,9 @@ def run(): epoch_str = max(epoch_str, 1) # global_step = (epoch_str - 1) * len(train_loader) global_step = int( - utils.get_steps(utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth")) + utils.get_steps( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth") + ) ) logger.info( f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"