Apply black formatter

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

View File

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

View File

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

View File

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

View File

@@ -67,7 +67,9 @@ def intersperse(lst: list[Any], item: Any) -> list[Any]:
return result
def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4) -> torch.Tensor:
def slice_segments(
x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4
) -> torch.Tensor:
"""
テンソルからセグメントをスライスする
@@ -85,7 +87,9 @@ def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4
return torch.gather(x, 2, gather_indices)
def rand_slice_segments(x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4) -> tuple[torch.Tensor, torch.Tensor]:
def rand_slice_segments(
x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4
) -> tuple[torch.Tensor, torch.Tensor]:
"""
ランダムなセグメントをスライスする
@@ -121,7 +125,9 @@ def subsequent_mask(length: int) -> torch.Tensor:
@torch.jit.script # type: ignore
def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor) -> torch.Tensor:
def fused_add_tanh_sigmoid_multiply(
input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor
) -> torch.Tensor:
"""
加算、tanh、sigmoid の活性化関数を組み合わせた演算を行う
@@ -141,7 +147,9 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor
return acts
def sequence_mask(length: torch.Tensor, max_length: Optional[int] = None) -> torch.Tensor:
def sequence_mask(
length: torch.Tensor, max_length: Optional[int] = None
) -> torch.Tensor:
"""
シーケンスマスクを生成する
@@ -180,7 +188,11 @@ def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
return path
def clip_grad_value_(parameters: Union[torch.Tensor, list[torch.Tensor]], clip_value: Optional[float], norm_type: float = 2.0) -> float:
def clip_grad_value_(
parameters: Union[torch.Tensor, list[torch.Tensor]],
clip_value: Optional[float],
norm_type: float = 2.0,
) -> float:
"""
勾配の値をクリップする

View File

@@ -9,8 +9,14 @@ from style_bert_vits2.models import commons
from style_bert_vits2.models import utils
from style_bert_vits2.models.hyper_parameters import HyperParameters
from style_bert_vits2.models.models import SynthesizerTrn
from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra
from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature
from style_bert_vits2.models.models_jp_extra import (
SynthesizerTrn as SynthesizerTrnJPExtra,
)
from style_bert_vits2.nlp import (
clean_text,
cleaned_text_to_sequence,
extract_bert_feature,
)
from style_bert_vits2.nlp.symbols import SYMBOLS
@@ -18,69 +24,71 @@ def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters):
if version.endswith("JP-Extra"):
logger.info("Using JP-Extra model")
net_g = SynthesizerTrnJPExtra(
n_vocab = len(SYMBOLS),
spec_channels = hps.data.filter_length // 2 + 1,
segment_size = hps.train.segment_size // hps.data.hop_length,
n_speakers = hps.data.n_speakers,
n_vocab=len(SYMBOLS),
spec_channels=hps.data.filter_length // 2 + 1,
segment_size=hps.train.segment_size // hps.data.hop_length,
n_speakers=hps.data.n_speakers,
# hps.model 以下のすべての値を引数に渡す
use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas = hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder,
use_duration_discriminator = hps.model.use_duration_discriminator,
use_wavlm_discriminator = hps.model.use_wavlm_discriminator,
inter_channels = hps.model.inter_channels,
hidden_channels = hps.model.hidden_channels,
filter_channels = hps.model.filter_channels,
n_heads = hps.model.n_heads,
n_layers = hps.model.n_layers,
kernel_size = hps.model.kernel_size,
p_dropout = hps.model.p_dropout,
resblock = hps.model.resblock,
resblock_kernel_sizes = hps.model.resblock_kernel_sizes,
resblock_dilation_sizes = hps.model.resblock_dilation_sizes,
upsample_rates = hps.model.upsample_rates,
upsample_initial_channel = hps.model.upsample_initial_channel,
upsample_kernel_sizes = hps.model.upsample_kernel_sizes,
n_layers_q = hps.model.n_layers_q,
use_spectral_norm = hps.model.use_spectral_norm,
gin_channels = hps.model.gin_channels,
slm = hps.model.slm,
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
use_duration_discriminator=hps.model.use_duration_discriminator,
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
inter_channels=hps.model.inter_channels,
hidden_channels=hps.model.hidden_channels,
filter_channels=hps.model.filter_channels,
n_heads=hps.model.n_heads,
n_layers=hps.model.n_layers,
kernel_size=hps.model.kernel_size,
p_dropout=hps.model.p_dropout,
resblock=hps.model.resblock,
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
upsample_rates=hps.model.upsample_rates,
upsample_initial_channel=hps.model.upsample_initial_channel,
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
n_layers_q=hps.model.n_layers_q,
use_spectral_norm=hps.model.use_spectral_norm,
gin_channels=hps.model.gin_channels,
slm=hps.model.slm,
).to(device)
else:
logger.info("Using normal model")
net_g = SynthesizerTrn(
n_vocab = len(SYMBOLS),
spec_channels = hps.data.filter_length // 2 + 1,
segment_size = hps.train.segment_size // hps.data.hop_length,
n_vocab=len(SYMBOLS),
spec_channels=hps.data.filter_length // 2 + 1,
segment_size=hps.train.segment_size // hps.data.hop_length,
n_speakers=hps.data.n_speakers,
# hps.model 以下のすべての値を引数に渡す
use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas = hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder,
use_duration_discriminator = hps.model.use_duration_discriminator,
use_wavlm_discriminator = hps.model.use_wavlm_discriminator,
inter_channels = hps.model.inter_channels,
hidden_channels = hps.model.hidden_channels,
filter_channels = hps.model.filter_channels,
n_heads = hps.model.n_heads,
n_layers = hps.model.n_layers,
kernel_size = hps.model.kernel_size,
p_dropout = hps.model.p_dropout,
resblock = hps.model.resblock,
resblock_kernel_sizes = hps.model.resblock_kernel_sizes,
resblock_dilation_sizes = hps.model.resblock_dilation_sizes,
upsample_rates = hps.model.upsample_rates,
upsample_initial_channel = hps.model.upsample_initial_channel,
upsample_kernel_sizes = hps.model.upsample_kernel_sizes,
n_layers_q = hps.model.n_layers_q,
use_spectral_norm = hps.model.use_spectral_norm,
gin_channels = hps.model.gin_channels,
slm = hps.model.slm,
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
use_duration_discriminator=hps.model.use_duration_discriminator,
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
inter_channels=hps.model.inter_channels,
hidden_channels=hps.model.hidden_channels,
filter_channels=hps.model.filter_channels,
n_heads=hps.model.n_heads,
n_layers=hps.model.n_layers,
kernel_size=hps.model.kernel_size,
p_dropout=hps.model.p_dropout,
resblock=hps.model.resblock,
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
upsample_rates=hps.model.upsample_rates,
upsample_initial_channel=hps.model.upsample_initial_channel,
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
n_layers_q=hps.model.n_layers_q,
use_spectral_norm=hps.model.use_spectral_norm,
gin_channels=hps.model.gin_channels,
slm=hps.model.slm,
).to(device)
net_g.state_dict()
_ = net_g.eval()
if model_path.endswith(".pth") or model_path.endswith(".pt"):
_ = utils.checkpoints.load_checkpoint(model_path, net_g, None, skip_optimizer=True)
_ = utils.checkpoints.load_checkpoint(
model_path, net_g, None, skip_optimizer=True
)
elif model_path.endswith(".safetensors"):
_ = utils.safetensors.load_safetensors(model_path, net_g, True)
else:
@@ -102,8 +110,8 @@ def get_text(
norm_text, phone, tone, word2ph = clean_text(
text,
language_str,
use_jp_extra = use_jp_extra,
raise_yomi_error = False,
use_jp_extra=use_jp_extra,
raise_yomi_error=False,
)
if given_tone is not None:
if len(given_tone) != len(phone):

View File

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

View File

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

View File

@@ -88,7 +88,9 @@ class DDSConv(nn.Module):
Dialted and Depth-Separable Convolution
"""
def __init__(self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0) -> None:
def __init__(
self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0
) -> None:
super().__init__()
self.channels = channels
self.kernel_size = kernel_size
@@ -117,7 +119,9 @@ class DDSConv(nn.Module):
self.norms_1.append(LayerNorm(channels))
self.norms_2.append(LayerNorm(channels))
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
def forward(
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
) -> torch.Tensor:
if g is not None:
x = x + g
for i in range(self.n_layers):
@@ -184,7 +188,13 @@ class WN(torch.nn.Module):
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight")
self.res_skip_layers.append(res_skip_layer)
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, **kwargs: Any) -> torch.Tensor:
def forward(
self,
x: torch.Tensor,
x_mask: torch.Tensor,
g: Optional[torch.Tensor] = None,
**kwargs: Any,
) -> torch.Tensor:
output = torch.zeros_like(x)
n_channels_tensor = torch.IntTensor([self.hidden_channels])
@@ -221,7 +231,12 @@ class WN(torch.nn.Module):
class ResBlock1(torch.nn.Module):
def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int, int] = (1, 3, 5)) -> None:
def __init__(
self,
channels: int,
kernel_size: int = 3,
dilation: tuple[int, int, int] = (1, 3, 5),
) -> None:
super(ResBlock1, self).__init__()
self.convs1 = nn.ModuleList(
[
@@ -295,7 +310,9 @@ class ResBlock1(torch.nn.Module):
)
self.convs2.apply(commons.init_weights)
def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
def forward(
self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None
) -> torch.Tensor:
for c1, c2 in zip(self.convs1, self.convs2):
xt = F.leaky_relu(x, LRELU_SLOPE)
if x_mask is not None:
@@ -318,7 +335,9 @@ class ResBlock1(torch.nn.Module):
class ResBlock2(torch.nn.Module):
def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)) -> None:
def __init__(
self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)
) -> None:
super(ResBlock2, self).__init__()
self.convs = nn.ModuleList(
[
@@ -346,7 +365,9 @@ class ResBlock2(torch.nn.Module):
)
self.convs.apply(commons.init_weights)
def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
def forward(
self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None
) -> torch.Tensor:
for c in self.convs:
xt = F.leaky_relu(x, LRELU_SLOPE)
if x_mask is not None:

View File

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

View File

@@ -39,12 +39,14 @@ def piecewise_rational_quadratic_transform(
min_bin_width=min_bin_width,
min_bin_height=min_bin_height,
min_derivative=min_derivative,
**spline_kwargs # type: ignore
**spline_kwargs, # type: ignore
)
return outputs, logabsdet
def searchsorted(bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
def searchsorted(
bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6
) -> torch.Tensor:
bin_locations[..., -1] += eps
return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1

View File

@@ -107,7 +107,9 @@ def plot_spectrogram_to_numpy(spectrogram: NDArray[Any]) -> NDArray[Any]:
return data
def plot_alignment_to_numpy(alignment: NDArray[Any], info: Optional[str] = None) -> NDArray[Any]:
def plot_alignment_to_numpy(
alignment: NDArray[Any], info: Optional[str] = None
) -> NDArray[Any]:
"""
指定されたアライメントを画像データに変換する
@@ -163,7 +165,9 @@ def load_wav_to_torch(full_path: Union[str, Path]) -> tuple[torch.FloatTensor, i
return torch.FloatTensor(data.astype(np.float32)), sampling_rate
def load_filepaths_and_text(filename: Union[str, Path], split: str = "|") -> list[list[str]]:
def load_filepaths_and_text(
filename: Union[str, Path], split: str = "|"
) -> list[list[str]]:
"""
指定されたファイルからファイルパスとテキストを読み込む
@@ -180,7 +184,9 @@ def load_filepaths_and_text(filename: Union[str, Path], split: str = "|") -> lis
return filepaths_and_text
def get_logger(model_dir_path: Union[str, Path], filename: str = "train.log") -> logging.Logger:
def get_logger(
model_dir_path: Union[str, Path], filename: str = "train.log"
) -> logging.Logger:
"""
ロガーを取得する

View File

@@ -14,7 +14,7 @@ def load_checkpoint(
model: torch.nn.Module,
optimizer: Optional[torch.optim.Optimizer] = None,
skip_optimizer: bool = False,
for_infer: bool = False
for_infer: bool = False,
) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]:
"""
指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。
@@ -107,7 +107,9 @@ def save_checkpoint(
iteration (int): イテレーション数
checkpoint_path (Union[str, Path]): 保存先のパス
"""
logger.info(f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}")
logger.info(
f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}"
)
if hasattr(model, "module"):
state_dict = model.module.state_dict()
else:
@@ -123,7 +125,11 @@ def save_checkpoint(
)
def clean_checkpoints(model_dir_path: Union[str, Path] = "logs/44k/", n_ckpts_to_keep: int = 2, sort_by_time: bool = True) -> None:
def clean_checkpoints(
model_dir_path: Union[str, Path] = "logs/44k/",
n_ckpts_to_keep: int = 2,
sort_by_time: bool = True,
) -> None:
"""
指定されたディレクトリから古いチェックポイントを削除して空き容量を確保する
@@ -172,7 +178,9 @@ def clean_checkpoints(model_dir_path: Union[str, Path] = "logs/44k/", n_ckpts_to
[del_routine(fn) for fn in to_del]
def get_latest_checkpoint_path(model_dir_path: Union[str, Path], regex: str = "G_*.pth") -> str:
def get_latest_checkpoint_path(
model_dir_path: Union[str, Path], regex: str = "G_*.pth"
) -> str:
"""
指定されたディレクトリから最新のチェックポイントのパスを取得する

View File

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

View File

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

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():
pinyin = single_rep_map[pinyin[0]] + pinyin[1:]
assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (pinyin, seg, raw_pinyin)
assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (
pinyin,
seg,
raw_pinyin,
)
phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ")
word2ph.append(len(phone))
@@ -125,7 +129,7 @@ if __name__ == "__main__":
text = normalize_text(text)
print(text)
phones, tones, word2ph = g2p(text)
bert = extract_bert_feature(text, word2ph, 'cuda')
bert = extract_bert_feature(text, word2ph, "cuda")
print(phones, tones, word2ph, bert.shape)

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -18,7 +18,11 @@ from fastapi import HTTPException
from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes
from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data
from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import (
MAX_PRIORITY,
MIN_PRIORITY,
part_of_speech_data,
)
# root_dir = engine_root()
# save_dir = get_save_dir()
@@ -26,9 +30,13 @@ from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIO
# if not save_dir.is_dir():
# save_dir.mkdir(parents=True)
default_dict_path = DEFAULT_USER_DICT_DIR / "default.csv" # VOICEVOXデフォルト辞書ファイルのパス
default_dict_path = (
DEFAULT_USER_DICT_DIR / "default.csv"
) # VOICEVOXデフォルト辞書ファイルのパス
user_dict_path = DEFAULT_USER_DICT_DIR / "user_dict.json" # ユーザー辞書ファイルのパス
compiled_dict_path = DEFAULT_USER_DICT_DIR / "user.dic" # コンパイル済み辞書ファイルのパス
compiled_dict_path = (
DEFAULT_USER_DICT_DIR / "user.dic"
) # コンパイル済み辞書ファイルのパス
# # 同時書き込みの制御

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

View File

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

View File

@@ -9,27 +9,28 @@ class StrEnum(str, enum.Enum):
def __new__(cls, *values: str) -> "StrEnum":
"values must already be of type `str`"
if len(values) > 3:
raise TypeError('too many arguments for str(): %r' % (values, ))
raise TypeError("too many arguments for str(): %r" % (values,))
if len(values) == 1:
# it must be a string
if not isinstance(values[0], str): # type: ignore
raise TypeError('%r is not a string' % (values[0], ))
raise TypeError("%r is not a string" % (values[0],))
if len(values) >= 2:
# check that encoding argument is a string
if not isinstance(values[1], str): # type: ignore
raise TypeError('encoding must be a string, not %r' % (values[1], ))
raise TypeError("encoding must be a string, not %r" % (values[1],))
if len(values) == 3:
# check that errors argument is a string
if not isinstance(values[2], str): # type: ignore
raise TypeError('errors must be a string, not %r' % (values[2]))
raise TypeError("errors must be a string, not %r" % (values[2]))
value = str(*values)
member = str.__new__(cls, value)
member._value_ = value
return member
@staticmethod
def _generate_next_value_(name: str, start: int, count: int, last_values: list[str]) -> str:
def _generate_next_value_(
name: str, start: int, count: int, last_values: list[str]
) -> str:
"""
Return the lower-cased version of the member name.
"""

View File

@@ -6,7 +6,9 @@ from style_bert_vits2.logging import logger
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[bool, str]:
def run_script_with_log(
cmd: list[str], ignore_warning: bool = False
) -> tuple[bool, str]:
"""
指定されたコマンドを実行し、そのログを記録する。
@@ -21,10 +23,10 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b
logger.info(f"Running: {' '.join(cmd)}")
result = subprocess.run(
[sys.executable] + cmd,
stdout = SAFE_STDOUT,
stderr = subprocess.PIPE,
text = True,
encoding = "utf-8",
stdout=SAFE_STDOUT,
stderr=subprocess.PIPE,
text=True,
encoding="utf-8",
)
if result.returncode != 0:
logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}")
@@ -37,7 +39,9 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b
return True, ""
def second_elem_of(original_function: Callable[..., tuple[Any, Any]]) -> Callable[..., Any]:
def second_elem_of(
original_function: Callable[..., tuple[Any, Any]]
) -> Callable[..., Any]:
"""
与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する。