from torch.utils.data import ConcatDataset from ditto.data_loaders.mtat import MTATDataset from ditto.data_loaders.pond5 import Pond5Dataset from ditto.data_loaders.ytm import YTMDataset from ditto.data_loaders.genius_lyric import GeniusLyricDataset from ditto.data_loaders.cleaned_concat import CleanedConcatDataset from ditto.data_loaders.cleaned_ytmsd import CleanedYTMSDDataset from ditto.data_loaders.artist_ytmsd import ArtistYTMSDDataset def get_dataset(dataset, split, num_samples=-1, input_length_s=30.0): if dataset == "mtat": return MTATDataset( split=split, num_samples=num_samples, input_length_s=input_length_s ) elif dataset == "pond5": return Pond5Dataset( split=split, num_samples=num_samples, input_length_s=input_length_s ) elif dataset == "genius_lyric": return GeniusLyricDataset( split=split, num_samples=num_samples, input_length_s=input_length_s ) elif dataset == "concat": datasets = [ Pond5Dataset( split=split, num_samples=num_samples, input_length_s=input_length_s ), YTMDataset( split=split, num_samples=num_samples, input_length_s=input_length_s ), GeniusLyricDataset( split=split, num_samples=num_samples, input_length_s=input_length_s, is_concat=True, ), ] return ConcatDataset(datasets) elif dataset == "cleaned_concat": return CleanedConcatDataset( split=split, num_samples=num_samples, input_length_s=input_length_s ) elif dataset == "cleaned_ytmsd": return CleanedYTMSDDataset( split=split, num_samples=num_samples, input_length_s=input_length_s ) elif dataset == "artist_ytmsd": return ArtistYTMSDDataset( split=split, num_samples=num_samples, input_length_s=input_length_s ) else: print("%s dataset is not supported yet" % dataset)