Artrajz commited on
Commit
282b2a6
·
1 Parent(s): a5fbcbd

Update Bert-VITS2 Extra

Browse files
bert_vits2/bert/Erlangshen-MegatronBert-1.3B-Chinese/config.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {"vocab_size": 21248, "hidden_size": 2048, "num_hidden_layers": 24, "num_attention_heads": 8, "hidden_act": "gelu_new", "intermediate_size": 8192, "hidden_dropout_prob": 0.1, "attention_probs_dropout_prob": 0.1, "max_position_embeddings": 512, "type_vocab_size": 2, "initializer_range": 0.02, "layer_norm_eps": 1e-12, "gradient_checkpointing": false, "position_embedding_type": "absolute", "use_cache": false, "model_type": "megatron-bert"}
bert_vits2/bert/Erlangshen-MegatronBert-1.3B-Chinese/vocab.txt ADDED
The diff for this file is too large to render. See raw diff
 
bert_vits2/bert_vits2.py CHANGED
@@ -1,5 +1,6 @@
1
  import logging
2
 
 
3
  import torch
4
 
5
  from bert_vits2 import commons
@@ -26,6 +27,7 @@ class Bert_VITS2:
26
  self.sampling_rate = self.hps_ms.data.sampling_rate
27
 
28
  self.bert_model_names = {}
 
29
  self.ja_bert_dim = 1024
30
  self.num_tones = num_tones
31
 
@@ -35,6 +37,7 @@ class Bert_VITS2:
35
  self.bert_extra_str_map = {"zh": "", "ja": "", "en": ""}
36
  self.hps_ms.model.emotion_embedding = None
37
  if self.version in ["1.0", "1.0.0", "1.0.1"]:
 
38
  self.symbols = symbols_legacy
39
  self.hps_ms.model.n_layers_trans_flow = 3
40
  self.lang = getattr(self.hps_ms.data, "lang", ["zh"])
@@ -43,6 +46,7 @@ class Bert_VITS2:
43
  self.text_extra_str_map.update({"zh": "_v100"})
44
 
45
  elif self.version in ["1.1.0-transition"]:
 
46
  self.hps_ms.model.n_layers_trans_flow = 3
47
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja"])
48
  self.ja_bert_dim = 768
@@ -52,6 +56,7 @@ class Bert_VITS2:
52
  self.bert_extra_str_map.update({"ja": "_v111"})
53
 
54
  elif self.version in ["1.1", "1.1.0", "1.1.1"]:
 
55
  self.hps_ms.model.n_layers_trans_flow = 6
56
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja"])
57
  self.ja_bert_dim = 768
@@ -61,6 +66,7 @@ class Bert_VITS2:
61
  self.bert_extra_str_map.update({"ja": "_v111"})
62
 
63
  elif self.version in ["2.0", "2.0.0", "2.0.1", "2.0.2"]:
 
64
  self.hps_ms.model.n_layers_trans_flow = 4
65
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
66
  self.num_tones = num_tones
@@ -70,6 +76,7 @@ class Bert_VITS2:
70
  self.bert_extra_str_map.update({"ja": "_v200", "en": "_v200"})
71
 
72
  elif self.version in ["2.1", "2.1.0"]:
 
73
  self.hps_ms.model.n_layers_trans_flow = 4
74
  self.hps_ms.model.emotion_embedding = 1
75
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
@@ -78,6 +85,7 @@ class Bert_VITS2:
78
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
79
 
80
  elif self.version in ["2.2", "2.2.0"]:
 
81
  self.hps_ms.model.n_layers_trans_flow = 4
82
  self.hps_ms.model.emotion_embedding = 2
83
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
@@ -85,20 +93,31 @@ class Bert_VITS2:
85
  if "ja" in self.lang: self.bert_model_names.update({"ja": "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM"})
86
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
87
  elif self.version in ["2.3", "2.3.0"]:
 
88
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
89
  self.num_tones = num_tones
90
  self.text_extra_str_map.update({"en": "_v230"})
91
  if "ja" in self.lang: self.bert_model_names.update({"ja": "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM"})
92
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
 
 
 
 
 
 
 
 
 
93
  else:
94
  logging.debug("Version information not found. Loaded as the newest version: v2.3.")
 
95
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
96
  self.num_tones = num_tones
97
  self.text_extra_str_map.update({"en": "_v230"})
98
  if "ja" in self.lang: self.bert_model_names.update({"ja": "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM"})
99
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
100
 
101
- if "zh" in self.lang:
102
  self.bert_model_names.update({"zh": "CHINESE_ROBERTA_WWM_EXT_LARGE"})
103
 
104
  self._symbol_to_id = {s: i for i, s in enumerate(self.symbols)}
@@ -108,7 +127,7 @@ class Bert_VITS2:
108
  def load_model(self, model_handler):
