Apply black formatter
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user