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
