Apply black formatter
This commit is contained in:
4
.vscode/settings.json
vendored
4
.vscode/settings.json
vendored
@@ -19,4 +19,8 @@
|
||||
"reportUnusedFunction": "none",
|
||||
"reportUnusedVariable": "information",
|
||||
},
|
||||
"[python]": {
|
||||
"editor.defaultFormatter": "ms-python.black-formatter",
|
||||
"editor.formatOnType": true,
|
||||
},
|
||||
}
|
||||
7
app.py
7
app.py
@@ -51,4 +51,9 @@ with gr.Blocks(theme=GRADIO_THEME) as app:
|
||||
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,
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
勾配の値をクリップする
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
ロガーを取得する
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
指定されたディレクトリから最新のチェックポイントのパスを取得する
|
||||
|
||||
|
||||
@@ -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 に変換する
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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]]:
|
||||
"""
|
||||
例: …私は、、そう思う。
|
||||
|
||||
@@ -34,11 +34,13 @@ def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, i
|
||||
"""
|
||||
|
||||
# 子音の集合
|
||||
CONSONANTS = set([
|
||||
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]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
) # コンパイル済み辞書ファイルのパス
|
||||
|
||||
|
||||
# # 同時書き込みの制御
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
|
||||
@@ -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 番目の要素のみを返す関数を生成する。
|
||||
|
||||
|
||||
@@ -5,15 +5,15 @@ from style_bert_vits2.constants import BASE_DIR, Languages
|
||||
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:
|
||||
|
||||
# jvnv-F2-jp モデルを探す
|
||||
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:
|
||||
|
||||
@@ -23,21 +23,21 @@ def synthesize(device: str = 'cpu'):
|
||||
sample_rate, audio_data = model.infer(
|
||||
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
|
||||
# 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ)
|
||||
language = Languages.JP,
|
||||
language=Languages.JP,
|
||||
# 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0)
|
||||
speaker_id = 0,
|
||||
speaker_id=0,
|
||||
# 感情表現の強さ (0.0 〜 1.0)
|
||||
sdp_ratio = 0.4,
|
||||
sdp_ratio=0.4,
|
||||
# スタイル (Neutral, Happy など)
|
||||
style = style,
|
||||
style=style,
|
||||
# スタイルの強さ (0.0 〜 100.0)
|
||||
style_weight = 6.0,
|
||||
style_weight=6.0,
|
||||
)
|
||||
|
||||
# 音声データを保存
|
||||
(BASE_DIR / 'tests/wavs').mkdir(exist_ok=True, parents=True)
|
||||
wav_file_path = BASE_DIR / f'tests/wavs/{style}.wav'
|
||||
with open(wav_file_path, 'wb') as f:
|
||||
(BASE_DIR / "tests/wavs").mkdir(exist_ok=True, parents=True)
|
||||
wav_file_path = BASE_DIR / f"tests/wavs/{style}.wav"
|
||||
with open(wav_file_path, "wb") as f:
|
||||
wavfile.write(f, sample_rate, audio_data)
|
||||
|
||||
# 音声データが保存されたことを確認
|
||||
@@ -48,8 +48,8 @@ def synthesize(device: str = 'cpu'):
|
||||
|
||||
|
||||
def test_synthesize_cpu():
|
||||
synthesize(device='cpu')
|
||||
synthesize(device="cpu")
|
||||
|
||||
|
||||
def test_synthesize_cuda():
|
||||
synthesize(device='cuda')
|
||||
synthesize(device="cuda")
|
||||
|
||||
48
train_ms.py
48
train_ms.py
@@ -281,28 +281,28 @@ def run():
|
||||
mas_noise_scale_initial=mas_noise_scale_initial,
|
||||
noise_scale_delta=noise_scale_delta,
|
||||
# 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,
|
||||
).cuda(local_rank)
|
||||
|
||||
if getattr(hps.train, "freeze_ZH_bert", False):
|
||||
@@ -389,7 +389,9 @@ def run():
|
||||
epoch_str = max(epoch_str, 1)
|
||||
# global_step = (epoch_str - 1) * len(train_loader)
|
||||
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(
|
||||
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
|
||||
|
||||
@@ -288,28 +288,28 @@ def run():
|
||||
mas_noise_scale_initial=mas_noise_scale_initial,
|
||||
noise_scale_delta=noise_scale_delta,
|
||||
# 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,
|
||||
).cuda(local_rank)
|
||||
if getattr(hps.train, "freeze_JP_bert", False):
|
||||
logger.info("Freezing (JP) bert encoder !!!")
|
||||
@@ -383,7 +383,9 @@ def run():
|
||||
if net_dur_disc is not None:
|
||||
try:
|
||||
_, _, 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,
|
||||
optim_dur_disc,
|
||||
skip_optimizer=hps.train.skip_optimizer,
|
||||
@@ -396,12 +398,16 @@ def run():
|
||||
print("Initialize dur_disc")
|
||||
if net_wd is not None:
|
||||
try:
|
||||
_, optim_wd, wd_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
|
||||
utils.checkpoints.get_latest_checkpoint_path(model_dir, "WD_*.pth"),
|
||||
_, optim_wd, wd_resume_lr, epoch_str = (
|
||||
utils.checkpoints.load_checkpoint(
|
||||
utils.checkpoints.get_latest_checkpoint_path(
|
||||
model_dir, "WD_*.pth"
|
||||
),
|
||||
net_wd,
|
||||
optim_wd,
|
||||
skip_optimizer=hps.train.skip_optimizer,
|
||||
)
|
||||
)
|
||||
if not optim_wd.param_groups[0].get("initial_lr"):
|
||||
optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr
|
||||
except:
|
||||
@@ -430,7 +436,9 @@ def run():
|
||||
epoch_str = max(epoch_str, 1)
|
||||
# global_step = (epoch_str - 1) * len(train_loader)
|
||||
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(
|
||||
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
|
||||
|
||||
Reference in New Issue
Block a user