{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import boto3\n",
    "import tempfile\n",
    "import os\n",
    "import numpy as np\n",
    "import statistics\n",
    "from tqdm import tqdm"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TIMESTAMP = \"2025_02_26-16_36_11\"\n",
    "DATA_FOLDER = \"modal_runs\"\n",
    "OUTPUT_FOLDER = \"scores\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def cosine_similarity(a, b):\n",
    "    dot_product = np.dot(a, b)\n",
    "\n",
    "    magnitude_a = np.sqrt(np.dot(a, a))\n",
    "    magnitude_b = np.sqrt(np.dot(b, b))\n",
    "\n",
    "    return dot_product / (magnitude_a * magnitude_b)\n",
    "\n",
    "\n",
    "def get_task_scores(mappings, task):\n",
    "    s3 = boto3.client(\"s3\")\n",
    "    bucket = \"suno-data-uploads\"\n",
    "    folder = f\"tasks/feature_eval/cover_persona/{TIMESTAMP}/\"\n",
    "\n",
    "    scores = {}\n",
    "    model = None\n",
    "    for source_id, covers in tqdm(mappings.items()):\n",
    "        scores[source_id] = []\n",
    "        with tempfile.NamedTemporaryFile(suffix=\".npz\") as temp_file:\n",
    "            s3.download_file(\n",
    "                bucket,\n",
    "                os.path.join(folder, f\"{source_id}_{task}_ditto.npz\"),\n",
    "                temp_file.name,\n",
    "            )\n",
    "            data = np.load(temp_file.name)\n",
    "            source_task = data[\"task\"]\n",
    "            source_ditto = data[\"embedding\"]\n",
    "            assert source_task == task\n",
    "            if model is None:\n",
    "                model = data[\"model\"]\n",
    "        for cover_id in covers:\n",
    "            with tempfile.NamedTemporaryFile(suffix=\".npz\") as temp_file:\n",
    "                s3.download_file(\n",
    "                    bucket,\n",
    "                    os.path.join(folder, f\"{cover_id}_{task}_ditto.npz\"),\n",
    "                    temp_file.name,\n",
    "                )\n",
    "\n",
    "                data = np.load(temp_file.name)\n",
    "                child_task = data[\"task\"]\n",
    "                child_ditto = data[\"embedding\"]\n",
    "                assert child_task == task\n",
    "            score = cosine_similarity(source_ditto, child_ditto)\n",
    "            scores[source_id].append(score)\n",
    "\n",
    "    return scores, task, model\n",
    "\n",
    "\n",
    "def summarize_scores(scores):\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",
    "        print(f\"{key}: average: {avg_score} min: {min_score} max: {max_score}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "if not os.path.exists(OUTPUT_FOLDER):\n",
    "    os.makedirs(OUTPUT_FOLDER)\n",
    "\n",
    "cover_path = os.path.join(DATA_FOLDER, f\"cover_mappings_{TIMESTAMP}.json\")\n",
    "score_cover = os.path.exists(cover_path)\n",
    "artist_path = os.path.join(DATA_FOLDER, f\"artist_mappings_{TIMESTAMP}.json\")\n",
    "score_artist = os.path.exists(artist_path)\n",
    "\n",
    "if score_cover:\n",
    "    with open(cover_path) as f:\n",
    "        cover_mappings = json.load(f)\n",
    "\n",
    "    cover_scores, c_task, c_model = get_task_scores(cover_mappings, \"self_sim\")\n",
    "\n",
    "if score_artist:\n",
    "    with open(artist_path) as f:\n",
    "        artist_mappings = json.load(f)\n",
    "\n",
    "    artist_scores, a_task, a_model = get_task_scores(artist_mappings, \"artist_vox_sim\")"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_clean",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.10.15"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
