Apply black formatter
This commit is contained in:
4
.vscode/settings.json
vendored
4
.vscode/settings.json
vendored
@@ -19,4 +19,8 @@
|
|||||||
"reportUnusedFunction": "none",
|
"reportUnusedFunction": "none",
|
||||||
"reportUnusedVariable": "information",
|
"reportUnusedVariable": "information",
|
||||||
},
|
},
|
||||||
|
"[python]": {
|
||||||
|
"editor.defaultFormatter": "ms-python.black-formatter",
|
||||||
|
"editor.formatOnType": true,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
7
app.py
7
app.py
@@ -51,4 +51,9 @@ with gr.Blocks(theme=GRADIO_THEME) as app:
|
|||||||
create_merge_app(model_holder=model_holder)
|
create_merge_app(model_holder=model_holder)
|
||||||
|
|
||||||
|
|
||||||
app.launch(server_name=args.host, server_port=args.port, inbrowser=not args.no_autolaunch, share=args.share)
|
app.launch(
|
||||||
|
server_name=args.host,
|
||||||
|
server_port=args.port,
|
||||||
|
inbrowser=not args.no_autolaunch,
|
||||||
|
share=args.share,
|
||||||
|
)
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ VERSION = "2.4"
|
|||||||
# Style-Bert-VITS2 のベースディレクトリ
|
# Style-Bert-VITS2 のベースディレクトリ
|
||||||
BASE_DIR = Path(__file__).parent.parent
|
BASE_DIR = Path(__file__).parent.parent
|
||||||
|
|
||||||
|
|
||||||
# 利用可能な言語
|
# 利用可能な言語
|
||||||
## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない
|
## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない
|
||||||
class Languages(StrEnum):
|
class Languages(StrEnum):
|
||||||
@@ -16,6 +17,7 @@ class Languages(StrEnum):
|
|||||||
EN = "EN"
|
EN = "EN"
|
||||||
ZH = "ZH"
|
ZH = "ZH"
|
||||||
|
|
||||||
|
|
||||||
# 言語ごとのデフォルトの BERT トークナイザーのパス
|
# 言語ごとのデフォルトの BERT トークナイザーのパス
|
||||||
DEFAULT_BERT_TOKENIZER_PATHS = {
|
DEFAULT_BERT_TOKENIZER_PATHS = {
|
||||||
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ logger.remove()
|
|||||||
# Add a new handler
|
# Add a new handler
|
||||||
logger.add(
|
logger.add(
|
||||||
SAFE_STDOUT,
|
SAFE_STDOUT,
|
||||||
format = "<g>{time:MM-DD HH:mm:ss}</g> |<lvl>{level:^8}</lvl>| {file}:{line} | {message}",
|
format="<g>{time:MM-DD HH:mm:ss}</g> |<lvl>{level:^8}</lvl>| {file}:{line} | {message}",
|
||||||
backtrace = True,
|
backtrace=True,
|
||||||
diagnose = True,
|
diagnose=True,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -24,7 +24,9 @@ class LayerNorm(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
@torch.jit.script # type: ignore
|
@torch.jit.script # type: ignore
|
||||||
def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor, n_channels: list[int]) -> torch.Tensor:
|
def fused_add_tanh_sigmoid_multiply(
|
||||||
|
input_a: torch.Tensor, input_b: torch.Tensor, n_channels: list[int]
|
||||||
|
) -> torch.Tensor:
|
||||||
n_channels_int = n_channels[0]
|
n_channels_int = n_channels[0]
|
||||||
in_act = input_a + input_b
|
in_act = input_a + input_b
|
||||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||||
@@ -44,7 +46,7 @@ class Encoder(nn.Module):
|
|||||||
p_dropout: float = 0.0,
|
p_dropout: float = 0.0,
|
||||||
window_size: int = 4,
|
window_size: int = 4,
|
||||||
isflow: bool = True,
|
isflow: bool = True,
|
||||||
**kwargs: Any
|
**kwargs: Any,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -99,7 +101,9 @@ class Encoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
|
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
|
||||||
x = x * x_mask
|
x = x * x_mask
|
||||||
for i in range(self.n_layers):
|
for i in range(self.n_layers):
|
||||||
@@ -131,7 +135,7 @@ class Decoder(nn.Module):
|
|||||||
p_dropout: float = 0.0,
|
p_dropout: float = 0.0,
|
||||||
proximal_bias: bool = False,
|
proximal_bias: bool = False,
|
||||||
proximal_init: bool = True,
|
proximal_init: bool = True,
|
||||||
**kwargs: Any
|
**kwargs: Any,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_channels = hidden_channels
|
self.hidden_channels = hidden_channels
|
||||||
@@ -180,7 +184,13 @@ class Decoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, h: torch.Tensor, h_mask: torch.Tensor) -> torch.Tensor:
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
h: torch.Tensor,
|
||||||
|
h_mask: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: decoder input
|
x: decoder input
|
||||||
h: encoder output
|
h: encoder output
|
||||||
@@ -262,7 +272,9 @@ class MultiHeadAttention(nn.Module):
|
|||||||
assert self.conv_q.bias is not None
|
assert self.conv_q.bias is not None
|
||||||
self.conv_k.bias.copy_(self.conv_q.bias)
|
self.conv_k.bias.copy_(self.conv_q.bias)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, c: torch.Tensor, attn_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, c: torch.Tensor, attn_mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
q = self.conv_q(x)
|
q = self.conv_q(x)
|
||||||
k = self.conv_k(c)
|
k = self.conv_k(c)
|
||||||
v = self.conv_v(c)
|
v = self.conv_v(c)
|
||||||
@@ -329,7 +341,9 @@ class MultiHeadAttention(nn.Module):
|
|||||||
) # [b, n_h, t_t, d_k] -> [b, d, t_t]
|
) # [b, n_h, t_t, d_k] -> [b, d, t_t]
|
||||||
return output, p_attn
|
return output, p_attn
|
||||||
|
|
||||||
def _matmul_with_relative_values(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
|
def _matmul_with_relative_values(
|
||||||
|
self, x: torch.Tensor, y: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: [b, h, l, m]
|
x: [b, h, l, m]
|
||||||
y: [h or 1, m, d]
|
y: [h or 1, m, d]
|
||||||
@@ -338,7 +352,9 @@ class MultiHeadAttention(nn.Module):
|
|||||||
ret = torch.matmul(x, y.unsqueeze(0))
|
ret = torch.matmul(x, y.unsqueeze(0))
|
||||||
return ret
|
return ret
|
||||||
|
|
||||||
def _matmul_with_relative_keys(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
|
def _matmul_with_relative_keys(
|
||||||
|
self, x: torch.Tensor, y: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
x: [b, h, l, d]
|
x: [b, h, l, d]
|
||||||
y: [h or 1, m, d]
|
y: [h or 1, m, d]
|
||||||
@@ -347,7 +363,9 @@ class MultiHeadAttention(nn.Module):
|
|||||||
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
|
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
|
||||||
return ret
|
return ret
|
||||||
|
|
||||||
def _get_relative_embeddings(self, relative_embeddings: torch.Tensor, length: int) -> torch.Tensor:
|
def _get_relative_embeddings(
|
||||||
|
self, relative_embeddings: torch.Tensor, length: int
|
||||||
|
) -> torch.Tensor:
|
||||||
assert self.window_size is not None
|
assert self.window_size is not None
|
||||||
2 * self.window_size + 1 # type: ignore
|
2 * self.window_size + 1 # type: ignore
|
||||||
# Pad first before slice to avoid using cond ops.
|
# Pad first before slice to avoid using cond ops.
|
||||||
|
|||||||
@@ -67,7 +67,9 @@ def intersperse(lst: list[Any], item: Any) -> list[Any]:
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4) -> torch.Tensor:
|
def slice_segments(
|
||||||
|
x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
テンソルからセグメントをスライスする
|
テンソルからセグメントをスライスする
|
||||||
|
|
||||||
@@ -85,7 +87,9 @@ def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4
|
|||||||
return torch.gather(x, 2, gather_indices)
|
return torch.gather(x, 2, gather_indices)
|
||||||
|
|
||||||
|
|
||||||
def rand_slice_segments(x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4) -> tuple[torch.Tensor, torch.Tensor]:
|
def rand_slice_segments(
|
||||||
|
x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
"""
|
"""
|
||||||
ランダムなセグメントをスライスする
|
ランダムなセグメントをスライスする
|
||||||
|
|
||||||
@@ -121,7 +125,9 @@ def subsequent_mask(length: int) -> torch.Tensor:
|
|||||||
|
|
||||||
|
|
||||||
@torch.jit.script # type: ignore
|
@torch.jit.script # type: ignore
|
||||||
def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor) -> torch.Tensor:
|
def fused_add_tanh_sigmoid_multiply(
|
||||||
|
input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
加算、tanh、sigmoid の活性化関数を組み合わせた演算を行う
|
加算、tanh、sigmoid の活性化関数を組み合わせた演算を行う
|
||||||
|
|
||||||
@@ -141,7 +147,9 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor
|
|||||||
return acts
|
return acts
|
||||||
|
|
||||||
|
|
||||||
def sequence_mask(length: torch.Tensor, max_length: Optional[int] = None) -> torch.Tensor:
|
def sequence_mask(
|
||||||
|
length: torch.Tensor, max_length: Optional[int] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
シーケンスマスクを生成する
|
シーケンスマスクを生成する
|
||||||
|
|
||||||
@@ -180,7 +188,11 @@ def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
def clip_grad_value_(parameters: Union[torch.Tensor, list[torch.Tensor]], clip_value: Optional[float], norm_type: float = 2.0) -> float:
|
def clip_grad_value_(
|
||||||
|
parameters: Union[torch.Tensor, list[torch.Tensor]],
|
||||||
|
clip_value: Optional[float],
|
||||||
|
norm_type: float = 2.0,
|
||||||
|
) -> float:
|
||||||
"""
|
"""
|
||||||
勾配の値をクリップする
|
勾配の値をクリップする
|
||||||
|
|
||||||
|
|||||||
@@ -9,8 +9,14 @@ from style_bert_vits2.models import commons
|
|||||||
from style_bert_vits2.models import utils
|
from style_bert_vits2.models import utils
|
||||||
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
||||||
from style_bert_vits2.models.models import SynthesizerTrn
|
from style_bert_vits2.models.models import SynthesizerTrn
|
||||||
from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra
|
from style_bert_vits2.models.models_jp_extra import (
|
||||||
from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
|
)
|
||||||
|
from style_bert_vits2.nlp import (
|
||||||
|
clean_text,
|
||||||
|
cleaned_text_to_sequence,
|
||||||
|
extract_bert_feature,
|
||||||
|
)
|
||||||
from style_bert_vits2.nlp.symbols import SYMBOLS
|
from style_bert_vits2.nlp.symbols import SYMBOLS
|
||||||
|
|
||||||
|
|
||||||
@@ -18,69 +24,71 @@ def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters):
|
|||||||
if version.endswith("JP-Extra"):
|
if version.endswith("JP-Extra"):
|
||||||
logger.info("Using JP-Extra model")
|
logger.info("Using JP-Extra model")
|
||||||
net_g = SynthesizerTrnJPExtra(
|
net_g = SynthesizerTrnJPExtra(
|
||||||
n_vocab = len(SYMBOLS),
|
n_vocab=len(SYMBOLS),
|
||||||
spec_channels = hps.data.filter_length // 2 + 1,
|
spec_channels=hps.data.filter_length // 2 + 1,
|
||||||
segment_size = hps.train.segment_size // hps.data.hop_length,
|
segment_size=hps.train.segment_size // hps.data.hop_length,
|
||||||
n_speakers = hps.data.n_speakers,
|
n_speakers=hps.data.n_speakers,
|
||||||
# hps.model 以下のすべての値を引数に渡す
|
# hps.model 以下のすべての値を引数に渡す
|
||||||
use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder,
|
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
|
||||||
use_noise_scaled_mas = hps.model.use_noise_scaled_mas,
|
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
|
||||||
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder,
|
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
|
||||||
use_duration_discriminator = hps.model.use_duration_discriminator,
|
use_duration_discriminator=hps.model.use_duration_discriminator,
|
||||||
use_wavlm_discriminator = hps.model.use_wavlm_discriminator,
|
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
|
||||||
inter_channels = hps.model.inter_channels,
|
inter_channels=hps.model.inter_channels,
|
||||||
hidden_channels = hps.model.hidden_channels,
|
hidden_channels=hps.model.hidden_channels,
|
||||||
filter_channels = hps.model.filter_channels,
|
filter_channels=hps.model.filter_channels,
|
||||||
n_heads = hps.model.n_heads,
|
n_heads=hps.model.n_heads,
|
||||||
n_layers = hps.model.n_layers,
|
n_layers=hps.model.n_layers,
|
||||||
kernel_size = hps.model.kernel_size,
|
kernel_size=hps.model.kernel_size,
|
||||||
p_dropout = hps.model.p_dropout,
|
p_dropout=hps.model.p_dropout,
|
||||||
resblock = hps.model.resblock,
|
resblock=hps.model.resblock,
|
||||||
resblock_kernel_sizes = hps.model.resblock_kernel_sizes,
|
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
|
||||||
resblock_dilation_sizes = hps.model.resblock_dilation_sizes,
|
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
|
||||||
upsample_rates = hps.model.upsample_rates,
|
upsample_rates=hps.model.upsample_rates,
|
||||||
upsample_initial_channel = hps.model.upsample_initial_channel,
|
upsample_initial_channel=hps.model.upsample_initial_channel,
|
||||||
upsample_kernel_sizes = hps.model.upsample_kernel_sizes,
|
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
|
||||||
n_layers_q = hps.model.n_layers_q,
|
n_layers_q=hps.model.n_layers_q,
|
||||||
use_spectral_norm = hps.model.use_spectral_norm,
|
use_spectral_norm=hps.model.use_spectral_norm,
|
||||||
gin_channels = hps.model.gin_channels,
|
gin_channels=hps.model.gin_channels,
|
||||||
slm = hps.model.slm,
|
slm=hps.model.slm,
|
||||||
).to(device)
|
).to(device)
|
||||||
else:
|
else:
|
||||||
logger.info("Using normal model")
|
logger.info("Using normal model")
|
||||||
net_g = SynthesizerTrn(
|
net_g = SynthesizerTrn(
|
||||||
n_vocab = len(SYMBOLS),
|
n_vocab=len(SYMBOLS),
|
||||||
spec_channels = hps.data.filter_length // 2 + 1,
|
spec_channels=hps.data.filter_length // 2 + 1,
|
||||||
segment_size = hps.train.segment_size // hps.data.hop_length,
|
segment_size=hps.train.segment_size // hps.data.hop_length,
|
||||||
n_speakers=hps.data.n_speakers,
|
n_speakers=hps.data.n_speakers,
|
||||||
# hps.model 以下のすべての値を引数に渡す
|
# hps.model 以下のすべての値を引数に渡す
|
||||||
use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder,
|
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
|
||||||
use_noise_scaled_mas = hps.model.use_noise_scaled_mas,
|
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
|
||||||
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder,
|
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
|
||||||
use_duration_discriminator = hps.model.use_duration_discriminator,
|
use_duration_discriminator=hps.model.use_duration_discriminator,
|
||||||
use_wavlm_discriminator = hps.model.use_wavlm_discriminator,
|
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
|
||||||
inter_channels = hps.model.inter_channels,
|
inter_channels=hps.model.inter_channels,
|
||||||
hidden_channels = hps.model.hidden_channels,
|
hidden_channels=hps.model.hidden_channels,
|
||||||
filter_channels = hps.model.filter_channels,
|
filter_channels=hps.model.filter_channels,
|
||||||
n_heads = hps.model.n_heads,
|
n_heads=hps.model.n_heads,
|
||||||
n_layers = hps.model.n_layers,
|
n_layers=hps.model.n_layers,
|
||||||
kernel_size = hps.model.kernel_size,
|
kernel_size=hps.model.kernel_size,
|
||||||
p_dropout = hps.model.p_dropout,
|
p_dropout=hps.model.p_dropout,
|
||||||
resblock = hps.model.resblock,
|
resblock=hps.model.resblock,
|
||||||
resblock_kernel_sizes = hps.model.resblock_kernel_sizes,
|
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
|
||||||
resblock_dilation_sizes = hps.model.resblock_dilation_sizes,
|
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
|
||||||
upsample_rates = hps.model.upsample_rates,
|
upsample_rates=hps.model.upsample_rates,
|
||||||
upsample_initial_channel = hps.model.upsample_initial_channel,
|
upsample_initial_channel=hps.model.upsample_initial_channel,
|
||||||
upsample_kernel_sizes = hps.model.upsample_kernel_sizes,
|
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
|
||||||
n_layers_q = hps.model.n_layers_q,
|
n_layers_q=hps.model.n_layers_q,
|
||||||
use_spectral_norm = hps.model.use_spectral_norm,
|
use_spectral_norm=hps.model.use_spectral_norm,
|
||||||
gin_channels = hps.model.gin_channels,
|
gin_channels=hps.model.gin_channels,
|
||||||
slm = hps.model.slm,
|
slm=hps.model.slm,
|
||||||
).to(device)
|
).to(device)
|
||||||
net_g.state_dict()
|
net_g.state_dict()
|
||||||
_ = net_g.eval()
|
_ = net_g.eval()
|
||||||
if model_path.endswith(".pth") or model_path.endswith(".pt"):
|
if model_path.endswith(".pth") or model_path.endswith(".pt"):
|
||||||
_ = utils.checkpoints.load_checkpoint(model_path, net_g, None, skip_optimizer=True)
|
_ = utils.checkpoints.load_checkpoint(
|
||||||
|
model_path, net_g, None, skip_optimizer=True
|
||||||
|
)
|
||||||
elif model_path.endswith(".safetensors"):
|
elif model_path.endswith(".safetensors"):
|
||||||
_ = utils.safetensors.load_safetensors(model_path, net_g, True)
|
_ = utils.safetensors.load_safetensors(model_path, net_g, True)
|
||||||
else:
|
else:
|
||||||
@@ -102,8 +110,8 @@ def get_text(
|
|||||||
norm_text, phone, tone, word2ph = clean_text(
|
norm_text, phone, tone, word2ph = clean_text(
|
||||||
text,
|
text,
|
||||||
language_str,
|
language_str,
|
||||||
use_jp_extra = use_jp_extra,
|
use_jp_extra=use_jp_extra,
|
||||||
raise_yomi_error = False,
|
raise_yomi_error=False,
|
||||||
)
|
)
|
||||||
if given_tone is not None:
|
if given_tone is not None:
|
||||||
if len(given_tone) != len(phone):
|
if len(given_tone) != len(phone):
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
filter_channels: int,
|
filter_channels: int,
|
||||||
kernel_size: int,
|
kernel_size: int,
|
||||||
p_dropout: float,
|
p_dropout: float,
|
||||||
gin_channels: int = 0
|
gin_channels: int = 0,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -330,7 +330,9 @@ class DurationPredictor(nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
g = torch.detach(g)
|
g = torch.detach(g)
|
||||||
@@ -582,7 +584,9 @@ class Generator(torch.nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = self.conv_pre(x)
|
x = self.conv_pre(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
x = x + self.cond(g)
|
x = x + self.cond(g)
|
||||||
@@ -613,7 +617,13 @@ class Generator(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class DiscriminatorP(torch.nn.Module):
|
class DiscriminatorP(torch.nn.Module):
|
||||||
def __init__(self, period: int, kernel_size: int = 5, stride: int = 3, use_spectral_norm: bool = False) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
period: int,
|
||||||
|
kernel_size: int = 5,
|
||||||
|
stride: int = 3,
|
||||||
|
use_spectral_norm: bool = False,
|
||||||
|
) -> None:
|
||||||
super(DiscriminatorP, self).__init__()
|
super(DiscriminatorP, self).__init__()
|
||||||
self.period = period
|
self.period = period
|
||||||
self.use_spectral_norm = use_spectral_norm
|
self.use_spectral_norm = use_spectral_norm
|
||||||
@@ -736,7 +746,9 @@ class MultiPeriodDiscriminator(torch.nn.Module):
|
|||||||
self,
|
self,
|
||||||
y: torch.Tensor,
|
y: torch.Tensor,
|
||||||
y_hat: torch.Tensor,
|
y_hat: torch.Tensor,
|
||||||
) -> tuple[list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]]:
|
) -> tuple[
|
||||||
|
list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]
|
||||||
|
]:
|
||||||
y_d_rs = []
|
y_d_rs = []
|
||||||
y_d_gs = []
|
y_d_gs = []
|
||||||
fmap_rs = []
|
fmap_rs = []
|
||||||
@@ -787,7 +799,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.proj = nn.Linear(128, gin_channels)
|
self.proj = nn.Linear(128, gin_channels)
|
||||||
|
|
||||||
def forward(self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
N = inputs.size(0)
|
N = inputs.size(0)
|
||||||
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
||||||
for conv in self.convs:
|
for conv in self.convs:
|
||||||
@@ -805,7 +819,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
|
|
||||||
return self.proj(out.squeeze(0))
|
return self.proj(out.squeeze(0))
|
||||||
|
|
||||||
def calculate_channels(self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int) -> int:
|
def calculate_channels(
|
||||||
|
self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int
|
||||||
|
) -> int:
|
||||||
for i in range(n_convs):
|
for i in range(n_convs):
|
||||||
L = (L - kernel_size + 2 * pad) // stride + 1
|
L = (L - kernel_size + 2 * pad) // stride + 1
|
||||||
return L
|
return L
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ class DurationDiscriminator(nn.Module): # vits2
|
|||||||
filter_channels: int,
|
filter_channels: int,
|
||||||
kernel_size: int,
|
kernel_size: int,
|
||||||
p_dropout: float,
|
p_dropout: float,
|
||||||
gin_channels: int = 0
|
gin_channels: int = 0,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -313,7 +313,9 @@ class DurationPredictor(nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
self.cond = nn.Conv1d(gin_channels, in_channels, 1)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = torch.detach(x)
|
x = torch.detach(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
g = torch.detach(g)
|
g = torch.detach(g)
|
||||||
@@ -587,7 +589,9 @@ class Generator(torch.nn.Module):
|
|||||||
if gin_channels != 0:
|
if gin_channels != 0:
|
||||||
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
x = self.conv_pre(x)
|
x = self.conv_pre(x)
|
||||||
if g is not None:
|
if g is not None:
|
||||||
x = x + self.cond(g)
|
x = x + self.cond(g)
|
||||||
@@ -618,7 +622,13 @@ class Generator(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class DiscriminatorP(torch.nn.Module):
|
class DiscriminatorP(torch.nn.Module):
|
||||||
def __init__(self, period: int, kernel_size: int = 5, stride: int = 3, use_spectral_norm: bool = False) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
period: int,
|
||||||
|
kernel_size: int = 5,
|
||||||
|
stride: int = 3,
|
||||||
|
use_spectral_norm: bool = False,
|
||||||
|
) -> None:
|
||||||
super(DiscriminatorP, self).__init__()
|
super(DiscriminatorP, self).__init__()
|
||||||
self.period = period
|
self.period = period
|
||||||
self.use_spectral_norm = use_spectral_norm
|
self.use_spectral_norm = use_spectral_norm
|
||||||
@@ -741,7 +751,9 @@ class MultiPeriodDiscriminator(torch.nn.Module):
|
|||||||
self,
|
self,
|
||||||
y: torch.Tensor,
|
y: torch.Tensor,
|
||||||
y_hat: torch.Tensor,
|
y_hat: torch.Tensor,
|
||||||
) -> tuple[list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]]:
|
) -> tuple[
|
||||||
|
list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]
|
||||||
|
]:
|
||||||
y_d_rs = []
|
y_d_rs = []
|
||||||
y_d_gs = []
|
y_d_gs = []
|
||||||
fmap_rs = []
|
fmap_rs = []
|
||||||
@@ -845,7 +857,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
)
|
)
|
||||||
self.proj = nn.Linear(128, gin_channels)
|
self.proj = nn.Linear(128, gin_channels)
|
||||||
|
|
||||||
def forward(self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
N = inputs.size(0)
|
N = inputs.size(0)
|
||||||
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
|
||||||
for conv in self.convs:
|
for conv in self.convs:
|
||||||
@@ -863,7 +877,9 @@ class ReferenceEncoder(nn.Module):
|
|||||||
|
|
||||||
return self.proj(out.squeeze(0))
|
return self.proj(out.squeeze(0))
|
||||||
|
|
||||||
def calculate_channels(self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int) -> int:
|
def calculate_channels(
|
||||||
|
self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int
|
||||||
|
) -> int:
|
||||||
for i in range(n_convs):
|
for i in range(n_convs):
|
||||||
L = (L - kernel_size + 2 * pad) // stride + 1
|
L = (L - kernel_size + 2 * pad) // stride + 1
|
||||||
return L
|
return L
|
||||||
|
|||||||
@@ -88,7 +88,9 @@ class DDSConv(nn.Module):
|
|||||||
Dialted and Depth-Separable Convolution
|
Dialted and Depth-Separable Convolution
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0) -> None:
|
def __init__(
|
||||||
|
self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0
|
||||||
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.channels = channels
|
self.channels = channels
|
||||||
self.kernel_size = kernel_size
|
self.kernel_size = kernel_size
|
||||||
@@ -117,7 +119,9 @@ class DDSConv(nn.Module):
|
|||||||
self.norms_1.append(LayerNorm(channels))
|
self.norms_1.append(LayerNorm(channels))
|
||||||
self.norms_2.append(LayerNorm(channels))
|
self.norms_2.append(LayerNorm(channels))
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
if g is not None:
|
if g is not None:
|
||||||
x = x + g
|
x = x + g
|
||||||
for i in range(self.n_layers):
|
for i in range(self.n_layers):
|
||||||
@@ -184,7 +188,13 @@ class WN(torch.nn.Module):
|
|||||||
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight")
|
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight")
|
||||||
self.res_skip_layers.append(res_skip_layer)
|
self.res_skip_layers.append(res_skip_layer)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, **kwargs: Any) -> torch.Tensor:
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
x_mask: torch.Tensor,
|
||||||
|
g: Optional[torch.Tensor] = None,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> torch.Tensor:
|
||||||
output = torch.zeros_like(x)
|
output = torch.zeros_like(x)
|
||||||
n_channels_tensor = torch.IntTensor([self.hidden_channels])
|
n_channels_tensor = torch.IntTensor([self.hidden_channels])
|
||||||
|
|
||||||
@@ -221,7 +231,12 @@ class WN(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class ResBlock1(torch.nn.Module):
|
class ResBlock1(torch.nn.Module):
|
||||||
def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int, int] = (1, 3, 5)) -> None:
|
def __init__(
|
||||||
|
self,
|
||||||
|
channels: int,
|
||||||
|
kernel_size: int = 3,
|
||||||
|
dilation: tuple[int, int, int] = (1, 3, 5),
|
||||||
|
) -> None:
|
||||||
super(ResBlock1, self).__init__()
|
super(ResBlock1, self).__init__()
|
||||||
self.convs1 = nn.ModuleList(
|
self.convs1 = nn.ModuleList(
|
||||||
[
|
[
|
||||||
@@ -295,7 +310,9 @@ class ResBlock1(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
self.convs2.apply(commons.init_weights)
|
self.convs2.apply(commons.init_weights)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
for c1, c2 in zip(self.convs1, self.convs2):
|
for c1, c2 in zip(self.convs1, self.convs2):
|
||||||
xt = F.leaky_relu(x, LRELU_SLOPE)
|
xt = F.leaky_relu(x, LRELU_SLOPE)
|
||||||
if x_mask is not None:
|
if x_mask is not None:
|
||||||
@@ -318,7 +335,9 @@ class ResBlock1(torch.nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class ResBlock2(torch.nn.Module):
|
class ResBlock2(torch.nn.Module):
|
||||||
def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)) -> None:
|
def __init__(
|
||||||
|
self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)
|
||||||
|
) -> None:
|
||||||
super(ResBlock2, self).__init__()
|
super(ResBlock2, self).__init__()
|
||||||
self.convs = nn.ModuleList(
|
self.convs = nn.ModuleList(
|
||||||
[
|
[
|
||||||
@@ -346,7 +365,9 @@ class ResBlock2(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
self.convs.apply(commons.init_weights)
|
self.convs.apply(commons.init_weights)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
def forward(
|
||||||
|
self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
for c in self.convs:
|
for c in self.convs:
|
||||||
xt = F.leaky_relu(x, LRELU_SLOPE)
|
xt = F.leaky_relu(x, LRELU_SLOPE)
|
||||||
if x_mask is not None:
|
if x_mask is not None:
|
||||||
|
|||||||
@@ -40,8 +40,8 @@ def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|||||||
numba.int32[::1],
|
numba.int32[::1],
|
||||||
numba.int32[::1],
|
numba.int32[::1],
|
||||||
),
|
),
|
||||||
nopython = True,
|
nopython=True,
|
||||||
nogil = True,
|
nogil=True,
|
||||||
) # type: ignore
|
) # type: ignore
|
||||||
def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None:
|
def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -39,12 +39,14 @@ def piecewise_rational_quadratic_transform(
|
|||||||
min_bin_width=min_bin_width,
|
min_bin_width=min_bin_width,
|
||||||
min_bin_height=min_bin_height,
|
min_bin_height=min_bin_height,
|
||||||
min_derivative=min_derivative,
|
min_derivative=min_derivative,
|
||||||
**spline_kwargs # type: ignore
|
**spline_kwargs, # type: ignore
|
||||||
)
|
)
|
||||||
return outputs, logabsdet
|
return outputs, logabsdet
|
||||||
|
|
||||||
|
|
||||||
def searchsorted(bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
|
def searchsorted(
|
||||||
|
bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6
|
||||||
|
) -> torch.Tensor:
|
||||||
bin_locations[..., -1] += eps
|
bin_locations[..., -1] += eps
|
||||||
return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1
|
return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1
|
||||||
|
|
||||||
|
|||||||
@@ -107,7 +107,9 @@ def plot_spectrogram_to_numpy(spectrogram: NDArray[Any]) -> NDArray[Any]:
|
|||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
def plot_alignment_to_numpy(alignment: NDArray[Any], info: Optional[str] = None) -> NDArray[Any]:
|
def plot_alignment_to_numpy(
|
||||||
|
alignment: NDArray[Any], info: Optional[str] = None
|
||||||
|
) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
指定されたアライメントを画像データに変換する
|
指定されたアライメントを画像データに変換する
|
||||||
|
|
||||||
@@ -163,7 +165,9 @@ def load_wav_to_torch(full_path: Union[str, Path]) -> tuple[torch.FloatTensor, i
|
|||||||
return torch.FloatTensor(data.astype(np.float32)), sampling_rate
|
return torch.FloatTensor(data.astype(np.float32)), sampling_rate
|
||||||
|
|
||||||
|
|
||||||
def load_filepaths_and_text(filename: Union[str, Path], split: str = "|") -> list[list[str]]:
|
def load_filepaths_and_text(
|
||||||
|
filename: Union[str, Path], split: str = "|"
|
||||||
|
) -> list[list[str]]:
|
||||||
"""
|
"""
|
||||||
指定されたファイルからファイルパスとテキストを読み込む
|
指定されたファイルからファイルパスとテキストを読み込む
|
||||||
|
|
||||||
@@ -180,7 +184,9 @@ def load_filepaths_and_text(filename: Union[str, Path], split: str = "|") -> lis
|
|||||||
return filepaths_and_text
|
return filepaths_and_text
|
||||||
|
|
||||||
|
|
||||||
def get_logger(model_dir_path: Union[str, Path], filename: str = "train.log") -> logging.Logger:
|
def get_logger(
|
||||||
|
model_dir_path: Union[str, Path], filename: str = "train.log"
|
||||||
|
) -> logging.Logger:
|
||||||
"""
|
"""
|
||||||
ロガーを取得する
|
ロガーを取得する
|
||||||
|
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ def load_checkpoint(
|
|||||||
model: torch.nn.Module,
|
model: torch.nn.Module,
|
||||||
optimizer: Optional[torch.optim.Optimizer] = None,
|
optimizer: Optional[torch.optim.Optimizer] = None,
|
||||||
skip_optimizer: bool = False,
|
skip_optimizer: bool = False,
|
||||||
for_infer: bool = False
|
for_infer: bool = False,
|
||||||
) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]:
|
) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]:
|
||||||
"""
|
"""
|
||||||
指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。
|
指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。
|
||||||
@@ -107,7 +107,9 @@ def save_checkpoint(
|
|||||||
iteration (int): イテレーション数
|
iteration (int): イテレーション数
|
||||||
checkpoint_path (Union[str, Path]): 保存先のパス
|
checkpoint_path (Union[str, Path]): 保存先のパス
|
||||||
"""
|
"""
|
||||||
logger.info(f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}")
|
logger.info(
|
||||||
|
f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}"
|
||||||
|
)
|
||||||
if hasattr(model, "module"):
|
if hasattr(model, "module"):
|
||||||
state_dict = model.module.state_dict()
|
state_dict = model.module.state_dict()
|
||||||
else:
|
else:
|
||||||
@@ -123,7 +125,11 @@ def save_checkpoint(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def clean_checkpoints(model_dir_path: Union[str, Path] = "logs/44k/", n_ckpts_to_keep: int = 2, sort_by_time: bool = True) -> None:
|
def clean_checkpoints(
|
||||||
|
model_dir_path: Union[str, Path] = "logs/44k/",
|
||||||
|
n_ckpts_to_keep: int = 2,
|
||||||
|
sort_by_time: bool = True,
|
||||||
|
) -> None:
|
||||||
"""
|
"""
|
||||||
指定されたディレクトリから古いチェックポイントを削除して空き容量を確保する
|
指定されたディレクトリから古いチェックポイントを削除して空き容量を確保する
|
||||||
|
|
||||||
@@ -172,7 +178,9 @@ def clean_checkpoints(model_dir_path: Union[str, Path] = "logs/44k/", n_ckpts_to
|
|||||||
[del_routine(fn) for fn in to_del]
|
[del_routine(fn) for fn in to_del]
|
||||||
|
|
||||||
|
|
||||||
def get_latest_checkpoint_path(model_dir_path: Union[str, Path], regex: str = "G_*.pth") -> str:
|
def get_latest_checkpoint_path(
|
||||||
|
model_dir_path: Union[str, Path], regex: str = "G_*.pth"
|
||||||
|
) -> str:
|
||||||
"""
|
"""
|
||||||
指定されたディレクトリから最新のチェックポイントのパスを取得する
|
指定されたディレクトリから最新のチェックポイントのパスを取得する
|
||||||
|
|
||||||
|
|||||||
@@ -74,16 +74,19 @@ def clean_text(
|
|||||||
if language == Languages.JP:
|
if language == Languages.JP:
|
||||||
from style_bert_vits2.nlp.japanese.g2p import g2p
|
from style_bert_vits2.nlp.japanese.g2p import g2p
|
||||||
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||||
|
|
||||||
norm_text = normalize_text(text)
|
norm_text = normalize_text(text)
|
||||||
phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error)
|
phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error)
|
||||||
elif language == Languages.EN:
|
elif language == Languages.EN:
|
||||||
from style_bert_vits2.nlp.english.g2p import g2p
|
from style_bert_vits2.nlp.english.g2p import g2p
|
||||||
from style_bert_vits2.nlp.english.normalizer import normalize_text
|
from style_bert_vits2.nlp.english.normalizer import normalize_text
|
||||||
|
|
||||||
norm_text = normalize_text(text)
|
norm_text = normalize_text(text)
|
||||||
phones, tones, word2ph = g2p(norm_text)
|
phones, tones, word2ph = g2p(norm_text)
|
||||||
elif language == Languages.ZH:
|
elif language == Languages.ZH:
|
||||||
from style_bert_vits2.nlp.chinese.g2p import g2p
|
from style_bert_vits2.nlp.chinese.g2p import g2p
|
||||||
from style_bert_vits2.nlp.chinese.normalizer import normalize_text
|
from style_bert_vits2.nlp.chinese.normalizer import normalize_text
|
||||||
|
|
||||||
norm_text = normalize_text(text)
|
norm_text = normalize_text(text)
|
||||||
phones, tones, word2ph = g2p(norm_text)
|
phones, tones, word2ph = g2p(norm_text)
|
||||||
else:
|
else:
|
||||||
@@ -92,7 +95,9 @@ def clean_text(
|
|||||||
return norm_text, phones, tones, word2ph
|
return norm_text, phones, tones, word2ph
|
||||||
|
|
||||||
|
|
||||||
def cleaned_text_to_sequence(cleaned_phones: list[str], tones: list[int], language: Languages) -> tuple[list[int], list[int], list[int]]:
|
def cleaned_text_to_sequence(
|
||||||
|
cleaned_phones: list[str], tones: list[int], language: Languages
|
||||||
|
) -> tuple[list[int], list[int], list[int]]:
|
||||||
"""
|
"""
|
||||||
テキスト文字列を、テキスト内の記号に対応する一連の ID に変換する
|
テキスト文字列を、テキスト内の記号に対応する一連の ID に変換する
|
||||||
|
|
||||||
|
|||||||
@@ -30,7 +30,9 @@ from style_bert_vits2.logging import logger
|
|||||||
__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
||||||
|
|
||||||
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||||
__loaded_tokenizers: dict[Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]] = {}
|
__loaded_tokenizers: dict[
|
||||||
|
Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]
|
||||||
|
] = {}
|
||||||
|
|
||||||
|
|
||||||
def load_model(
|
def load_model(
|
||||||
@@ -63,18 +65,24 @@ def load_model(
|
|||||||
|
|
||||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \
|
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
||||||
f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model."
|
language
|
||||||
|
].exists(), f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model."
|
||||||
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
||||||
|
|
||||||
# BERT モデルをロードし、辞書に格納して返す
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
## 英語のみ DebertaV2Model でロードする必要がある
|
## 英語のみ DebertaV2Model でロードする必要がある
|
||||||
if language == Languages.EN:
|
if language == Languages.EN:
|
||||||
model = cast(DebertaV2Model, DebertaV2Model.from_pretrained(pretrained_model_name_or_path))
|
model = cast(
|
||||||
|
DebertaV2Model,
|
||||||
|
DebertaV2Model.from_pretrained(pretrained_model_name_or_path),
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path)
|
model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path)
|
||||||
__loaded_models[language] = model
|
__loaded_models[language] = model
|
||||||
logger.info(f"Loaded the {language} BERT model from {pretrained_model_name_or_path}")
|
logger.info(
|
||||||
|
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
||||||
|
)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
@@ -109,8 +117,9 @@ def load_tokenizer(
|
|||||||
|
|
||||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \
|
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
||||||
f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model."
|
language
|
||||||
|
].exists(), f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model."
|
||||||
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
||||||
|
|
||||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
@@ -120,7 +129,9 @@ def load_tokenizer(
|
|||||||
else:
|
else:
|
||||||
tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path)
|
tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path)
|
||||||
__loaded_tokenizers[language] = tokenizer
|
__loaded_tokenizers[language] = tokenizer
|
||||||
logger.info(f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}")
|
logger.info(
|
||||||
|
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
||||||
|
)
|
||||||
|
|
||||||
return tokenizer
|
return tokenizer
|
||||||
|
|
||||||
|
|||||||
@@ -95,7 +95,11 @@ def __g2p(segments: list[str]) -> tuple[list[str], list[int], list[int]]:
|
|||||||
if pinyin[0] in single_rep_map.keys():
|
if pinyin[0] in single_rep_map.keys():
|
||||||
pinyin = single_rep_map[pinyin[0]] + pinyin[1:]
|
pinyin = single_rep_map[pinyin[0]] + pinyin[1:]
|
||||||
|
|
||||||
assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (pinyin, seg, raw_pinyin)
|
assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (
|
||||||
|
pinyin,
|
||||||
|
seg,
|
||||||
|
raw_pinyin,
|
||||||
|
)
|
||||||
phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ")
|
phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ")
|
||||||
word2ph.append(len(phone))
|
word2ph.append(len(phone))
|
||||||
|
|
||||||
@@ -125,7 +129,7 @@ if __name__ == "__main__":
|
|||||||
text = normalize_text(text)
|
text = normalize_text(text)
|
||||||
print(text)
|
print(text)
|
||||||
phones, tones, word2ph = g2p(text)
|
phones, tones, word2ph = g2p(text)
|
||||||
bert = extract_bert_feature(text, word2ph, 'cuda')
|
bert = extract_bert_feature(text, word2ph, "cuda")
|
||||||
|
|
||||||
print(phones, tones, word2ph, bert.shape)
|
print(phones, tones, word2ph, bert.shape)
|
||||||
|
|
||||||
|
|||||||
@@ -121,7 +121,9 @@ def __expand_number(m: re.Match[str]) -> str:
|
|||||||
else:
|
else:
|
||||||
return __INFLECT.number_to_words(
|
return __INFLECT.number_to_words(
|
||||||
num, andword="", zero="oh", group=2 # type: ignore
|
num, andword="", zero="oh", group=2 # type: ignore
|
||||||
).replace(", ", " ") # type: ignore
|
).replace(
|
||||||
|
", ", " "
|
||||||
|
) # type: ignore
|
||||||
else:
|
else:
|
||||||
return __INFLECT.number_to_words(num, andword="") # type: ignore
|
return __INFLECT.number_to_words(num, andword="") # type: ignore
|
||||||
|
|
||||||
|
|||||||
@@ -10,9 +10,7 @@ from style_bert_vits2.nlp.symbols import PUNCTUATIONS
|
|||||||
|
|
||||||
|
|
||||||
def g2p(
|
def g2p(
|
||||||
norm_text: str,
|
norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False
|
||||||
use_jp_extra: bool = True,
|
|
||||||
raise_yomi_error: bool = False
|
|
||||||
) -> tuple[list[str], list[int], list[int]]:
|
) -> tuple[list[str], list[int], list[int]]:
|
||||||
"""
|
"""
|
||||||
他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、
|
他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、
|
||||||
@@ -93,8 +91,7 @@ def g2p(
|
|||||||
|
|
||||||
|
|
||||||
def text_to_sep_kata(
|
def text_to_sep_kata(
|
||||||
norm_text: str,
|
norm_text: str, raise_yomi_error: bool = False
|
||||||
raise_yomi_error: bool = False
|
|
||||||
) -> tuple[list[str], list[str]]:
|
) -> tuple[list[str], list[str]]:
|
||||||
"""
|
"""
|
||||||
`normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、
|
`normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、
|
||||||
@@ -212,7 +209,9 @@ def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]:
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def __pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]:
|
def __pyopenjtalk_g2p_prosody(
|
||||||
|
text: str, drop_unvoiced_vowels: bool = True
|
||||||
|
) -> list[str]:
|
||||||
"""
|
"""
|
||||||
ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。
|
ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。
|
||||||
ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py
|
ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py
|
||||||
@@ -414,8 +413,7 @@ def __kata_to_phoneme_list(text: str) -> list[str]:
|
|||||||
|
|
||||||
|
|
||||||
def __align_tones(
|
def __align_tones(
|
||||||
phones_with_punct: list[str],
|
phones_with_punct: list[str], phone_tone_list: list[tuple[str, int]]
|
||||||
phone_tone_list: list[tuple[str, int]]
|
|
||||||
) -> list[tuple[str, int]]:
|
) -> list[tuple[str, int]]:
|
||||||
"""
|
"""
|
||||||
例: …私は、、そう思う。
|
例: …私は、、そう思う。
|
||||||
|
|||||||
@@ -34,11 +34,13 @@ def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, i
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
# 子音の集合
|
# 子音の集合
|
||||||
CONSONANTS = set([
|
CONSONANTS = set(
|
||||||
|
[
|
||||||
consonant
|
consonant
|
||||||
for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values()
|
for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values()
|
||||||
if consonant is not None
|
if consonant is not None
|
||||||
])
|
]
|
||||||
|
)
|
||||||
|
|
||||||
phone_tone = phone_tone[1:] # 最初の("_", 0)を無視
|
phone_tone = phone_tone[1:] # 最初の("_", 0)を無視
|
||||||
phones = [phone for phone, _ in phone_tone]
|
phones = [phone for phone, _ in phone_tone]
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ def run_frontend(text: str) -> list[dict[str, Any]]:
|
|||||||
else:
|
else:
|
||||||
# without worker
|
# without worker
|
||||||
import pyopenjtalk
|
import pyopenjtalk
|
||||||
|
|
||||||
return pyopenjtalk.run_frontend(text)
|
return pyopenjtalk.run_frontend(text)
|
||||||
|
|
||||||
|
|
||||||
@@ -36,6 +37,7 @@ def make_label(njd_features: Any) -> list[str]:
|
|||||||
else:
|
else:
|
||||||
# without worker
|
# without worker
|
||||||
import pyopenjtalk
|
import pyopenjtalk
|
||||||
|
|
||||||
return pyopenjtalk.make_label(njd_features)
|
return pyopenjtalk.make_label(njd_features)
|
||||||
|
|
||||||
|
|
||||||
@@ -45,6 +47,7 @@ def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None) -
|
|||||||
else:
|
else:
|
||||||
# without worker
|
# without worker
|
||||||
import pyopenjtalk
|
import pyopenjtalk
|
||||||
|
|
||||||
pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab)
|
pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab)
|
||||||
|
|
||||||
|
|
||||||
@@ -54,6 +57,7 @@ def update_global_jtalk_with_user_dict(path: str) -> None:
|
|||||||
else:
|
else:
|
||||||
# without worker
|
# without worker
|
||||||
import pyopenjtalk
|
import pyopenjtalk
|
||||||
|
|
||||||
pyopenjtalk.update_global_jtalk_with_user_dict(path)
|
pyopenjtalk.update_global_jtalk_with_user_dict(path)
|
||||||
|
|
||||||
|
|
||||||
@@ -63,6 +67,7 @@ def unset_user_dict() -> None:
|
|||||||
else:
|
else:
|
||||||
# without worker
|
# without worker
|
||||||
import pyopenjtalk
|
import pyopenjtalk
|
||||||
|
|
||||||
pyopenjtalk.unset_user_dict()
|
pyopenjtalk.unset_user_dict()
|
||||||
|
|
||||||
|
|
||||||
@@ -102,7 +107,12 @@ def initialize_worker(port: int = WORKER_PORT) -> None:
|
|||||||
else:
|
else:
|
||||||
# align with Windows behavior
|
# align with Windows behavior
|
||||||
# start_new_session is same as specifying setsid in preexec_fn
|
# start_new_session is same as specifying setsid in preexec_fn
|
||||||
subprocess.Popen(args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, start_new_session=True)
|
subprocess.Popen(
|
||||||
|
args,
|
||||||
|
stdout=subprocess.DEVNULL,
|
||||||
|
stderr=subprocess.DEVNULL,
|
||||||
|
start_new_session=True,
|
||||||
|
)
|
||||||
|
|
||||||
# wait until server listening
|
# wait until server listening
|
||||||
count = 0
|
count = 0
|
||||||
|
|||||||
@@ -2,12 +2,15 @@ import socket
|
|||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import RequestType, receive_data, send_data
|
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import (
|
||||||
|
RequestType,
|
||||||
|
receive_data,
|
||||||
|
send_data,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class WorkerClient:
|
class WorkerClient:
|
||||||
""" pyopenjtalk worker client """
|
"""pyopenjtalk worker client"""
|
||||||
|
|
||||||
|
|
||||||
def __init__(self, port: int) -> None:
|
def __init__(self, port: int) -> None:
|
||||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||||
@@ -16,19 +19,15 @@ class WorkerClient:
|
|||||||
sock.connect((socket.gethostname(), port))
|
sock.connect((socket.gethostname(), port))
|
||||||
self.sock = sock
|
self.sock = sock
|
||||||
|
|
||||||
|
|
||||||
def __enter__(self) -> "WorkerClient":
|
def __enter__(self) -> "WorkerClient":
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
||||||
def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
|
def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
|
||||||
self.close()
|
self.close()
|
||||||
|
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
self.sock.close()
|
self.sock.close()
|
||||||
|
|
||||||
|
|
||||||
def dispatch_pyopenjtalk(self, func: str, *args: Any, **kwargs: Any) -> Any:
|
def dispatch_pyopenjtalk(self, func: str, *args: Any, **kwargs: Any) -> Any:
|
||||||
data = {
|
data = {
|
||||||
"request-type": RequestType.PYOPENJTALK,
|
"request-type": RequestType.PYOPENJTALK,
|
||||||
@@ -43,7 +42,6 @@ class WorkerClient:
|
|||||||
logger.trace(f"client received response: {response}")
|
logger.trace(f"client received response: {response}")
|
||||||
return response.get("return")
|
return response.get("return")
|
||||||
|
|
||||||
|
|
||||||
def status(self) -> int:
|
def status(self) -> int:
|
||||||
data = {"request-type": RequestType.STATUS}
|
data = {"request-type": RequestType.STATUS}
|
||||||
logger.trace(f"client sends request: {data}")
|
logger.trace(f"client sends request: {data}")
|
||||||
@@ -53,7 +51,6 @@ class WorkerClient:
|
|||||||
logger.trace(f"client received response: {response}")
|
logger.trace(f"client received response: {response}")
|
||||||
return cast(int, response.get("client-count"))
|
return cast(int, response.get("client-count"))
|
||||||
|
|
||||||
|
|
||||||
def quit_server(self) -> None:
|
def quit_server(self) -> None:
|
||||||
data = {"request-type": RequestType.QUIT_SERVER}
|
data = {"request-type": RequestType.QUIT_SERVER}
|
||||||
logger.trace(f"client sends request: {data}")
|
logger.trace(f"client sends request: {data}")
|
||||||
|
|||||||
@@ -26,14 +26,12 @@ PYOPENJTALK_FUNC_DICT = {
|
|||||||
|
|
||||||
|
|
||||||
class WorkerServer:
|
class WorkerServer:
|
||||||
""" pyopenjtalk worker server """
|
"""pyopenjtalk worker server"""
|
||||||
|
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self.client_count: int = 0
|
self.client_count: int = 0
|
||||||
self.quit: bool = False
|
self.quit: bool = False
|
||||||
|
|
||||||
|
|
||||||
def handle_request(self, request: dict[str, Any]) -> dict[str, Any]:
|
def handle_request(self, request: dict[str, Any]) -> dict[str, Any]:
|
||||||
request_type = None
|
request_type = None
|
||||||
try:
|
try:
|
||||||
@@ -70,7 +68,6 @@ class WorkerServer:
|
|||||||
|
|
||||||
return response
|
return response
|
||||||
|
|
||||||
|
|
||||||
def start_server(self, port: int, no_client_timeout: int = 30) -> None:
|
def start_server(self, port: int, no_client_timeout: int = 30) -> None:
|
||||||
logger.info("start pyopenjtalk worker server")
|
logger.info("start pyopenjtalk worker server")
|
||||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket:
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket:
|
||||||
|
|||||||
@@ -18,7 +18,11 @@ from fastapi import HTTPException
|
|||||||
from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR
|
from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR
|
||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes
|
from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes
|
||||||
from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data
|
from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import (
|
||||||
|
MAX_PRIORITY,
|
||||||
|
MIN_PRIORITY,
|
||||||
|
part_of_speech_data,
|
||||||
|
)
|
||||||
|
|
||||||
# root_dir = engine_root()
|
# root_dir = engine_root()
|
||||||
# save_dir = get_save_dir()
|
# save_dir = get_save_dir()
|
||||||
@@ -26,9 +30,13 @@ from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIO
|
|||||||
# if not save_dir.is_dir():
|
# if not save_dir.is_dir():
|
||||||
# save_dir.mkdir(parents=True)
|
# save_dir.mkdir(parents=True)
|
||||||
|
|
||||||
default_dict_path = DEFAULT_USER_DICT_DIR / "default.csv" # VOICEVOXデフォルト辞書ファイルのパス
|
default_dict_path = (
|
||||||
|
DEFAULT_USER_DICT_DIR / "default.csv"
|
||||||
|
) # VOICEVOXデフォルト辞書ファイルのパス
|
||||||
user_dict_path = DEFAULT_USER_DICT_DIR / "user_dict.json" # ユーザー辞書ファイルのパス
|
user_dict_path = DEFAULT_USER_DICT_DIR / "user_dict.json" # ユーザー辞書ファイルのパス
|
||||||
compiled_dict_path = DEFAULT_USER_DICT_DIR / "user.dic" # コンパイル済み辞書ファイルのパス
|
compiled_dict_path = (
|
||||||
|
DEFAULT_USER_DICT_DIR / "user.dic"
|
||||||
|
) # コンパイル済み辞書ファイルのパス
|
||||||
|
|
||||||
|
|
||||||
# # 同時書き込みの制御
|
# # 同時書き込みの制御
|
||||||
|
|||||||
@@ -25,7 +25,9 @@ from style_bert_vits2.constants import (
|
|||||||
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
from style_bert_vits2.models.hyper_parameters import HyperParameters
|
||||||
from style_bert_vits2.models.infer import get_net_g, infer
|
from style_bert_vits2.models.infer import get_net_g, infer
|
||||||
from style_bert_vits2.models.models import SynthesizerTrn
|
from style_bert_vits2.models.models import SynthesizerTrn
|
||||||
from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra
|
from style_bert_vits2.models.models_jp_extra import (
|
||||||
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.voice import adjust_voice
|
from style_bert_vits2.voice import adjust_voice
|
||||||
|
|
||||||
@@ -36,7 +38,6 @@ class TTSModel:
|
|||||||
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。
|
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model_path: Path,
|
model_path: Path,
|
||||||
@@ -59,7 +60,9 @@ class TTSModel:
|
|||||||
self.config_path: Path = config_path
|
self.config_path: Path = config_path
|
||||||
self.style_vec_path: Path = style_vec_path
|
self.style_vec_path: Path = style_vec_path
|
||||||
self.device: str = device
|
self.device: str = device
|
||||||
self.hyper_parameters: HyperParameters = HyperParameters.load_from_json(self.config_path)
|
self.hyper_parameters: HyperParameters = HyperParameters.load_from_json(
|
||||||
|
self.config_path
|
||||||
|
)
|
||||||
self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id
|
self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id
|
||||||
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()}
|
||||||
|
|
||||||
@@ -82,19 +85,17 @@ class TTSModel:
|
|||||||
|
|
||||||
self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
||||||
|
|
||||||
|
|
||||||
def load(self) -> None:
|
def load(self) -> None:
|
||||||
"""
|
"""
|
||||||
音声合成モデルをデバイスにロードする。
|
音声合成モデルをデバイスにロードする。
|
||||||
"""
|
"""
|
||||||
self.__net_g = get_net_g(
|
self.__net_g = get_net_g(
|
||||||
model_path = str(self.model_path),
|
model_path=str(self.model_path),
|
||||||
version = self.hyper_parameters.version,
|
version=self.hyper_parameters.version,
|
||||||
device = self.device,
|
device=self.device,
|
||||||
hps = self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
|
def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
スタイルベクトルを取得する。
|
スタイルベクトルを取得する。
|
||||||
@@ -111,8 +112,9 @@ class TTSModel:
|
|||||||
style_vec = mean + (style_vec - mean) * weight
|
style_vec = mean + (style_vec - mean) * weight
|
||||||
return style_vec
|
return style_vec
|
||||||
|
|
||||||
|
def __get_style_vector_from_audio(
|
||||||
def __get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]:
|
self, audio_path: str, weight: float = 1.0
|
||||||
|
) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
音声からスタイルベクトルを推論する。
|
音声からスタイルベクトルを推論する。
|
||||||
|
|
||||||
@@ -126,8 +128,10 @@ class TTSModel:
|
|||||||
# スタイルベクトルを取得するための推論モデルを初期化
|
# スタイルベクトルを取得するための推論モデルを初期化
|
||||||
if self.__style_vector_inference is None:
|
if self.__style_vector_inference is None:
|
||||||
self.__style_vector_inference = pyannote.audio.Inference(
|
self.__style_vector_inference = pyannote.audio.Inference(
|
||||||
model = pyannote.audio.Model.from_pretrained("pyannote/wespeaker-voxceleb-resnet34-LM"),
|
model=pyannote.audio.Model.from_pretrained(
|
||||||
window = "whole",
|
"pyannote/wespeaker-voxceleb-resnet34-LM"
|
||||||
|
),
|
||||||
|
window="whole",
|
||||||
)
|
)
|
||||||
self.__style_vector_inference.to(torch.device(self.device))
|
self.__style_vector_inference.to(torch.device(self.device))
|
||||||
|
|
||||||
@@ -137,7 +141,6 @@ class TTSModel:
|
|||||||
xvec = mean + (xvec - mean) * weight
|
xvec = mean + (xvec - mean) * weight
|
||||||
return xvec
|
return xvec
|
||||||
|
|
||||||
|
|
||||||
def infer(
|
def infer(
|
||||||
self,
|
self,
|
||||||
text: str,
|
text: str,
|
||||||
@@ -209,20 +212,20 @@ class TTSModel:
|
|||||||
if not line_split:
|
if not line_split:
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
audio = infer(
|
audio = infer(
|
||||||
text = text,
|
text=text,
|
||||||
sdp_ratio = sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale = noise,
|
noise_scale=noise,
|
||||||
noise_scale_w = noise_w,
|
noise_scale_w=noise_w,
|
||||||
length_scale = length,
|
length_scale=length,
|
||||||
sid = speaker_id,
|
sid=speaker_id,
|
||||||
language = language,
|
language=language,
|
||||||
hps = self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
net_g = self.__net_g,
|
net_g=self.__net_g,
|
||||||
device = self.device,
|
device=self.device,
|
||||||
assist_text = assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight = assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
style_vec = style_vector,
|
style_vec=style_vector,
|
||||||
given_tone = given_tone,
|
given_tone=given_tone,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
texts = text.split("\n")
|
texts = text.split("\n")
|
||||||
@@ -232,19 +235,19 @@ class TTSModel:
|
|||||||
for i, t in enumerate(texts):
|
for i, t in enumerate(texts):
|
||||||
audios.append(
|
audios.append(
|
||||||
infer(
|
infer(
|
||||||
text = t,
|
text=t,
|
||||||
sdp_ratio = sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale = noise,
|
noise_scale=noise,
|
||||||
noise_scale_w = noise_w,
|
noise_scale_w=noise_w,
|
||||||
length_scale = length,
|
length_scale=length,
|
||||||
sid = speaker_id,
|
sid=speaker_id,
|
||||||
language = language,
|
language=language,
|
||||||
hps = self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
net_g = self.__net_g,
|
net_g=self.__net_g,
|
||||||
device = self.device,
|
device=self.device,
|
||||||
assist_text = assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight = assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
style_vec = style_vector,
|
style_vec=style_vector,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
if i != len(texts) - 1:
|
if i != len(texts) - 1:
|
||||||
@@ -253,10 +256,10 @@ class TTSModel:
|
|||||||
logger.info("Audio data generated successfully")
|
logger.info("Audio data generated successfully")
|
||||||
if not (pitch_scale == 1.0 and intonation_scale == 1.0):
|
if not (pitch_scale == 1.0 and intonation_scale == 1.0):
|
||||||
_, audio = adjust_voice(
|
_, audio = adjust_voice(
|
||||||
fs = self.hyper_parameters.data.sampling_rate,
|
fs=self.hyper_parameters.data.sampling_rate,
|
||||||
wave = audio,
|
wave=audio,
|
||||||
pitch_scale = pitch_scale,
|
pitch_scale=pitch_scale,
|
||||||
intonation_scale = intonation_scale,
|
intonation_scale=intonation_scale,
|
||||||
)
|
)
|
||||||
with warnings.catch_warnings():
|
with warnings.catch_warnings():
|
||||||
warnings.simplefilter("ignore")
|
warnings.simplefilter("ignore")
|
||||||
@@ -277,7 +280,6 @@ class TTSModelHolder:
|
|||||||
model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。
|
model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
def __init__(self, model_root_dir: Path, device: str) -> None:
|
def __init__(self, model_root_dir: Path, device: str) -> None:
|
||||||
"""
|
"""
|
||||||
Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。
|
Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。
|
||||||
@@ -308,7 +310,6 @@ class TTSModelHolder:
|
|||||||
self.models_info: list[TTSModelInfo] = []
|
self.models_info: list[TTSModelInfo] = []
|
||||||
self.refresh()
|
self.refresh()
|
||||||
|
|
||||||
|
|
||||||
def refresh(self) -> None:
|
def refresh(self) -> None:
|
||||||
"""
|
"""
|
||||||
音声合成モデルの一覧を更新する。
|
音声合成モデルの一覧を更新する。
|
||||||
@@ -342,13 +343,14 @@ class TTSModelHolder:
|
|||||||
styles = list(style2id.keys())
|
styles = list(style2id.keys())
|
||||||
spk2id: dict[str, int] = hyper_parameters.data.spk2id
|
spk2id: dict[str, int] = hyper_parameters.data.spk2id
|
||||||
speakers = list(spk2id.keys())
|
speakers = list(spk2id.keys())
|
||||||
self.models_info.append(TTSModelInfo(
|
self.models_info.append(
|
||||||
name = model_dir.name,
|
TTSModelInfo(
|
||||||
files = [str(f) for f in model_files],
|
name=model_dir.name,
|
||||||
styles = styles,
|
files=[str(f) for f in model_files],
|
||||||
speakers = speakers,
|
styles=styles,
|
||||||
))
|
speakers=speakers,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def get_model(self, model_name: str, model_path_str: str) -> TTSModel:
|
def get_model(self, model_name: str, model_path_str: str) -> TTSModel:
|
||||||
"""
|
"""
|
||||||
@@ -370,16 +372,17 @@ class TTSModelHolder:
|
|||||||
raise ValueError(f"Model file `{model_path}` is not found")
|
raise ValueError(f"Model file `{model_path}` is not found")
|
||||||
if self.current_model is None or self.current_model.model_path != model_path:
|
if self.current_model is None or self.current_model.model_path != model_path:
|
||||||
self.current_model = TTSModel(
|
self.current_model = TTSModel(
|
||||||
model_path = model_path,
|
model_path=model_path,
|
||||||
config_path = self.root_dir / model_name / "config.json",
|
config_path=self.root_dir / model_name / "config.json",
|
||||||
style_vec_path = self.root_dir / model_name / "style_vectors.npy",
|
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
||||||
device = self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
return self.current_model
|
return self.current_model
|
||||||
|
|
||||||
|
def get_model_for_gradio(
|
||||||
def get_model_for_gradio(self, model_name: str, model_path_str: str) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]:
|
self, model_name: str, model_path_str: str
|
||||||
|
) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]:
|
||||||
model_path = Path(model_path_str)
|
model_path = Path(model_path_str)
|
||||||
if model_name not in self.model_files_dict:
|
if model_name not in self.model_files_dict:
|
||||||
raise ValueError(f"Model `{model_name}` is not found")
|
raise ValueError(f"Model `{model_name}` is not found")
|
||||||
@@ -398,10 +401,10 @@ class TTSModelHolder:
|
|||||||
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
||||||
)
|
)
|
||||||
self.current_model = TTSModel(
|
self.current_model = TTSModel(
|
||||||
model_path = model_path,
|
model_path=model_path,
|
||||||
config_path = self.root_dir / model_name / "config.json",
|
config_path=self.root_dir / model_name / "config.json",
|
||||||
style_vec_path = self.root_dir / model_name / "style_vectors.npy",
|
style_vec_path=self.root_dir / model_name / "style_vectors.npy",
|
||||||
device = self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
speakers = list(self.current_model.spk2id.keys())
|
speakers = list(self.current_model.spk2id.keys())
|
||||||
styles = list(self.current_model.style2id.keys())
|
styles = list(self.current_model.style2id.keys())
|
||||||
@@ -411,13 +414,13 @@ class TTSModelHolder:
|
|||||||
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown:
|
def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown:
|
||||||
model_files = self.model_files_dict[model_name]
|
model_files = self.model_files_dict[model_name]
|
||||||
return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore
|
return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore
|
||||||
|
|
||||||
|
def update_model_names_for_gradio(
|
||||||
def update_model_names_for_gradio(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]:
|
self,
|
||||||
|
) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]:
|
||||||
self.refresh()
|
self.refresh()
|
||||||
initial_model_name = self.model_names[0]
|
initial_model_name = self.model_names[0]
|
||||||
initial_model_files = self.model_files_dict[initial_model_name]
|
initial_model_files = self.model_files_dict[initial_model_name]
|
||||||
|
|||||||
@@ -8,40 +8,35 @@ class StdoutWrapper(TextIO):
|
|||||||
`sys.stdout` wrapper for both Google Colab and local environment.
|
`sys.stdout` wrapper for both Google Colab and local environment.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self.temp_file = tempfile.NamedTemporaryFile(
|
self.temp_file = tempfile.NamedTemporaryFile(
|
||||||
mode="w+", delete=False, encoding="utf-8"
|
mode="w+", delete=False, encoding="utf-8"
|
||||||
)
|
)
|
||||||
self.original_stdout = sys.stdout
|
self.original_stdout = sys.stdout
|
||||||
|
|
||||||
|
|
||||||
def write(self, message: str) -> int:
|
def write(self, message: str) -> int:
|
||||||
result = self.temp_file.write(message)
|
result = self.temp_file.write(message)
|
||||||
self.temp_file.flush()
|
self.temp_file.flush()
|
||||||
print(message, end="", file=self.original_stdout)
|
print(message, end="", file=self.original_stdout)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def flush(self) -> None:
|
def flush(self) -> None:
|
||||||
self.temp_file.flush()
|
self.temp_file.flush()
|
||||||
|
|
||||||
|
|
||||||
def read(self, n: int = -1) -> str:
|
def read(self, n: int = -1) -> str:
|
||||||
self.temp_file.seek(0)
|
self.temp_file.seek(0)
|
||||||
return self.temp_file.read(n)
|
return self.temp_file.read(n)
|
||||||
|
|
||||||
|
|
||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
self.temp_file.close()
|
self.temp_file.close()
|
||||||
|
|
||||||
|
|
||||||
def fileno(self) -> int:
|
def fileno(self) -> int:
|
||||||
return self.temp_file.fileno()
|
return self.temp_file.fileno()
|
||||||
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import google.colab # type: ignore
|
import google.colab # type: ignore
|
||||||
|
|
||||||
SAFE_STDOUT = StdoutWrapper()
|
SAFE_STDOUT = StdoutWrapper()
|
||||||
except ImportError:
|
except ImportError:
|
||||||
SAFE_STDOUT = sys.stdout
|
SAFE_STDOUT = sys.stdout
|
||||||
|
|||||||
@@ -9,27 +9,28 @@ class StrEnum(str, enum.Enum):
|
|||||||
def __new__(cls, *values: str) -> "StrEnum":
|
def __new__(cls, *values: str) -> "StrEnum":
|
||||||
"values must already be of type `str`"
|
"values must already be of type `str`"
|
||||||
if len(values) > 3:
|
if len(values) > 3:
|
||||||
raise TypeError('too many arguments for str(): %r' % (values, ))
|
raise TypeError("too many arguments for str(): %r" % (values,))
|
||||||
if len(values) == 1:
|
if len(values) == 1:
|
||||||
# it must be a string
|
# it must be a string
|
||||||
if not isinstance(values[0], str): # type: ignore
|
if not isinstance(values[0], str): # type: ignore
|
||||||
raise TypeError('%r is not a string' % (values[0], ))
|
raise TypeError("%r is not a string" % (values[0],))
|
||||||
if len(values) >= 2:
|
if len(values) >= 2:
|
||||||
# check that encoding argument is a string
|
# check that encoding argument is a string
|
||||||
if not isinstance(values[1], str): # type: ignore
|
if not isinstance(values[1], str): # type: ignore
|
||||||
raise TypeError('encoding must be a string, not %r' % (values[1], ))
|
raise TypeError("encoding must be a string, not %r" % (values[1],))
|
||||||
if len(values) == 3:
|
if len(values) == 3:
|
||||||
# check that errors argument is a string
|
# check that errors argument is a string
|
||||||
if not isinstance(values[2], str): # type: ignore
|
if not isinstance(values[2], str): # type: ignore
|
||||||
raise TypeError('errors must be a string, not %r' % (values[2]))
|
raise TypeError("errors must be a string, not %r" % (values[2]))
|
||||||
value = str(*values)
|
value = str(*values)
|
||||||
member = str.__new__(cls, value)
|
member = str.__new__(cls, value)
|
||||||
member._value_ = value
|
member._value_ = value
|
||||||
return member
|
return member
|
||||||
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _generate_next_value_(name: str, start: int, count: int, last_values: list[str]) -> str:
|
def _generate_next_value_(
|
||||||
|
name: str, start: int, count: int, last_values: list[str]
|
||||||
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Return the lower-cased version of the member name.
|
Return the lower-cased version of the member name.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -6,7 +6,9 @@ from style_bert_vits2.logging import logger
|
|||||||
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
|
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
|
||||||
|
|
||||||
|
|
||||||
def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[bool, str]:
|
def run_script_with_log(
|
||||||
|
cmd: list[str], ignore_warning: bool = False
|
||||||
|
) -> tuple[bool, str]:
|
||||||
"""
|
"""
|
||||||
指定されたコマンドを実行し、そのログを記録する。
|
指定されたコマンドを実行し、そのログを記録する。
|
||||||
|
|
||||||
@@ -21,10 +23,10 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b
|
|||||||
logger.info(f"Running: {' '.join(cmd)}")
|
logger.info(f"Running: {' '.join(cmd)}")
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
[sys.executable] + cmd,
|
[sys.executable] + cmd,
|
||||||
stdout = SAFE_STDOUT,
|
stdout=SAFE_STDOUT,
|
||||||
stderr = subprocess.PIPE,
|
stderr=subprocess.PIPE,
|
||||||
text = True,
|
text=True,
|
||||||
encoding = "utf-8",
|
encoding="utf-8",
|
||||||
)
|
)
|
||||||
if result.returncode != 0:
|
if result.returncode != 0:
|
||||||
logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}")
|
logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}")
|
||||||
@@ -37,7 +39,9 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b
|
|||||||
return True, ""
|
return True, ""
|
||||||
|
|
||||||
|
|
||||||
def second_elem_of(original_function: Callable[..., tuple[Any, Any]]) -> Callable[..., Any]:
|
def second_elem_of(
|
||||||
|
original_function: Callable[..., tuple[Any, Any]]
|
||||||
|
) -> Callable[..., Any]:
|
||||||
"""
|
"""
|
||||||
与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する。
|
与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する。
|
||||||
|
|
||||||
|
|||||||
@@ -5,15 +5,15 @@ from style_bert_vits2.constants import BASE_DIR, Languages
|
|||||||
from style_bert_vits2.tts_model import TTSModelHolder
|
from style_bert_vits2.tts_model import TTSModelHolder
|
||||||
|
|
||||||
|
|
||||||
def synthesize(device: str = 'cpu'):
|
def synthesize(device: str = "cpu"):
|
||||||
|
|
||||||
# 音声合成モデルが配置されていれば、音声合成を実行
|
# 音声合成モデルが配置されていれば、音声合成を実行
|
||||||
model_holder = TTSModelHolder(BASE_DIR / 'model_assets', device)
|
model_holder = TTSModelHolder(BASE_DIR / "model_assets", device)
|
||||||
if len(model_holder.models_info) > 0:
|
if len(model_holder.models_info) > 0:
|
||||||
|
|
||||||
# jvnv-F2-jp モデルを探す
|
# jvnv-F2-jp モデルを探す
|
||||||
for model_info in model_holder.models_info:
|
for model_info in model_holder.models_info:
|
||||||
if model_info.name == 'jvnv-F2-jp':
|
if model_info.name == "jvnv-F2-jp":
|
||||||
# すべてのスタイルに対して音声合成を実行
|
# すべてのスタイルに対して音声合成を実行
|
||||||
for style in model_info.styles:
|
for style in model_info.styles:
|
||||||
|
|
||||||
@@ -23,21 +23,21 @@ def synthesize(device: str = 'cpu'):
|
|||||||
sample_rate, audio_data = model.infer(
|
sample_rate, audio_data = model.infer(
|
||||||
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
|
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
|
||||||
# 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ)
|
# 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ)
|
||||||
language = Languages.JP,
|
language=Languages.JP,
|
||||||
# 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0)
|
# 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0)
|
||||||
speaker_id = 0,
|
speaker_id=0,
|
||||||
# 感情表現の強さ (0.0 〜 1.0)
|
# 感情表現の強さ (0.0 〜 1.0)
|
||||||
sdp_ratio = 0.4,
|
sdp_ratio=0.4,
|
||||||
# スタイル (Neutral, Happy など)
|
# スタイル (Neutral, Happy など)
|
||||||
style = style,
|
style=style,
|
||||||
# スタイルの強さ (0.0 〜 100.0)
|
# スタイルの強さ (0.0 〜 100.0)
|
||||||
style_weight = 6.0,
|
style_weight=6.0,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 音声データを保存
|
# 音声データを保存
|
||||||
(BASE_DIR / 'tests/wavs').mkdir(exist_ok=True, parents=True)
|
(BASE_DIR / "tests/wavs").mkdir(exist_ok=True, parents=True)
|
||||||
wav_file_path = BASE_DIR / f'tests/wavs/{style}.wav'
|
wav_file_path = BASE_DIR / f"tests/wavs/{style}.wav"
|
||||||
with open(wav_file_path, 'wb') as f:
|
with open(wav_file_path, "wb") as f:
|
||||||
wavfile.write(f, sample_rate, audio_data)
|
wavfile.write(f, sample_rate, audio_data)
|
||||||
|
|
||||||
# 音声データが保存されたことを確認
|
# 音声データが保存されたことを確認
|
||||||
@@ -48,8 +48,8 @@ def synthesize(device: str = 'cpu'):
|
|||||||
|
|
||||||
|
|
||||||
def test_synthesize_cpu():
|
def test_synthesize_cpu():
|
||||||
synthesize(device='cpu')
|
synthesize(device="cpu")
|
||||||
|
|
||||||
|
|
||||||
def test_synthesize_cuda():
|
def test_synthesize_cuda():
|
||||||
synthesize(device='cuda')
|
synthesize(device="cuda")
|
||||||
|
|||||||
48
train_ms.py
48
train_ms.py
@@ -281,28 +281,28 @@ def run():
|
|||||||
mas_noise_scale_initial=mas_noise_scale_initial,
|
mas_noise_scale_initial=mas_noise_scale_initial,
|
||||||
noise_scale_delta=noise_scale_delta,
|
noise_scale_delta=noise_scale_delta,
|
||||||
# hps.model 以下のすべての値を引数に渡す
|
# hps.model 以下のすべての値を引数に渡す
|
||||||
use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder,
|
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
|
||||||
use_noise_scaled_mas = hps.model.use_noise_scaled_mas,
|
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
|
||||||
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder,
|
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
|
||||||
use_duration_discriminator = hps.model.use_duration_discriminator,
|
use_duration_discriminator=hps.model.use_duration_discriminator,
|
||||||
use_wavlm_discriminator = hps.model.use_wavlm_discriminator,
|
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
|
||||||
inter_channels = hps.model.inter_channels,
|
inter_channels=hps.model.inter_channels,
|
||||||
hidden_channels = hps.model.hidden_channels,
|
hidden_channels=hps.model.hidden_channels,
|
||||||
filter_channels = hps.model.filter_channels,
|
filter_channels=hps.model.filter_channels,
|
||||||
n_heads = hps.model.n_heads,
|
n_heads=hps.model.n_heads,
|
||||||
n_layers = hps.model.n_layers,
|
n_layers=hps.model.n_layers,
|
||||||
kernel_size = hps.model.kernel_size,
|
kernel_size=hps.model.kernel_size,
|
||||||
p_dropout = hps.model.p_dropout,
|
p_dropout=hps.model.p_dropout,
|
||||||
resblock = hps.model.resblock,
|
resblock=hps.model.resblock,
|
||||||
resblock_kernel_sizes = hps.model.resblock_kernel_sizes,
|
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
|
||||||
resblock_dilation_sizes = hps.model.resblock_dilation_sizes,
|
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
|
||||||
upsample_rates = hps.model.upsample_rates,
|
upsample_rates=hps.model.upsample_rates,
|
||||||
upsample_initial_channel = hps.model.upsample_initial_channel,
|
upsample_initial_channel=hps.model.upsample_initial_channel,
|
||||||
upsample_kernel_sizes = hps.model.upsample_kernel_sizes,
|
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
|
||||||
n_layers_q = hps.model.n_layers_q,
|
n_layers_q=hps.model.n_layers_q,
|
||||||
use_spectral_norm = hps.model.use_spectral_norm,
|
use_spectral_norm=hps.model.use_spectral_norm,
|
||||||
gin_channels = hps.model.gin_channels,
|
gin_channels=hps.model.gin_channels,
|
||||||
slm = hps.model.slm,
|
slm=hps.model.slm,
|
||||||
).cuda(local_rank)
|
).cuda(local_rank)
|
||||||
|
|
||||||
if getattr(hps.train, "freeze_ZH_bert", False):
|
if getattr(hps.train, "freeze_ZH_bert", False):
|
||||||
@@ -389,7 +389,9 @@ def run():
|
|||||||
epoch_str = max(epoch_str, 1)
|
epoch_str = max(epoch_str, 1)
|
||||||
# global_step = (epoch_str - 1) * len(train_loader)
|
# global_step = (epoch_str - 1) * len(train_loader)
|
||||||
global_step = int(
|
global_step = int(
|
||||||
utils.get_steps(utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth"))
|
utils.get_steps(
|
||||||
|
utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth")
|
||||||
|
)
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
|
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
|
||||||
|
|||||||
@@ -288,28 +288,28 @@ def run():
|
|||||||
mas_noise_scale_initial=mas_noise_scale_initial,
|
mas_noise_scale_initial=mas_noise_scale_initial,
|
||||||
noise_scale_delta=noise_scale_delta,
|
noise_scale_delta=noise_scale_delta,
|
||||||
# hps.model 以下のすべての値を引数に渡す
|
# hps.model 以下のすべての値を引数に渡す
|
||||||
use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder,
|
use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,
|
||||||
use_noise_scaled_mas = hps.model.use_noise_scaled_mas,
|
use_noise_scaled_mas=hps.model.use_noise_scaled_mas,
|
||||||
use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder,
|
use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,
|
||||||
use_duration_discriminator = hps.model.use_duration_discriminator,
|
use_duration_discriminator=hps.model.use_duration_discriminator,
|
||||||
use_wavlm_discriminator = hps.model.use_wavlm_discriminator,
|
use_wavlm_discriminator=hps.model.use_wavlm_discriminator,
|
||||||
inter_channels = hps.model.inter_channels,
|
inter_channels=hps.model.inter_channels,
|
||||||
hidden_channels = hps.model.hidden_channels,
|
hidden_channels=hps.model.hidden_channels,
|
||||||
filter_channels = hps.model.filter_channels,
|
filter_channels=hps.model.filter_channels,
|
||||||
n_heads = hps.model.n_heads,
|
n_heads=hps.model.n_heads,
|
||||||
n_layers = hps.model.n_layers,
|
n_layers=hps.model.n_layers,
|
||||||
kernel_size = hps.model.kernel_size,
|
kernel_size=hps.model.kernel_size,
|
||||||
p_dropout = hps.model.p_dropout,
|
p_dropout=hps.model.p_dropout,
|
||||||
resblock = hps.model.resblock,
|
resblock=hps.model.resblock,
|
||||||
resblock_kernel_sizes = hps.model.resblock_kernel_sizes,
|
resblock_kernel_sizes=hps.model.resblock_kernel_sizes,
|
||||||
resblock_dilation_sizes = hps.model.resblock_dilation_sizes,
|
resblock_dilation_sizes=hps.model.resblock_dilation_sizes,
|
||||||
upsample_rates = hps.model.upsample_rates,
|
upsample_rates=hps.model.upsample_rates,
|
||||||
upsample_initial_channel = hps.model.upsample_initial_channel,
|
upsample_initial_channel=hps.model.upsample_initial_channel,
|
||||||
upsample_kernel_sizes = hps.model.upsample_kernel_sizes,
|
upsample_kernel_sizes=hps.model.upsample_kernel_sizes,
|
||||||
n_layers_q = hps.model.n_layers_q,
|
n_layers_q=hps.model.n_layers_q,
|
||||||
use_spectral_norm = hps.model.use_spectral_norm,
|
use_spectral_norm=hps.model.use_spectral_norm,
|
||||||
gin_channels = hps.model.gin_channels,
|
gin_channels=hps.model.gin_channels,
|
||||||
slm = hps.model.slm,
|
slm=hps.model.slm,
|
||||||
).cuda(local_rank)
|
).cuda(local_rank)
|
||||||
if getattr(hps.train, "freeze_JP_bert", False):
|
if getattr(hps.train, "freeze_JP_bert", False):
|
||||||
logger.info("Freezing (JP) bert encoder !!!")
|
logger.info("Freezing (JP) bert encoder !!!")
|
||||||
@@ -383,7 +383,9 @@ def run():
|
|||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
try:
|
try:
|
||||||
_, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
|
_, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
|
||||||
utils.checkpoints.get_latest_checkpoint_path(model_dir, "DUR_*.pth"),
|
utils.checkpoints.get_latest_checkpoint_path(
|
||||||
|
model_dir, "DUR_*.pth"
|
||||||
|
),
|
||||||
net_dur_disc,
|
net_dur_disc,
|
||||||
optim_dur_disc,
|
optim_dur_disc,
|
||||||
skip_optimizer=hps.train.skip_optimizer,
|
skip_optimizer=hps.train.skip_optimizer,
|
||||||
@@ -396,12 +398,16 @@ def run():
|
|||||||
print("Initialize dur_disc")
|
print("Initialize dur_disc")
|
||||||
if net_wd is not None:
|
if net_wd is not None:
|
||||||
try:
|
try:
|
||||||
_, optim_wd, wd_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
|
_, optim_wd, wd_resume_lr, epoch_str = (
|
||||||
utils.checkpoints.get_latest_checkpoint_path(model_dir, "WD_*.pth"),
|
utils.checkpoints.load_checkpoint(
|
||||||
|
utils.checkpoints.get_latest_checkpoint_path(
|
||||||
|
model_dir, "WD_*.pth"
|
||||||
|
),
|
||||||
net_wd,
|
net_wd,
|
||||||
optim_wd,
|
optim_wd,
|
||||||
skip_optimizer=hps.train.skip_optimizer,
|
skip_optimizer=hps.train.skip_optimizer,
|
||||||
)
|
)
|
||||||
|
)
|
||||||
if not optim_wd.param_groups[0].get("initial_lr"):
|
if not optim_wd.param_groups[0].get("initial_lr"):
|
||||||
optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr
|
optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr
|
||||||
except:
|
except:
|
||||||
@@ -430,7 +436,9 @@ def run():
|
|||||||
epoch_str = max(epoch_str, 1)
|
epoch_str = max(epoch_str, 1)
|
||||||
# global_step = (epoch_str - 1) * len(train_loader)
|
# global_step = (epoch_str - 1) * len(train_loader)
|
||||||
global_step = int(
|
global_step = int(
|
||||||
utils.get_steps(utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth"))
|
utils.get_steps(
|
||||||
|
utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth")
|
||||||
|
)
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
|
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
|
||||||
|
|||||||
Reference in New Issue
Block a user