import numpy as np
import os
from tqdm import tqdm
import torch
from suno_utils.audio import Audio
import json
import random
from suno_utils.tasks.ditto import Ditto

if __name__ == "__main__":
    avialbe_device = f"cuda:{os.environ['CUDA_VISIBLE_DEVICES']}"
    print(avialbe_device)

    with open("/home/tony/Data/Preference/7b_v2/7b_before_recode_20240412_s3.json", "r") as fp:
        all_s3_ids = json.load(fp)
    print(f"total jobs are, {len(all_s3_ids)}")
    
    my_ditto = Ditto(
        music_encoder_name="musicfm_concat",
        latent_dim=128,
        model_path="/home/tony/Data/Ditto/ditto.pt",
        music_encoder_path="/home/tony/Data/Ditto/music_encoder.pt",
        is_flash=False
    )
    my_ditto = my_ditto.eval().cuda()
    print("finish loading ditto")

    def get_ditto_music_emb(test_id):
        s3_url = f"s3://suno-data-uploads/studio/uploads/{test_id}.mp3"
        start = 0
        dur = 120
        audio = Audio.from_s3(s3_url, n_channels=1, sample_rate=24000)
        audio = audio.get_segment(from_s=start, to_s=start + dur)
        wav = torch.tensor(audio.array_float).unsqueeze(0).cuda()
        emb = my_ditto.music_to_latent(wav)[0].detach().cpu().numpy()
        return emb

    def get_ditto_text_emb(test_text):
        emb = my_ditto.text_to_latent("[CLS]" + test_text)[0].detach().cpu().numpy()
        return emb
    
    def get_ditto_and_save(test_id):
        output_path = f"/app/suno/data/dpo/ditto_npz/{test_id}.npy"
        if os.path.exists(output_path):
            return
        with open(output_path, "w") as fp:
            fp.write("")
        try:
            emb = get_ditto_music_emb(test_id)
            np.save(output_path, emb)
        except Exception as e:
            print(f"failed {test_id} with {str(e)}")
            os.remove(output_path)

    print("start working")
    random.shuffle(all_s3_ids)
    for filename in tqdm(all_s3_ids):
        get_ditto_and_save(filename)
    print("finish working")