import os
import torch
import argparse
import torchaudio
import numpy as np
import pandas as pd
import pyloudnorm as pyln
import multiprocessing as mp

from tqdm import tqdm
from suno_utils.audio import Audio

from suno_ear.system import EarSystem
from suno_boost.utils import apply_normalization

if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--csv_path",
        default="/home/tony/Data/Preference/7b_v0/interesting_clips_v3_processed.csv",
        type=str,
    )
    parser.add_argument(
        "--output_dir",
        default="/app/suno/christian/data/interesting_clips_v3_processed",
    )
    parser.add_argument("--num_examples", type=int, default=-1)
    args = parser.parse_args()

    if not os.path.isdir(args.output_dir):
        os.makedirs(args.output_dir)

    # create func
    def download_from_s3_and_save(s3_id: str):
        filepath = os.path.join(args.output_dir, s3_id + ".mp3")

        if os.path.isfile(filepath):
            return
        else:
            try:
                pos_audio = Audio.from_s3(
                    f"s3://suno-data-uploads/studio/uploads/{s3_id}.mp3"
                )
            except:
                return

        pos_audio.write_mp3(filepath)

    # load csv
    df = pd.read_csv(args.csv_path)
    print(df.shape)

    # get pairs
    df = df.sort_values(by=["request_id", "preference"])
    print(df.head())

    s3_ids = []
    for idx in tqdm(np.arange(0, len(df), 1)):
        # idx = np.random.randint(0, len(df))
        # idx += idx % 2
        row = df.iloc[idx]
        # pos_row = df.iloc[idx + 1]

        s3_ids.append(row["s3_id"])

    print(len(s3_ids))

    with mp.Pool(32) as pool:
        pool.map(download_from_s3_and_save, s3_ids)
