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

@@ -19,4 +19,8 @@
"reportUnusedFunction": "none", "reportUnusedFunction": "none",
"reportUnusedVariable": "information", "reportUnusedVariable": "information",
}, },
"[python]": {
"editor.defaultFormatter": "ms-python.black-formatter",
"editor.formatOnType": true,
},
} }

7
app.py
View File

@@ -51,4 +51,9 @@ with gr.Blocks(theme=GRADIO_THEME) as app:
create_merge_app(model_holder=model_holder) 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,
)

View File

@@ -9,6 +9,7 @@ VERSION = "2.4"
# Style-Bert-VITS2 のベースディレクトリ # Style-Bert-VITS2 のベースディレクトリ
BASE_DIR = Path(__file__).parent.parent BASE_DIR = Path(__file__).parent.parent
# 利用可能な言語 # 利用可能な言語
## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない ## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない
class Languages(StrEnum): class Languages(StrEnum):
@@ -16,6 +17,7 @@ class Languages(StrEnum):
EN = "EN" EN = "EN"
ZH = "ZH" ZH = "ZH"
# 言語ごとのデフォルトの BERT トークナイザーのパス # 言語ごとのデフォルトの BERT トークナイザーのパス
DEFAULT_BERT_TOKENIZER_PATHS = { DEFAULT_BERT_TOKENIZER_PATHS = {
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm", Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",

View File

@@ -9,7 +9,7 @@ logger.remove()
# Add a new handler # Add a new handler
logger.add( logger.add(
SAFE_STDOUT, SAFE_STDOUT,
format = "<g>{time:MM-DD HH:mm:ss}</g> |<lvl>{level:^8}</lvl>| {file}:{line} | {message}", format="<g>{time:MM-DD HH:mm:ss}</g> |<lvl>{level:^8}</lvl>| {file}:{line} | {message}",
backtrace = True, backtrace=True,
diagnose = True, diagnose=True,
) )

View File

@@ -24,7 +24,9 @@ class LayerNorm(nn.Module):
@torch.jit.script # type: ignore @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] n_channels_int = n_channels[0]
in_act = input_a + input_b in_act = input_a + input_b
t_act = torch.tanh(in_act[:, :n_channels_int, :]) t_act = torch.tanh(in_act[:, :n_channels_int, :])
@@ -44,7 +46,7 @@ class Encoder(nn.Module):
p_dropout: float = 0.0, p_dropout: float = 0.0,
window_size: int = 4, window_size: int = 4,
isflow: bool = True, isflow: bool = True,
**kwargs: Any **kwargs: Any,
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_channels = hidden_channels self.hidden_channels = hidden_channels
@@ -99,7 +101,9 @@ class Encoder(nn.Module):
) )
self.norm_layers_2.append(LayerNorm(hidden_channels)) 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) attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
x = x * x_mask x = x * x_mask
for i in range(self.n_layers): for i in range(self.n_layers):
@@ -131,7 +135,7 @@ class Decoder(nn.Module):
p_dropout: float = 0.0, p_dropout: float = 0.0,
proximal_bias: bool = False, proximal_bias: bool = False,
proximal_init: bool = True, proximal_init: bool = True,
**kwargs: Any **kwargs: Any,
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_channels = hidden_channels self.hidden_channels = hidden_channels
@@ -180,7 +184,13 @@ class Decoder(nn.Module):
) )
self.norm_layers_2.append(LayerNorm(hidden_channels)) 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 x: decoder input
h: encoder output h: encoder output
@@ -262,7 +272,9 @@ class MultiHeadAttention(nn.Module):
assert self.conv_q.bias is not None assert self.conv_q.bias is not None
self.conv_k.bias.copy_(self.conv_q.bias) 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) q = self.conv_q(x)
k = self.conv_k(c) k = self.conv_k(c)
v = self.conv_v(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] ) # [b, n_h, t_t, d_k] -> [b, d, t_t]
return output, p_attn 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] x: [b, h, l, m]
y: [h or 1, m, d] y: [h or 1, m, d]
@@ -338,7 +352,9 @@ class MultiHeadAttention(nn.Module):
ret = torch.matmul(x, y.unsqueeze(0)) ret = torch.matmul(x, y.unsqueeze(0))
return ret 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] x: [b, h, l, d]
y: [h or 1, m, 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)) ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
return ret 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 assert self.window_size is not None
2 * self.window_size + 1 # type: ignore 2 * self.window_size + 1 # type: ignore
# Pad first before slice to avoid using cond ops. # Pad first before slice to avoid using cond ops.

View File