109
  self.model_handler = model_handler
110
 
111
- if self.version in ["2.3", "2.3.0"]:
112
  self.net_g = SynthesizerTrn_v230(
113
  len(symbols),
114
  self.hps_ms.data.filter_length // 2 + 1,
@@ -125,6 +144,7 @@ class Bert_VITS2:
125
  symbols=self.symbols,
126
  ja_bert_dim=self.ja_bert_dim,
127
  num_tones=self.num_tones,
 
128
  **self.hps_ms.model).to(self.device)
129
  _ = self.net_g.eval()
130
  bert_vits2_utils.load_checkpoint(self.model_path, self.net_g, None, skip_optimizer=True, version=self.version)
@@ -156,7 +176,7 @@ class Bert_VITS2:
156
  del word2ph
157
  assert bert.shape[-1] == len(phone), phone
158
 
159
- if language_str == "zh":
160
  zh_bert = bert
161
  ja_bert = torch.zeros(self.ja_bert_dim, len(phone))
162
  en_bert = torch.zeros(1024, len(phone))
@@ -180,7 +200,7 @@ class Bert_VITS2:
180
  language = torch.LongTensor(language)
181
  return zh_bert, ja_bert, en_bert, phone, tone, language
182
 
183
- def get_emo_(self, reference_audio, emotion):
184
  if reference_audio:
