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)