@@ -67,7 +67,9 @@ def intersperse(lst: list[Any], item: Any) -> list[Any]:
return result 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) 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 @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 の活性化関数を組み合わせた演算を行う 加算、tanh、sigmoid の活性化関数を組み合わせた演算を行う
@@ -141,7 +147,9 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor
return acts 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 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:
""" """
勾配の値をクリップする 勾配の値をクリップする

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 import utils
from style_bert_vits2.models.hyper_parameters import HyperParameters 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 (
from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature 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 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"): if version.endswith("JP-Extra"):
logger.info("Using JP-Extra model") logger.info("Using JP-Extra model")
net_g = SynthesizerTrnJPExtra( net_g = SynthesizerTrnJPExtra(
n_vocab = len(SYMBOLS), n_vocab=len(SYMBOLS),
spec_channels = hps.data.filter_length // 2 + 1, spec_channels=hps.data.filter_length // 2 + 1,
segment_size = 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_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas = hps.model.use_noise_scaled_mas, use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
use_duration_discriminator = hps.model.use_duration_discriminator, use_duration_discriminator=hps.model.use_duration_discriminator,
use_wavlm_discriminator = hps.model.use_wavlm_discriminator, use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
inter_channels = hps.model.inter_channels, inter_channels=hps.model.inter_channels,
hidden_channels = hps.model.hidden_channels, hidden_channels=hps.model.hidden_channels,
filter_channels = hps.model.filter_channels, filter_channels=hps.model.filter_channels,
n_heads = hps.model.n_heads, n_heads=hps.model.n_heads,
n_layers = hps.model.n_layers, n_layers=hps.model.n_layers,
kernel_size = hps.model.kernel_size, kernel_size=hps.model.kernel_size,
p_dropout = hps.model.p_dropout, p_dropout=hps.model.p_dropout,
resblock = hps.model.resblock, resblock=hps.model.resblock,
resblock_kernel_sizes = hps.model.resblock_kernel_sizes, resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
resblock_dilation_sizes = hps.model.resblock_dilation_sizes, resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
upsample_rates = hps.model.upsample_rates, upsample_rates=hps.model.upsample_rates,
upsample_initial_channel = hps.model.upsample_initial_channel, upsample_initial_channel=hps.model.upsample_initial_channel,
upsample_kernel_sizes = hps.model.upsample_kernel_sizes, upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
n_layers_q = hps.model.n_layers_q, n_layers_q=hps.model.n_layers_q,
use_spectral_norm = hps.model.use_spectral_norm, use_spectral_norm=hps.model.use_spectral_norm,
gin_channels = hps.model.gin_channels, gin_channels=hps.model.gin_channels,
slm = hps.model.slm, 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(
n_vocab = len(SYMBOLS), n_vocab=len(SYMBOLS),
spec_channels = hps.data.filter_length // 2 + 1, spec_channels=hps.data.filter_length // 2 + 1,
segment_size = 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_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas = hps.model.use_noise_scaled_mas, use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
use_duration_discriminator = hps.model.use_duration_discriminator, use_duration_discriminator=hps.model.use_duration_discriminator,
use_wavlm_discriminator = hps.model.use_wavlm_discriminator, use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
inter_channels = hps.model.inter_channels, inter_channels=hps.model.inter_channels,
hidden_channels = hps.model.hidden_channels, hidden_channels=hps.model.hidden_channels,
filter_channels = hps.model.filter_channels, filter_channels=hps.model.filter_channels,
n_heads = hps.model.n_heads, n_heads=hps.model.n_heads,
n_layers = hps.model.n_layers, n_layers=hps.model.n_layers,
kernel_size = hps.model.kernel_size, kernel_size=hps.model.kernel_size,
p_dropout = hps.model.p_dropout, p_dropout=hps.model.p_dropout,
resblock = hps.model.resblock, resblock=hps.model.resblock,
resblock_kernel_sizes = hps.model.resblock_kernel_sizes, resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
resblock_dilation_sizes = hps.model.resblock_dilation_sizes, resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
upsample_rates = hps.model.upsample_rates, upsample_rates=hps.model.upsample_rates,
upsample_initial_channel = hps.model.upsample_initial_channel, upsample_initial_channel=hps.model.upsample_initial_channel,
upsample_kernel_sizes = hps.model.upsample_kernel_sizes, upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
n_layers_q = hps.model.n_layers_q, n_layers_q=hps.model.n_layers_q,
use_spectral_norm = hps.model.use_spectral_norm, use_spectral_norm=hps.model.use_spectral_norm,
gin_channels = hps.model.gin_channels, gin_channels=hps.model.gin_channels,
slm = hps.model.slm, slm=hps.model.slm,
).to(device) ).to(device)
net_g.state_dict() net_g.state_dict()
_ = net_g.eval() _ = net_g.eval()
if model_path.endswith(".pth") or model_path.endswith(".pt"): 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"): elif model_path.endswith(".safetensors"):
_ = utils.safetensors.load_safetensors(model_path, net_g, True) _ = utils.safetensors.load_safetensors(model_path, net_g, True)
else: else:
@@ -102,8 +110,8 @@ def get_text(
norm_text, phone, tone, word2ph = clean_text( norm_text, phone, tone, word2ph = clean_text(
text, text,
language_str, language_str,
use_jp_extra = use_jp_extra, use_jp_extra=use_jp_extra,
raise_yomi_error = False, raise_yomi_error=False,
) )
if given_tone is not None: if given_tone is not None:
if len(given_tone) != len(phone): if len(given_tone) != len(phone):

View File

@@ -21,7 +21,7 @@ class DurationDiscriminator(nn.Module): # vits2
filter_channels: int, filter_channels: int,
kernel_size: int, kernel_size: int,
p_dropout: float, p_dropout: float,
gin_channels: int = 0 gin_channels: int = 0,
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -330,7 +330,9 @@ class DurationPredictor(nn.Module):
if gin_channels != 0: if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, in_channels, 1) 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) x = torch.detach(x)
if g is not None: if g is not None:
g = torch.detach(g) g = torch.detach(g)
@@ -582,7 +584,9 @@ class Generator(torch.nn.Module):
if gin_channels != 0: if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) 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) x = self.conv_pre(x)
if g is not None: if g is not None:
x = x + self.cond(g) x = x + self.cond(g)
@@ -613,7 +617,13 @@ class Generator(torch.nn.Module):
class DiscriminatorP(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__() super(DiscriminatorP, self).__init__()
self.period = period self.period = period
self.use_spectral_norm = use_spectral_norm self.use_spectral_norm = use_spectral_norm
@@ -736,7 +746,9 @@ class MultiPeriodDiscriminator(torch.nn.Module):
self, self,
y: torch.Tensor, y: torch.Tensor,
y_hat: 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_rs = []
y_d_gs = [] y_d_gs = []
fmap_rs = [] fmap_rs = []
@@ -787,7 +799,9 @@ class ReferenceEncoder(nn.Module):
) )
self.proj = nn.Linear(128, gin_channels) 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) N = inputs.size(0)
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs] out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
for conv in self.convs: for conv in self.convs:
@@ -805,7 +819,9 @@ class ReferenceEncoder(nn.Module):
return self.proj(out.squeeze(0)) 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): for i in range(n_convs):
L = (L - kernel_size + 2 * pad) // stride + 1 L = (L - kernel_size + 2 * pad) // stride + 1
return L return L

