| 51 | self.vocoder.remove_weight_norm() |
| 52 | |
| 53 | def forward(self, inputs: Dict[str, Tensor]) -> Dict[str, Tensor]: |
| 54 | target_wav_path = inputs['target_wav'] |
| 55 | source_wav_path = inputs['source_wav'] |
| 56 | save_wav_path = inputs['save_path'] |
| 57 | |
| 58 | with torch.no_grad(): |
| 59 | source_enc = self.encoder.inference(source_wav_path).to( |
| 60 | self.device) |
| 61 | |
| 62 | spk_emb = self.spk_emb.forward(target_wav_path).to(self.device) |
| 63 | |
| 64 | style_mc = self.encoder.get_feats(target_wav_path).to(self.device) |
| 65 | |
| 66 | coded_sp_converted_norm = self.converter(source_enc, spk_emb, |
| 67 | style_mc) |
| 68 | |
| 69 | wav = self.vocoder(coded_sp_converted_norm.permute([0, 2, 1])) |
| 70 | if os.path.exists(save_wav_path): |
| 71 | sf.write(save_wav_path, |
| 72 | wav.flatten().cpu().data.numpy(), 16000) |
| 73 | |
| 74 | return wav.flatten().cpu().data.numpy() |