CoolFace
Apppublic

jone/Music_Source_Separation

sourceHugging Faceupdated 4y agoView on Hugging Face
3likes
samplers.py189 linesDownload Raw Back to data
1import pickle2from typing import Dict, List, NoReturn3 4import numpy as np5import torch.distributed as dist6 7 8class SegmentSampler:9    def __init__(10        self,11        indexes_path: str,12        segment_samples: int,13        mixaudio_dict: Dict,14        batch_size: int,15        steps_per_epoch: int,16        random_seed=1234,17    ):18        r"""Sample training indexes of sources.19 20        Args:21            indexes_path: str, path of indexes dict22            segment_samplers: int23            mixaudio_dict, dict, including hyper-parameters for mix-audio data24                augmentation, e.g., {'voclas': 2, 'accompaniment': 2}25            batch_size: int26            steps_per_epoch: int, #steps_per_epoch is called an `epoch`27            random_seed: int28        """29        self.segment_samples = segment_samples30        self.mixaudio_dict = mixaudio_dict31        self.batch_size = batch_size32        self.steps_per_epoch = steps_per_epoch33 34        self.meta_dict = pickle.load(open(indexes_path, "rb"))35        # E.g., {36        #     'vocals': [37        #         {'hdf5_path': 'songA.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 0, 'end_sample': 132300},38        #         {'hdf5_path': 'songB.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 4410, 'end_sample': 445410},39        #         ...40        #     ],41        #     'accompaniment': [42        #         {'hdf5_path': 'songA.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 0, 'end_sample': 132300},43        #         {'hdf5_path': 'songB.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 4410, 'end_sample': 445410},44        #         ...45        #     ]46        # }47 48        self.source_types = self.meta_dict.keys()49        # E.g., ['vocals', 'accompaniment']50 51        self.pointers_dict = {source_type: 0 for source_type in self.source_types}52        # E.g., {'vocals': 0, 'accompaniment': 0}53 54        self.indexes_dict = {55            source_type: np.arange(len(self.meta_dict[source_type]))56            for source_type in self.source_types57        }58        # E.g. {59        #     'vocals': [0, 1, ..., 225751],60        #     'accompaniment': [0, 1, ..., 225751]61        # }62 63        self.random_state = np.random.RandomState(random_seed)64 65        # Shuffle indexes.66        for source_type in self.source_types:67            self.random_state.shuffle(self.indexes_dict[source_type])68            print("{}: {}".format(source_type, len(self.indexes_dict[source_type])))69 70    def __iter__(self) -> List[Dict]:71        r"""Yield a batch of meta info.72 73        Returns:74            batch_meta_list: (batch_size,) e.g., when mix-audio is 2, looks like [75                {'vocals': [76                    {'hdf5_path': 'songA.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 13406400, 'end_sample': 13538700},77                    {'hdf5_path': 'songB.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 4440870, 'end_sample': 4573170}]78                'accompaniment': [79                    {'hdf5_path': 'songE.h5', 'key_in_hdf5': 'accompaniment', 'begin_sample': 14579460, 'end_sample': 14711760},80                    {'hdf5_path': 'songF.h5', 'key_in_hdf5': 'accompaniment', 'begin_sample': 3995460, 'end_sample': 4127760}]81                }82                ...83            ]84        """85        batch_size = self.batch_size86 87        while True:88            batch_meta_dict = {source_type: [] for source_type in self.source_types}89 90            for source_type in self.source_types:91                # E.g., ['vocals', 'accompaniment']92 93                # Loop until get a mini-batch.94                while len(batch_meta_dict[source_type]) != batch_size:95 96                    largest_index = (97                        len(self.indexes_dict[source_type])98                        - self.mixaudio_dict[source_type]99                    )100                    # E.g., 225750 = 225752 - 2101 102                    if self.pointers_dict[source_type] > largest_index:103 104                        # Reset pointer, and shuffle indexes.105                        self.pointers_dict[source_type] = 0106                        self.random_state.shuffle(self.indexes_dict[source_type])107 108                    source_metas = []109                    mix_audios_num = self.mixaudio_dict[source_type]110 111                    for _ in range(mix_audios_num):112 113                        pointer = self.pointers_dict[source_type]114                        # E.g., 1115 116                        index = self.indexes_dict[source_type][pointer]117                        # E.g., 12231118 119                        self.pointers_dict[source_type] += 1120 121                        source_meta = self.meta_dict[source_type][index]122                        # E.g., ['song_A.h5', 198450, 330750]123 124                        # source_metas.append(new_source_meta)125                        source_metas.append(source_meta)126 127                    batch_meta_dict[source_type].append(source_metas)128            # When mix-audio is 2, batch_meta_dict looks like: {129            #     'vocals': [130            #         [{'hdf5_path': 'songA.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 13406400, 'end_sample': 13538700},131            #          {'hdf5_path': 'songB.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 4440870, 'end_sample': 4573170}],132            #         [{'hdf5_path': 'songC.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 1186290, 'end_sample': 1318590},133            #          {'hdf5_path': 'songD.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 8462790, 'end_sample': 8595090}]134            #     ]135            #     'accompaniment': [136            #         [{'hdf5_path': 'songE.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 24232950, 'end_sample': 24365250},137            #          {'hdf5_path': 'songF.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 1569960, 'end_sample': 1702260}],138            #         [{'hdf5_path': 'songG.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 2795940, 'end_sample': 2928240},139            #          {'hdf5_path': 'songH.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 10923570, 'end_sample': 11055870}]140            #     ]141            # }142 143            batch_meta_list = [144                {145                    source_type: batch_meta_dict[source_type][i]146                    for source_type in self.source_types147                }148                for i in range(batch_size)149            ]150            # When mix-audio is 2, batch_meta_list looks like: [151            #     {'vocals': [152            #         {'hdf5_path': 'songA.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 13406400, 'end_sample': 13538700},153            #         {'hdf5_path': 'songB.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 4440870, 'end_sample': 4573170}]154            #      'accompaniment': [155            #         {'hdf5_path': 'songE.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 14579460, 'end_sample': 14711760},156            #         {'hdf5_path': 'songF.h5', 'key_in_hdf5': 'vocals', 'begin_sample': 3995460, 'end_sample': 4127760}]157            #     }158            #     ...159            # ]160 161            yield batch_meta_list162 163    def __len__(self) -> int:164        return self.steps_per_epoch165 166    def state_dict(self) -> Dict:167        state = {'pointers_dict': self.pointers_dict, 'indexes_dict': self.indexes_dict}168        return state169 170    def load_state_dict(self, state) -> NoReturn:171        self.pointers_dict = state['pointers_dict']172        self.indexes_dict = state['indexes_dict']173 174 175class DistributedSamplerWrapper:176    def __init__(self, sampler):177        r"""Distributed wrapper of sampler."""178        self.sampler = sampler179 180    def __iter__(self):181        num_replicas = dist.get_world_size()182        rank = dist.get_rank()183 184        for indices in self.sampler:185            yield indices[rank::num_replicas]186 187    def __len__(self) -> int:188        return len(self.sampler)189