185
  emo = torch.from_numpy(
186
  get_emo(reference_audio, self.model_handler.emotion_model,
@@ -191,6 +211,17 @@ class Bert_VITS2:
191
 
192
  return emo
193
 
 
 
 
 
 
 
 
 
 
 
 
194
  def _infer(self, id, phones, tones, lang_ids, zh_bert, ja_bert, en_bert, sdp_ratio, noise, noisew, length,
195
  emo=None):
196
  with torch.no_grad():
@@ -225,18 +256,11 @@ class Bert_VITS2:
225
  zh_bert, ja_bert, en_bert, phones, tones, lang_ids = self.get_text(text, lang, self.hps_ms, style_text,
226
  style_weigth)
227
 
 
228
  if self.hps_ms.model.emotion_embedding == 1:
229
- emo = self.get_emo_(reference_audio, emotion).to(self.device).unsqueeze(0)
230
  elif self.hps_ms.model.emotion_embedding == 2:
231
- if isinstance(text_prompt, str):
232
- emo = get_clap_text_feature(text_prompt, self.model_handler.clap_model,
233
- self.model_handler.clap_processor, self.device)
234
- else:
235
- emo = get_clap_audio_feature(reference_audio, self.model_handler.clap_model,
236
- self.model_handler.clap_processor, self.device)
237
- emo = torch.squeeze(emo, dim=1).unsqueeze(0)
238
- else:
239
- emo = None
240
 
241
  return self._infer(id, phones, tones, lang_ids, zh_bert, ja_bert, en_bert, sdp_ratio, noise, noisew, length,
242
  emo)
@@ -245,54 +269,49 @@ class Bert_VITS2:
245
  text_prompt=None, style_text=None, style_weigth=0.7, **kwargs):
246
  sentences_list = split_by_language(text, self.lang)
247
 
 
248
  if self.hps_ms.model.emotion_embedding == 1:
249
- emo = self.get_emo_(reference_audio, emotion).to(self.device).unsqueeze(0)
250
  elif self.hps_ms.model.emotion_embedding == 2:
251
- if isinstance(text_prompt, str):
252
- emo = get_clap_text_feature(text_prompt, self.model_handler.clap_model,
253
- self.model_handler.clap_processor, self.device)
254
- else:
255
- emo = get_clap_audio_feature(reference_audio, self.model_handler.clap_model,
256
- self.model_handler.clap_processor, self.device)
257
- emo = torch.squeeze(emo, dim=1).unsqueeze(0)
258
- else:
259
- emo = None
260
 
261
- tmp_phones, tmp_tones, tmp_lang_ids, tmp_zh_bert, tmp_ja_bert, tmp_en_bert = [], [], [], [], [], []
262
 
263
  for idx, (_text, lang) in enumerate(sentences_list):
264
  skip_start = idx != 0
265
  skip_end = idx != len(sentences_list) - 1
266
- zh_bert, ja_bert, en_bert, phones, tones, lang_ids = self.get_text(_text, lang, self.hps_ms, style_text,
267
- style_weigth)
 
268
  if skip_start:
269
- phones = phones[3:]
270
- tones = tones[3:]
271
- lang_ids = lang_ids[3:]
272
- zh_bert = zh_bert[:, 3:]
273
- ja_bert = ja_bert[:, 3:]
274
- en_bert = en_bert[:, 3:]
275
  if skip_end:
276
- phones = phones[:-2]
277
- tones = tones[:-2]
278
- lang_ids = lang_ids[:-2]
279
- zh_bert = zh_bert[:, :-2]
280
- ja_bert = ja_bert[:, :-2]
281
- en_bert = en_bert[:, :-2]
282
-
283
- tmp_phones.append(phones)
284
- tmp_tones.append(tones)
285
- tmp_lang_ids.append(lang_ids)
286
- tmp_zh_bert.append(zh_bert)
287
- tmp_ja_bert.append(ja_bert)
288
- tmp_en_bert.append(en_bert)
289
-
290
- zh_bert = torch.concatenate(tmp_zh_bert, dim=1)
291
- ja_bert = torch.concatenate(tmp_ja_bert, dim=1)
292
- en_bert = torch.concatenate(tmp_en_bert, dim=1)
293
- phones = torch.concatenate(tmp_phones, dim=0)
294
- tones = torch.concatenate(tmp_tones, dim=0)
295
- lang_ids = torch.concatenate(tmp_lang_ids, dim=0)
 
296
  audio = self._infer(id, phones, tones, lang_ids, zh_bert, ja_bert, en_bert, sdp_ratio, noise,
297
  noisew, length, emo)
298
 
 
1
  import logging
2
 
3
+ import numpy as np
4
  import torch
5
 
6
  from bert_vits2 import commons
 
27
  self.sampling_rate = self.hps_ms.data.sampling_rate
28
 
29
  self.bert_model_names = {}
30
+ self.zh_bert_extra = False
31
  self.ja_bert_dim = 1024
32
  self.num_tones = num_tones
33
 
 
37
  self.bert_extra_str_map = {"zh": "", "ja": "", "en": ""}
38
  self.hps_ms.model.emotion_embedding = None
39
  if self.version in ["1.0", "1.0.0", "1.0.1"]:
40
+ self.version = "1.0"
41
  self.symbols = symbols_legacy
42
  self.hps_ms.model.n_layers_trans_flow = 3
43
  self.lang = getattr(self.hps_ms.data, "lang", ["zh"])
 
46
  self.text_extra_str_map.update({"zh": "_v100"})
47
 
48
  elif self.version in ["1.1.0-transition"]:
49
+ self.version = "1.1.0-transition"
50
  self.hps_ms.model.n_layers_trans_flow = 3
51
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja"])
52
  self.ja_bert_dim = 768
 
56
  self.bert_extra_str_map.update({"ja": "_v111"})
57
 
58
  elif self.version in ["1.1", "1.1.0", "1.1.1"]:
59
+ self.version = "1.1"
60
  self.hps_ms.model.n_layers_trans_flow = 6
61
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja"])
62
  self.ja_bert_dim = 768
 
66
  self.bert_extra_str_map.update({"ja": "_v111"})
67
 
68
  elif self.version in ["2.0", "2.0.0", "2.0.1", "2.0.2"]:
69
+ self.version = "2.0"
70
  self.hps_ms.model.n_layers_trans_flow = 4
71
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
72
  self.num_tones = num_tones
 
76
  self.bert_extra_str_map.update({"ja": "_v200", "en": "_v200"})
77
 
78
  elif self.version in ["2.1", "2.1.0"]:
79
+ self.version = "2.1"
80
  self.hps_ms.model.n_layers_trans_flow = 4
81
  self.hps_ms.model.emotion_embedding = 1
82
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
 
85
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
86
 
87
  elif self.version in ["2.2", "2.2.0"]:
88
+ self.version = "2.2"
89
  self.hps_ms.model.n_layers_trans_flow = 4
90
  self.hps_ms.model.emotion_embedding = 2
91
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
 
93
  if "ja" in self.lang: self.bert_model_names.update({"ja": "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM"})
94
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
95
  elif self.version in ["2.3", "2.3.0"]:
96
+ self.version = "2.3"
97
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
98
  self.num_tones = num_tones
99
  self.text_extra_str_map.update({"en": "_v230"})
100
  if "ja" in self.lang: self.bert_model_names.update({"ja": "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM"})
101
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
102
+ elif self.version.lower().replace("-", "_") in ["extra", "zh_clap"]:
103
+ self.version = "extra"
104
+ self.hps_ms.model.emotion_embedding = 2
105
+ self.hps_ms.model.n_layers_trans_flow = 6
106
+ self.lang = ["zh"]
107
+ self.num_tones = num_tones
108
+ self.zh_bert_extra = True
109
+ self.bert_model_names.update({"zh": "Erlangshen-MegatronBert-1.3B-Chinese"})
110
+ self.bert_extra_str_map.update({"zh": "_extra"})
111
  else:
112
  logging.debug("Version information not found. Loaded as the newest version: v2.3.")
113
+ self.version = "2.3"
114
  self.lang = getattr(self.hps_ms.data, "lang", ["zh", "ja", "en"])
115
  self.num_tones = num_tones
116
  self.text_extra_str_map.update({"en": "_v230"})
117
  if "ja" in self.lang: self.bert_model_names.update({"ja": "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM"})
118
  if "en" in self.lang: self.bert_model_names.update({"en": "DEBERTA_V3_LARGE"})
119
 
120
+ if "zh" in self.lang and "zh" not in self.bert_model_names.keys():
121
  self.bert_model_names.update({"zh": "CHINESE_ROBERTA_WWM_EXT_LARGE"})
122
 
123
  self._symbol_to_id = {s: i for i, s in enumerate(self.symbols)}
 
127
  def load_model(self, model_handler):
128
  self.model_handler = model_handler
129
 
130
+ if self.version == "2.3":
131
  self.net_g = SynthesizerTrn_v230(
132
  len(symbols),
133
  self.hps_ms.data.filter_length // 2 + 1,
 
144
  symbols=self.symbols,
145
  ja_bert_dim=self.ja_bert_dim,
146
  num_tones=self.num_tones,
147
+ zh_bert_extra=self.zh_bert_extra,
148
  **self.hps_ms.model).to(self.device)
149
  _ = self.net_g.eval()
150
  bert_vits2_utils.load_checkpoint(self.model_path, self.net_g, None, skip_optimizer=True, version=self.version)
 
176
  del word2ph
177
  assert bert.shape[-1] == len(phone), phone
178
 
179
+ if language_str == "zh" or self.zh_bert_extra:
180
  zh_bert = bert
181
  ja_bert = torch.zeros(self.ja_bert_dim, len(phone))
182
  en_bert = torch.zeros(1024, len(phone))
 
200
  language = torch.LongTensor(language)
201
  return zh_bert, ja_bert, en_bert, phone, tone, language
202
 
203
+ def _get_emo(self, reference_audio, emotion):
204
  if reference_audio:
205
  emo = torch.from_numpy(
206
  get_emo(reference_audio, self.model_handler.emotion_model,
 
211
 
212
  return emo
213
 
214
+ def _get_clap(self, reference_audio, text_prompt):
215
+ if isinstance(reference_audio, np.ndarray):
216
+ emo = get_clap_audio_feature(reference_audio, self.model_handler.clap_model,
217
+ self.model_handler.clap_processor, self.device)
218
+ else:
219
+ if text_prompt is None: text_prompt = ""
220
+ emo = get_clap_text_feature(text_prompt, self.model_handler.clap_model,
221
+ self.model_handler.clap_processor, self.device)
222
+ emo = torch.squeeze(emo, dim=1).unsqueeze(0)
223
+ return emo
224
+
225
  def _infer(self, id, phones, tones, lang_ids, zh_bert, ja_bert, en_bert, sdp_ratio, noise, noisew, length,
226
  emo=None):
227
  with torch.no_grad():
 
256
  zh_bert, ja_bert, en_bert, phones, tones, lang_ids = self.get_text(text, lang, self.hps_ms, style_text,
257
  style_weigth)
258
 
259
+ emo = None
260
  if self.hps_ms.model.emotion_embedding == 1:
261
+ emo = self._get_emo(reference_audio, emotion).to(self.device).unsqueeze(0)
262
  elif self.hps_ms.model.emotion_embedding == 2:
263
+ emo = self._get_clap(reference_audio, text_prompt)
 
 
 
 
 
 
 
 
264
 
265
  return self._infer(id, phones, tones, lang_ids, zh_bert, ja_bert, en_bert, sdp_ratio, noise, noisew, length,
266
  emo)
 
269
  text_prompt=None, style_text=None, style_weigth=0.7, **kwargs):
270
  sentences_list = split_by_language(text, self.lang)
271
 
272
+ emo = None
273
  if self.hps_ms.model.emotion_embedding == 1:
274
+ emo = self._get_emo(reference_audio, emotion).to(self.device).unsqueeze(0)
275
  elif self.hps_ms.model.emotion_embedding == 2:
276
+ emo = self._get_clap(reference_audio, text_prompt)
 
 
 
 
 
 
 
 
277
 
278
+ phones, tones, lang_ids, zh_bert, ja_bert, en_bert = [], [], [], [], [], []
279
 
280
  for idx, (_text, lang) in enumerate(sentences_list):
281
  skip_start = idx != 0
282
  skip_end = idx != len(sentences_list) - 1
283
+ _zh_bert, _ja_bert, _en_bert, _phones, _tones, _lang_ids = self.get_text(_text, lang, self.hps_ms,
284
+ style_text, style_weigth)
285
+
286
  if skip_start:
287
+ _phones = _phones[3:]
288
+ _tones = _tones[3:]
289
+ _lang_ids = _lang_ids[3:]
290
+ _zh_bert = _zh_bert[:, 3:]
291
+ _ja_bert = _ja_bert[:, 3:]
292
+ _en_bert = _en_bert[:, 3:]
293
  if skip_end:
294
+ _phones = _phones[:-2]
295
+ _tones = _tones[:-2]
296
+ _lang_ids = _lang_ids[:-2]
297
+ _zh_bert = _zh_bert[:, :-2]
298
+ _ja_bert = _ja_bert[:, :-2]
299
+ _en_bert = _en_bert[:, :-2]
300
+
301
+ phones.append(_phones)
302
+ tones.append(_tones)
303
+ lang_ids.append(_lang_ids)
304
+ zh_bert.append(_zh_bert)
305
+ ja_bert.append(_ja_bert)
306
+ en_bert.append(_en_bert)
307
+
308
+ zh_bert = torch.cat(zh_bert, dim=1)
309
+ ja_bert = torch.cat(ja_bert, dim=1)
310
+ en_bert = torch.cat(en_bert, dim=1)
311
+ phones = torch.cat(phones, dim=0)
312
+ tones = torch.cat(tones, dim=0)
313
+ lang_ids = torch.cat(lang_ids, dim=0)
314
+
315
  audio = self._infer(id, phones, tones, lang_ids, zh_bert, ja_bert, en_bert, sdp_ratio, noise,
316
  noisew, length, emo)
317
 
bert_vits2/model_handler.py CHANGED
@@ -13,6 +13,7 @@ from bert_vits2.text.japanese_bert import get_bert_feature as ja_bert
13
  from bert_vits2.text.japanese_bert_v111 import get_bert_feature as ja_bert_v111
14
  from bert_vits2.text.japanese_bert_v200 import get_bert_feature as ja_bert_v200
15
  from bert_vits2.text.english_bert_mock_v200 import get_bert_feature as en_bert_v200
 
16
 
17
 
18
  class ModelHandler:
@@ -53,6 +54,10 @@ class ModelHandler:
53
  "CLAP_HTSAT_FUSED": [
54
  "https://huggingface.co/laion/clap-htsat-fused/resolve/main/pytorch_model.bin?download=true",
55
  "https://hf-mirror.com/laion/clap-htsat-fused/resolve/main/pytorch_model.bin?download=true",
 
 
 
 
56
  ]
57
  }
58
 
@@ -66,6 +71,7 @@ class ModelHandler:
66
  "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM": "bf0dab8ad87bd7c22e85ec71e04f2240804fda6d33196157d6b5923af6ea1201",
67
  "WAV2VEC2_LARGE_ROBUST_12_FT_EMOTION_MSP_DIM": "176d9d1ce29a8bddbab44068b9c1c194c51624c7f1812905e01355da58b18816",
68
  "CLAP_HTSAT_FUSED": "1ed5d0215d887551ddd0a49ce7311b21429ebdf1e6a129d4e68f743357225253",
 
69
  }
