codejin/diffsingerkr
6
1import torch2import numpy as np3import logging, yaml, os, sys, argparse, math4import matplotlib.pyplot as plt5from tqdm import tqdm6from librosa import griffinlim7 8from Modules.Modules import DiffSinger9from Datasets import Inference_Dataset as Dataset, Inference_Collater as Collater10from meldataset import spectral_de_normalize_torch11from Arg_Parser import Recursive_Parse12 13import matplotlib as mpl14# 유니코드 깨짐현상 해결15mpl.rcParams['axes.unicode_minus'] = False16# 나눔고딕 폰트 적용17plt.rcParams["font.family"] = 'NanumGothic'18 19logging.basicConfig(20 level=logging.INFO, stream=sys.stdout,21 format= '%(asctime)s (%(module)s:%(lineno)d) %(levelname)s: %(message)s'22 )23 24class Inferencer:25 def __init__(26 self,27 hp_path: str,28 checkpoint_path: str,29 batch_size= 130 ):31 self.device = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')32 33 self.hp = Recursive_Parse(yaml.load(34 open(hp_path, encoding='utf-8'),35 Loader=yaml.Loader36 ))37 38 self.model = DiffSinger(self.hp).to(self.device)39 if self.hp.Feature_Type == 'Mel':40 self.vocoder = torch.jit.load('vocoder.pts', map_location='cpu').to(self.device)41 42 if self.hp.Feature_Type == 'Spectrogram':43 self.feature_range_info_dict = yaml.load(open(self.hp.Spectrogram_Range_Info_Path), Loader=yaml.Loader)44 if self.hp.Feature_Type == 'Mel':45 self.feature_range_info_dict = yaml.load(open(self.hp.Mel_Range_Info_Path), Loader=yaml.Loader)46 self.index_singer_dict = {47 value: key48 for key, value in yaml.load(open(self.hp.Singer_Info_Path), Loader=yaml.Loader).items()49 }50 51 if self.hp.Feature_Type == 'Spectrogram':52 self.feature_size = self.hp.Sound.N_FFT // 2 + 153 elif self.hp.Feature_Type == 'Mel':54 self.feature_size = self.hp.Sound.Mel_Dim55 else:56 raise ValueError('Unknown feature type: {}'.format(self.hp.Feature_Type))57 58 self.Load_Checkpoint(checkpoint_path)59 self.batch_size = batch_size60 61 def Dataset_Generate(self, message_times_list, lyrics, notes, singers, genres):62 token_dict = yaml.load(open(self.hp.Token_Path), Loader=yaml.Loader)63 singer_info_dict = yaml.load(open(self.hp.Singer_Info_Path), Loader=yaml.Loader)64 genre_info_dict = yaml.load(open(self.hp.Genre_Info_Path), Loader=yaml.Loader)65 66 return torch.utils.data.DataLoader(67 dataset= Dataset(68 token_dict= token_dict,69 singer_info_dict= singer_info_dict,70 genre_info_dict= genre_info_dict,71 durations= message_times_list,72 lyrics= lyrics,73 notes= notes,74 singers= singers,75 genres= genres,76 sample_rate= self.hp.Sound.Sample_Rate,77 frame_shift= self.hp.Sound.Frame_Shift,78 equality_duration= self.hp.Duration.Equality,79 consonant_duration= self.hp.Duration.Consonant_Duration80 ),81 shuffle= False,82 collate_fn= Collater(83 token_dict= token_dict84 ),85 batch_size= self.batch_size,86 num_workers= 0,87 pin_memory= True88 )89 90 def Load_Checkpoint(self, path):91 state_dict = torch.load(path, map_location= 'cpu')92 self.model.load_state_dict(state_dict['Model']['DiffSVS']) 93 self.steps = state_dict['Steps']94 95 self.model.eval()96 97 logging.info('Checkpoint loaded at {} steps.'.format(self.steps))98 99 @torch.inference_mode()100 def Inference_Step(self, tokens, notes, durations, lengths, singers, genres, singer_labels, ddim_steps):101 tokens = tokens.to(self.device, non_blocking=True)102 notes = notes.to(self.device, non_blocking=True)103 durations = durations.to(self.device, non_blocking=True)104 lengths = lengths.to(self.device, non_blocking=True)105 singers = singers.to(self.device, non_blocking=True)106 genres = genres.to(self.device, non_blocking=True)107 108 linear_predictions, diffusion_predictions, _, _ = self.model(109 tokens= tokens,110 notes= notes,111 durations= durations,112 lengths= lengths,113 genres= genres,114 singers= singers,115 ddim_steps= ddim_steps116 )117 linear_predictions = linear_predictions.clamp(-1.0, 1.0)118 diffusion_predictions = diffusion_predictions.clamp(-1.0, 1.0)119 120 linear_prediction_list, diffusion_prediction_list = [], []121 for linear_prediction, diffusion_prediction, singer in zip(linear_predictions, diffusion_predictions, singer_labels):122 feature_max = self.feature_range_info_dict[singer]['Max']123 feature_min = self.feature_range_info_dict[singer]['Min']124 linear_prediction_list.append((linear_prediction + 1.0) / 2.0 * (feature_max - feature_min) + feature_min)125 diffusion_prediction_list.append((diffusion_prediction + 1.0) / 2.0 * (feature_max - feature_min) + feature_min)126 linear_predictions = torch.stack(linear_prediction_list, dim= 0)127 diffusion_predictions = torch.stack(diffusion_prediction_list, dim= 0)128 129 if self.hp.Feature_Type == 'Mel':130 audios = self.vocoder(diffusion_predictions)131 if audios.ndim == 1: # This is temporal because of the vocoder problem.132 audios = audios.unsqueeze(0)133 audios = [134 audio[:min(length * self.hp.Sound.Frame_Shift, audio.size(0))].cpu().numpy()135 for audio, length in zip(audios, lengths)136 ]137 elif self.hp.Feature_Type == 'Spectrogram':138 audios = []139 for prediction, length in zip(140 diffusion_predictions,141 lengths142 ):143 prediction = spectral_de_normalize_torch(prediction).cpu().numpy()144 audio = griffinlim(prediction)[:min(prediction.size(1), length) * self.hp.Sound.Frame_Shift]145 audio = (audio / np.abs(audio).max() * 32767.5).astype(np.int16)146 audios.append(audio)147 148 return audios149 150 def Inference_Epoch(self, message_times_list, lyrics, notes, singers, genres, ddim_steps= None, use_tqdm= True):151 dataloader = self.Dataset_Generate(152 message_times_list= message_times_list,153 lyrics= lyrics,154 notes= notes,155 singers= singers,156 genres= genres157 )158 if use_tqdm:159 dataloader = tqdm(160 dataloader,161 desc='[Inference]',162 total= math.ceil(len(dataloader.dataset) / self.batch_size)163 )164 audios = []165 for tokens, notes, durations, lengths, singers, genres, singer_labels, lyrics in dataloader:166 audios.extend(self.Inference_Step(tokens, notes, durations, lengths, singers, genres, singer_labels, ddim_steps))167 168 return audios