codejin/diffsingerkr
6
1from argparse import Namespace2import torch3import numpy as np4import pickle, os, logging5from typing import Dict, List, Optional6import hgtk7 8from Pattern_Generator import Convert_Feature_Based_Music, Expand_by_Duration9 10def Decompose(syllable: str): 11 onset, nucleus, coda = hgtk.letter.decompose(syllable)12 coda += '_'13 14 return onset, nucleus, coda15 16def Lyric_to_Token(lyric: List[str], token_dict: Dict[str, int]):17 return [18 token_dict[letter]19 for letter in list(lyric)20 ]21 22def Token_Stack(tokens: List[List[int]], token_dict: Dict[str, int], max_length: Optional[int]= None):23 max_token_length = max_length or max([len(token) for token in tokens])24 tokens = np.stack(25 [np.pad(token[:max_token_length], [0, max_token_length - len(token[:max_token_length])], constant_values= token_dict['<X>']) for token in tokens],26 axis= 027 )28 return tokens29 30def Note_Stack(notes: List[List[int]], max_length: Optional[int]= None):31 max_note_length = max_length or max([len(note) for note in notes])32 notes = np.stack(33 [np.pad(note[:max_note_length], [0, max_note_length - len(note[:max_note_length])], constant_values= 0) for note in notes],34 axis= 035 )36 return notes37 38def Duration_Stack(durations: List[List[int]], max_length: Optional[int]= None):39 max_duration_length = max_length or max([len(duration) for duration in durations])40 durations = np.stack(41 [np.pad(duration[:max_duration_length], [0, max_duration_length - len(duration[:max_duration_length])], constant_values= 0) for duration in durations],42 axis= 043 )44 return durations45 46def Feature_Stack(features: List[np.array], max_length: Optional[int]= None):47 max_feature_length = max_length or max([feature.shape[0] for feature in features])48 features = np.stack(49 [np.pad(feature, [[0, max_feature_length - feature.shape[0]], [0, 0]], constant_values= -1.0) for feature in features],50 axis= 051 )52 return features53 54def Log_F0_Stack(log_f0s: List[np.array], max_length: int= None):55 max_log_f0_length = max_length or max([len(log_f0) for log_f0 in log_f0s])56 log_f0s = np.stack(57 [np.pad(log_f0, [0, max_log_f0_length - len(log_f0)], constant_values= 0.0) for log_f0 in log_f0s],58 axis= 059 )60 return log_f0s61 62class Inference_Dataset(torch.utils.data.Dataset):63 def __init__(64 self,65 token_dict: Dict[str, int], 66 singer_info_dict: Dict[str, int],67 genre_info_dict: Dict[str, int],68 durations: List[List[float]],69 lyrics: List[List[str]],70 notes: List[List[int]],71 singers: List[str],72 genres: List[str],73 sample_rate: int,74 frame_shift: int,75 equality_duration: bool= False,76 consonant_duration: int= 377 ):78 super().__init__()79 self.token_dict = token_dict80 self.singer_info_dict = singer_info_dict81 self.genre_info_dict = genre_info_dict82 self.equality_duration = equality_duration83 self.consonant_duration = consonant_duration84 85 self.patterns = []86 for index, (duration, lyric, note, singer, genre) in enumerate(zip(durations, lyrics, notes, singers, genres)):87 if not singer in self.singer_info_dict.keys():88 logging.warn('The singer \'{}\' is incorrect. The pattern \'{}\' is ignoired.'.format(singer, index))89 continue90 if not genre in self.genre_info_dict.keys():91 logging.warn('The genre \'{}\' is incorrect. The pattern \'{}\' is ignoired.'.format(genre, index))92 continue93 94 music = [x for x in zip(duration, lyric, note)] 95 singer_label = singer96 text = lyric97 98 lyric, note, duration = Convert_Feature_Based_Music(99 music= music,100 sample_rate= sample_rate,101 frame_shift= frame_shift,102 consonant_duration= consonant_duration,103 equality_duration= equality_duration104 )105 lyric_expand, note_expand, duration_expand = Expand_by_Duration(lyric, note, duration)106 107 singer = self.singer_info_dict[singer]108 genre = self.genre_info_dict[genre]109 110 self.patterns.append((lyric_expand, note_expand, duration_expand, singer, genre, singer_label, text))111 112 def __getitem__(self, idx):113 lyric, note, duration, singer, genre, singer_label, text = self.patterns[idx]114 115 return Lyric_to_Token(lyric, self.token_dict), note, duration, singer, genre, singer_label, text116 117 def __len__(self):118 return len(self.patterns)119 120class Inference_Collater:121 def __init__(self,122 token_dict: Dict[str, int]123 ):124 self.token_dict = token_dict125 126 def __call__(self, batch):127 tokens, notes, durations, singers, genres, singer_labels, lyrics = zip(*batch)128 129 lengths = np.array([len(token) for token in tokens])130 131 max_length = max(lengths)132 133 tokens = Token_Stack(tokens, self.token_dict, max_length)134 notes = Note_Stack(notes, max_length) 135 durations = Duration_Stack(durations, max_length)136 137 tokens = torch.LongTensor(tokens) # [Batch, Time]138 notes = torch.LongTensor(notes) # [Batch, Time]139 durations = torch.LongTensor(durations) # [Batch, Time]140 lengths = torch.LongTensor(lengths) # [Batch]141 singers = torch.LongTensor(singers) # [Batch]142 genres = torch.LongTensor(genres) # [Batch]143 144 lyrics = [''.join([(x if x != '<X>' else ' ') for x in lyric]) for lyric in lyrics]145 146 return tokens, notes, durations, lengths, singers, genres, singer_labels, lyrics