{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import boto3\n",
    "import statistics\n",
    "from typing import Dict, List, Any\n",
    "import matplotlib.pyplot as plt\n",
    "import numpy as np\n",
    "import math\n",
    "import pandas as pd\n",
    "from tqdm import tqdm\n",
    "\n",
    "from suno_utils.audio import Audio"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "def load_json_from_s3(bucket_name: str, key: str) -> Dict[str, Any]:\n",
    "    \"\"\"Load JSON data from S3 bucket\"\"\"\n",
    "    s3 = boto3.client(\"s3\")\n",
    "    try:\n",
    "        response = s3.get_object(Bucket=bucket_name, Key=key)\n",
    "        content = response[\"Body\"].read().decode(\"utf-8\")\n",
    "        return json.loads(content)\n",
    "    except Exception as e:\n",
    "        # print(f\"Error loading {key}: {e}\")\n",
    "        return None\n",
    "\n",
    "\n",
    "def process_wer_data(\n",
    "    ids_data: Dict[\n",
    "        str,\n",
    "        List[str],\n",
    "    ],\n",
    "    bucket_name: str = \"suno-data-uploads\",\n",
    "    s3_prefix: str = \"tasks/feature_eval/cover_persona/2025_07_11-16_20_50/\",\n",
    ") -> Dict[str, Any]:\n",
    "    \"\"\"Process WER data for all files\"\"\"\n",
    "\n",
    "    results = {}\n",
    "    flat_data = []\n",
    "    all_wers = []  # Collect all WER values for overall stats\n",
    "    all_data = []  # Collect all data for overall stats\n",
    "\n",
    "    for group_id, file_ids in tqdm(ids_data.items()):\n",
    "        group_wers = []\n",
    "        group_data = []\n",
    "\n",
    "        for file_id in file_ids:\n",
    "            # Construct S3 key\n",
    "            s3_key = f\"{s3_prefix}{file_id}_infill_wer.json\"\n",
    "\n",
    "            # Load JSON from S3\n",
    "            data = load_json_from_s3(bucket_name, s3_key)\n",
    "\n",
    "            if data and \"wer\" in data:\n",
    "                data[\"s3_id\"] = file_id\n",
    "                flat_data.append(data)\n",
    "\n",
    "                group_wers.append(data[\"wer\"])\n",
    "                group_data.append(data)\n",
    "\n",
    "                # Add to overall collections\n",
    "                all_wers.append(data[\"wer\"])\n",
    "                all_data.append(data)\n",
    "            else:\n",
    "                pass\n",
    "                # print(f\"Missing or invalid data for {file_id}\")\n",
    "\n",
    "        # Calculate statistics for this group\n",
    "        if group_wers:\n",
    "            results[group_id] = {\n",
    "                \"wers\": group_wers,\n",
    "                \"mean_wer\": statistics.mean(group_wers),\n",
    "                \"median_wer\": statistics.median(group_wers),\n",
    "                \"min_wer\": min(group_wers),\n",
    "                \"max_wer\": max(group_wers),\n",
    "                \"std_wer\": statistics.stdev(group_wers) if len(group_wers) > 1 else 0,\n",
    "                \"count\": len(group_wers),\n",
    "                \"data\": group_data,  # Include full data if needed\n",
    "            }\n",
    "\n",
    "    # Add overall statistics as \"all\" entry\n",
    "    if all_wers:\n",
    "        results[\"all\"] = {\n",
    "            \"wers\": all_wers,\n",
    "            \"mean_wer\": statistics.mean(all_wers),\n",
    "            \"median_wer\": statistics.median(all_wers),\n",
    "            \"min_wer\": min(all_wers),\n",
    "            \"max_wer\": max(all_wers),\n",
    "            \"std_wer\": statistics.stdev(all_wers) if len(all_wers) > 1 else 0,\n",
    "            \"count\": len(all_wers),\n",
    "            \"data\": all_data,\n",
    "        }\n",
    "\n",
    "    return results, flat_data\n",
    "\n",
    "\n",
    "def get_infill_data(timestamp):\n",
    "    infill_path = f\"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/infill_mappings_{timestamp}.json\"\n",
    "    with open(infill_path, \"r\") as f:\n",
    "        ids_data = json.load(f)\n",
    "    wer_results, flat_data = process_wer_data(\n",
    "        ids_data, s3_prefix=f\"tasks/feature_eval/cover_persona/{timestamp}/\"\n",
    "    )\n",
    "    return wer_results, flat_data\n",
    "\n",
    "\n",
    "def get_dur_success(flat_data):\n",
    "    df = pd.DataFrame(flat_data)\n",
    "    counts = df[\"duration_matches\"].value_counts()\n",
    "    print(f\"{len(df)} items\")\n",
    "    percentages = counts / len(df)\n",
    "    return percentages\n",
    "\n",
    "\n",
    "def get_closest_bucket(infill_dur):\n",
    "    buckets = [8, 15, 25]\n",
    "    bucket_distance = {\n",
    "        bucket_dur: abs(infill_dur - bucket_dur) for bucket_dur in buckets\n",
    "    }\n",
    "    min_key = min(bucket_distance, key=bucket_distance.get)\n",
    "    return min_key\n",
    "\n",
    "\n",
    "def get_wer_by_duration(flat_data):\n",
    "    df = pd.DataFrame(flat_data)\n",
    "    df[\"duration_bucket\"] = df[\"infill_dur_s\"].apply(get_closest_bucket)\n",
    "    return df.groupby(\"duration_bucket\")[\"wer\"].mean()\n",
    "\n",
    "\n",
    "def plot_multi_model_wer_grouped(\n",
    "    models_data: Dict[str, Dict[str, Any]],\n",
    "    include_all: bool = True,\n",
    "    figsize: tuple = (16, 8),\n",
    "):\n",
    "    # Get all unique group IDs (genres)\n",
    "    all_group_ids = set()\n",
    "    for model_data in models_data.values():\n",
    "        all_group_ids.update(model_data.keys())\n",
    "\n",
    "    if not include_all and \"all\" in all_group_ids:\n",
    "        all_group_ids.remove(\"all\")\n",
    "\n",
    "    group_ids = sorted(list(all_group_ids))\n",
    "    model_names = list(models_data.keys())\n",
    "\n",
    "    fig, ax = plt.subplots(figsize=figsize)\n",
    "\n",
    "    # Calculate bar positions\n",
    "    n_groups = len(group_ids)\n",
    "    n_models = len(model_names)\n",
    "\n",
    "    # Width of each individual bar\n",
    "    bar_width = 0.8 / n_models\n",
    "\n",
    "    # Colors for each model\n",
    "    colors = plt.cm.Set3(np.linspace(0, 1, n_models))\n",
    "\n",
    "    # X positions for each group\n",
    "    group_positions = np.arange(n_groups)\n",
    "\n",
    "    # Plot bars for each model\n",
    "    for model_idx, model_name in enumerate(model_names):\n",
    "        means = []\n",
    "        standard_errors = []\n",
    "\n",
    "        # Get data for each group for this model\n",
    "        for group_id in group_ids:\n",
    "            if group_id in models_data[model_name]:\n",
    "                wers = models_data[model_name][group_id][\"wers\"]\n",
    "                mean_wer = statistics.mean(wers)\n",
    "                se = (\n",
    "                    statistics.stdev(wers) / math.sqrt(len(wers))\n",
    "                    if len(wers) > 1\n",
    "                    else 0\n",
    "                )\n",
    "            else:\n",
    "                mean_wer = 0\n",
    "                se = 0\n",
    "\n",
    "            means.append(mean_wer)\n",
    "            standard_errors.append(se)\n",
    "\n",
    "        # Calculate x positions for this model's bars\n",
    "        # Center the group of bars around each group position\n",
    "        offset = (model_idx - (n_models - 1) / 2) * bar_width\n",
    "        x_positions = group_positions + offset\n",
    "\n",
    "        # Plot bars for this model\n",
    "        bars = ax.bar(\n",
    "            x_positions,\n",
    "            means,\n",
    "            bar_width,\n",
    "            yerr=standard_errors,\n",
    "            capsize=1,\n",
    "            label=model_name,\n",
    "            alpha=0.8,\n",
    "            color=colors[model_idx],\n",
    "            error_kw={\"ecolor\": \"black\", \"capthick\": 1, \"elinewidth\": 1},\n",
    "        )\n",
    "\n",
    "        # Add value labels on bars\n",
    "        for j, (x_pos, mean, se) in enumerate(zip(x_positions, means, standard_errors)):\n",
    "            if mean > 0:  # Only label if there's data\n",
    "                ax.text(\n",
    "                    x_pos,\n",
    "                    mean + se + 0.005,\n",
    "                    f\"{mean:.3f}\",\n",
    "                    ha=\"center\",\n",
    "                    va=\"bottom\",\n",
    "                    fontsize=8,\n",
    "                    fontweight=\"bold\",\n",
    "                )\n",
    "\n",
    "    # Customize the plot\n",
    "    ax.set_xlabel(\"Genre\", fontsize=12)\n",
    "    ax.set_ylabel(\"Mean WER\", fontsize=12)\n",
    "    ax.set_title(\"WER Comparison Across Models by Genre\\n\", fontsize=14)\n",
    "\n",
    "    # Set x-axis ticks and labels\n",
    "    ax.set_xticks(group_positions)\n",
    "    ax.set_xticklabels(\n",
    "        [id[:12] + \"...\" if len(id) > 15 else id for id in group_ids],\n",
    "        rotation=45,\n",
    "        ha=\"right\",\n",
    "    )\n",
    "\n",
    "    # Add legend\n",
    "    ax.legend(bbox_to_anchor=(1.05, 1), loc=\"upper left\")\n",
    "\n",
    "    # Add grid\n",
    "    ax.grid(True, alpha=0.3, axis=\"y\")\n",
    "\n",
    "    # Set y-axis to start from 0\n",
    "    ax.set_ylim(bottom=0)\n",
    "\n",
    "    plt.tight_layout()\n",
    "\n",
    "    plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "metadata": {},
   "outputs": [],
   "source": [
    "timestamp = \"2025_07_14-20_38_56\"\n",
    "results_30b, flat_30b = get_infill_data(timestamp)\n",
    "print(get_dur_success(flat_30b))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "np.savez(\"ç\", **results_30b)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "timestamp = \"2025_07_14-19_32_41\"\n",
    "results_auk, flat_auk = get_infill_data(timestamp)\n",
    "print(get_dur_success(flat_auk))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5",
   "metadata": {},
   "outputs": [],
   "source": [
    "np.savez(\n",
    "    \"/app2/suno/data/ditto_evals/baselines/auk_2025_07_15-19_14_15_infill.npz\",\n",
    "    **results_auk,\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "6",
   "metadata": {},
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "metadata": {},
   "outputs": [],
   "source": [
    "timestamp = \"2025_07_15-18_13_06\"\n",
    "results_blue_base, flat_blue_base = get_infill_data(timestamp)\n",
    "print(get_dur_success(flat_blue_base))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8",
   "metadata": {},
   "outputs": [],
   "source": [
    "np.savez(\n",
    "    \"/app2/suno/data/ditto_evals/baselines/bluejay_2025_07_15-19_14_15_infill.npz\",\n",
    "    **results_blue_base,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9",
   "metadata": {},
   "outputs": [],
   "source": [
    "timestamp = \"2025_07_15-19_14_15\"\n",
    "results_blue, flat_blue = get_infill_data(timestamp)\n",
    "print(get_dur_success(flat_blue))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10",
   "metadata": {},
   "outputs": [],
   "source": [
    "np.savez(\n",
    "    \"/app2/suno/data/ditto_evals/baselines/bluejay_2025_07_15-19_14_15_infill.npz\",\n",
    "    **results_blue,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "11",
   "metadata": {},
   "outputs": [],
   "source": [
    "results_blue[\"r&b\"].keys()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "12",
   "metadata": {},
   "outputs": [],
   "source": [
    "model_results = {\n",
    "    \"30b\": results_30b,\n",
    "    \"auk\": results_auk,\n",
    "    \"bluejay\": results_blue_base,\n",
    "    \"bluejay alt\": results_blue,\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13",
   "metadata": {},
   "outputs": [],
   "source": [
    "plot_multi_model_wer_grouped(model_results)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(get_wer_by_duration(flat_30b))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "15",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(get_wer_by_duration(flat_auk))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "16",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(get_wer_by_duration(flat_blue_base))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "17",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(get_wer_by_duration(flat_blue))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "18",
   "metadata": {},
   "source": [
    "### Test an Example"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "19",
   "metadata": {},
   "outputs": [],
   "source": [
    "test_id = \"25eea987-8409-4595-91d8-5f40e4e07b25\"\n",
    "json_data = load_json_from_s3(\n",
    "    \"suno-data-uploads\",\n",
    "    f\"tasks/feature_eval/cover_persona/2025_07_15-14_13_43/{test_id}_infill_wer.json\",\n",
    ")  # 2025_07_14-16_40_05"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "20",
   "metadata": {},
   "outputs": [],
   "source": [
    "audio = Audio.from_s3(test_id)\n",
    "audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "21",
   "metadata": {},
   "outputs": [],
   "source": [
    "start = json_data[\"infill_context_start_s\"]\n",
    "end = json_data[\"infill_context_end_s\"]\n",
    "source = json_data[\"source\"]\n",
    "lyrics = json_data[\"stripped_expected\"]\n",
    "lyrics_actual = json_data[\"stripped_result\"]\n",
    "print(lyrics)\n",
    "print(lyrics_actual)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "22",
   "metadata": {},
   "outputs": [],
   "source": [
    "source_audio = Audio.from_s3(source).get_segment(from_s=start, to_s=end)\n",
    "source_audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "23",
   "metadata": {},
   "outputs": [],
   "source": [
    "for key in json_data.keys():\n",
    "    if \"lyrics\" not in key:\n",
    "        print(f\"{key}: {json_data[key]}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "24",
   "metadata": {},
   "outputs": [],
   "source": [
    "json_data[\"infill_lyrics\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "25",
   "metadata": {},
   "outputs": [],
   "source": [
    "json_data[\"context_lyrics\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "26",
   "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
}