70
  self.model_path = {
71
  "CHINESE_ROBERTA_WWM_EXT_LARGE": os.path.join(config.ABS_PATH,
@@ -78,11 +84,13 @@ class ModelHandler:
78
  "bert_vits2/bert/deberta-v2-large-japanese-char-wwm"),
79
  "WAV2VEC2_LARGE_ROBUST_12_FT_EMOTION_MSP_DIM": os.path.join(config.ABS_PATH,
80
  "bert_vits2/emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim"),
81
- "CLAP_HTSAT_FUSED": os.path.join(config.ABS_PATH, "bert_vits2/emotional/clap-htsat-fused")
 
 
82
  }
83
 
84
  self.lang_bert_func_map = {"zh": zh_bert, "en": en_bert, "ja": ja_bert, "ja_v111": ja_bert_v111,
85
- "ja_v200": ja_bert_v200, "en_v200": en_bert_v200}
86
 
87
  self.bert_models = {} # Value: (tokenizer, model, reference_count)
88
  self.emotion = None
 
13
  from bert_vits2.text.japanese_bert_v111 import get_bert_feature as ja_bert_v111
14
  from bert_vits2.text.japanese_bert_v200 import get_bert_feature as ja_bert_v200
15
  from bert_vits2.text.english_bert_mock_v200 import get_bert_feature as en_bert_v200