View File

@@ -21,7 +21,7 @@ class DurationDiscriminator(nn.Module): # vits2
filter_channels: int, filter_channels: int,
kernel_size: int, kernel_size: int,
p_dropout: float, p_dropout: float,
gin_channels: int = 0 gin_channels: int = 0,
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -313,7 +313,9 @@ class DurationPredictor(nn.Module):
if gin_channels != 0: if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, in_channels, 1) 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) x = torch.detach(x)
if g is not None: if g is not None:
g = torch.detach(g) g = torch.detach(g)
@@ -587,7 +589,9 @@ class Generator(torch.nn.Module):
if gin_channels != 0: if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) 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) x = self.conv_pre(x)
if g is not None: if g is not None:
x = x + self.cond(g) x = x + self.cond(g)
@@ -618,7 +622,13 @@ class Generator(torch.nn.Module):
class DiscriminatorP(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__() super(DiscriminatorP, self).__init__()
self.period = period self.period = period
self.use_spectral_norm = use_spectral_norm self.use_spectral_norm = use_spectral_norm
@@ -741,7 +751,9 @@ class MultiPeriodDiscriminator(torch.nn.Module):
self, self,
y: torch.Tensor, y: torch.Tensor,
y_hat: 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_rs = []
y_d_gs = [] y_d_gs = []
fmap_rs = [] fmap_rs = []
@@ -845,7 +857,9 @@ class ReferenceEncoder(nn.Module):
) )
self.proj = nn.Linear(128, gin_channels) 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) N = inputs.size(0)
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs] out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
for conv in self.convs: for conv in self.convs:
@@ -863,7 +877,9 @@ class ReferenceEncoder(nn.Module):
return self.proj(out.squeeze(0)) 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): for i in range(n_convs):
L = (L - kernel_size + 2 * pad) // stride + 1 L = (L - kernel_size + 2 * pad) // stride + 1
return L return L

View File

