Team Ai
Apppublic

codejin/diffsingerkr

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