jone/Music_Source_Separation
3
1from typing import Dict, List, NoReturn, Optional2 3import h5py4import librosa5import numpy as np6import torch7from pytorch_lightning.core.datamodule import LightningDataModule8 9from bytesep.data.samplers import DistributedSamplerWrapper10from bytesep.utils import int16_to_float3211 12 13class DataModule(LightningDataModule):14 def __init__(15 self,16 train_sampler: object,17 train_dataset: object,18 num_workers: int,19 distributed: bool,20 ):21 r"""Data module.22 23 Args:24 train_sampler: Sampler object25 train_dataset: Dataset object26 num_workers: int27 distributed: bool28 """29 super().__init__()30 self._train_sampler = train_sampler31 self.train_dataset = train_dataset32 self.num_workers = num_workers33 self.distributed = distributed34 35 def setup(self, stage: Optional[str] = None) -> NoReturn:36 r"""called on every device."""37 38 # SegmentSampler is used for selecting segments for training.39 # On multiple devices, each SegmentSampler samples a part of mini-batch40 # data.41 if self.distributed:42 self.train_sampler = DistributedSamplerWrapper(self._train_sampler)43 44 else:45 self.train_sampler = self._train_sampler46 47 def train_dataloader(self) -> torch.utils.data.DataLoader:48 r"""Get train loader."""49 train_loader = torch.utils.data.DataLoader(50 dataset=self.train_dataset,51 batch_sampler=self.train_sampler,52 collate_fn=collate_fn,53 num_workers=self.num_workers,54 pin_memory=True,55 )56 57 return train_loader58 59 60class Dataset:61 def __init__(self, augmentor: object, segment_samples: int):62 r"""Used for getting data according to a meta.63 64 Args:65 augmentor: Augmentor class66 segment_samples: int67 """68 self.augmentor = augmentor69 self.segment_samples = segment_samples70 71 def __getitem__(self, meta: Dict) -> Dict:72 r"""Return data according to a meta. E.g., an input meta looks like: {73 'vocals': [['song_A.h5', 6332760, 6465060], ['song_B.h5', 198450, 330750]],74 'accompaniment': [['song_C.h5', 24232920, 24365250], ['song_D.h5', 1569960, 1702260]]}.75 }76 77 Then, vocals segments of song_A and song_B will be mixed (mix-audio augmentation).78 Accompaniment segments of song_C and song_B will be mixed (mix-audio augmentation).79 Finally, mixture is created by summing vocals and accompaniment.80 81 Args:82 meta: dict, e.g., {83 'vocals': [['song_A.h5', 6332760, 6465060], ['song_B.h5', 198450, 330750]],84 'accompaniment': [['song_C.h5', 24232920, 24365250], ['song_D.h5', 1569960, 1702260]]}85 }86 87 Returns:88 data_dict: dict, e.g., {89 'vocals': (channels, segments_num),90 'accompaniment': (channels, segments_num),91 'mixture': (channels, segments_num),92 }93 """94 source_types = meta.keys()95 data_dict = {}96 97 for source_type in source_types:98 # E.g., ['vocals', 'bass', ...]99 100 waveforms = [] # Audio segments to be mix-audio augmented.101 102 for m in meta[source_type]:103 # E.g., {104 # 'hdf5_path': '.../song_A.h5',105 # 'key_in_hdf5': 'vocals',106 # 'begin_sample': '13406400',107 # 'end_sample': 13538700,108 # }109 110 hdf5_path = m['hdf5_path']111 key_in_hdf5 = m['key_in_hdf5']112 bgn_sample = m['begin_sample']113 end_sample = m['end_sample']114 115 with h5py.File(hdf5_path, 'r') as hf:116 117 if source_type == 'audioset':118 index_in_hdf5 = m['index_in_hdf5']119 waveform = int16_to_float32(120 hf['waveform'][index_in_hdf5][bgn_sample:end_sample]121 )122 waveform = waveform[None, :]123 else:124 waveform = int16_to_float32(125 hf[key_in_hdf5][:, bgn_sample:end_sample]126 )127 128 if self.augmentor:129 waveform = self.augmentor(waveform, source_type)130 131 waveform = librosa.util.fix_length(132 waveform, size=self.segment_samples, axis=1133 )134 # (channels_num, segments_num)135 136 waveforms.append(waveform)137 # E.g., waveforms: [(channels_num, audio_samples), (channels_num, audio_samples)]138 139 # mix-audio augmentation140 data_dict[source_type] = np.sum(waveforms, axis=0)141 # data_dict[source_type]: (channels_num, audio_samples)142 143 # data_dict looks like: {144 # 'voclas': (channels_num, audio_samples),145 # 'accompaniment': (channels_num, audio_samples)146 # }147 148 # Mix segments from different sources.149 mixture = np.sum(150 [data_dict[source_type] for source_type in source_types], axis=0151 )152 data_dict['mixture'] = mixture153 # shape: (channels_num, audio_samples)154 155 return data_dict156 157 158def collate_fn(list_data_dict: List[Dict]) -> Dict:159 r"""Collate mini-batch data to inputs and targets for training.160 161 Args:162 list_data_dict: e.g., [163 {'vocals': (channels_num, segment_samples),164 'accompaniment': (channels_num, segment_samples),165 'mixture': (channels_num, segment_samples)166 },167 {'vocals': (channels_num, segment_samples),168 'accompaniment': (channels_num, segment_samples),169 'mixture': (channels_num, segment_samples)170 },171 ...]172 173 Returns:174 data_dict: e.g. {175 'vocals': (batch_size, channels_num, segment_samples),176 'accompaniment': (batch_size, channels_num, segment_samples),177 'mixture': (batch_size, channels_num, segment_samples)178 }179 """180 data_dict = {}181 182 for key in list_data_dict[0].keys():183 data_dict[key] = torch.Tensor(184 np.array([data_dict[key] for data_dict in list_data_dict])185 )186 187 return data_dict188 