{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import numpy as np\n",
    "from tqdm import tqdm\n",
    "import matplotlib.pyplot as plt\n",
    "from IPython.display import Audio, display"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TASK = \"vox_artist_sim\"\n",
    "\n",
    "DATA_PATH = f\"/app/suno/sara/ditto_v2_{TASK}/metas/*.jsonl\"\n",
    "DITTO_PATH = f\"/app/suno/sara/ditto_v2_{TASK}/*.npz\"\n",
    "COVER_PATH = \"/home/sara/glockenspiel/metas_v0_cover.jsonl\"\n",
    "OUT_FILE = f\"/app/suno/sara/ditto_v2_{TASK}_cover_scores.jsonl\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "all_scores = []\n",
    "\n",
    "with open(OUT_FILE, \"r\", encoding=\"utf-8\") as file:\n",
    "    for line in tqdm(file):\n",
    "        try:\n",
    "            sample = json.loads(line)\n",
    "            all_scores.append(sample[f\"score_{TASK}\"])\n",
    "        except json.JSONDecodeError as e:\n",
    "            print(f\"Error decoding JSON: {e}\")\n",
    "\n",
    "print(len(all_scores))\n",
    "all_scores = np.array(all_scores)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def plot_histogram(data, bins=30):\n",
    "    plt.figure(figsize=(10, 6))\n",
    "\n",
    "    # Plot histogram\n",
    "    plt.hist(data, bins=bins, color=\"skyblue\", edgecolor=\"black\")\n",
    "\n",
    "    # Add mean and median lines\n",
    "    mean_val = np.mean(data)\n",
    "    median_val = np.median(data)\n",
    "    plt.axvline(\n",
    "        mean_val,\n",
    "        color=\"red\",\n",
    "        linestyle=\"dashed\",\n",
    "        linewidth=1,\n",
    "        label=f\"Mean: {mean_val:.2f}\",\n",
    "    )\n",
    "    plt.axvline(\n",
    "        median_val,\n",
    "        color=\"green\",\n",
    "        linestyle=\"dashed\",\n",
    "        linewidth=1,\n",
    "        label=f\"Median: {median_val:.2f}\",\n",
    "    )\n",
    "\n",
    "    # Add labels and title\n",
    "    plt.xlabel(f\"Ditto {TASK} Scores\")\n",
    "    plt.ylabel(\"Counts\")\n",
    "    plt.title(f\"Cosine Similarity of Cover Pairs, {TASK}\")\n",
    "    plt.legend()\n",
    "\n",
    "    # Add grid for better readability\n",
    "    plt.grid(True, alpha=0.3)\n",
    "\n",
    "    # Adjust layout to prevent label cutoff\n",
    "    plt.tight_layout()\n",
    "\n",
    "    return plt"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plot = plot_histogram(all_scores)\n",
    "plot.show()"
   ]
  },
  {
   "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 IPython.display import Audio, display\n",
    "from IPython.display import clear_output"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "TIMESTAMP = \"2025_01_30-18_21_08\"\n",
    "DATA_FOLDER = \"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs\"\n",
    "\n",
    "with open(os.path.join(DATA_FOLDER, f\"cover_mappings_{TIMESTAMP}.json\")) as f:\n",
    "    cover_mappings = json.load(f)\n",
    "with open(os.path.join(DATA_FOLDER, f\"artist_mappings_{TIMESTAMP}.json\")) as f:\n",
    "    artist_mappings = json.load(f)\n",
    "\n",
    "s3 = boto3.client(\"s3\")\n",
    "bucket = \"suno-data-uploads\"\n",
    "folder = f\"tasks/feature_eval/cover_persona/{TIMESTAMP}/\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "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_task_scores(mappings):\n",
    "    scores = {}\n",
    "    task = None\n",
    "    model = None\n",
    "    for source_id, covers in cover_mappings.items():\n",
    "        scores[source_id] = []\n",
    "        with tempfile.NamedTemporaryFile(suffix=\".npz\") as temp_file:\n",
    "            s3.download_file(\n",
    "                bucket, os.path.join(folder, f\"{source_id}_ditto.npz\"), temp_file.name\n",
    "            )\n",
    "            data = np.load(temp_file.name)\n",
    "            source_ditto = data[\"embedding\"]\n",
    "            if task is None:\n",
    "                task = data[\"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}_ditto.npz\"),\n",
    "                    temp_file.name,\n",
    "                )\n",
    "                child_ditto = np.load(temp_file.name)[\"embedding\"]\n",
    "            score = cosine_similarity(source_ditto, child_ditto)\n",
    "            if score > 0.9 or score < 0.2:\n",
    "                with tempfile.NamedTemporaryFile(suffix=\".mp3\") as temp_file:\n",
    "                    print(source_id)\n",
    "                    s3.download_file(\n",
    "                        bucket, os.path.join(folder, f\"{cover_id}.mp3\"), temp_file.name\n",
    "                    )\n",
    "                    display(Audio(temp_file.name))\n",
    "                    clear_output()\n",
    "            scores[source_id].append(score)\n",
    "\n",
    "    return scores\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": [
    "cover_scores = get_task_scores(cover_mappings)\n",
    "summarize_scores(cover_scores)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "artist_scores = get_task_scores(artist_mappings)\n",
    "summarize_scores(artist_scores)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "Audio.dp"
   ]
  }
 ],
 "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
}
