Remove: remove currently unused code in style_bert_vits2/models/commons.py

This commit is contained in:
tsukumi
2024-03-07 08:46:13 +00:00
parent bffd5a67bb
commit e57cfbf072
3 changed files with 25 additions and 151 deletions

View File

@@ -3,7 +3,6 @@
コードと完全に一致している保証はない。あくまで参考程度とすること。 コードと完全に一致している保証はない。あくまで参考程度とすること。
""" """
import math
import torch import torch
from torch.nn import functional as F from torch.nn import functional as F
from typing import Any from typing import Any
@@ -68,54 +67,6 @@ def intersperse(lst: list[Any], item: Any) -> list[Any]:
return result return result
def kl_divergence(m_p: torch.Tensor, logs_p: torch.Tensor, m_q: torch.Tensor, logs_q: torch.Tensor) -> torch.Tensor:
"""
2つの正規分布間の KL ダイバージェンスを計算する
Args:
m_p (torch.Tensor): P の平均
logs_p (torch.Tensor): P の対数標準偏差
m_q (torch.Tensor): Q の平均
logs_q (torch.Tensor): Q の対数標準偏差
Returns:
torch.Tensor: KL ダイバージェンスの値。
"""
kl = (logs_q - logs_p) - 0.5
kl += (
0.5 * (torch.exp(2.0 * logs_p) + ((m_p - m_q) ** 2)) * torch.exp(-2.0 * logs_q)
)
return kl
def rand_gumbel(shape: torch.Size) -> torch.Tensor:
"""
Gumbel 分布からサンプリングし、オーバーフローを防ぐ
Args:
shape (torch.Size): サンプルの形状
Returns:
torch.Tensor: Gumbel 分布からのサンプル
"""
uniform_samples = torch.rand(shape) * 0.99998 + 0.00001
return -torch.log(-torch.log(uniform_samples))
def rand_gumbel_like(x: torch.Tensor) -> torch.Tensor:
"""
引数と同じ形状のテンソルで、Gumbel 分布からサンプリングする
Args:
x (torch.Tensor): 形状を基にするテンソル
Returns:
torch.Tensor: Gumbel 分布からのサンプル
"""
g = rand_gumbel(x.size()).to(dtype=x.dtype, device=x.device)
return g
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:
""" """
テンソルからセグメントをスライスする テンソルからセグメントをスライスする
@@ -155,69 +106,6 @@ def rand_slice_segments(x: torch.Tensor, x_lengths: torch.Tensor | None = None,
return ret, ids_str return ret, ids_str
def get_timing_signal_1d(length: int, channels: int, min_timescale: float = 1.0, max_timescale: float = 1.0e4) -> torch.Tensor:
"""
1D タイミング信号を取得する
Args:
length (int): シグナルの長さ
channels (int): シグナルのチャネル数
min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0)
max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4)
Returns:
torch.Tensor: タイミング信号
"""
position = torch.arange(length, dtype=torch.float)
num_timescales = channels // 2
log_timescale_increment = math.log(float(max_timescale) / float(min_timescale)) / (
num_timescales - 1
)
inv_timescales = min_timescale * torch.exp(
torch.arange(num_timescales, dtype=torch.float) * -log_timescale_increment
)
scaled_time = position.unsqueeze(0) * inv_timescales.unsqueeze(1)
signal = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], 0)
signal = F.pad(signal, [0, 0, 0, channels % 2])
signal = signal.view(1, channels, length)
return signal
def add_timing_signal_1d(x: torch.Tensor, min_timescale: float = 1.0, max_timescale: float = 1.0e4) -> torch.Tensor:
"""
1D タイミング信号をテンソルに追加する
Args:
x (torch.Tensor): 入力テンソル
min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0)
max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4)
Returns:
torch.Tensor: タイミング信号が追加されたテンソル
"""
b, channels, length = x.size()
signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale)
return x + signal.to(dtype=x.dtype, device=x.device)
def cat_timing_signal_1d(x: torch.Tensor, min_timescale: float = 1.0, max_timescale: float = 1.0e4, axis: int = 1) -> torch.Tensor:
"""
1D タイミング信号をテンソルに連結する
Args:
x (torch.Tensor): 入力テンソル
min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0)
max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4)
axis (int, optional): 連結する軸 (デフォルト: 1)
Returns:
torch.Tensor: タイミング信号が連結されたテンソル
"""
b, channels, length = x.size()
signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale)
return torch.cat([x, signal.to(dtype=x.dtype, device=x.device)], axis)
def subsequent_mask(length: int) -> torch.Tensor: def subsequent_mask(length: int) -> torch.Tensor:
""" """
後続のマスクを生成する 後続のマスクを生成する
@@ -253,20 +141,6 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor
return acts return acts
def shift_1d(x: torch.Tensor) -> torch.Tensor:
"""
与えられたテンソルを 1D でシフトする
Args:
x (torch.Tensor): シフトするテンソル
Returns:
torch.Tensor: シフトされたテンソル
"""
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1]
return x
def sequence_mask(length: torch.Tensor, max_length: int | None = None) -> torch.Tensor: def sequence_mask(length: torch.Tensor, max_length: int | None = None) -> torch.Tensor:
""" """
シーケンスマスクを生成する シーケンスマスクを生成する

