jone/Music_Source_Separation
3
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 