From e57cfbf072c2a34dcfd22201ba9c8c67632f9122 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 08:46:13 +0000 Subject: [PATCH] Remove: remove currently unused code in style_bert_vits2/models/commons.py --- style_bert_vits2/models/commons.py | 126 ------------------ .../text_processing/japanese/g2p.py | 4 +- .../text_processing/japanese/normalizer.py | 46 +++---- 3 files changed, 25 insertions(+), 151 deletions(-) diff --git a/style_bert_vits2/models/commons.py b/style_bert_vits2/models/commons.py index 064ef5f..969ed36 100644 --- a/style_bert_vits2/models/commons.py +++ b/style_bert_vits2/models/commons.py @@ -3,7 +3,6 @@ コードと完全に一致している保証はない。あくまで参考程度とすること。 """ -import math import torch from torch.nn import functional as F from typing import Any @@ -68,54 +67,6 @@ def intersperse(lst: list[Any], item: Any) -> list[Any]: 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: """ テンソルからセグメントをスライスする @@ -155,69 +106,6 @@ def rand_slice_segments(x: torch.Tensor, x_lengths: torch.Tensor | None = None, 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: """ 後続のマスクを生成する @@ -253,20 +141,6 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor 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: """ シーケンスマスクを生成する diff --git a/style_bert_vits2/text_processing/japanese/g2p.py b/style_bert_vits2/text_processing/japanese/g2p.py index 7968751..c04d078 100644 --- a/style_bert_vits2/text_processing/japanese/g2p.py +++ b/style_bert_vits2/text_processing/japanese/g2p.py @@ -171,7 +171,7 @@ def __g2phone_tone_wo_punct(text: str) -> 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}") result: 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 -def pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]: +def __pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]: """ ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。 ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py diff --git a/style_bert_vits2/text_processing/japanese/normalizer.py b/style_bert_vits2/text_processing/japanese/normalizer.py index 92c2e87..8276338 100644 --- a/style_bert_vits2/text_processing/japanese/normalizer.py +++ b/style_bert_vits2/text_processing/japanese/normalizer.py @@ -48,29 +48,6 @@ def normalize_text(text: str) -> str: 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: """ 句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalk で読みが取得できるもののみ残す: @@ -159,3 +136,26 @@ def replace_punctuation(text: str) -> str: ) 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