Team Ai
Apppublic

codejin/diffsingerkr

sourceHugging Facemitupdated 4y agoView on Hugging Face
6likes
Inference.py168 linesDownload Raw Back to root
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