16
+ from bert_vits2.text.chinese_bert_extra import get_bert_feature as zh_bert_extra
17
 
18
 
19
  class ModelHandler:
 
54
  "CLAP_HTSAT_FUSED": [
55
  "https://huggingface.co/laion/clap-htsat-fused/resolve/main/pytorch_model.bin?download=true",
56
  "https://hf-mirror.com/laion/clap-htsat-fused/resolve/main/pytorch_model.bin?download=true",
57
+ ],
58
+ "Erlangshen-MegatronBert-1.3B-Chinese": [
59
+ "https://huggingface.co/IDEA-CCNL/Erlangshen-UniMC-MegatronBERT-1.3B-Chinese/resolve/main/pytorch_model.bin",
60
+ "https://hf-mirror.com/IDEA-CCNL/Erlangshen-UniMC-MegatronBERT-1.3B-Chinese/resolve/main/pytorch_model.bin",
61
  ]
62
  }
63
 
 
71
  "DEBERTA_V2_LARGE_JAPANESE_CHAR_WWM": "bf0dab8ad87bd7c22e85ec71e04f2240804fda6d33196157d6b5923af6ea1201",
72
  "WAV2VEC2_LARGE_ROBUST_12_FT_EMOTION_MSP_DIM": "176d9d1ce29a8bddbab44068b9c1c194c51624c7f1812905e01355da58b18816",
