Apply black formatter

This commit is contained in:
litagin02
2024-03-11 09:47:47 +09:00
parent 42ee7d7608
commit c776c08235
31 changed files with 463 additions and 298 deletions

View File

@@ -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):