Refactor: replace utils.HParams with HyperParameters Pydantic model

HyperParameters is largely a drop-in replacement for utils.HParams, which ensures type safety for hyper-parameters.
This commit is contained in:
tsukumi
2024-03-08 15:52:37 +00:00
parent 7f0b252806
commit a84783a6cc
11 changed files with 190 additions and 221 deletions

View File

@@ -8,7 +8,7 @@ from tqdm import tqdm
from config import config from config import config
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.models import commons 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.nlp import cleaned_text_to_sequence, extract_bert_feature from style_bert_vits2.nlp import cleaned_text_to_sequence, extract_bert_feature
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
@@ -62,7 +62,7 @@ if __name__ == "__main__":
) )
args, _ = parser.parse_known_args() args, _ = parser.parse_known_args()
config_path = args.config config_path = args.config
hps = utils.get_hparams_from_file(config_path) hps = HyperParameters.load_from_json(config_path)
lines = [] lines = []
with open(hps.data.training_files, encoding="utf-8") as f: with open(hps.data.training_files, encoding="utf-8") as f:
lines.extend(f.readlines()) lines.extend(f.readlines())

View File

@@ -11,6 +11,7 @@ from config import config
from mel_processing import mel_spectrogram_torch, spectrogram_torch from mel_processing import mel_spectrogram_torch, spectrogram_torch
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.models import commons from style_bert_vits2.models import commons
from style_bert_vits2.models.hyper_parameters import HyperParametersData
from style_bert_vits2.models.utils import load_filepaths_and_text, load_wav_to_torch from style_bert_vits2.models.utils import load_filepaths_and_text, load_wav_to_torch
from style_bert_vits2.nlp import cleaned_text_to_sequence from style_bert_vits2.nlp import cleaned_text_to_sequence
@@ -24,7 +25,7 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
3) computes spectrograms from audio files. 3) computes spectrograms from audio files.
""" """
def __init__(self, audiopaths_sid_text, hparams): def __init__(self, audiopaths_sid_text: str, hparams: HyperParametersData):
self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text) self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text)
self.max_wav_value = hparams.max_wav_value self.max_wav_value = hparams.max_wav_value
self.sampling_rate = hparams.sampling_rate self.sampling_rate = hparams.sampling_rate

View File

@@ -1,16 +1,16 @@
""" """
Style-Bert-VITS2 モデルのハイパーパラメータを表す Pydantic モデル。 Style-Bert-VITS2 モデルのハイパーパラメータを表す Pydantic モデル。
デフォルト値は configs/configs_jp_extra.json 内の定義と同一で、 デフォルト値は configs/configs_jp_extra.json 内の定義と概ね同一で、
万が一ロードした config.json に存在しないキーがあった際のフェイルセーフとして適用される。 万が一ロードした config.json に存在しないキーがあった際のフェイルセーフとして適用される。
""" """
from pathlib import Path from pathlib import Path
from typing import Optional, Union from typing import Optional, Union
from pydantic import BaseModel from pydantic import BaseModel, ConfigDict
class __HyperParametersTrain(BaseModel): class HyperParametersTrain(BaseModel):
log_interval: int = 200 log_interval: int = 200
eval_interval: int = 1000 eval_interval: int = 1000
seed: int = 42 seed: int = 42
@@ -36,7 +36,8 @@ class __HyperParametersTrain(BaseModel):
freeze_style: bool = False freeze_style: bool = False
freeze_decoder: bool = False freeze_decoder: bool = False
class __HyperParametersData(BaseModel):
class HyperParametersData(BaseModel):
use_jp_extra: bool = True use_jp_extra: bool = True
training_files: str = "Data/dummy/train.list" training_files: str = "Data/dummy/train.list"
validation_files: str = "Data/dummy/val.list" validation_files: str = "Data/dummy/val.list"
@@ -59,7 +60,8 @@ class __HyperParametersData(BaseModel):
"Neutral": 0, "Neutral": 0,
} }
class __HyperParametersModel(BaseModel):
class HyperParametersModel(BaseModel):
use_spk_conditioned_encoder: bool = True use_spk_conditioned_encoder: bool = True
use_noise_scaled_mas: bool = True use_noise_scaled_mas: bool = True
use_mel_posterior_encoder: bool = False use_mel_posterior_encoder: bool = False
@@ -93,12 +95,21 @@ class __HyperParametersModel(BaseModel):
"initial_channel": 64 "initial_channel": 64
} }
class HyperParameters(BaseModel): class HyperParameters(BaseModel):
version: str = "2.0-JP-Extra"
model_name: str = 'dummy' model_name: str = 'dummy'
train: __HyperParametersTrain version: str = "2.0-JP-Extra"
data: __HyperParametersData train: HyperParametersTrain
model: __HyperParametersModel data: HyperParametersData
model: HyperParametersModel
# 以下は学習時にのみ動的に設定されるパラメータ (通常 config.json には存在しない)
model_dir: Optional[str] = None
speedup: bool = False
repo_id: Optional[str] = None
# model_ 以下を Pydantic の保護対象から除外する
model_config = ConfigDict(protected_namespaces=())
@staticmethod @staticmethod
@@ -112,5 +123,6 @@ class HyperParameters(BaseModel):
Returns: Returns:
HyperParameters: ハイパーパラメータ HyperParameters: ハイパーパラメータ
""" """
with open(json_path, "r") as f: with open(json_path, "r") as f:
return HyperParameters.model_validate_json(f.read()) return HyperParameters.model_validate_json(f.read())