@@ -88,7 +88,9 @@ class DDSConv(nn.Module):
Dialted and Depth-Separable Convolution 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__() super().__init__()
self.channels = channels self.channels = channels
self.kernel_size = kernel_size self.kernel_size = kernel_size
@@ -117,7 +119,9 @@ class DDSConv(nn.Module):
self.norms_1.append(LayerNorm(channels)) self.norms_1.append(LayerNorm(channels))
self.norms_2.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: if g is not None:
x = x + g x = x + g
for i in range(self.n_layers): 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") res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight")
self.res_skip_layers.append(res_skip_layer) 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) output = torch.zeros_like(x)
n_channels_tensor = torch.IntTensor([self.hidden_channels]) n_channels_tensor = torch.IntTensor([self.hidden_channels])
@@ -221,7 +231,12 @@ class WN(torch.nn.Module):
class ResBlock1(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__() super(ResBlock1, self).__init__()
self.convs1 = nn.ModuleList( self.convs1 = nn.ModuleList(
[ [
@@ -295,7 +310,9 @@ class ResBlock1(torch.nn.Module):
) )
self.convs2.apply(commons.init_weights) 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): for c1, c2 in zip(self.convs1, self.convs2):
xt = F.leaky_relu(x, LRELU_SLOPE) xt = F.leaky_relu(x, LRELU_SLOPE)
if x_mask is not None: if x_mask is not None:
@@ -318,7 +335,9 @@ class ResBlock1(torch.nn.Module):
class ResBlock2(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__() super(ResBlock2, self).__init__()
self.convs = nn.ModuleList( self.convs = nn.ModuleList(
[ [
@@ -346,7 +365,9 @@ class ResBlock2(torch.nn.Module):
) )
self.convs.apply(commons.init_weights) 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: for c in self.convs:
xt = F.leaky_relu(x, LRELU_SLOPE) xt = F.leaky_relu(x, LRELU_SLOPE)
if x_mask is not None: if x_mask is not None:

View File

@@ -40,8 +40,8 @@ def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
numba.int32[::1], numba.int32[::1],
numba.int32[::1], numba.int32[::1],
), ),
nopython = True, nopython=True,
nogil = True, nogil=True,
) # type: ignore ) # type: ignore
def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None: def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None:
""" """

View File

@@ -39,12 +39,14 @@ def piecewise_rational_quadratic_transform(
min_bin_width=min_bin_width, min_bin_width=min_bin_width,
min_bin_height=min_bin_height, min_bin_height=min_bin_height,
min_derivative=min_derivative, min_derivative=min_derivative,
**spline_kwargs # type: ignore **spline_kwargs, # type: ignore
) )
return outputs, logabsdet 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 bin_locations[..., -1] += eps
return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1 return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1

View File

@@ -107,7 +107,9 @@ def plot_spectrogram_to_numpy(spectrogram: NDArray[Any]) -> NDArray[Any]:
return data 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 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 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:
""" """
ロガーを取得する ロガーを取得する

View File

@@ -14,7 +14,7 @@ def load_checkpoint(
model: torch.nn.Module, model: torch.nn.Module,
optimizer: Optional[torch.optim.Optimizer] = None, optimizer: Optional[torch.optim.Optimizer] = None,
skip_optimizer: bool = False, skip_optimizer: bool = False,
for_infer: bool = False for_infer: bool = False,
) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]: ) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]:
""" """
指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。 指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。
@@ -107,7 +107,9 @@ def save_checkpoint(
iteration (int): イテレーション数 iteration (int): イテレーション数
checkpoint_path (Union[str, Path]): 保存先のパス 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"): if hasattr(model, "module"):
state_dict = model.module.state_dict() state_dict = model.module.state_dict()
else: 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] [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:
""" """
指定されたディレクトリから最新のチェックポイントのパスを取得する 指定されたディレクトリから最新のチェックポイントのパスを取得する

View File

@@ -74,16 +74,19 @@ def clean_text(
if language == Languages.JP: if language == Languages.JP:
from style_bert_vits2.nlp.japanese.g2p import g2p from style_bert_vits2.nlp.japanese.g2p import g2p
from style_bert_vits2.nlp.japanese.normalizer import normalize_text from style_bert_vits2.nlp.japanese.normalizer import normalize_text
norm_text = normalize_text(text) norm_text = normalize_text(text)
phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error) phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error)
elif language == Languages.EN: elif language == Languages.EN:
from style_bert_vits2.nlp.english.g2p import g2p from style_bert_vits2.nlp.english.g2p import g2p
from style_bert_vits2.nlp.english.normalizer import normalize_text from style_bert_vits2.nlp.english.normalizer import normalize_text
norm_text = normalize_text(text) norm_text = normalize_text(text)
phones, tones, word2ph = g2p(norm_text) phones, tones, word2ph = g2p(norm_text)
elif language == Languages.ZH: elif language == Languages.ZH:
from style_bert_vits2.nlp.chinese.g2p import g2p from style_bert_vits2.nlp.chinese.g2p import g2p
from style_bert_vits2.nlp.chinese.normalizer import normalize_text from style_bert_vits2.nlp.chinese.normalizer import normalize_text
norm_text = normalize_text(text) norm_text = normalize_text(text)
phones, tones, word2ph = g2p(norm_text) phones, tones, word2ph = g2p(norm_text)
else: else:
@@ -92,7 +95,9 @@ def clean_text(
return norm_text, phones, tones, word2ph 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 に変換する テキスト文字列を、テキスト内の記号に対応する一連の ID に変換する

View File

@@ -30,7 +30,9 @@ from style_bert_vits2.logging import logger
__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {} __loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書 # 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
__loaded_tokenizers: dict[Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]] = {} __loaded_tokenizers: dict[
Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]
] = {}
def load_model( def load_model(
@@ -63,18 +65,24 @@ def load_model(
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
if pretrained_model_name_or_path is None: if pretrained_model_name_or_path is None:
assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \ assert DEFAULT_BERT_TOKENIZER_PATHS[
f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model." 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]) pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
# BERT モデルをロードし、辞書に格納して返す # BERT モデルをロードし、辞書に格納して返す
## 英語のみ DebertaV2Model でロードする必要がある ## 英語のみ DebertaV2Model でロードする必要がある
if language == Languages.EN: 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: else:
model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path) model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path)
__loaded_models[language] = model __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 return model
@@ -109,8 +117,9 @@ def load_tokenizer(
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
if pretrained_model_name_or_path is None: if pretrained_model_name_or_path is None:
assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \ assert DEFAULT_BERT_TOKENIZER_PATHS[
f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." 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]) pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
# BERT トークナイザーをロードし、辞書に格納して返す # BERT トークナイザーをロードし、辞書に格納して返す
@@ -120,7 +129,9 @@ def load_tokenizer(
else: else:
tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path) tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path)
__loaded_tokenizers[language] = tokenizer __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 return tokenizer

View File

@@ -95,7 +95,11 @@ def __g2p(segments: list[str]) -> tuple[list[str], list[int], list[int]]:
if pinyin[0] in single_rep_map.keys(): if pinyin[0] in single_rep_map.keys():
pinyin = single_rep_map[pinyin[0]] + pinyin[1:] 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(" ") phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ")
word2ph.append(len(phone)) word2ph.append(len(phone))
@@ -125,7 +129,7 @@ if __name__ == "__main__":
text = normalize_text(text) text = normalize_text(text)
print(text) print(text)
phones, tones, word2ph = g2p(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) print(phones, tones, word2ph, bert.shape)

View File

@@ -121,7 +121,9 @@ def __expand_number(m: re.Match[str]) -> str:
else: else:
return __INFLECT.number_to_words( return __INFLECT.number_to_words(
num, andword="", zero="oh", group=2 # type: ignore num, andword="", zero="oh", group=2 # type: ignore
).replace(", ", " ") # type: ignore ).replace(
", ", " "
) # type: ignore
else: else:
return __INFLECT.number_to_words(num, andword="") # type: ignore return __INFLECT.number_to_words(num, andword="") # type: ignore

View File

@@ -10,9 +10,7 @@ from style_bert_vits2.nlp.symbols import PUNCTUATIONS
def g2p( def g2p(
norm_text: str, norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False
use_jp_extra: bool = True,
raise_yomi_error: bool = False
) -> tuple[list[str], list[int], list[int]]: ) -> tuple[list[str], list[int], list[int]]:
""" """
他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、 他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、
@@ -93,8 +91,7 @@ def g2p(
def text_to_sep_kata( def text_to_sep_kata(
norm_text: str, norm_text: str, raise_yomi_error: bool = False
raise_yomi_error: bool = False
) -> tuple[list[str], list[str]]: ) -> tuple[list[str], list[str]]:
""" """
`normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、 `normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、
@@ -212,7 +209,9 @@ def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]:
return result 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」なことに注意。 ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。
ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py 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( def __align_tones(
phones_with_punct: list[str], phones_with_punct: list[str], phone_tone_list: list[tuple[str, int]]
phone_tone_list: list[tuple[str, int]]
) -> list[tuple[str, int]]: ) -> list[tuple[str, int]]:
""" """
例: …私は、、そう思う。 例: …私は、、そう思う。

View File

@@ -34,11 +34,13 @@ def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, i
""" """
# 子音の集合 # 子音の集合
CONSONANTS = set([ CONSONANTS = set(
[
consonant consonant
for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values() for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values()
if consonant is not None if consonant is not None
]) ]
)
phone_tone = phone_tone[1:] # 最初の("_", 0)を無視 phone_tone = phone_tone[1:] # 最初の("_", 0)を無視
phones = [phone for phone, _ in phone_tone] phones = [phone for phone, _ in phone_tone]

View File

@@ -25,6 +25,7 @@ def run_frontend(text: str) -> list[dict[str, Any]]:
else: else:
# without worker # without worker
import pyopenjtalk import pyopenjtalk
return pyopenjtalk.run_frontend(text) return pyopenjtalk.run_frontend(text)
@@ -36,6 +37,7 @@ def make_label(njd_features: Any) -> list[str]:
else: else:
# without worker # without worker
import pyopenjtalk import pyopenjtalk
return pyopenjtalk.make_label(njd_features) 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: else:
# without worker # without worker
import pyopenjtalk import pyopenjtalk
pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab) pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab)
@@ -54,6 +57,7 @@ def update_global_jtalk_with_user_dict(path: str) -> None:
else: else:
# without worker # without worker
import pyopenjtalk import pyopenjtalk
pyopenjtalk.update_global_jtalk_with_user_dict(path) pyopenjtalk.update_global_jtalk_with_user_dict(path)
@@ -63,6 +67,7 @@ def unset_user_dict() -> None:
else: else:
# without worker # without worker
import pyopenjtalk import pyopenjtalk
pyopenjtalk.unset_user_dict() pyopenjtalk.unset_user_dict()
@@ -102,7 +107,12 @@ def initialize_worker(port: int = WORKER_PORT) -> None:
else: else:
# align with Windows behavior # align with Windows behavior
# start_new_session is same as specifying setsid in preexec_fn # 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 # wait until server listening
count = 0 count = 0

View File

@@ -2,12 +2,15 @@ import socket
from typing import Any, cast from typing import Any, cast
from style_bert_vits2.logging import logger 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: class WorkerClient:
""" pyopenjtalk worker client """ """pyopenjtalk worker client"""
def __init__(self, port: int) -> None: def __init__(self, port: int) -> None:
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
@@ -16,19 +19,15 @@ class WorkerClient:
sock.connect((socket.gethostname(), port)) sock.connect((socket.gethostname(), port))
self.sock = sock self.sock = sock
def __enter__(self) -> "WorkerClient": def __enter__(self) -> "WorkerClient":
return self return self
def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None: def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
self.close() self.close()
def close(self) -> None: def close(self) -> None:
self.sock.close() self.sock.close()
def dispatch_pyopenjtalk(self, func: str, *args: Any, **kwargs: Any) -> Any: def dispatch_pyopenjtalk(self, func: str, *args: Any, **kwargs: Any) -> Any:
data = { data = {
"request-type": RequestType.PYOPENJTALK, "request-type": RequestType.PYOPENJTALK,
@@ -43,7 +42,6 @@ class WorkerClient:
logger.trace(f"client received response: {response}") logger.trace(f"client received response: {response}")
return response.get("return") return response.get("return")
def status(self) -> int: def status(self) -> int:
data = {"request-type": RequestType.STATUS} data = {"request-type": RequestType.STATUS}
logger.trace(f"client sends request: {data}") logger.trace(f"client sends request: {data}")
@@ -53,7 +51,6 @@ class WorkerClient:
logger.trace(f"client received response: {response}") logger.trace(f"client received response: {response}")
return cast(int, response.get("client-count")) return cast(int, response.get("client-count"))
def quit_server(self) -> None: def quit_server(self) -> None:
data = {"request-type": RequestType.QUIT_SERVER} data = {"request-type": RequestType.QUIT_SERVER}
logger.trace(f"client sends request: {data}") logger.trace(f"client sends request: {data}")

View File

@@ -26,14 +26,12 @@ PYOPENJTALK_FUNC_DICT = {
class WorkerServer: class WorkerServer:
""" pyopenjtalk worker server """ """pyopenjtalk worker server"""
def __init__(self) -> None: def __init__(self) -> None:
self.client_count: int = 0 self.client_count: int = 0
self.quit: bool = False self.quit: bool = False
def handle_request(self, request: dict[str, Any]) -> dict[str, Any]: def handle_request(self, request: dict[str, Any]) -> dict[str, Any]:
request_type = None request_type = None
try: try:
@@ -70,7 +68,6 @@ class WorkerServer:
return response return response
def start_server(self, port: int, no_client_timeout: int = 30) -> None: def start_server(self, port: int, no_client_timeout: int = 30) -> None:
logger.info("start pyopenjtalk worker server") logger.info("start pyopenjtalk worker server")
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket: with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket:

View File

@@ -18,7 +18,11 @@ from fastapi import HTTPException
from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR 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 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.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() # root_dir = engine_root()
# save_dir = get_save_dir() # 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(): # if not save_dir.is_dir():
# save_dir.mkdir(parents=True) # 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" # ユーザー辞書ファイルのパス 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"
) # コンパイル済み辞書ファイルのパス
# # 同時書き込みの制御 # # 同時書き込みの制御

View File

@@ -25,7 +25,9 @@ from style_bert_vits2.constants import (
from style_bert_vits2.models.hyper_parameters import HyperParameters 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
from style_bert_vits2.voice import adjust_voice from style_bert_vits2.voice import adjust_voice
@@ -36,7 +38,6 @@ class TTSModel:
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。 モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。
""" """
def __init__( def __init__(
self, self,
model_path: Path, model_path: Path,
@@ -59,7 +60,9 @@ class TTSModel:
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.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.spk2id: dict[str, int] = self.hyper_parameters.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()}
@@ -82,19 +85,17 @@ class TTSModel:
self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
def load(self) -> None: def load(self) -> None:
""" """
音声合成モデルをデバイスにロードする。 音声合成モデルをデバイスにロードする。
""" """
self.__net_g = get_net_g( self.__net_g = get_net_g(
model_path = str(self.model_path), model_path=str(self.model_path),
version = self.hyper_parameters.version, version=self.hyper_parameters.version,
device = self.device, device=self.device,
hps = self.hyper_parameters, hps=self.hyper_parameters,
) )
def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: 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 style_vec = mean + (style_vec - mean) * weight
return style_vec return style_vec
def __get_style_vector_from_audio(
def __get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]: self, audio_path: str, weight: float = 1.0
) -> NDArray[Any]:
""" """
音声からスタイルベクトルを推論する。 音声からスタイルベクトルを推論する。
@@ -126,8 +128,10 @@ class TTSModel:
# スタイルベクトルを取得するための推論モデルを初期化 # スタイルベクトルを取得するための推論モデルを初期化
if self.__style_vector_inference is None: if self.__style_vector_inference is None:
self.__style_vector_inference = pyannote.audio.Inference( self.__style_vector_inference = pyannote.audio.Inference(
model = pyannote.audio.Model.from_pretrained("pyannote/wespeaker-voxceleb-resnet34-LM"), model=pyannote.audio.Model.from_pretrained(
window = "whole", "pyannote/wespeaker-voxceleb-resnet34-LM"
),
window="whole",
) )
self.__style_vector_inference.to(torch.device(self.device)) self.__style_vector_inference.to(torch.device(self.device))
@@ -137,7 +141,6 @@ class TTSModel:
xvec = mean + (xvec - mean) * weight xvec = mean + (xvec - mean) * weight
return xvec return xvec
def infer( def infer(
self, self,
text: str, text: str,
@@ -209,20 +212,20 @@ class TTSModel:
if not line_split: if not line_split:
with torch.no_grad(): with torch.no_grad():
audio = infer( audio = infer(
text = text, text=text,
sdp_ratio = sdp_ratio, sdp_ratio=sdp_ratio,
noise_scale = noise, noise_scale=noise,
noise_scale_w = noise_w, noise_scale_w=noise_w,
length_scale = length, length_scale=length,
sid = speaker_id, sid=speaker_id,
language = language, language=language,
hps = self.hyper_parameters, hps=self.hyper_parameters,
net_g = self.__net_g, net_g=self.__net_g,
device = self.device, device=self.device,
assist_text = assist_text, assist_text=assist_text,
assist_text_weight = assist_text_weight, assist_text_weight=assist_text_weight,
style_vec = style_vector, style_vec=style_vector,
given_tone = given_tone, given_tone=given_tone,
) )
else: else:
texts = text.split("\n") texts = text.split("\n")
@@ -232,19 +235,19 @@ class TTSModel:
for i, t in enumerate(texts): for i, t in enumerate(texts):
audios.append( audios.append(
infer( infer(
text = t, text=t,
sdp_ratio = sdp_ratio, sdp_ratio=sdp_ratio,
noise_scale = noise, noise_scale=noise,
noise_scale_w = noise_w, noise_scale_w=noise_w,
length_scale = length, length_scale=length,
sid = speaker_id, sid=speaker_id,
language = language, language=language,
hps = self.hyper_parameters, hps=self.hyper_parameters,
net_g = self.__net_g, net_g=self.__net_g,
device = self.device, device=self.device,
assist_text = assist_text, assist_text=assist_text,
assist_text_weight = assist_text_weight, assist_text_weight=assist_text_weight,
style_vec = style_vector, style_vec=style_vector,
) )
) )
if i != len(texts) - 1: if i != len(texts) - 1:
@@ -253,10 +256,10 @@ class TTSModel:
logger.info("Audio data generated successfully") logger.info("Audio data generated successfully")
if not (pitch_scale == 1.0 and intonation_scale == 1.0): if not (pitch_scale == 1.0 and intonation_scale == 1.0):
_, audio = adjust_voice( _, audio = adjust_voice(
fs = self.hyper_parameters.data.sampling_rate, fs=self.hyper_parameters.data.sampling_rate,
wave = audio, wave=audio,
pitch_scale = pitch_scale, pitch_scale=pitch_scale,
intonation_scale = intonation_scale, intonation_scale=intonation_scale,
) )
with warnings.catch_warnings(): with warnings.catch_warnings():
warnings.simplefilter("ignore") warnings.simplefilter("ignore")
@@ -277,7 +280,6 @@ class TTSModelHolder:
model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。 model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。
""" """
def __init__(self, model_root_dir: Path, device: str) -> None: def __init__(self, model_root_dir: Path, device: str) -> None:
""" """
Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。 Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。
@@ -308,7 +310,6 @@ class TTSModelHolder:
self.models_info: list[TTSModelInfo] = [] self.models_info: list[TTSModelInfo] = []
self.refresh() self.refresh()
def refresh(self) -> None: def refresh(self) -> None:
""" """
音声合成モデルの一覧を更新する。 音声合成モデルの一覧を更新する。
@@ -342,13 +343,14 @@ class TTSModelHolder:
styles = list(style2id.keys()) styles = list(style2id.keys())
spk2id: dict[str, int] = hyper_parameters.data.spk2id spk2id: dict[str, int] = hyper_parameters.data.spk2id
speakers = list(spk2id.keys()) speakers = list(spk2id.keys())
self.models_info.append(TTSModelInfo( self.models_info.append(
name = model_dir.name, TTSModelInfo(
files = [str(f) for f in model_files], name=model_dir.name,
styles = styles, files=[str(f) for f in model_files],
speakers = speakers, styles=styles,
)) speakers=speakers,
)
)
def get_model(self, model_name: str, model_path_str: str) -> TTSModel: 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") raise ValueError(f"Model file `{model_path}` is not found")
if self.current_model is None or self.current_model.model_path != model_path: if self.current_model is None or self.current_model.model_path != model_path:
self.current_model = TTSModel( self.current_model = TTSModel(
model_path = model_path, model_path=model_path,
config_path = self.root_dir / model_name / "config.json", config_path=self.root_dir / model_name / "config.json",
style_vec_path = self.root_dir / model_name / "style_vectors.npy", style_vec_path=self.root_dir / model_name / "style_vectors.npy",
device = self.device, device=self.device,
) )
return self.current_model return self.current_model
def get_model_for_gradio(
def get_model_for_gradio(self, model_name: str, model_path_str: str) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: self, model_name: str, model_path_str: str
) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]:
model_path = Path(model_path_str) model_path = Path(model_path_str)
if model_name not in self.model_files_dict: if model_name not in self.model_files_dict:
raise ValueError(f"Model `{model_name}` is not found") raise ValueError(f"Model `{model_name}` is not found")
@@ -398,10 +401,10 @@ class TTSModelHolder:
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
) )
self.current_model = TTSModel( self.current_model = TTSModel(
model_path = model_path, model_path=model_path,
config_path = self.root_dir / model_name / "config.json", config_path=self.root_dir / model_name / "config.json",
style_vec_path = self.root_dir / model_name / "style_vectors.npy", style_vec_path=self.root_dir / model_name / "style_vectors.npy",
device = self.device, device=self.device,
) )
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())
@@ -411,13 +414,13 @@ class TTSModelHolder:
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
) )
def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown: def update_model_files_for_gradio(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]) # type: ignore return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore
def update_model_names_for_gradio(
def update_model_names_for_gradio(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: 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]

View File

@@ -8,40 +8,35 @@ class StdoutWrapper(TextIO):
`sys.stdout` wrapper for both Google Colab and local environment. `sys.stdout` wrapper for both Google Colab and local environment.
""" """
def __init__(self) -> None: def __init__(self) -> None:
self.temp_file = tempfile.NamedTemporaryFile( self.temp_file = tempfile.NamedTemporaryFile(
mode="w+", delete=False, encoding="utf-8" mode="w+", delete=False, encoding="utf-8"
) )
self.original_stdout = sys.stdout self.original_stdout = sys.stdout
def write(self, message: str) -> int: def write(self, message: str) -> int:
result = self.temp_file.write(message) result = self.temp_file.write(message)
self.temp_file.flush() self.temp_file.flush()
print(message, end="", file=self.original_stdout) print(message, end="", file=self.original_stdout)
return result return result
def flush(self) -> None: def flush(self) -> None:
self.temp_file.flush() self.temp_file.flush()
def read(self, n: int = -1) -> str: def read(self, n: int = -1) -> str:
self.temp_file.seek(0) self.temp_file.seek(0)
return self.temp_file.read(n) return self.temp_file.read(n)
def close(self) -> None: def close(self) -> None:
self.temp_file.close() self.temp_file.close()
def fileno(self) -> int: def fileno(self) -> int:
return self.temp_file.fileno() return self.temp_file.fileno()
try: try:
import google.colab # type: ignore import google.colab # type: ignore
SAFE_STDOUT = StdoutWrapper() SAFE_STDOUT = StdoutWrapper()
except ImportError: except ImportError:
SAFE_STDOUT = sys.stdout SAFE_STDOUT = sys.stdout

View File

@@ -9,27 +9,28 @@ class StrEnum(str, enum.Enum):
def __new__(cls, *values: str) -> "StrEnum": def __new__(cls, *values: str) -> "StrEnum":
"values must already be of type `str`" "values must already be of type `str`"
if len(values) > 3: 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: if len(values) == 1:
# it must be a string # it must be a string
if not isinstance(values[0], str): # type: ignore 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: if len(values) >= 2:
# check that encoding argument is a string # check that encoding argument is a string
if not isinstance(values[1], str): # type: ignore 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: if len(values) == 3:
# check that errors argument is a string # check that errors argument is a string
if not isinstance(values[2], str): # type: ignore 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) value = str(*values)
member = str.__new__(cls, value) member = str.__new__(cls, value)
member._value_ = value member._value_ = value
return member return member
@staticmethod @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. Return the lower-cased version of the member name.
""" """

View File

@@ -6,7 +6,9 @@ from style_bert_vits2.logging import logger
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT 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)}") logger.info(f"Running: {' '.join(cmd)}")
result = subprocess.run( result = subprocess.run(
[sys.executable] + cmd, [sys.executable] + cmd,
stdout = SAFE_STDOUT, stdout=SAFE_STDOUT,
stderr = subprocess.PIPE, stderr=subprocess.PIPE,
text = True, text=True,
encoding = "utf-8", encoding="utf-8",
) )
if result.returncode != 0: if result.returncode != 0:
logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}") 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, "" 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 番目の要素のみを返す関数を生成する。 与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する。

