Spaces:
Sleeping
Sleeping
Update Bert-VITS2 num_tones
Browse files- bert_vits2/bert_vits2.py +7 -0
- bert_vits2/models.py +6 -3
- bert_vits2/text/__init__.py +9 -0
- bert_vits2/text/symbols.py +9 -1
bert_vits2/bert_vits2.py
CHANGED
|
@@ -27,12 +27,14 @@ class Bert_VITS2:
|
|
| 27 |
self.bert_model_names = {"zh": "CHINESE_ROBERTA_WWM_EXT_LARGE"}
|
| 28 |
self.ja_bert_dim = 1024
|
| 29 |
self.ja_extra_str = ""
|
|
|
|
| 30 |
|
| 31 |
if self.version in ["1.0", "1.0.0", "1.0.1"]:
|
| 32 |
self.symbols = symbols_legacy
|
| 33 |
self.hps_ms.model.n_layers_trans_flow = 3
|
| 34 |
self.lang = ["zh"]
|
| 35 |
self.ja_bert_dim = 768
|
|
|
|
| 36 |
|
| 37 |
elif self.version in ["1.1.0-transition"]:
|
| 38 |
self.hps_ms.model.n_layers_trans_flow = 3
|
|
@@ -40,6 +42,7 @@ class Bert_VITS2:
|
|
| 40 |
self.bert_model_names["ja"] = "BERT_BASE_JAPANESE_V3"
|
| 41 |
self.ja_bert_dim = 768
|
| 42 |
self.ja_extra_str = "_v111"
|
|
|
|
| 43 |
|
| 44 |
elif self.version in ["1.1", "1.1.0", "1.1.1"]:
|
| 45 |
self.hps_ms.model.n_layers_trans_flow = 6
|
|
@@ -47,12 +50,15 @@ class Bert_VITS2:
|
|
| 47 |
self.bert_model_names["ja"] = "BERT_BASE_JAPANESE_V3"
|
| 48 |
self.ja_bert_dim = 768
|
| 49 |
self.ja_extra_str = "_v111"
|
|
|
|
| 50 |
|
| 51 |
elif self.version in ["2.0", "2.0.0"]:
|
| 52 |
self.hps_ms.model.n_layers_trans_flow = 4
|
| 53 |
self.bert_model_names = {"zh": "CHINESE_ROBERTA_WWM_EXT_LARGE",
|
| 54 |
"ja": "DEBERTA_V2_LARGE_JAPANESE",
|
| 55 |
"en": "DEBERTA_V3_LARGE"}
|
|
|
|
|
|
|
| 56 |
|
| 57 |
# self.bert_handler = BertHandler(self.lang)
|
| 58 |
|
|
@@ -69,6 +75,7 @@ class Bert_VITS2:
|
|
| 69 |
n_speakers=self.hps_ms.data.n_speakers,
|
| 70 |
symbols=self.symbols,
|
| 71 |
ja_bert_dim=self.ja_bert_dim,
|
|
|
|
| 72 |
**self.hps_ms.model).to(self.device)
|
| 73 |
_ = self.net_g.eval()
|
| 74 |
bert_vits2_utils.load_checkpoint(self.model_path, self.net_g, None, skip_optimizer=True, version=self.version)
|
|
|
|
| 27 |
self.bert_model_names = {"zh": "CHINESE_ROBERTA_WWM_EXT_LARGE"}
|
| 28 |
self.ja_bert_dim = 1024
|
| 29 |
self.ja_extra_str = ""
|
| 30 |
+
self.num_tones = num_tones
|
| 31 |
|
| 32 |
if self.version in ["1.0", "1.0.0", "1.0.1"]:
|
| 33 |
self.symbols = symbols_legacy
|
| 34 |
self.hps_ms.model.n_layers_trans_flow = 3
|
| 35 |
self.lang = ["zh"]
|
| 36 |
self.ja_bert_dim = 768
|
| 37 |
+
self.num_tones = num_tones_v111
|
| 38 |
|
| 39 |
elif self.version in ["1.1.0-transition"]:
|
| 40 |
self.hps_ms.model.n_layers_trans_flow = 3
|
|
|
|
| 42 |
self.bert_model_names["ja"] = "BERT_BASE_JAPANESE_V3"
|
| 43 |
self.ja_bert_dim = 768
|
| 44 |
self.ja_extra_str = "_v111"
|
| 45 |
+
self.num_tones = num_tones_v111
|
| 46 |
|
| 47 |
elif self.version in ["1.1", "1.1.0", "1.1.1"]:
|
| 48 |
self.hps_ms.model.n_layers_trans_flow = 6
|
|
|
|
| 50 |
self.bert_model_names["ja"] = "BERT_BASE_JAPANESE_V3"
|
| 51 |
self.ja_bert_dim = 768
|
| 52 |
self.ja_extra_str = "_v111"
|
| 53 |
+
self.num_tones = num_tones_v111
|
| 54 |
|
| 55 |
elif self.version in ["2.0", "2.0.0"]:
|
| 56 |
self.hps_ms.model.n_layers_trans_flow = 4
|
| 57 |
self.bert_model_names = {"zh": "CHINESE_ROBERTA_WWM_EXT_LARGE",
|
| 58 |
"ja": "DEBERTA_V2_LARGE_JAPANESE",
|
| 59 |
"en": "DEBERTA_V3_LARGE"}
|
| 60 |
+
self.num_tones = num_tones
|
| 61 |
+
|
| 62 |
|
| 63 |
# self.bert_handler = BertHandler(self.lang)
|
| 64 |
|
|
|
|
| 75 |
n_speakers=self.hps_ms.data.n_speakers,
|
| 76 |
symbols=self.symbols,
|
| 77 |
ja_bert_dim=self.ja_bert_dim,
|
| 78 |
+
num_tones=self.num_tones,
|
| 79 |
**self.hps_ms.model).to(self.device)
|
| 80 |
_ = self.net_g.eval()
|
| 81 |
bert_vits2_utils.load_checkpoint(self.model_path, self.net_g, None, skip_optimizer=True, version=self.version)
|
bert_vits2/models.py
CHANGED
|
@@ -11,7 +11,7 @@ from torch.nn import Conv1d, ConvTranspose1d, AvgPool1d, Conv2d
|
|
| 11 |
from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm
|
| 12 |
|
| 13 |
from bert_vits2.commons import init_weights, get_padding
|
| 14 |
-
from bert_vits2.text import
|
| 15 |
|
| 16 |
|
| 17 |
class DurationDiscriminator(nn.Module): # vits2
|
|
@@ -258,7 +258,8 @@ class TextEncoder(nn.Module):
|
|
| 258 |
p_dropout,
|
| 259 |
gin_channels=0,
|
| 260 |
symbols=None,
|
| 261 |
-
ja_bert_dim=1024
|
|
|
|
| 262 |
super().__init__()
|
| 263 |
self.n_vocab = n_vocab
|
| 264 |
self.out_channels = out_channels
|
|
@@ -601,6 +602,7 @@ class SynthesizerTrn(nn.Module):
|
|
| 601 |
use_transformer_flow=True,
|
| 602 |
symbols=None,
|
| 603 |
ja_bert_dim=1024,
|
|
|
|
| 604 |
**kwargs):
|
| 605 |
|
| 606 |
super().__init__()
|
|
@@ -641,7 +643,8 @@ class SynthesizerTrn(nn.Module):
|
|
| 641 |
p_dropout,
|
| 642 |
gin_channels=self.enc_gin_channels,
|
| 643 |
symbols=symbols,
|
| 644 |
-
ja_bert_dim=ja_bert_dim
|
|
|
|
| 645 |
)
|
| 646 |
self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates,
|
| 647 |
upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)
|
|
|
|
| 11 |
from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm
|
| 12 |
|
| 13 |
from bert_vits2.commons import init_weights, get_padding
|
| 14 |
+
from bert_vits2.text import num_languages
|
| 15 |
|
| 16 |
|
| 17 |
class DurationDiscriminator(nn.Module): # vits2
|
|
|
|
| 258 |
p_dropout,
|
| 259 |
gin_channels=0,
|
| 260 |
symbols=None,
|
| 261 |
+
ja_bert_dim=1024,
|
| 262 |
+
num_tones=None):
|
| 263 |
super().__init__()
|
| 264 |
self.n_vocab = n_vocab
|
| 265 |
self.out_channels = out_channels
|
|
|
|
| 602 |
use_transformer_flow=True,
|
| 603 |
symbols=None,
|
| 604 |
ja_bert_dim=1024,
|
| 605 |
+
num_tones=None,
|
| 606 |
**kwargs):
|
| 607 |
|
| 608 |
super().__init__()
|
|
|
|
| 643 |
p_dropout,
|
| 644 |
gin_channels=self.enc_gin_channels,
|
| 645 |
symbols=symbols,
|
| 646 |
+
ja_bert_dim=ja_bert_dim,
|
| 647 |
+
num_tones=num_tones
|
| 648 |
)
|
| 649 |
self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates,
|
| 650 |
upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)
|
bert_vits2/text/__init__.py
CHANGED
|
@@ -2,6 +2,15 @@ from bert_vits2.text.symbols import *
|
|
| 2 |
from bert_vits2.text.bert_handler import BertHandler
|
| 3 |
|
| 4 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
def cleaned_text_to_sequence(cleaned_text, tones, language, _symbol_to_id):
|
| 6 |
"""Converts a string of text to a sequence of IDs corresponding to the symbols in the text.
|
| 7 |
Args:
|
|
|
|
| 2 |
from bert_vits2.text.bert_handler import BertHandler
|
| 3 |
|
| 4 |
|
| 5 |
+
def cleaned_text_to_sequence_v111(cleaned_text, tones, language, _symbol_to_id):
|
| 6 |
+
"""version <= 1.1.1"""
|
| 7 |
+
phones = [_symbol_to_id[symbol] for symbol in cleaned_text]
|
| 8 |
+
tone_start = language_tone_start_map_v111[language]
|
| 9 |
+
tones = [i + tone_start for i in tones]
|
| 10 |
+
lang_id = language_id_map[language]
|
| 11 |
+
lang_ids = [lang_id for i in phones]
|
| 12 |
+
return phones, tones, lang_ids
|
| 13 |
+
|
| 14 |
def cleaned_text_to_sequence(cleaned_text, tones, language, _symbol_to_id):
|
| 15 |
"""Converts a string of text to a sequence of IDs corresponding to the symbols in the text.
|
| 16 |
Args:
|
bert_vits2/text/symbols.py
CHANGED
|
@@ -120,7 +120,8 @@ ja_symbols = [
|
|
| 120 |
"z",
|
| 121 |
"zy",
|
| 122 |
]
|
| 123 |
-
|
|
|
|
| 124 |
|
| 125 |
# English
|
| 126 |
en_symbols = [
|
|
@@ -176,12 +177,19 @@ symbols_legacy = [pad] + normal_symbols_legacy + pu_symbols
|
|
| 176 |
sil_phonemes_ids_legacy = [symbols_legacy.index(i) for i in pu_symbols]
|
| 177 |
|
| 178 |
# combine all tones
|
|
|
|
| 179 |
num_tones = num_zh_tones + num_ja_tones + num_en_tones
|
| 180 |
|
| 181 |
# language maps
|
| 182 |
language_id_map = {"zh": 0, "ja": 1, "en": 2}
|
| 183 |
num_languages = len(language_id_map.keys())
|
| 184 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 185 |
language_tone_start_map = {
|
| 186 |
"zh": 0,
|
| 187 |
"ja": num_zh_tones,
|
|
|
|
| 120 |
"z",
|
| 121 |
"zy",
|
| 122 |
]
|
| 123 |
+
num_ja_tones_v111 = 1
|
| 124 |
+
num_ja_tones = 2
|
| 125 |
|
| 126 |
# English
|
| 127 |
en_symbols = [
|
|
|
|
| 177 |
sil_phonemes_ids_legacy = [symbols_legacy.index(i) for i in pu_symbols]
|
| 178 |
|
| 179 |
# combine all tones
|
| 180 |
+
num_tones_v111 = num_zh_tones + num_ja_tones_v111 + num_en_tones
|
| 181 |
num_tones = num_zh_tones + num_ja_tones + num_en_tones
|
| 182 |
|
| 183 |
# language maps
|
| 184 |
language_id_map = {"zh": 0, "ja": 1, "en": 2}
|
| 185 |
num_languages = len(language_id_map.keys())
|
| 186 |
|
| 187 |
+
language_tone_start_map_v111 = {
|
| 188 |
+
"zh": 0,
|
| 189 |
+
"ja": num_zh_tones,
|
| 190 |
+
"en": num_zh_tones + num_ja_tones_v111,
|
| 191 |
+
}
|
| 192 |
+
|
| 193 |
language_tone_start_map = {
|
| 194 |
"zh": 0,
|
| 195 |
"ja": num_zh_tones,
|