CoolFace
Apppublic

jone/Music_Source_Separation

sourceHugging Faceupdated 4y agoView on Hugging Face
3likes
data_modules.py188 linesDownload Raw Back to data
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