73
  "CLAP_HTSAT_FUSED": "1ed5d0215d887551ddd0a49ce7311b21429ebdf1e6a129d4e68f743357225253",
74
+ "Erlangshen-MegatronBert-1.3B-Chinese": "3456bb8f2c7157985688a4cb5cecdb9e229cb1dcf785b01545c611462ffe3579",
75
  }
76
  self.model_path = {
77
  "CHINESE_ROBERTA_WWM_EXT_LARGE": os.path.join(config.ABS_PATH,
 
84
  "bert_vits2/bert/deberta-v2-large-japanese-char-wwm"),
85
  "WAV2VEC2_LARGE_ROBUST_12_FT_EMOTION_MSP_DIM": os.path.join(config.ABS_PATH,
86
  "bert_vits2/emotional/wav2vec2-large-robust-12-ft-emotion-msp-dim"),
87
+ "CLAP_HTSAT_FUSED": os.path.join(config.ABS_PATH, "bert_vits2/emotional/clap-htsat-fused"),
88
+ "Erlangshen-MegatronBert-1.3B-Chinese": os.path.join(config.ABS_PATH,
89
+ "bert_vits2/bert/Erlangshen-MegatronBert-1.3B-Chinese"),
90
  }
91
 
92
  self.lang_bert_func_map = {"zh": zh_bert, "en": en_bert, "ja": ja_bert, "ja_v111": ja_bert_v111,
93
+ "ja_v200": ja_bert_v200, "en_v200": en_bert_v200, "zh_extra": zh_bert_extra}
94
 
95
  self.bert_models = {} # Value: (tokenizer, model, reference_count)
96
  self.emotion = None
bert_vits2/models.py CHANGED
@@ -287,6 +287,7 @@ class TextEncoder(nn.Module):
287
  ja_bert_dim=1024,
288
  num_tones=None,
289
  emotion_embedding=1,
 
290
  ):
291
  super().__init__()
292
  self.n_vocab = n_vocab
@@ -305,9 +306,13 @@ class TextEncoder(nn.Module):
305
  self.language_emb = nn.Embedding(num_languages, hidden_channels)
306
  nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels ** -0.5)
307
  self.bert_proj = nn.Conv1d(1024, hidden_channels, 1)
 
 
 
308
  self.ja_bert_proj = nn.Conv1d(ja_bert_dim, hidden_channels, 1)
309
  self.en_bert_proj = nn.Conv1d(1024, hidden_channels, 1)
310
  self.emotion_embedding = emotion_embedding
 
311
  if self.emotion_embedding == 1:
312
  self.emo_proj = nn.Linear(1024, 1024)
313
  self.emo_quantizer = VectorQuantize(
@@ -356,20 +361,24 @@ class TextEncoder(nn.Module):
356
  self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
357
 
358
  def forward(self, x, x_lengths, tone, language, zh_bert, ja_bert, en_bert, emo=None, sid=None, g=None):
359
- zh_bert_emb = self.bert_proj(zh_bert).transpose(1, 2)
360
- ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2)
361
- en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2)
362
- x = self.emb(x) + self.tone_emb(tone) + self.language_emb(language) + zh_bert_emb + ja_bert_emb + en_bert_emb
 
 
 
363
 
 
364
  if self.emotion_embedding == 1:
365
- emo = emo.to(zh_bert_emb.device)
366
  if emo.size(-1) == 1024:
367
  emo_emb = self.emo_proj(emo.unsqueeze(1))
368
  emo_emb_ = []
369
  for i in range(emo_emb.size(0)):
370
  temp_emo_emb, _, _ = self.emo_quantizer(
371
- emo_emb[i].unsqueeze(0)
372
- )
373
  emo_emb_.append(temp_emo_emb)
374
  emo_emb = torch.cat(emo_emb_, dim=0).to(emo_emb.device)
375
  else:
@@ -694,6 +703,7 @@ class SynthesizerTrn(nn.Module):
694
  ja_bert_dim=1024,
695
  num_tones=None,
696
  emotion_embedding=False,
 
697
  **kwargs):
698
 
