Artrajz commited on
Commit
b5172da
·
1 Parent(s): 684250c

Update Bert-VITS2 num_tones

Browse files
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 num_tones, num_languages
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
- num_ja_tones = 1
 
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,