{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import numpy as np\n",
    "import os\n",
    "import tempfile\n",
    "import noisereduce as nr\n",
    "\n",
    "from suno_utils.tasks import ss_vad\n",
    "from suno_utils.gpt import chirp_v2_5 as chirp_v3\n",
    "from suno_utils.models.ditto_v2.ditto_v2 import Ditto\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.worker.loader import get_latents\n",
    "from suno_utils.tasks.dac_vae_100hz_peaq import (\n",
    "    decode_stream_to_full_audio as decode_vae_to_audio,\n",
    ")\n",
    "from suno_utils.tasks.dac_vae_100hz_peaq import preload_models as preload_vae_models\n",
    "from suno_utils.audio.conversion import trim_audio_silence"
   ]
  },
  {
   "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\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "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().cuda()\n",
    "print(\"Finish loading models\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "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",
    ")\n",
    "\n",
    "preload_vae_models(VAE_S3_PATH)\n",
    "print(\"Finish loading models\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Persona Vocals, Pop Rock"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def encode_audio(audio_path, task, start=0, duration=30):\n",
    "    if os.path.exists(audio_path):\n",
    "        audio = Audio.from_file(audio_path, n_channels=1, sample_rate=24000)\n",
    "    else:\n",
    "        s3_url = f\"s3://suno-data-uploads/studio/uploads/{audio_path}.mp3\"\n",
    "        audio = Audio.from_s3(s3_url, n_channels=1, sample_rate=24000)\n",
    "\n",
    "    audio.play()\n",
    "    print(f\"Starting ditto encoding for {audio_path}, duration {audio.duration_s}.\")\n",
    "\n",
    "    audio = audio.get_segment(from_s=start, to_s=start + duration)\n",
    "    wav = torch.tensor(audio.array_float).unsqueeze(0).cuda()\n",
    "    embedding = ditto.music_to_latent(wav, task=task)[0].detach().cpu().numpy()\n",
    "\n",
    "    return embedding\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 get_stems(s3_id):\n",
    "    with tempfile.TemporaryDirectory() as td:\n",
    "        vae_latents = get_latents(os.path.join(td, f\"{s3_id}_vae.npz\"), s3_id)\n",
    "\n",
    "        n_default_stride_tokens = 60 * 25  # very long duration -- 1 min\n",
    "        assert vae_latents is not None\n",
    "        audio = decode_vae_to_audio(\n",
    "            vae_latents,\n",
    "            n_stride_tokens=min(n_default_stride_tokens, vae_latents.shape[0] - 5),\n",
    "        )\n",
    "        print(f\"{s3_id}: decoded vae latents, {audio.duration_s}s\")\n",
    "\n",
    "        vocals_array = ss_vad.encode(audio)\n",
    "        assert isinstance(vocals_array, np.ndarray)\n",
    "        if np.abs(vocals_array).max() > 1:\n",
    "            vocals_array = vocals_array / np.abs(vocals_array).max()\n",
    "        audio = audio.convert(\n",
    "            sample_rate=ss_vad.SAMPLE_RATE,\n",
    "            byte_width=audio.byte_width,\n",
    "            n_channels=audio.n_channels,\n",
    "        )\n",
    "        print(\n",
    "            f\"Done with source separation: vocals shape = \"\n",
    "            f\"{vocals_array.shape}, audio shape = \"\n",
    "            f\"{audio.array_float.shape}\"\n",
    "        )\n",
    "        max_audio_length = min(vocals_array.shape[1], audio.array_float.shape[1])\n",
    "        instrumentals_array = (\n",
    "            audio.array_float[:, :max_audio_length] - vocals_array[:, :max_audio_length]\n",
    "        )\n",
    "        # normalize since it is possible to overflow\n",
    "        if np.abs(instrumentals_array).max() > 1:\n",
    "            instrumentals_array = (\n",
    "                instrumentals_array / np.abs(instrumentals_array).max()\n",
    "            )\n",
    "        vocals_audio = Audio.from_array_float(\n",
    "            vocals_array, sample_rate=ss_vad.SAMPLE_RATE\n",
    "        )\n",
    "        instrumentals_audio = Audio.from_array_float(\n",
    "            instrumentals_array, sample_rate=ss_vad.SAMPLE_RATE\n",
    "        )\n",
    "\n",
    "        return vocals_audio, instrumentals_audio\n",
    "\n",
    "\n",
    "def embed_vocals(s3_id, task, trim=True, denoise=False):\n",
    "    vocals, _ = get_stems(s3_id)\n",
    "\n",
    "    if trim:\n",
    "        vocals, start, end = trim_audio_silence(\n",
    "            vocals, silence_threshold=0.3, output_original_start_end_time=True\n",
    "        )\n",
    "        print(f\"Trimmed to {start}-{end}\")\n",
    "\n",
    "    if denoise:\n",
    "        print(\"Denoising\")\n",
    "        reduced_noise = nr.reduce_noise(\n",
    "            y=vocals.array_float,\n",
    "            sr=vocals.sample_rate,\n",
    "            n_std_thresh_stationary=1.5,  # noise threshold\n",
    "            stationary=False,\n",
    "            prop_decrease=0.2,  # 1.0 is most aggresive\n",
    "        )\n",
    "\n",
    "        vocals = Audio.from_array_float(reduced_noise, sample_rate=vocals.sample_rate)\n",
    "\n",
    "    with tempfile.TemporaryDirectory() as td:\n",
    "        temp_path = os.path.join(td, \"test.mp3\")\n",
    "        vocals.write_hq_mp3(temp_path)\n",
    "        embed = encode_audio(temp_path, task)\n",
    "\n",
    "    return embed\n",
    "\n",
    "\n",
    "def embed_full(s3_id, task):\n",
    "    return encode_audio(s3_id, task)\n",
    "\n",
    "\n",
    "def compare_vocals(s3_a, s3_b, task, trim=True, denoise=False):\n",
    "    embed_a = embed_vocals(s3_a, task, trim=trim, denoise=denoise)\n",
    "    embed_b = embed_vocals(s3_b, task, trim=trim, denoise=denoise)\n",
    "    return cosine_similarity(embed_a, embed_b)\n",
    "\n",
    "\n",
    "def compare_full(s3_a, s3_b, task):\n",
    "    embed_a = embed_full(s3_a, task)\n",
    "    embed_b = embed_full(s3_b, task)\n",
    "    return cosine_similarity(embed_a, embed_b)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Dad Rock Persona"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"artist_vox_sim\"  # self_sim #self_vox_sim #artist_sim #artist_vox_sim #album_sim #genre_sim #lyric_sim\n",
    "S3_A = \"ed1a70e7-150c-4189-933b-2eda8cfd6dd1\"  # dad rock persona source\n",
    "S3_B = \"de600531-a268-4a9e-8cc7-fdf3631c33f4\"  # failed persona, wrong gender\n",
    "\n",
    "score = compare_vocals(S3_A, S3_B, TASK)\n",
    "print(score)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "S3_B = \"8737624b-39f7-473f-98e4-d8cd91946b4e\"  # right gender but not a match\n",
    "score = compare_vocals(S3_A, S3_B, TASK)\n",
    "print(score)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "S3_B = \"3ec2924a-8e8a-4364-94b8-d97e0392a597\"  # good match\n",
    "score = compare_vocals(S3_A, S3_B, TASK)\n",
    "print(score)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Fabrizio Persona, Opera -> Techno"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"artist_vox_sim\"\n",
    "FABRIZIO = \"1ec5f5ef-cdaa-455a-b258-4b09917a768e\"\n",
    "TECHNO_FABRIZIO = \"7fb9d1c8-0436-49d5-b06d-522a14b24280\"\n",
    "\n",
    "score = compare_vocals(FABRIZIO, TECHNO_FABRIZIO, TASK)\n",
    "print(score)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Cover, Classical -> Techno"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"self_sim\"\n",
    "S3_ID = \"6778d4a9-60c2-41ac-abe8-7184aa83075e\"\n",
    "S3_ID_COVER = \"4461dff1-107b-4cdf-8709-6bd48801ca9e\"  # its basically the same\n",
    "\n",
    "score = compare_full(S3_ID, S3_ID_COVER, TASK)\n",
    "print(score)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"self_sim\"\n",
    "S3_ID = \"42d40038-5eb8-485f-aaa3-27e31d446bbf\"\n",
    "S3_ID_COVER = \"be401ddc-c081-4ba0-8958-be6d4b3d6c05\"  # more obscure failure, timbre isn't the same\n",
    "\n",
    "score = compare_full(S3_ID, S3_ID_COVER, TASK)\n",
    "print(score)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"self_sim\"\n",
    "S3_ID = \"7326a60b-77ab-4fa5-a0a1-21264f09b1bc\"\n",
    "S3_ID_COVER = \"68ba3068-eb1b-4e03-af91-4b31897631d4\"  # another exact match\n",
    "\n",
    "score = compare_full(S3_ID, S3_ID_COVER, TASK)\n",
    "print(score)"
   ]
  }
 ],
 "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
}