699
  super().__init__()
@@ -738,7 +748,8 @@ class SynthesizerTrn(nn.Module):
738
  symbols=symbols,
739
  ja_bert_dim=ja_bert_dim,
740
  num_tones=num_tones,
741
- emotion_embedding=self.emotion_embedding
 
742
  )
743
  self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates,
744
  upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)
@@ -769,8 +780,8 @@ class SynthesizerTrn(nn.Module):
769
  g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
770
  x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, zh_bert, ja_bert, en_bert, emo, sid, g=g)
771
  logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (sdp_ratio) + self.dp(x, x_mask,
772
- g=g) * (
773
- 1 - sdp_ratio)
774
  w = torch.exp(logw) * x_mask * length_scale
775
  w_ceil = torch.ceil(w)
776
  y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long()
 
287
  ja_bert_dim=1024,
288
  num_tones=None,
289
  emotion_embedding=1,
290
+ zh_bert_extra=False,
291
  ):
292
  super().__init__()
293
  self.n_vocab = n_vocab
 
306
  self.language_emb = nn.Embedding(num_languages, hidden_channels)
307
  nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels ** -0.5)
308
  self.bert_proj = nn.Conv1d(1024, hidden_channels, 1)
309
+ self.zh_bert_extra = zh_bert_extra
310
+ if self.zh_bert_extra:
311
+ self.bert_pre_proj = nn.Conv1d(2048, 1024, 1)
312
  self.ja_bert_proj = nn.Conv1d(ja_bert_dim, hidden_channels, 1)
313
  self.en_bert_proj = nn.Conv1d(1024, hidden_channels, 1)
314
  self.emotion_embedding = emotion_embedding
315
+
316
  if self.emotion_embedding == 1:
317
  self.emo_proj = nn.Linear(1024, 1024)
318
  self.emo_quantizer = VectorQuantize(
 
361
  self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
362
 
363
  def forward(self, x, x_lengths, tone, language, zh_bert, ja_bert, en_bert, emo=None, sid=None, g=None):
364
+ x = self.emb(x) + self.tone_emb(tone) + self.language_emb(language)
365
+
366
+ if self.zh_bert_extra:
367
+ zh_bert = self.bert_pre_proj(zh_bert)
368
+ x += self.bert_proj(zh_bert).transpose(1, 2)
369
+ x += self.ja_bert_proj(ja_bert).transpose(1, 2)
370
+ x += self.en_bert_proj(en_bert).transpose(1, 2)
371
 
372
+ x *= math.sqrt(self.hidden_channels) # [b, t, h]
373
  if self.emotion_embedding == 1:
374
+ # emo = emo.to(zh_bert_emb.device)
375
  if emo.size(-1) == 1024:
376
  emo_emb = self.emo_proj(emo.unsqueeze(1))
377
  emo_emb_ = []
378
  for i in range(emo_emb.size(0)):
379
  temp_emo_emb, _, _ = self.emo_quantizer(
380
+ emo_emb[i].unsqueeze(0).to(emo.device)
381
+ )
382
  emo_emb_.append(temp_emo_emb)
383
  emo_emb = torch.cat(emo_emb_, dim=0).to(emo_emb.device)
384
  else:
 
703
  ja_bert_dim=1024,
704
  num_tones=None,
705
  emotion_embedding=False,
706
+ zh_bert_extra=False,
707
  **kwargs):
708
 
709
  super().__init__()
 
748
  symbols=symbols,
749
  ja_bert_dim=ja_bert_dim,
750
  num_tones=num_tones,
751
+ emotion_embedding=self.emotion_embedding,
752
+ zh_bert_extra=zh_bert_extra,
753
  )
754
  self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates,
755
  upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)
 
780
  g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
781
  x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, zh_bert, ja_bert, en_bert, emo, sid, g=g)
782
  logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (sdp_ratio) + self.dp(x, x_mask,
783
+ g=g) * (
784
+ 1 - sdp_ratio)
785
  w = torch.exp(logw) * x_mask * length_scale
786
  w_ceil = torch.ceil(w)
787
  y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long()
bert_vits2/models_v230.py CHANGED
@@ -374,19 +374,13 @@ class TextEncoder(nn.Module):
374
  self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
375
 
376
  def forward(self, x, x_lengths, tone, language, zh_bert, ja_bert, en_bert, g=None):
377
- zh_bert_emb = self.bert_proj(zh_bert).transpose(1, 2)
378
- ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2)
379
- en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2)
380
- x = (
381
- self.emb(x)
382
- + self.tone_emb(tone)
383
- + self.language_emb(language)
384
- + zh_bert_emb
385
- + ja_bert_emb
386
- + en_bert_emb
387
- ) * math.sqrt(
388
- self.hidden_channels
389
- ) # [b, t, h]
390
  x = torch.transpose(x, 1, -1) # [b, h, t]