View File

@@ -5,15 +5,15 @@ from style_bert_vits2.constants import BASE_DIR, Languages
from style_bert_vits2.tts_model import TTSModelHolder 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: if len(model_holder.models_info) > 0:
# jvnv-F2-jp モデルを探す # jvnv-F2-jp モデルを探す
for model_info in model_holder.models_info: 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: for style in model_info.styles:
@@ -23,21 +23,21 @@ def synthesize(device: str = 'cpu'):
sample_rate, audio_data = model.infer( sample_rate, audio_data = model.infer(
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。", "あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
# 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ) # 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ)
language = Languages.JP, language=Languages.JP,
# 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0) # 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0)
speaker_id = 0, speaker_id=0,
# 感情表現の強さ (0.0 〜 1.0) # 感情表現の強さ (0.0 〜 1.0)
sdp_ratio = 0.4, sdp_ratio=0.4,
# スタイル (Neutral, Happy など) # スタイル (Neutral, Happy など)
style = style, style=style,
# スタイルの強さ (0.0 〜 100.0) # スタイルの強さ (0.0 〜 100.0)
style_weight = 6.0, style_weight=6.0,
) )
# 音声データを保存 # 音声データを保存
(BASE_DIR / 'tests/wavs').mkdir(exist_ok=True, parents=True) (BASE_DIR / "tests/wavs").mkdir(exist_ok=True, parents=True)
wav_file_path = BASE_DIR / f'tests/wavs/{style}.wav' wav_file_path = BASE_DIR / f"tests/wavs/{style}.wav"
with open(wav_file_path, 'wb') as f: with open(wav_file_path, "wb") as f:
wavfile.write(f, sample_rate, audio_data) wavfile.write(f, sample_rate, audio_data)
# 音声データが保存されたことを確認 # 音声データが保存されたことを確認
@@ -48,8 +48,8 @@ def synthesize(device: str = 'cpu'):
def test_synthesize_cpu(): def test_synthesize_cpu():
synthesize(device='cpu') synthesize(device="cpu")
def test_synthesize_cuda(): def test_synthesize_cuda():
synthesize(device='cuda') synthesize(device="cuda")

