{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import tempfile\n",
    "import torchaudio\n",
    "import boto3\n",
    "import json\n",
    "from dataclasses import dataclass\n",
    "from typing import Optional, List\n",
    "import numpy as np\n",
    "import torch\n",
    "import os\n",
    "import statistics\n",
    "import librosa\n",
    "\n",
    "from suno_utils.gpt import chirp_v2_5 as chirp_v3\n",
    "from suno_utils.diffusion import generation as diffusion_gen\n",
    "from suno_utils.gpt.engine import Engine\n",
    "from suno_utils.gpt.generation_engine import make_request, semantic_codes\n",
    "\n",
    "from suno_utils.tasks.upsample_engine import (\n",
    "    UpsampleEngine,\n",
    "    Request,\n",
    "    DiffusionGenerationConfig,\n",
    ")\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.worker.loader import _get_compatible_tokens\n",
    "\n",
    "from suno_utils.tasks import ss_vad\n",
    "from suno_utils.models.ditto_v2.ditto_v2 import Ditto\n",
    "from suno_utils.tasks.dac_vae_100hz_peaq import preload_models as preload_vae_models\n",
    "from suno_utils.gpt.generation import CfgGenerationConfig"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "s3_client = boto3.client(\"s3\")\n",
    "S3_BUCKET = \"suno-data-uploads\"\n",
    "USE_COMPILE = False\n",
    "MAX_STREAMS = 10\n",
    "N_BATCH = 4\n",
    "PARALLEL = 2\n",
    "MIN_CHUNK_SIZE = 25 * 30\n",
    "MAX_DURATION = 10\n",
    "NOISE_CUTOFF = 15\n",
    "SEP_SAMPLE_RATE = 44100\n",
    "ENCODER_RATE = 25\n",
    "\n",
    "DIFFUSION_CKPT_PATH = (\n",
    "    \"s3://suno-data/tony/tmp/diff/dit_v2_dpo_t2_v1_3k.pt\"  # 30b_t6 diffusion default\n",
    ")\n",
    "\n",
    "diffusion_gen.preload_models(\n",
    "    dit_model_filepath=DIFFUSION_CKPT_PATH, compile=USE_COMPILE, device=\"cuda:1\"\n",
    ")\n",
    "diff_engine = UpsampleEngine(min_chunk_size=MIN_CHUNK_SIZE)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "MODEL_NAME = \"model_30b_t5_v15_b40\"  # \"model_30b_t5_v1\"\n",
    "GPT_CKPT_PATH = f\"s3://suno-data/tony/tmp/{MODEL_NAME}.pt\"  # \"s3://suno-data/tony/tmp/model_30b_t5_v15_b40.pt\" #30b_t6\n",
    "MODEL_NAME = \"cover_dpo\"\n",
    "GPT_CKPT_PATH = \"/app/suno/checkpoints/2025-01-19_05-27-38/last_ckpt_infer.pt\"\n",
    "\n",
    "chirp_v3.preload_models(gpt_ckpt_path=GPT_CKPT_PATH)\n",
    "\n",
    "gpt_engine = Engine(\n",
    "    GPT_CKPT_PATH,\n",
    "    chirp_v3.TOKENIZER_PATH,\n",
    "    max_sequences=MAX_STREAMS,\n",
    "    compile=USE_COMPILE,\n",
    "    max_length_s=MAX_DURATION,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "DITTO_S3_PATH = \"s3://suno-data/minz/models/ditto_v2_epoch_57.pt\"\n",
    "DITTO_EMBEDDING_DIM = 128\n",
    "VAE_S3_PATH = \"s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth\"\n",
    "\n",
    "ditto_path = chirp_v3._get_model_if_needed(DITTO_S3_PATH)\n",
    "ditto = Ditto(\n",
    "    latent_dim=DITTO_EMBEDDING_DIM,\n",
    "    model_path=ditto_path,\n",
    "    is_flash=False,\n",
    "    is_serving=True,\n",
    ")\n",
    "\n",
    "ditto = ditto.eval().to(\"cuda:1\")\n",
    "\n",
    "ss_vad_config_path = chirp_v3._get_model_if_needed(ss_vad.YAML_PATH)\n",
    "ss_vad_model_path = chirp_v3._get_model_if_needed(ss_vad.MODEL_PATH)\n",
    "ss_vad.preload_models(\n",
    "    checkpoint_filepath=ss_vad_model_path,\n",
    "    config_path=ss_vad_config_path,\n",
    "    device=\"cuda:1\",\n",
    ")\n",
    "\n",
    "preload_vae_models(VAE_S3_PATH, device=\"cuda:1\")\n",
    "print(\"Finish loading models\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "DATA_PATH = \"test_task.json\"\n",
    "\n",
    "with open(DATA_PATH, \"r\") as file:\n",
    "    data = json.load(file)\n",
    "\n",
    "test_covers = data[\"cover\"]\n",
    "test_artist = data[\"artist\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "@dataclass\n",
    "class Prompt:\n",
    "    lyrics: str\n",
    "    tags: str\n",
    "    s3_id: str\n",
    "    test_tags: Optional[str] = None\n",
    "    test_lyrics: Optional[str] = None\n",
    "    neg_tags: Optional[str] = None\n",
    "    cover_arr: Optional[np.ndarray] = None\n",
    "    artist_arr: Optional[np.ndarray] = None\n",
    "\n",
    "\n",
    "def load_prompt_arr(audio_id):\n",
    "    with tempfile.NamedTemporaryFile(suffix=\".npz\") as temp_file:\n",
    "        target_version = \"4.0\"\n",
    "        tokens = _get_compatible_tokens(temp_file.name, audio_id, target_version)\n",
    "    return tokens[: (ENCODER_RATE * 120)]\n",
    "\n",
    "\n",
    "def run_gpt(prompt: Prompt):\n",
    "    gen_config = chirp_v3.GenerationConfig(\n",
    "        text=prompt.test_lyrics or prompt.lyrics,\n",
    "        text_tags=prompt.test_tags or prompt.tags,\n",
    "        text_neg_tags=prompt.neg_tags,\n",
    "        cfg_coef=1.2,\n",
    "        artist_arr=prompt.artist_arr,\n",
    "        cover_arr=prompt.cover_arr,\n",
    "        n_batch=1,\n",
    "        min_text_offset=0,\n",
    "        min_eos_p=0.1,\n",
    "        use_whisper=False,\n",
    "        n_repeat_tags=3,\n",
    "        n_repeat_neg_tags=3,\n",
    "        max_gen_duration_s=MAX_DURATION,\n",
    "        cfg_streams=[\n",
    "            CfgGenerationConfig(stream_type=\"tag\", weight=2.0, max_steps=25 * 30),\n",
    "            CfgGenerationConfig(stream_type=\"neg_tag\", weight=-1.0, max_steps=25 * 30),\n",
    "            # CfgGenerationConfig(stream_type=\"artist\", weight=1.5, max_steps=25*48)\n",
    "        ],\n",
    "    )\n",
    "\n",
    "    requests = [\n",
    "        make_request(f\"{i}\", gen_config, gpt_engine.model.config, gpt_engine.tokenizer)\n",
    "        for i in range(N_BATCH)\n",
    "    ]\n",
    "\n",
    "    gpt_outputs = []\n",
    "    for idx in range(0, len(requests), PARALLEL):\n",
    "        inner_requests = requests[idx : (idx + PARALLEL)]\n",
    "        for job in gpt_engine.run_request(inner_requests, tqdm_enabled=True):\n",
    "            stream = gpt_engine.token_generator(job)\n",
    "            semantics = list(semantic_codes(stream, gpt_engine.model.config))\n",
    "            gpt_outputs.append(semantics.copy())\n",
    "\n",
    "    return gpt_outputs\n",
    "\n",
    "\n",
    "def run_diffusion(prompt: Prompt, gpt_outputs: List):\n",
    "    diffusion_config = DiffusionGenerationConfig(\n",
    "        lyrics=prompt.test_lyrics or prompt.lyrics, tags=prompt.test_tags or prompt.tags\n",
    "    )\n",
    "\n",
    "    diff_outputs = []\n",
    "    for x in gpt_outputs:\n",
    "        request = Request(\n",
    "            id=f\"{x}\",\n",
    "            generation_config=diffusion_config,\n",
    "            tokens=x,\n",
    "            input_tokens_finished=True,\n",
    "        )\n",
    "        job = diff_engine.run_request(request, tqdm_enabled=True)\n",
    "        full_audio = Audio.concatenate(job.generated_audios)\n",
    "        latents = np.concatenate(job.vae_latents, axis=0)\n",
    "        diff_outputs.append((full_audio, latents))\n",
    "\n",
    "    return diff_outputs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def load_audio(audio_path, duration=MAX_DURATION):\n",
    "    s3 = boto3.client(\"s3\")\n",
    "\n",
    "    if os.path.exists(audio_path):\n",
    "        waveform, sr = torchaudio.load(audio_path)\n",
    "    else:\n",
    "        with tempfile.NamedTemporaryFile(suffix=\".mp3\") as temp_file:\n",
    "            s3.download_file(\n",
    "                \"suno-data-uploads\", f\"studio/uploads/{audio_path}.mp3\", temp_file.name\n",
    "            )\n",
    "            waveform, sr = torchaudio.load(temp_file.name)\n",
    "\n",
    "    waveform = torch.mean(waveform, dim=0).unsqueeze(0)[:, : sr * duration]\n",
    "    if sr != SEP_SAMPLE_RATE:\n",
    "        resampler = torchaudio.transforms.Resample(sr, SEP_SAMPLE_RATE)\n",
    "        waveform = resampler(waveform)\n",
    "\n",
    "    return waveform\n",
    "\n",
    "\n",
    "def load_audio_from_gen(audio, duration=MAX_DURATION, trim=True):\n",
    "    with tempfile.NamedTemporaryFile(suffix=\".mp3\") as temp_file:\n",
    "        audio.get_segment(to_s=MAX_DURATION).write_hq_mp3(temp_file.name)\n",
    "        waveform = load_audio(temp_file.name, duration=duration)\n",
    "    return waveform\n",
    "\n",
    "\n",
    "def cosine_similarity(a, b):\n",
    "    # Calculate dot product\n",
    "    dot_product = np.dot(a, b)\n",
    "\n",
    "    # Calculate magnitudes\n",
    "    magnitude_a = np.sqrt(np.dot(a, a))\n",
    "    magnitude_b = np.sqrt(np.dot(b, b))\n",
    "\n",
    "    # Calculate cosine similarity\n",
    "    return dot_product / (magnitude_a * magnitude_b)\n",
    "\n",
    "\n",
    "def load_vocals(audio_path, duration=MAX_DURATION, trim=True):\n",
    "    waveform = load_audio(audio_path, duration=duration)\n",
    "    vocals = ss_vad.encode(waveform)\n",
    "    if trim:\n",
    "        vocals, _ = librosa.effects.trim(vocals, top_db=NOISE_CUTOFF)\n",
    "    vocals = torch.from_numpy(vocals)\n",
    "    return vocals\n",
    "\n",
    "\n",
    "def load_vocals_from_gen(audio, duration=MAX_DURATION, trim=True):\n",
    "    with tempfile.NamedTemporaryFile(suffix=\".mp3\") as temp_file:\n",
    "        audio.get_segment(to_s=MAX_DURATION).write_hq_mp3(temp_file.name)\n",
    "        waveform = load_audio(temp_file.name, duration=duration)\n",
    "    vocals = ss_vad.encode(waveform)\n",
    "    if trim:\n",
    "        vocals, _ = librosa.effects.trim(vocals, top_db=NOISE_CUTOFF)\n",
    "    vocals = torch.from_numpy(vocals)\n",
    "    return vocals\n",
    "\n",
    "\n",
    "def embed_waveform(waveform, task):\n",
    "    return ditto.music_to_latent(waveform, task=task)[0].detach().cpu().numpy()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"artist_vox_sim\"\n",
    "TRIM = True\n",
    "\n",
    "scores = {}\n",
    "for artist in test_artist:\n",
    "    prompt = Prompt(**artist)\n",
    "    prompt.artist_arr = load_prompt_arr(prompt.s3_id)\n",
    "\n",
    "    gpt_outputs = run_gpt(prompt)\n",
    "    diff_outputs = run_diffusion(prompt, gpt_outputs)\n",
    "\n",
    "    S3_ID = prompt.s3_id\n",
    "    source_wav = load_vocals(S3_ID)\n",
    "    source_embed = embed_waveform(source_wav, TASK)\n",
    "    scores[S3_ID] = []\n",
    "    for full_audio, _ in diff_outputs:\n",
    "        artist_wav = load_vocals_from_gen(full_audio)\n",
    "        artist_embed = embed_waveform(artist_wav, TASK)\n",
    "        score = cosine_similarity(source_embed, artist_embed)\n",
    "        scores[S3_ID].append(score)\n",
    "\n",
    "for key in scores:\n",
    "    avg_score = statistics.mean(scores[key])\n",
    "    min_score = min(scores[key])\n",
    "    max_score = max(scores[key])\n",
    "\n",
    "    print(f\"{key}: average: {avg_score} min: {min_score} max: {max_score}\")\n",
    "\n",
    "np.savez(f\"{MODEL_NAME}_{TASK}_artist_UPDATED.npz\", **scores)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"self_sim\"\n",
    "\n",
    "scores = {}\n",
    "for artist in test_artist:\n",
    "    prompt = Prompt(**artist)\n",
    "    prompt.artist_arr = load_prompt_arr(prompt.s3_id)\n",
    "    print(f\"Generating {prompt.s3_id} with {prompt.test_tags}\")\n",
    "\n",
    "    gpt_outputs = run_gpt(prompt)\n",
    "    diff_outputs = run_diffusion(prompt, gpt_outputs)\n",
    "\n",
    "    S3_ID = prompt.s3_id\n",
    "    source_wav = load_audio(S3_ID)\n",
    "    source_embed = embed_waveform(source_wav, TASK)\n",
    "    scores[S3_ID] = []\n",
    "    for full_audio, _ in diff_outputs:\n",
    "        artist_wav = load_audio_from_gen(full_audio)\n",
    "        artist_embed = embed_waveform(artist_wav, TASK)\n",
    "        score = cosine_similarity(source_embed, artist_embed)\n",
    "        scores[S3_ID].append(score)\n",
    "\n",
    "for key in scores:\n",
    "    avg_score = statistics.mean(scores[key])\n",
    "    min_score = min(scores[key])\n",
    "    max_score = max(scores[key])\n",
    "\n",
    "    print(f\"{key}: average: {avg_score} min: {min_score} max: {max_score}\")\n",
    "\n",
    "np.savez(f\"{MODEL_NAME}_{TASK}_cover_UPDATED.npz\", **scores)"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_clean",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.10.15"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