View File

@@ -1,34 +1,81 @@
from typing import Any, cast, Optional, Union
import torch import torch
from typing import Optional from numpy.typing import NDArray
from style_bert_vits2.constants import Languages from style_bert_vits2.constants import Languages
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.models import commons from style_bert_vits2.models import commons
from style_bert_vits2.models import utils 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 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.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature
from style_bert_vits2.nlp.symbols import SYMBOLS from style_bert_vits2.nlp.symbols import SYMBOLS
def get_net_g(model_path: str, version: str, device: str, hps): def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters):
if version.endswith("JP-Extra"): if version.endswith("JP-Extra"):
logger.info("Using JP-Extra model") logger.info("Using JP-Extra model")
net_g = SynthesizerTrnJPExtra( net_g = SynthesizerTrnJPExtra(
len(SYMBOLS), n_vocab = len(SYMBOLS),
hps.data.filter_length // 2 + 1, spec_channels = hps.data.filter_length // 2 + 1,
hps.train.segment_size // hps.data.hop_length, segment_size = hps.train.segment_size // hps.data.hop_length,
n_speakers=hps.data.n_speakers, n_speakers = hps.data.n_speakers,
**hps.model, # 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,
).to(device) ).to(device)
else: else:
logger.info("Using normal model") logger.info("Using normal model")
net_g = SynthesizerTrn( net_g = SynthesizerTrn(
len(SYMBOLS), n_vocab = len(SYMBOLS),
hps.data.filter_length // 2 + 1, spec_channels = hps.data.filter_length // 2 + 1,
hps.train.segment_size // hps.data.hop_length, segment_size = hps.train.segment_size // hps.data.hop_length,
n_speakers=hps.data.n_speakers, n_speakers=hps.data.n_speakers,
**hps.model, # 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,
).to(device) ).to(device)
net_g.state_dict() net_g.state_dict()
_ = net_g.eval() _ = net_g.eval()
@@ -44,7 +91,7 @@ def get_net_g(model_path: str, version: str, device: str, hps):
def get_text( def get_text(
text: str, text: str,
language_str: Languages, language_str: Languages,
hps, hps: HyperParameters,
device: str, device: str,
assist_text: Optional[str] = None, assist_text: Optional[str] = None,
assist_text_weight: float = 0.7, assist_text_weight: float = 0.7,
@@ -111,15 +158,15 @@ def get_text(
def infer( def infer(
text: str, text: str,
style_vec, style_vec: NDArray[Any],
sdp_ratio: float, sdp_ratio: float,
noise_scale: float, noise_scale: float,
noise_scale_w: float, noise_scale_w: float,
length_scale: float, length_scale: float,
sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id
language: Languages, language: Languages,
hps, hps: HyperParameters,
net_g, net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra],
device: str, device: str,
skip_start: bool = False, skip_start: bool = False,
skip_end: bool = False, skip_end: bool = False,
@@ -159,25 +206,25 @@ def infer(
ja_bert = ja_bert.to(device).unsqueeze(0) ja_bert = ja_bert.to(device).unsqueeze(0)
en_bert = en_bert.to(device).unsqueeze(0) en_bert = en_bert.to(device).unsqueeze(0)
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
style_vec = torch.from_numpy(style_vec).to(device).unsqueeze(0) style_vec_tensor = torch.from_numpy(style_vec).to(device).unsqueeze(0)
del phones del phones
sid_tensor = torch.LongTensor([sid]).to(device) sid_tensor = torch.LongTensor([sid]).to(device)
if is_jp_extra: if is_jp_extra:
output = net_g.infer( output = cast(SynthesizerTrnJPExtra, net_g).infer(
x_tst, x_tst,
x_tst_lengths, x_tst_lengths,
sid_tensor, sid_tensor,
tones, tones,
lang_ids, lang_ids,
ja_bert, ja_bert,
style_vec=style_vec, style_vec=style_vec_tensor,
sdp_ratio=sdp_ratio, sdp_ratio=sdp_ratio,
noise_scale=noise_scale, noise_scale=noise_scale,
noise_scale_w=noise_scale_w, noise_scale_w=noise_scale_w,
length_scale=length_scale, length_scale=length_scale,
) )
else: else:
output = net_g.infer( output = cast(SynthesizerTrn, net_g).infer(
x_tst, x_tst,
x_tst_lengths, x_tst_lengths,
sid_tensor, sid_tensor,
@@ -186,7 +233,7 @@ def infer(
bert, bert,
ja_bert, ja_bert,
en_bert, en_bert,
style_vec=style_vec, style_vec=style_vec_tensor,
sdp_ratio=sdp_ratio, sdp_ratio=sdp_ratio,
noise_scale=noise_scale, noise_scale=noise_scale,
noise_scale_w=noise_scale_w, noise_scale_w=noise_scale_w,
@@ -209,110 +256,5 @@ def infer(
return audio return audio
def infer_multilang(
text: str,
style_vec,
sdp_ratio: float,
noise_scale: float,
noise_scale_w: float,
length_scale: float,
sid: int,
language: Languages,
hps,
net_g,
device: str,
skip_start: bool = False,
skip_end: bool = False,
):
bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], []
# emo = get_emo_(reference_audio, emotion, sid)
# if isinstance(reference_audio, np.ndarray):
# emo = get_clap_audio_feature(reference_audio, device)
# else:
# emo = get_clap_text_feature(emotion, device)
# emo = torch.squeeze(emo, dim=1)
for idx, (txt, lang) in enumerate(zip(text, language)):
_skip_start = (idx != 0) or (skip_start and idx == 0)
_skip_end = (idx != len(language) - 1) or skip_end
(
temp_bert,
temp_ja_bert,
temp_en_bert,
temp_phones,
temp_tones,
temp_lang_ids,
) = get_text(txt, lang, hps, device) # type: ignore
if _skip_start:
temp_bert = temp_bert[:, 3:]
temp_ja_bert = temp_ja_bert[:, 3:]
temp_en_bert = temp_en_bert[:, 3:]
temp_phones = temp_phones[3:]
temp_tones = temp_tones[3:]
temp_lang_ids = temp_lang_ids[3:]
if _skip_end:
temp_bert = temp_bert[:, :-2]
temp_ja_bert = temp_ja_bert[:, :-2]
temp_en_bert = temp_en_bert[:, :-2]
temp_phones = temp_phones[:-2]
temp_tones = temp_tones[:-2]
temp_lang_ids = temp_lang_ids[:-2]
bert.append(temp_bert)
ja_bert.append(temp_ja_bert)
en_bert.append(temp_en_bert)
phones.append(temp_phones)
tones.append(temp_tones)
lang_ids.append(temp_lang_ids)
bert = torch.concatenate(bert, dim=1)
ja_bert = torch.concatenate(ja_bert, dim=1)
en_bert = torch.concatenate(en_bert, dim=1)
phones = torch.concatenate(phones, dim=0)
tones = torch.concatenate(tones, dim=0)
lang_ids = torch.concatenate(lang_ids, dim=0)
with torch.no_grad():
x_tst = phones.to(device).unsqueeze(0)
tones = tones.to(device).unsqueeze(0)
lang_ids = lang_ids.to(device).unsqueeze(0)
bert = bert.to(device).unsqueeze(0)
ja_bert = ja_bert.to(device).unsqueeze(0)
en_bert = en_bert.to(device).unsqueeze(0)
# emo = emo.to(device).unsqueeze(0)
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
del phones
speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device)
audio = (
net_g.infer(
x_tst,
x_tst_lengths,
speakers,
tones,
lang_ids,
bert,
ja_bert,
en_bert,
style_vec=style_vec,
sdp_ratio=sdp_ratio,
noise_scale=noise_scale,
noise_scale_w=noise_scale_w,
length_scale=length_scale,
)[0][0, 0]
.data.cpu()
.float()
.numpy()
)
del (
x_tst,
tones,
lang_ids,
bert,
x_tst_lengths,
speakers,
ja_bert,
en_bert,
) # , emo
if torch.cuda.is_available():
torch.cuda.empty_cache()
return audio
class InvalidToneError(ValueError): class InvalidToneError(ValueError):
pass pass

View File

@@ -983,10 +983,10 @@ class SynthesizerTrn(nn.Module):
en_bert, en_bert,
style_vec, style_vec,
noise_scale=0.667, noise_scale=0.667,
length_scale=1, length_scale=1.0,
noise_scale_w=0.8, noise_scale_w=0.8,
max_len=None, max_len=None,
sdp_ratio=0, sdp_ratio=0.0,
y=None, y=None,
): ):
# x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) # x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert)

View File

@@ -1029,10 +1029,10 @@ class SynthesizerTrn(nn.Module):
bert, bert,
style_vec, style_vec,
noise_scale=0.667, noise_scale=0.667,
length_scale=1, length_scale=1.0,
noise_scale_w=0.8, noise_scale_w=0.8,
max_len=None, max_len=None,
sdp_ratio=0, sdp_ratio=0.0,
y=None, y=None,
): ):
# x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) # x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert)

View File

@@ -355,45 +355,3 @@ def check_git_hash(model_dir):
) )
else: else:
open(path, "w").write(cur_hash) open(path, "w").write(cur_hash)
def get_hparams_from_file(config_path):
# print("config_path: ", config_path)
with open(config_path, "r", encoding="utf-8") as f:
data = f.read()
config = json.loads(data)
hparams = HParams(**config)
return hparams
class HParams:
def __init__(self, **kwargs):
for k, v in kwargs.items():
if type(v) == dict:
v = HParams(**v)
self[k] = v
def keys(self):
return self.__dict__.keys()
def items(self):
return self.__dict__.items()
def values(self):
return self.__dict__.values()
def __len__(self):
return len(self.__dict__)
def __getitem__(self, key):
return getattr(self, key)
def __setitem__(self, key, value):
return setattr(self, key, value)
def __contains__(self, key):
return key in self.__dict__
def __repr__(self):
return self.__dict__.__repr__()

View File

@@ -1,11 +1,12 @@
import warnings import warnings
from pathlib import Path from pathlib import Path
from typing import Optional, Union from typing import Any, Optional, Union
import gradio as gr import gradio as gr
import numpy as np import numpy as np
import torch import torch
from gradio.processing_utils import convert_to_16_bit_wav from gradio.processing_utils import convert_to_16_bit_wav
from numpy.typing import NDArray
from style_bert_vits2.constants import ( from style_bert_vits2.constants import (
DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_ASSIST_TEXT_WEIGHT,
@@ -17,15 +18,22 @@ from style_bert_vits2.constants import (
DEFAULT_SPLIT_INTERVAL, DEFAULT_SPLIT_INTERVAL,
DEFAULT_STYLE, DEFAULT_STYLE,
DEFAULT_STYLE_WEIGHT, DEFAULT_STYLE_WEIGHT,
Languages,
) )
from style_bert_vits2.models import utils 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.infer import get_net_g, infer
from style_bert_vits2.models.models import SynthesizerTrn 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.logging import logger
def adjust_voice(fs, wave, pitch_scale, intonation_scale): def adjust_voice(
fs: int,
wave: NDArray[Any],
pitch_scale: float,
intonation_scale: float,
) -> tuple[int, NDArray[Any]]:
if pitch_scale == 1.0 and intonation_scale == 1.0: if pitch_scale == 1.0 and intonation_scale == 1.0:
# 初期値の場合は、音質劣化を避けるためにそのまま返す # 初期値の場合は、音質劣化を避けるためにそのまま返す
return fs, wave return fs, wave
@@ -37,15 +45,17 @@ def adjust_voice(fs, wave, pitch_scale, intonation_scale):
"pyworld is not installed. Please install it by `pip install pyworld`" "pyworld is not installed. Please install it by `pip install pyworld`"
) )
# pyworldf0を加工して合成 # pyworldf0 を加工して合成
# pyworldよりもよいのがあるかもしれないが…… # pyworld よりもよいのがあるかもしれないが……
## pyworld は Cython で書かれているが、スタブファイルがないため型補完が全く効かない…
wave = wave.astype(np.double) wave = wave.astype(np.double)
f0, t = pyworld.harvest(wave, fs)
# 質が高そうだしとりあえずharvestにしておく
sp = pyworld.cheaptrick(wave, f0, t, fs) # 質が高そうだしとりあえずharvestにしておく
ap = pyworld.d4c(wave, f0, t, fs) f0, t = pyworld.harvest(wave, fs) # type: ignore
sp = pyworld.cheaptrick(wave, f0, t, fs) # type: ignore
ap = pyworld.d4c(wave, f0, t, fs) # type: ignore
non_zero_f0 = [f for f in f0 if f != 0] non_zero_f0 = [f for f in f0 if f != 0]
f0_mean = sum(non_zero_f0) / len(non_zero_f0) f0_mean = sum(non_zero_f0) / len(non_zero_f0)
@@ -55,7 +65,7 @@ def adjust_voice(fs, wave, pitch_scale, intonation_scale):
continue continue
f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean) f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean)
wave = pyworld.synthesize(f0, sp, ap, fs) wave = pyworld.synthesize(f0, sp, ap, fs) # type: ignore
return fs, wave return fs, wave
@@ -67,7 +77,7 @@ class Model:
self.config_path: Path = config_path self.config_path: Path = config_path
self.style_vec_path: Path = style_vec_path self.style_vec_path: Path = style_vec_path
self.device: str = device self.device: str = device
self.hps: utils.HParams = utils.get_hparams_from_file(self.config_path) self.hps: HyperParameters = HyperParameters.load_from_json(self.config_path)
self.spk2id: dict[str, int] = self.hps.data.spk2id self.spk2id: dict[str, int] = self.hps.data.spk2id
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()} self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
@@ -81,7 +91,7 @@ class Model:
f"Number of styles ({self.num_styles}) does not match the number of style2id ({len(self.style2id)})" f"Number of styles ({self.num_styles}) does not match the number of style2id ({len(self.style2id)})"
) )
self.style_vectors: np.ndarray = np.load(self.style_vec_path) self.style_vectors: NDArray[Any] = np.load(self.style_vec_path)
if self.style_vectors.shape[0] != self.num_styles: if self.style_vectors.shape[0] != self.num_styles:
raise ValueError( raise ValueError(
f"The number of styles ({self.num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})" f"The number of styles ({self.num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})"
@@ -97,7 +107,7 @@ class Model:
hps=self.hps, hps=self.hps,
) )
def get_style_vector(self, style_id: int, weight: float = 1.0) -> np.ndarray: def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
mean = self.style_vectors[0] mean = self.style_vectors[0]
style_vec = self.style_vectors[style_id] style_vec = self.style_vectors[style_id]
style_vec = mean + (style_vec - mean) * weight style_vec = mean + (style_vec - mean) * weight
@@ -105,7 +115,7 @@ class Model:
def get_style_vector_from_audio( def get_style_vector_from_audio(
self, audio_path: str, weight: float = 1.0 self, audio_path: str, weight: float = 1.0
) -> np.ndarray: ) -> NDArray[Any]:
from style_gen import get_style_vector from style_gen import get_style_vector
xvec = get_style_vector(audio_path) xvec = get_style_vector(audio_path)
@@ -116,7 +126,7 @@ class Model:
def infer( def infer(
self, self,
text: str, text: str,
language: str = "JP", language: Languages = Languages.JP,
sid: int = 0, sid: int = 0,
reference_audio_path: Optional[str] = None, reference_audio_path: Optional[str] = None,
sdp_ratio: float = DEFAULT_SDP_RATIO, sdp_ratio: float = DEFAULT_SDP_RATIO,
@@ -133,7 +143,7 @@ class Model:
given_tone: Optional[list[int]] = None, given_tone: Optional[list[int]] = None,
pitch_scale: float = 1.0, pitch_scale: float = 1.0,
intonation_scale: float = 1.0, intonation_scale: float = 1.0,
) -> tuple[int, np.ndarray]: ) -> tuple[int, NDArray[Any]]:
logger.info(f"Start generating audio data from text:\n{text}") logger.info(f"Start generating audio data from text:\n{text}")
if language != "JP" and self.hps.version.endswith("JP-Extra"): if language != "JP" and self.hps.version.endswith("JP-Extra"):
raise ValueError( raise ValueError(
@@ -146,6 +156,7 @@ class Model:
if self.net_g is None: if self.net_g is None:
self.load_net_g() self.load_net_g()
assert self.net_g is not None
if reference_audio_path is None: if reference_audio_path is None:
style_id = self.style2id[style] style_id = self.style2id[style]
style_vector = self.get_style_vector(style_id, style_weight) style_vector = self.get_style_vector(style_id, style_weight)
@@ -246,19 +257,17 @@ class ModelHolder:
continue continue
self.model_files_dict[model_dir.name] = model_files self.model_files_dict[model_dir.name] = model_files
self.model_names.append(model_dir.name) self.model_names.append(model_dir.name)
hps = utils.get_hparams_from_file(config_path) hps = HyperParameters.load_from_json(config_path)
style2id: dict[str, int] = hps.data.style2id style2id: dict[str, int] = hps.data.style2id
styles = list(style2id.keys()) styles = list(style2id.keys())
spk2id: dict[str, int] = hps.data.spk2id spk2id: dict[str, int] = hps.data.spk2id
speakers = list(spk2id.keys()) speakers = list(spk2id.keys())
self.models_info.append( self.models_info.append({
{ "name": model_dir.name,
"name": model_dir.name, "files": [str(f) for f in model_files],
"files": [str(f) for f in model_files], "styles": styles,
"styles": styles, "speakers": speakers,
"speakers": speakers, })
}
)
def load_model(self, model_name: str, model_path_str: str): def load_model(self, model_name: str, model_path_str: str):
model_path = Path(model_path_str) model_path = Path(model_path_str)
@@ -291,9 +300,9 @@ class ModelHolder:
speakers = list(self.current_model.spk2id.keys()) speakers = list(self.current_model.spk2id.keys())
styles = list(self.current_model.style2id.keys()) styles = list(self.current_model.style2id.keys())
return ( return (
gr.Dropdown(choices=styles, value=styles[0]), gr.Dropdown(choices=styles, value=styles[0]), # type: ignore
gr.Button(interactive=True, value="音声合成"), gr.Button(interactive=True, value="音声合成"),
gr.Dropdown(choices=speakers, value=speakers[0]), gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
) )
self.current_model = Model( self.current_model = Model(
model_path=model_path, model_path=model_path,
@@ -304,21 +313,21 @@ class ModelHolder:
speakers = list(self.current_model.spk2id.keys()) speakers = list(self.current_model.spk2id.keys())
styles = list(self.current_model.style2id.keys()) styles = list(self.current_model.style2id.keys())
return ( return (
gr.Dropdown(choices=styles, value=styles[0]), gr.Dropdown(choices=styles, value=styles[0]), # type: ignore
gr.Button(interactive=True, value="音声合成"), gr.Button(interactive=True, value="音声合成"),
gr.Dropdown(choices=speakers, value=speakers[0]), gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
) )
def update_model_files_gr(self, model_name: str) -> gr.Dropdown: def update_model_files_gr(self, model_name: str) -> gr.Dropdown:
model_files = self.model_files_dict[model_name] model_files = self.model_files_dict[model_name]
return gr.Dropdown(choices=model_files, value=model_files[0]) return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore
def update_model_names_gr(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: def update_model_names_gr(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]:
self.refresh() self.refresh()
initial_model_name = self.model_names[0] initial_model_name = self.model_names[0]
initial_model_files = self.model_files_dict[initial_model_name] initial_model_files = self.model_files_dict[initial_model_name]
return ( return (
gr.Dropdown(choices=self.model_names, value=initial_model_name), gr.Dropdown(choices=self.model_names, value=initial_model_name), # type: ignore
gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]), gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]), # type: ignore
gr.Button(interactive=False), # For tts_button gr.Button(interactive=False), # For tts_button
) )

View File

@@ -8,6 +8,7 @@ from tqdm import tqdm
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.models import utils from style_bert_vits2.models import utils
from style_bert_vits2.models.hyper_parameters import HyperParameters
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
from config import config from config import config
@@ -72,7 +73,7 @@ if __name__ == "__main__":
config_path = args.config config_path = args.config
num_processes = args.num_processes num_processes = args.num_processes
hps = utils.get_hparams_from_file(config_path) hps = HyperParameters.load_from_json(config_path)
device = config.style_gen_config.device device = config.style_gen_config.device

View File

@@ -26,6 +26,7 @@ from mel_processing import mel_spectrogram_torch, spec_to_mel_torch
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.models import commons from style_bert_vits2.models import commons
from style_bert_vits2.models import utils from style_bert_vits2.models import utils
from style_bert_vits2.models.hyper_parameters import HyperParameters
from style_bert_vits2.models.models import ( from style_bert_vits2.models.models import (
DurationDiscriminator, DurationDiscriminator,
MultiPeriodDiscriminator, MultiPeriodDiscriminator,
@@ -130,7 +131,7 @@ def run():
local_rank = int(os.environ["LOCAL_RANK"]) local_rank = int(os.environ["LOCAL_RANK"])
n_gpus = dist.get_world_size() n_gpus = dist.get_world_size()
hps = utils.get_hparams_from_file(args.config) hps = HyperParameters.load_from_json(args.config)
# This is needed because we have to pass values to `train_and_evaluate()` # This is needed because we have to pass values to `train_and_evaluate()`
hps.model_dir = model_dir hps.model_dir = model_dir
hps.speedup = args.speedup hps.speedup = args.speedup
@@ -288,7 +289,29 @@ def run():
n_speakers=hps.data.n_speakers, n_speakers=hps.data.n_speakers,
mas_noise_scale_initial=mas_noise_scale_initial, mas_noise_scale_initial=mas_noise_scale_initial,
noise_scale_delta=noise_scale_delta, noise_scale_delta=noise_scale_delta,
**hps.model, # 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,
).cuda(local_rank) ).cuda(local_rank)
if getattr(hps.train, "freeze_ZH_bert", False): if getattr(hps.train, "freeze_ZH_bert", False):
@@ -547,7 +570,7 @@ def train_and_evaluate(
rank, rank,
local_rank, local_rank,
epoch, epoch,
hps, hps: HyperParameters,
nets, nets,
optims, optims,
schedulers, schedulers,

View File

@@ -26,6 +26,7 @@ from mel_processing import mel_spectrogram_torch, spec_to_mel_torch
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.models import commons from style_bert_vits2.models import commons
from style_bert_vits2.models import utils from style_bert_vits2.models import utils
from style_bert_vits2.models.hyper_parameters import HyperParameters
from style_bert_vits2.models.models_jp_extra import ( from style_bert_vits2.models.models_jp_extra import (
DurationDiscriminator, DurationDiscriminator,
MultiPeriodDiscriminator, MultiPeriodDiscriminator,
@@ -129,7 +130,7 @@ def run():
local_rank = int(os.environ["LOCAL_RANK"]) local_rank = int(os.environ["LOCAL_RANK"])
n_gpus = dist.get_world_size() n_gpus = dist.get_world_size()
hps = utils.get_hparams_from_file(args.config) hps = HyperParameters.load_from_json(args.config)
# This is needed because we have to pass values to `train_and_evaluate() # This is needed because we have to pass values to `train_and_evaluate()
hps.model_dir = model_dir hps.model_dir = model_dir
hps.speedup = args.speedup hps.speedup = args.speedup
@@ -298,7 +299,29 @@ def run():
n_speakers=hps.data.n_speakers, n_speakers=hps.data.n_speakers,
mas_noise_scale_initial=mas_noise_scale_initial, mas_noise_scale_initial=mas_noise_scale_initial,
noise_scale_delta=noise_scale_delta, noise_scale_delta=noise_scale_delta,
**hps.model, # 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,
).cuda(local_rank) ).cuda(local_rank)
if getattr(hps.train, "freeze_JP_bert", False): if getattr(hps.train, "freeze_JP_bert", False):
logger.info("Freezing (JP) bert encoder !!!") logger.info("Freezing (JP) bert encoder !!!")