View File

@@ -281,28 +281,28 @@ def run():
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_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas = hps.model.use_noise_scaled_mas, use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
use_duration_discriminator = hps.model.use_duration_discriminator, use_duration_discriminator=hps.model.use_duration_discriminator,
use_wavlm_discriminator = hps.model.use_wavlm_discriminator, use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
inter_channels = hps.model.inter_channels, inter_channels=hps.model.inter_channels,
hidden_channels = hps.model.hidden_channels, hidden_channels=hps.model.hidden_channels,
filter_channels = hps.model.filter_channels, filter_channels=hps.model.filter_channels,
n_heads = hps.model.n_heads, n_heads=hps.model.n_heads,
n_layers = hps.model.n_layers, n_layers=hps.model.n_layers,
kernel_size = hps.model.kernel_size, kernel_size=hps.model.kernel_size,
p_dropout = hps.model.p_dropout, p_dropout=hps.model.p_dropout,
resblock = hps.model.resblock, resblock=hps.model.resblock,
resblock_kernel_sizes = hps.model.resblock_kernel_sizes, resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
resblock_dilation_sizes = hps.model.resblock_dilation_sizes, resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
upsample_rates = hps.model.upsample_rates, upsample_rates=hps.model.upsample_rates,
upsample_initial_channel = hps.model.upsample_initial_channel, upsample_initial_channel=hps.model.upsample_initial_channel,
upsample_kernel_sizes = hps.model.upsample_kernel_sizes, upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
n_layers_q = hps.model.n_layers_q, n_layers_q=hps.model.n_layers_q,
use_spectral_norm = hps.model.use_spectral_norm, use_spectral_norm=hps.model.use_spectral_norm,
gin_channels = hps.model.gin_channels, gin_channels=hps.model.gin_channels,
slm = hps.model.slm, 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):
@@ -389,7 +389,9 @@ def run():
epoch_str = max(epoch_str, 1) epoch_str = max(epoch_str, 1)
# global_step = (epoch_str - 1) * len(train_loader) # global_step = (epoch_str - 1) * len(train_loader)
global_step = int( 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( logger.info(
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************" f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"

View File

@@ -288,28 +288,28 @@ def run():
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_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas = hps.model.use_noise_scaled_mas, use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
use_duration_discriminator = hps.model.use_duration_discriminator, use_duration_discriminator=hps.model.use_duration_discriminator,
use_wavlm_discriminator = hps.model.use_wavlm_discriminator, use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
inter_channels = hps.model.inter_channels, inter_channels=hps.model.inter_channels,
hidden_channels = hps.model.hidden_channels, hidden_channels=hps.model.hidden_channels,
filter_channels = hps.model.filter_channels, filter_channels=hps.model.filter_channels,
n_heads = hps.model.n_heads, n_heads=hps.model.n_heads,
n_layers = hps.model.n_layers, n_layers=hps.model.n_layers,
kernel_size = hps.model.kernel_size, kernel_size=hps.model.kernel_size,
p_dropout = hps.model.p_dropout, p_dropout=hps.model.p_dropout,
resblock = hps.model.resblock, resblock=hps.model.resblock,
resblock_kernel_sizes = hps.model.resblock_kernel_sizes, resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
resblock_dilation_sizes = hps.model.resblock_dilation_sizes, resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
upsample_rates = hps.model.upsample_rates, upsample_rates=hps.model.upsample_rates,
upsample_initial_channel = hps.model.upsample_initial_channel, upsample_initial_channel=hps.model.upsample_initial_channel,
upsample_kernel_sizes = hps.model.upsample_kernel_sizes, upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
n_layers_q = hps.model.n_layers_q, n_layers_q=hps.model.n_layers_q,
use_spectral_norm = hps.model.use_spectral_norm, use_spectral_norm=hps.model.use_spectral_norm,
gin_channels = hps.model.gin_channels, gin_channels=hps.model.gin_channels,
slm = hps.model.slm, 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 !!!")
@@ -383,7 +383,9 @@ def run():
if net_dur_disc is not None: if net_dur_disc is not None:
try: try:
_, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( _, _, 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, net_dur_disc,
optim_dur_disc, optim_dur_disc,
skip_optimizer=hps.train.skip_optimizer, skip_optimizer=hps.train.skip_optimizer,
@@ -396,12 +398,16 @@ def run():
print("Initialize dur_disc") print("Initialize dur_disc")
if net_wd is not None: if net_wd is not None:
try: try:
_, optim_wd, wd_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( _, optim_wd, wd_resume_lr, epoch_str = (
utils.checkpoints.get_latest_checkpoint_path(model_dir, "WD_*.pth"), utils.checkpoints.load_checkpoint(
utils.checkpoints.get_latest_checkpoint_path(
model_dir, "WD_*.pth"
),
net_wd, net_wd,
optim_wd, optim_wd,
skip_optimizer=hps.train.skip_optimizer, skip_optimizer=hps.train.skip_optimizer,
) )
)
if not optim_wd.param_groups[0].get("initial_lr"): if not optim_wd.param_groups[0].get("initial_lr"):
optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr
except: except:
@@ -430,7 +436,9 @@ def run():
epoch_str = max(epoch_str, 1) epoch_str = max(epoch_str, 1)
# global_step = (epoch_str - 1) * len(train_loader) # global_step = (epoch_str - 1) * len(train_loader)
global_step = int( 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( logger.info(
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************" f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"