391
  x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
392
  x.dtype
@@ -940,7 +934,7 @@ class SynthesizerTrn(nn.Module):
940
  sid,
941
  tone,
942
  language,
943
- bert,
944
  ja_bert,
945
  en_bert,
946
  noise_scale=0.667,
@@ -958,7 +952,7 @@ class SynthesizerTrn(nn.Module):
958
  else:
959
  g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
960
  x, m_p, logs_p, x_mask = self.enc_p(
961
- x, x_lengths, tone, language, bert, ja_bert, en_bert, g=g
962
  )
963
  logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (
964
  sdp_ratio
 
374
  self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
375
 
376
  def forward(self, x, x_lengths, tone, language, zh_bert, ja_bert, en_bert, g=None):
377
+ x = self.emb(x) + self.tone_emb(tone) + self.language_emb(language)
378
+
379
+ x +=self.bert_proj(zh_bert).transpose(1, 2)
380
+ x += self.ja_bert_proj(ja_bert).transpose(1, 2)
381
+ x += self.en_bert_proj(en_bert).transpose(1, 2)
382
+
383
+ x *= math.sqrt(self.hidden_channels) # [b, t, h]
 
 
 
 
 
 
384
  x = torch.transpose(x, 1, -1) # [b, h, t]
385
  x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
386
  x.dtype
 
934
  sid,
935
  tone,
936
  language,
937
+ zh_bert,
938
  ja_bert,
939
  en_bert,
940
  noise_scale=0.667,
 
952
  else:
953
  g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
954
  x, m_p, logs_p, x_mask = self.enc_p(
955
+ x, x_lengths, tone, language, zh_bert, ja_bert, en_bert, g=g
956
  )
957
  logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (
958
  sdp_ratio
bert_vits2/text/chinese_bert_extra.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+ from utils.config_manager import global_config
4
+
5
+
6
+ def get_bert_feature(text, word2ph, tokenizer, model, device=global_config.DEVICE, style_text=None, style_weight=0.7, **kwargs):
7
+ with torch.no_grad():
8
+ inputs = tokenizer(text, return_tensors='pt')
9
+ for i in inputs:
10
+ inputs[i] = inputs[i].to(device)
11
+ res = model(**inputs, output_hidden_states=True)
12
+ res = torch.nn.functional.normalize(
13
+ torch.cat(res["hidden_states"][-3:-2], -1)[0], dim=0
14
+ ).cpu()
15
+ if style_text:
16
+ style_inputs = tokenizer(style_text, return_tensors="pt")
17
+ for i in style_inputs:
18
+ style_inputs[i] = style_inputs[i].to(device)
19
+ style_res = model(**style_inputs, output_hidden_states=True)
20
+ style_res = torch.nn.functional.normalize(
21
+ torch.cat(style_res["hidden_states"][-3:-2], -1)[0], dim=0
22
+ ).cpu()
23
+ style_res_mean = style_res.mean(0)
24
+
25
+ assert len(word2ph) == len(text) + 2
26
+ word2phone = word2ph
27
+ phone_level_feature = []
28
+ for i in range(len(word2phone)):
29
+ if style_text:
30
+ repeat_feature = (
31
+ res[i].repeat(word2phone[i], 1) * (1 - style_weight)
32
+ + style_res_mean.repeat(word2phone[i], 1) * style_weight
33
+ )
34
+ else:
35
+ repeat_feature = res[i].repeat(word2phone[i], 1)
36
+ phone_level_feature.append(repeat_feature)
37
+
38
+ phone_level_feature = torch.cat(phone_level_feature, dim=0)
39
+
40
+ return phone_level_feature.T
41
+
42
+
43
+ if __name__ == '__main__':
44
+
45
+ word_level_feature = torch.rand(38, 2048) # 12个词,每个词2048维特征
46
+ word2phone = [1, 2, 1, 2, 2, 1, 2, 2, 1, 2, 2, 1, 2, 2, 2, 2, 2, 1, 1, 2, 2, 1, 2, 2, 2, 2, 1, 2, 2, 2, 2, 2, 1, 2,
47
+ 2, 2, 2, 1]
48
+
49
+ # 计算总帧数
50
+ total_frames = sum(word2phone)
51
+ print(word_level_feature.shape)
52
+ print(word2phone)
53
+ phone_level_feature = []
54
+ for i in range(len(word2phone)):
55
+ print(word_level_feature[i].shape)
56
+
57
+ # 对每个词重复word2phone[i]次
58
+ repeat_feature = word_level_feature[i].repeat(word2phone[i], 1)
59
+ phone_level_feature.append(repeat_feature)
60
+
61
+ phone_level_feature = torch.cat(phone_level_feature, dim=0)
62
+ print(phone_level_feature.shape) # torch.Size([36, 2048])