{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "78a8a11e",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.text import read_json\n",
    "from suno_utils.utils.clip import SunoClip\n",
    "from suno_utils.tasks.audio_features.vocal import VocalExtractor\n",
    "from tqdm import tqdm\n",
    "import matplotlib.pyplot as plt\n",
    "import numpy as np"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1cd0a28f",
   "metadata": {},
   "outputs": [],
   "source": [
    "MODEL_DATA = {\n",
    "    \"Prod V4\": \"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/genre_mappings_chirp-v4-h-s-32_no_aug_vocalist_first_2025_04_25-15_34_09.json\",\n",
    "    \"Auk T1\": \"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/genre_mappings_chirp-auk-t1_no_aug_vocalist_first_2025_04_25-15_54_20.json\",\n",
    "    \"Auk Rep 1 + Neg\": \"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/genre_mappings_chirp-auk-t1_rep_1_neg_gender_only_2025_04_25-16_11_58.json\",\n",
    "    \"Auk Rep 2 + Neg\": \"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/genre_mappings_chirp-auk-t1_rep2_neg1_2025_04_25-16_47_49.json\",\n",
    "    \"Auk Rep 3 + Neg\": \"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/genre_mappings_chirp-auk-t1_rep3_neg1_2025_04_25-16_29_42.json\",\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "65160824",
   "metadata": {},
   "outputs": [],
   "source": [
    "ve = VocalExtractor(device=\"cuda\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2a3de5c9",
   "metadata": {},
   "outputs": [],
   "source": [
    "def score_model(generations):\n",
    "    genre_scores = {}\n",
    "    print(f\"Running eval on {len(generations)} genres...\")\n",
    "    for genre in generations:\n",
    "        genre_scores[genre] = []\n",
    "        print(f\"Scoring {genre}\")\n",
    "        for gen in tqdm(generations[genre]):\n",
    "            s3_id = gen[\"s3_id\"]\n",
    "            tags = gen[\"tags\"]\n",
    "            gender = None\n",
    "            if \"female\" in tags:\n",
    "                gender = \"female\"\n",
    "            elif \"male\" in tags:\n",
    "                gender = \"male\"\n",
    "\n",
    "            clip = SunoClip(s3_id)\n",
    "            vocal_logits, vocal_tags = ve.extract(clip, threshold=0.7)\n",
    "            if len(vocal_tags) > 0:\n",
    "                if gender == vocal_tags[0]:\n",
    "                    genre_scores[genre].append(1)\n",
    "                else:\n",
    "                    print(genre, s3_id)\n",
    "                    genre_scores[genre].append(0)\n",
    "\n",
    "    return genre_scores"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cfd8ab54",
   "metadata": {},
   "outputs": [],
   "source": [
    "model_scores = {}\n",
    "for model_name, model_path in MODEL_DATA.items():\n",
    "    generations = read_json(model_path)\n",
    "    genre_scores = score_model(generations)\n",
    "    hit_rate = {genre: sum(scores) / len(scores) for genre, scores in genre_scores.items()}\n",
    "    model_scores[model_name] = hit_rate"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "108fc182",
   "metadata": {},
   "outputs": [],
   "source": [
    "def plot_genre_scores(data, title=\"Global Vocalist Accuracy by Genre\", figsize=(12, 8)):\n",
    "    # Get all unique genres\n",
    "    genres = set()\n",
    "    for model_scores in data.values():\n",
    "        genres.update(model_scores.keys())\n",
    "    genres = sorted(list(genres))\n",
    "\n",
    "    # Get all model names\n",
    "    models = list(data.keys())\n",
    "\n",
    "    # Set up the figure\n",
    "    fig, ax = plt.subplots(figsize=figsize)\n",
    "\n",
    "    # Set width of bars and positions\n",
    "    bar_width = 0.8 / len(models)\n",
    "    indices = np.arange(len(genres))\n",
    "\n",
    "    # Plot bars for each model\n",
    "    for i, model in enumerate(models):\n",
    "        model_scores = [data[model].get(genre, 0) for genre in genres]\n",
    "        position = indices - 0.4 + (i + 0.5) * bar_width\n",
    "        ax.bar(position, model_scores, bar_width, label=model)\n",
    "\n",
    "    # Add labels, title and legend\n",
    "    ax.set_xlabel(\"Genre\")\n",
    "    ax.set_ylabel(\"Score\")\n",
    "    ax.set_title(title)\n",
    "    ax.set_xticks(indices)\n",
    "    ax.set_xticklabels(genres)\n",
    "    ax.legend()\n",
    "\n",
    "    # Add grid lines for better readability\n",
    "    ax.grid(axis=\"y\", linestyle=\"--\", alpha=0.7)\n",
    "\n",
    "    # Adjust layout\n",
    "    plt.tight_layout()\n",
    "\n",
    "    return fig, ax"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "99da5dfc",
   "metadata": {},
   "outputs": [],
   "source": [
    "plot_genre_scores(model_scores)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "af0d8046",
   "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": 5
}
