(text, word2ph, bert_model, tokenizer, device)
| 292 | return phones,None,norm_text |
| 293 | |
| 294 | def get_bert_feature(text, word2ph, bert_model, tokenizer, device): |
| 295 | with torch.no_grad(): |
| 296 | inputs = tokenizer(text, return_tensors="pt") |
| 297 | for i in inputs: |
| 298 | inputs[i] = inputs[i].to(device) |
| 299 | res = bert_model(**inputs, output_hidden_states=True) |
| 300 | res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()[1:-1] |
| 301 | assert len(word2ph) == len(text) |
| 302 | phone_level_feature = [] |
| 303 | for i in range(len(word2ph)): |
| 304 | repeat_feature = res[i].repeat(word2ph[i], 1) |
| 305 | phone_level_feature.append(repeat_feature) |
| 306 | phone_level_feature = torch.cat(phone_level_feature, dim=0) |
| 307 | return phone_level_feature.T |
| 308 | |
| 309 | |
| 310 | def only_punc(text): |
nothing calls this directly
no outgoing calls
no test coverage detected