import os import torch import auraloss import torchaudio import numpy as np import pyloudnorm as pyln import matplotlib.pyplot as plt from tqdm import tqdm from ear.data import CorruptAudioDataset if __name__ == "__main__": os.makedirs("debug", exist_ok=True) audio_dir = "/app/suno/data/audio_2ch_24khz_lg/val" samples = 131072 sample_rate = 24000 corrupts = [ "stereo_to_mono", "overlay_copy", "drop_random_samples", "tanh_distortion", "clipping_distortion", "white_noise", ] corrupt_probs = [0.25, 0.1, 0.25, 0.25, 0.25, 0.25] ffmpeg_filters = [ "exciter", "bandpass", "lowpass", "highpass", "bitcrusher", "deesser", ] ffmpeg_filter_probs = [0.0, 0.1, 0.25, 0.25, 0.1, 0.1] meter = pyln.Meter(sample_rate) dataset = CorruptAudioDataset( audio_dir, samples, sample_rate, corrupts=corrupts, corrupt_probs=corrupt_probs, ffmpeg_filters=ffmpeg_filters, ffmpeg_filter_probs=ffmpeg_filter_probs, mp3_codec=0.5, target_loudness_lufs_db=-20.0, ) dataloader = torch.utils.data.DataLoader( dataset, batch_size=8, drop_last=True, num_workers=8, ) save_audio = False melstft = auraloss.freq.MelSTFTLoss( 24000, fft_size=4096, hop_size=1024, win_length=4096, reduction="none", ) min_quant_error = 0.0 max_quant_error = 8.0 num_quant_levels = 32 quant_labels = [] pref_labels = [] for bidx, batch in enumerate(tqdm(dataloader)): audio_in_a, audio_out_a, audio_in_b, audio_out_b = batch print(bidx, audio_in_a.shape, audio_out_a.shape) melstft_error_a = melstft( audio_in_a.mean(dim=1, keepdim=True), audio_out_a.mean(dim=1, keepdim=True), ).mean(dim=(1, 2)) melstft_error_b = melstft( audio_in_b.mean(dim=1, keepdim=True), audio_out_b.mean(dim=1, keepdim=True), ).mean(dim=(1, 2)) pref_label = (melstft_error_a > melstft_error_b).float() for label in pref_label: pref_labels.append(label.item()) print(np.mean(pref_labels)) quant_label = torch.abs(melstft_error_a - melstft_error_b).clamp( min_quant_error, max_quant_error, ) print(quant_label) bin_edges = torch.linspace( min_quant_error, max_quant_error, num_quant_levels + 1, ).type_as(quant_label) bin_indices = torch.bucketize(quant_label, bin_edges) - 1 # print(bin_edges) for index in bin_indices: quant_labels.append(index.item()) print(quant_labels) quant_label = torch.nn.functional.one_hot( bin_indices, num_classes=num_quant_levels ).float() smoothing_value = 0.6 smoothing_neighbor_value = 0.2 # Create smoothed vectors smoothed_quant_label = quant_label * smoothing_value # Add neighbor values for i, index in enumerate(bin_indices): if index > 0: smoothed_quant_label[i, index - 1] += smoothing_neighbor_value if index < num_quant_levels - 1: smoothed_quant_label[i, index + 1] += smoothing_neighbor_value # print(smoothed_quant_label) if save_audio: for item_idx in range(audio_in_a.shape[0]): outfilepath = os.path.join( "debug", f"{bidx}_{item_idx}_audio_out_a.wav" ) in_lufs_db = meter.integrated_loudness( audio_in_a[item_idx, ...].T.numpy() ) out_lufs_db = meter.integrated_loudness( audio_out_a[item_idx, ...].T.numpy() ) print(in_lufs_db, out_lufs_db) infilepath = os.path.join("debug", f"{bidx}_{item_idx}_audio_in_a.wav") torchaudio.save(outfilepath, audio_out_a[item_idx, ...], 24000) torchaudio.save(infilepath, audio_in_a[item_idx, ...], 24000) if bidx > 100: break plt.hist(quant_labels) plt.savefig("debug/quant_labels.png", dpi=300)