{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from itertools import chain\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def plot_model_performance(model_scores, model_names=None, title=None):\n",
    "    # Create figure and axis\n",
    "    fig, ax = plt.subplots(figsize=(10, 6))\n",
    "\n",
    "    # Create box plot\n",
    "    bp = ax.boxplot(\n",
    "        model_scores,\n",
    "        patch_artist=True,  # Fill boxes with color\n",
    "        notch=True,  # Add notches to show confidence interval around median\n",
    "        widths=0.7,\n",
    "    )  # Width of boxes\n",
    "\n",
    "    # Customize box colors\n",
    "    colors = [\n",
    "        \"blue\",\n",
    "        \"red\",\n",
    "        \"green\",\n",
    "        \"orange\",\n",
    "        \"purple\",\n",
    "        \"pink\",\n",
    "        \"grey\",\n",
    "        \"blue\",\n",
    "        \"red\",\n",
    "        \"green\",\n",
    "        \"orange\",\n",
    "        \"purple\",\n",
    "        \"pink\",\n",
    "        \"grey\",\n",
    "    ]\n",
    "    for idx, box in enumerate(bp[\"boxes\"]):\n",
    "        box.set_facecolor(colors[idx])\n",
    "        box.set_alpha(0.7)\n",
    "\n",
    "    # Customize other elements\n",
    "    plt.setp(bp[\"whiskers\"], color=\"gray\")\n",
    "    plt.setp(bp[\"caps\"], color=\"gray\")\n",
    "    plt.setp(bp[\"medians\"], color=\"black\", linewidth=2)\n",
    "    plt.setp(bp[\"fliers\"], marker=\"o\", markerfacecolor=\"gray\", alpha=0.5)\n",
    "\n",
    "    # Set labels and title\n",
    "    if model_names:\n",
    "        ax.set_xticklabels(model_names)\n",
    "    ax.set_ylabel(\"Ditto Score\")\n",
    "    if title is not None:\n",
    "        ax.set_title(title)\n",
    "\n",
    "    # Add grid for easier reading\n",
    "    ax.yaxis.grid(True, linestyle=\"--\", alpha=0.7)\n",
    "\n",
    "    # Remove top and right spines\n",
    "    ax.spines[\"top\"].set_visible(False)\n",
    "    ax.spines[\"right\"].set_visible(False)\n",
    "\n",
    "    # Add individual points with jitter\n",
    "    for i, data in enumerate(model_scores, 1):\n",
    "        # Create jittered x-positions\n",
    "        x = np.random.normal(i, 0.04, size=len(data))\n",
    "        ax.scatter(x, data, alpha=0.4, color=\"black\", s=30)\n",
    "\n",
    "    # Adjust layout\n",
    "    plt.tight_layout()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "COVER_FILES = {\n",
    "    \"Base Model\": \"/home/sara/glockenspiel/z_eval_data/feat_eval_cover_30b_base_self_sim_2025_01_30-16_09_29.npz\",\n",
    "    # \"4.5 CFG 2.5\": \"/home/sara/glockenspiel/z_scratch/model_45_full_cfg_2.5_self_sim_cover_120.npz\",\n",
    "    \"4.5\": \"/home/sara/glockenspiel/z_scratch/model_45_self_sim_cover_120.npz\",\n",
    "    \"4.5 4 Min\": \"/home/sara/glockenspiel/z_scratch/model_45_self_sim_cover_240.npz\",\n",
    "    \"4.5 Tag CFG 3\": \"/home/sara/glockenspiel/z_scratch/model_45_cfg_3_self_sim_cover_120.npz\",\n",
    "    \"Filtered 4.5\": \"/home/sara/glockenspiel/z_scratch/model_45_w_filter_self_sim_cover_120.npz\",\n",
    "    # \"Tony 1.27\": \"/home/sara/glockenspiel/z_eval_data/tony_latest_dpo_self_sim_cover_120.npz\",\n",
    "    \"Current Prod\": \"/home/sara/glockenspiel/suno_utils/task_eval/scores/feat_eval_cover_30b_t6_prod_self_sim_2025_01_30-16_42_06.npz\",\n",
    "    # \"Prod Tag CFG 2.5\": \"/home/sara/glockenspiel/suno_utils/task_eval/scores/feat_eval_cover_30b_t6_prod_cfg2.5_self_sim_2025_01_30-19_37_30.npz\",\n",
    "    \"Prod Tag CFG 3\": \"/home/sara/glockenspiel/suno_utils/task_eval/scores/feat_eval_cover_30b_t6_prod_cfg3_self_sim_2025_01_30-18_43_44.npz\",\n",
    "    # \"Prod Tag CFG 4\": \"/home/sara/glockenspiel/suno_utils/task_eval/scores/feat_eval_cover_30b_t6_prod_cfg4_self_sim_2025_01_30-18_21_08.npz\",\n",
    "    \"Training Data\": \"/home/sara/glockenspiel/z_eval_data/training_data_cover_self_sim.npz\",\n",
    "}\n",
    "\n",
    "ARTIST_FILES = {\n",
    "    \"Base Model\": \"/home/sara/glockenspiel/z_eval_data/feat_eval_artist_30b_base_artist_vox_sim_2025_01_30-16_09_29.npz\",\n",
    "    \"Base CFG\": \"/home/sara/glockenspiel/z_eval_data/model_30b_t5_v1_cfg_artist_vox_sim_artist_120_30.npz\",\n",
    "    \"Tony 1.27\": \"/home/sara/glockenspiel/z_eval_data/tony_latest_dpo_artist_vox_sim_artist_120_30.npz\",\n",
    "    \"Tony 1.27 CFG\": \"/home/sara/glockenspiel/z_eval_data/ftony_latest_dpo_artist_vox_sim_CFG_artist_120_30.npz\",\n",
    "    \"Current Prod\": \"/home/sara/glockenspiel/suno_utils/task_eval/scores/feat_eval_artist_30b_t6_prod_artist_vox_sim_2025_01_30-16_42_06.npz\",\n",
    "    \"Training Data\": \"/home/sara/glockenspiel/z_eval_data/training_data_artist_artist_vox_sim.npz\",\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def parse_data(files):\n",
    "    all_flat = []\n",
    "    names = []\n",
    "    for (\n",
    "        name,\n",
    "        loc,\n",
    "    ) in files.items():\n",
    "        names.append(name)\n",
    "        scores = np.load(loc)\n",
    "        flat_scores = list(chain(*scores.values()))\n",
    "        all_flat.append(flat_scores)\n",
    "\n",
    "    return names, all_flat\n",
    "\n",
    "\n",
    "def parse_best_cover(files):\n",
    "    best_per_source = []\n",
    "    names = []\n",
    "    for (\n",
    "        name,\n",
    "        loc,\n",
    "    ) in files.items():\n",
    "        names.append(name)\n",
    "        scores = np.load(loc)\n",
    "        source_best = []\n",
    "        for _, source_scores in scores.items():\n",
    "            source_best.append(min(source_scores))\n",
    "        best_per_source.append(source_best)\n",
    "\n",
    "    return names, best_per_source\n",
    "\n",
    "\n",
    "def parse_best_artist(files):\n",
    "    best_per_source = []\n",
    "    names = []\n",
    "    for (\n",
    "        name,\n",
    "        loc,\n",
    "    ) in files.items():\n",
    "        names.append(name)\n",
    "        scores = np.load(loc)\n",
    "        source_best = []\n",
    "        for _, source_scores in scores.items():\n",
    "            source_best.append(max(source_scores))\n",
    "        best_per_source.append(source_best)\n",
    "\n",
    "    return names, best_per_source"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "cover_names, cover_flat = parse_data(COVER_FILES)\n",
    "plot_model_performance(cover_flat, cover_names, \"Cover Self Similarity, All Trials\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "cover_names, best_covers = parse_best_cover(COVER_FILES)\n",
    "plot_model_performance(best_covers, cover_names, \"Cover Self Similarity, Best of 10\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "artist_names, artist_flat = parse_data(ARTIST_FILES)\n",
    "plot_model_performance(artist_flat, artist_names, \"Artist Vocal Similarity, All Trials\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "artist_names, best_artist = parse_best_artist(ARTIST_FILES)\n",
    "plot_model_performance(best_artist, artist_names, \"Artist Vocal Similarity, Best of 10\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "COVER_FILES = {\n",
    "    \"Base Model\": \"/home/sara/glockenspiel/model_30b_t5_v1_self_sim_cover_120.npz\",\n",
    "    \"Cover DPO\": \"/home/sara/glockenspiel/cover_dpo_self_sim.npz\",\n",
    "    \"Tony 1.27\": \"/home/sara/glockenspiel/tony_latest_dpo_self_sim_cover_120.npz\",\n",
    "    \"Training Data\": \"/home/sara/glockenspiel/training_data_cover_self_sim.npz\",\n",
    "}\n",
    "\n",
    "ARTIST_FILES = {\n",
    "    \"Base Model\": \"/home/sara/glockenspiel/model_30b_t5_v1_artist_vox_sim_artist_120_30.npz\",\n",
    "    \"Base CFG\": \"/home/sara/glockenspiel/model_30b_t5_v1_cfg_artist_vox_sim_artist_120_30.npz\",\n",
    "    \"Tony 1.27\": \"/home/sara/glockenspiel/tony_latest_dpo_artist_vox_sim_artist_120_30.npz\",\n",
    "    \"Tony 1.27 CFG\": \"/home/sara/glockenspiel/tony_latest_dpo_artist_vox_sim_CFG_artist_120_30.npz\",\n",
    "    \"Training Data\": \"/home/sara/glockenspiel/training_data_artist_artist_vox_sim.npz\",\n",
    "}"
   ]
  }
 ],
 "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
}
