{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import torchaudio\n",
    "import boto3\n",
    "import tempfile\n",
    "import torch\n",
    "import numpy as np\n",
    "import statistics\n",
    "import random\n",
    "from tqdm import tqdm\n",
    "from IPython.display import Audio, display\n",
    "import librosa\n",
    "\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.tasks import ss_vad\n",
    "from suno_utils.tasks.dac_vae_100hz_peaq import preload_models as preload_vae_models"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "DATA_PATH = \"/home/sara/glockenspiel/metas_v0_artist.jsonl\"\n",
    "\n",
    "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\")\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, config_path=ss_vad_config_path, device=\"cuda\"\n",
    ")\n",
    "\n",
    "preload_vae_models(VAE_S3_PATH, device=\"cuda\")\n",
    "print(\"Finish loading models\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "artists = {}\n",
    "\n",
    "S3 = \"s3_filepath\"\n",
    "ID = \"id\"\n",
    "ARTIST = \"artists\"\n",
    "EXAMPLES = 5000\n",
    "LYRICS = \"text\"\n",
    "\n",
    "lines_read = 0\n",
    "with open(DATA_PATH, \"r\", encoding=\"utf-8\") as file:\n",
    "    for line in file:\n",
    "        try:\n",
    "            sample = json.loads(line)\n",
    "            if LYRICS not in sample or len(sample[LYRICS]) < 50:\n",
    "                continue\n",
    "            if ARTIST not in sample:\n",
    "                continue\n",
    "            artist_ids = sample[ARTIST]\n",
    "            for artist_id in artist_ids:\n",
    "                if artist_id not in artists:\n",
    "                    artists[artist_id] = []\n",
    "                artists[artist_id].append(sample[S3])\n",
    "        except json.JSONDecodeError as e:\n",
    "            print(f\"Error decoding JSON: {e}\")\n",
    "        lines_read += 1\n",
    "        if EXAMPLES is not None and lines_read > EXAMPLES:\n",
    "            break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "s3 = boto3.client(\"s3\")\n",
    "stem_sample_rate = 44100\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_webm_audio_from_s3(s3_path):\n",
    "    s3 = boto3.client(\"s3\")\n",
    "    prefix = \"s3://\"\n",
    "    path = s3_path[len(prefix) :]\n",
    "    bucket, *key_parts = path.split(\"/\")\n",
    "    key = \"/\".join(key_parts)\n",
    "\n",
    "    with tempfile.NamedTemporaryFile(suffix=\".webm\") as temp_file:\n",
    "        s3.download_file(bucket, key, temp_file.name)\n",
    "        waveform, sr = torchaudio.load(temp_file.name)\n",
    "        waveform = torch.mean(waveform, dim=0).unsqueeze(0)[:, : sr * 120]\n",
    "        if sr != stem_sample_rate:\n",
    "            resampler = torchaudio.transforms.Resample(sr, stem_sample_rate)\n",
    "            waveform = resampler(waveform)\n",
    "\n",
    "        return waveform\n",
    "\n",
    "\n",
    "def get_vocal_embed(s3_id, task, trim=True):\n",
    "    waveform = load_webm_audio_from_s3(s3_id)\n",
    "    vocals = ss_vad.encode(waveform)\n",
    "    if trim:\n",
    "        vocals, _ = librosa.effects.trim(vocals, top_db=15)\n",
    "    vocals = torch.from_numpy(vocals)\n",
    "    vocals_embed = ditto.music_to_latent(vocals, task=task)[0].detach().cpu().numpy()\n",
    "    return vocals, vocals_embed"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"artist_vox_sim\"\n",
    "ITEMS_PER_ARIST = 10\n",
    "\n",
    "score_data = {}\n",
    "these_suck_maybe = []\n",
    "for parent_id, parent_s3 in tqdm(artists.items()):\n",
    "    num_items = len(parent_s3)\n",
    "    if num_items < 2:\n",
    "        continue\n",
    "    random_idx = random.sample(range(0, num_items), min(ITEMS_PER_ARIST * 2, num_items))\n",
    "    for idx in range(0, len(random_idx) - 1, 2):\n",
    "        artist_parent = parent_s3[random_idx[idx]]\n",
    "        parent_vocals, parent_embed = get_vocal_embed(artist_parent, task=TASK)\n",
    "        artist_child = parent_s3[random_idx[idx + 1]]\n",
    "        child_vocals, child_embed = get_vocal_embed(artist_child, task=TASK)\n",
    "\n",
    "        score = cosine_similarity(parent_embed, child_embed)\n",
    "\n",
    "        if parent_id not in score_data:\n",
    "            score_data[parent_id] = []\n",
    "        score_data[parent_id].append(score)\n",
    "        if score < 0.6:\n",
    "            print(\"Really far\")\n",
    "            display(Audio(parent_vocals.cpu().numpy(), rate=stem_sample_rate))\n",
    "            display(Audio(child_vocals.cpu().numpy(), rate=stem_sample_rate))\n",
    "\n",
    "\n",
    "np.savez(f\"training_data_artist_{TASK}.npz\", **score_data)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for key in score_data:\n",
    "    avg_score = statistics.mean(score_data[key])\n",
    "    min_score = min(score_data[key])\n",
    "    max_score = max(score_data[key])\n",
    "\n",
    "    print(f\"{key}: average: {avg_score} min: {min_score} max: {max_score}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "np.savez(f\"training_data_artist_{TASK}.npz\", **score_data)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "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
}
