Apply black formatter

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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