Spaces:
Sleeping
Sleeping
Update Bert-VITS2 Extra
Browse files- bert_vits2/bert/Erlangshen-MegatronBert-1.3B-Chinese/config.json +1 -0
- bert_vits2/bert/Erlangshen-MegatronBert-1.3B-Chinese/vocab.txt +0 -0
- bert_vits2/bert_vits2.py +72 -53
- bert_vits2/model_handler.py +10 -2
- bert_vits2/models.py +21 -10
- bert_vits2/models_v230.py +9 -15
- bert_vits2/text/chinese_bert_extra.py +62 -0
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
|
| 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
|
| 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.
|
| 230 |
elif self.hps_ms.model.emotion_embedding == 2:
|
| 231 |
-
|
| 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.
|
| 250 |
elif self.hps_ms.model.emotion_embedding == 2:
|
| 251 |
-
|
| 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 |
-
|
| 262 |
|
| 263 |
for idx, (_text, lang) in enumerate(sentences_list):
|
| 264 |
skip_start = idx != 0
|
| 265 |
skip_end = idx != len(sentences_list) - 1
|
| 266 |
-
|
| 267 |
-
|
|
|
|
| 268 |
if skip_start:
|
| 269 |
-
|
| 270 |
-
|
| 271 |
-
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
if skip_end:
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
| 287 |
-
|
| 288 |
-
|
| 289 |
-
|
| 290 |
-
zh_bert = torch.
|
| 291 |
-
ja_bert = torch.
|
| 292 |
-
en_bert = torch.
|
| 293 |
-
phones = torch.
|
| 294 |
-
tones = torch.
|
| 295 |
-
lang_ids = torch.
|
|
|
|
| 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 |
-
|
| 360 |
-
|
| 361 |
-
|
| 362 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
|
| 773 |
-
|
| 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 |
-
|
| 378 |
-
|
| 379 |
-
|
| 380 |
-
x = (
|
| 381 |
-
|
| 382 |
-
|
| 383 |
-
|
| 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 |
-
|
| 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,
|
| 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])
|