Apply black formatter
This commit is contained in:
@@ -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:
|
||||
"""
|
||||
指定されたディレクトリから最新のチェックポイントのパスを取得する
|
||||
|
||||
|
||||
Reference in New Issue
Block a user