import torch import torchaudio import numpy as np from tqdm import tqdm from suno_boost.data import CorruptAudioDataset if __name__ == "__main__": audio_dir = "/app/suno/christian/data/codec_audio/genius_hq" samples = 262144 sample_rate = 48000 corrupts = [ "stereo_to_mono", ] corrupt_probs = [0.9] ffmpeg_filters = ["lowpass", "highpass"] ffmpeg_filter_probs = [0.1, 0.1] dataset = CorruptAudioDataset( audio_dir, samples, sample_rate, corrupts=corrupts, corrupt_probs=corrupt_probs, ffmpeg_filters=ffmpeg_filters, ffmpeg_filter_probs=ffmpeg_filter_probs, codec_names=["dac_2c_25x12"], mp3_codec=1.0, ) dataloader = torch.utils.data.DataLoader( dataset, batch_size=16, drop_last=True, num_workers=8, ) save_audio = False for bidx, batch in enumerate(tqdm(dataloader)): low_quality, high_quality = batch # print(bidx, low_quality.shape, high_quality.shape) # save one random output per batch if save_audio: elem_idx = np.random.randint(0, low_quality.shape[0]) low_quality_item = low_quality[elem_idx] high_quality_item = high_quality[elem_idx] # peak normalize # low_quality_item /= low_quality_item.abs().max() # high_quality_item /= high_quality_item.abs().max() torchaudio.save( f"debug/{bidx:03d}_low_quality.wav", low_quality_item, 48000 ) torchaudio.save( f"debug/{bidx:03d}_high_quality.wav", high_quality_item, 48000 ) if bidx > 10: break