View File

@@ -171,7 +171,7 @@ def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]:
list[tuple[str, int]]: 音素とアクセントのペアのリスト list[tuple[str, int]]: 音素とアクセントのペアのリスト
""" """
prosodies = pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True) prosodies = __pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True)
# logger.debug(f"prosodies: {prosodies}") # logger.debug(f"prosodies: {prosodies}")
result: list[tuple[str, int]] = [] result: list[tuple[str, int]] = []
current_phrase: list[tuple[str, int]] = [] current_phrase: list[tuple[str, int]] = []
@@ -212,7 +212,7 @@ 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

View File

@@ -48,29 +48,6 @@ def normalize_text(text: str) -> str:
return res return res
def __convert_numbers_to_words(text: str) -> str:
"""
記号や数字を日本語の文字表現に変換する。
Args:
text (str): 変換するテキスト
Returns:
str: 変換されたテキスト
"""
NUMBER_WITH_SEPARATOR_PATTERN = re.compile("[0-9]{1,3}(,[0-9]{3})+")
CURRENCY_MAP = {"$": "ドル", "¥": "", "£": "ポンド", "": "ユーロ"}
CURRENCY_PATTERN = re.compile(r"([$¥£€])([0-9.]*[0-9])")
NUMBER_PATTERN = re.compile(r"[0-9]+(\.[0-9]+)?")
res = NUMBER_WITH_SEPARATOR_PATTERN.sub(lambda m: m[0].replace(",", ""), text)
res = CURRENCY_PATTERN.sub(lambda m: m[2] + CURRENCY_MAP.get(m[1], m[1]), res)
res = NUMBER_PATTERN.sub(lambda m: num2words(m[0], lang="ja"), res)
return res
def replace_punctuation(text: str) -> str: def replace_punctuation(text: str) -> str:
""" """
句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalk で読みが取得できるもののみ残す: 句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalk で読みが取得できるもののみ残す:
@@ -159,3 +136,26 @@ def replace_punctuation(text: str) -> str:
) )
return replaced_text return replaced_text
def __convert_numbers_to_words(text: str) -> str:
"""
記号や数字を日本語の文字表現に変換する。
Args:
text (str): 変換するテキスト
Returns:
str: 変換されたテキスト
"""
NUMBER_WITH_SEPARATOR_PATTERN = re.compile("[0-9]{1,3}(,[0-9]{3})+")
CURRENCY_MAP = {"$": "ドル", "¥": "", "£": "ポンド", "": "ユーロ"}
CURRENCY_PATTERN = re.compile(r"([$¥£€])([0-9.]*[0-9])")
NUMBER_PATTERN = re.compile(r"[0-9]+(\.[0-9]+)?")
res = NUMBER_WITH_SEPARATOR_PATTERN.sub(lambda m: m[0].replace(",", ""), text)
res = CURRENCY_PATTERN.sub(lambda m: m[2] + CURRENCY_MAP.get(m[1], m[1]), res)
res = NUMBER_PATTERN.sub(lambda m: num2words(m[0], lang="ja"), res)
return res