{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import pandas as pd\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "import matplotlib.pyplot as plt\n",
    "import numpy as np\n",
    "from collections import defaultdict\n",
    "from suno_utils.audio.conversion import Audio\n",
    "from IPython.display import clear_output"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "import re\n",
    "\n",
    "pattern = r\"^([^_]+)_format_trimmed_([^_]+)$\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "metadata": {},
   "outputs": [],
   "source": [
    "labelbox_path = \"suno_utils/task_eval/labelbox/outputs/Export  project - stems-trimmed-preference - 6_18_2025.ndjson\"\n",
    "labelbox_data = read_jsonl(labelbox_path)\n",
    "\n",
    "metadata_filepath = (\n",
    "    \"suno_utils/task_eval/labelbox/outputs/metadata-more-stems-trimmed-20250617.json\"\n",
    ")\n",
    "metadata = json.load(open(metadata_filepath, \"r\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "metadata"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "model_names = [\"original\", \"suno\", \"lalal\"]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5",
   "metadata": {},
   "outputs": [],
   "source": [
    "seen_per = {name: 0 for name in model_names}\n",
    "wins_per = {name: 0 for name in model_names}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "results = []\n",
    "for row in labelbox_data:\n",
    "    row_id = row[\"data_row\"][\"global_key\"]\n",
    "\n",
    "    row_metadata = metadata[row_id]\n",
    "    source_a = re.match(pattern, row_metadata[\"source_a\"]).group(1)\n",
    "    source_b = re.match(pattern, row_metadata[\"source_b\"]).group(1)\n",
    "    song_name = row_metadata[\"song_name\"]\n",
    "    instrument = row_metadata[\"instrument\"]\n",
    "\n",
    "    projects = list(row[\"projects\"].keys())\n",
    "    assert len(projects) == 1\n",
    "    project_id = projects[0]\n",
    "    project = row[\"projects\"][project_id]\n",
    "    labels = project[\"labels\"]\n",
    "    for label in labels:  # represents each rating of the row, should be consensus count\n",
    "        classifications = label[\"annotations\"][\"classifications\"]\n",
    "        if len(classifications) != 1:\n",
    "            print(classifications)\n",
    "        for c in classifications:\n",
    "            question = c[\"name\"]\n",
    "            answer = c[\"radio_answer\"][\"value\"]\n",
    "\n",
    "            won = None\n",
    "            if answer == \"A\":\n",
    "                won = source_a\n",
    "            elif answer == \"B\":\n",
    "                won = source_b\n",
    "\n",
    "            results.append(\n",
    "                {\n",
    "                    \"row_id\": row_id,\n",
    "                    \"won\": won,\n",
    "                    \"source_a\": source_a,\n",
    "                    \"source_b\": source_b,\n",
    "                    \"song_name\": song_name,\n",
    "                    \"instrument\": instrument,\n",
    "                    \"question\": question,\n",
    "                    \"answer\": answer,\n",
    "                }\n",
    "            )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "metadata": {},
   "outputs": [],
   "source": [
    "df = pd.DataFrame(results)\n",
    "df.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8",
   "metadata": {},
   "outputs": [],
   "source": [
    "suno_v_lalal = df[df[[\"source_a\", \"source_b\"]].isin([\"original\"]).any(axis=1)]\n",
    "suno_v_lalal = suno_v_lalal[\n",
    "    suno_v_lalal[[\"source_a\", \"source_b\"]].isin([\"lalal\"]).any(axis=1)\n",
    "]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9",
   "metadata": {},
   "outputs": [],
   "source": [
    "suno_files = []\n",
    "lalal_files = []\n",
    "for row_id in suno_v_lalal[\"row_id\"]:\n",
    "    row_metadata = metadata[row_id]\n",
    "    source_a = row_metadata[\"source_a_fp\"]\n",
    "    source_b = row_metadata[\"source_b_fp\"]\n",
    "    if \"lalal\" in source_a:\n",
    "        suno_source = source_b\n",
    "        lalal_source = source_a\n",
    "    else:\n",
    "        suno_source = source_a\n",
    "        lalal_source = source_b\n",
    "    suno_files.append(suno_source)\n",
    "    lalal_files.append(lalal_source)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10",
   "metadata": {},
   "outputs": [],
   "source": [
    "lalal_wins = suno_v_lalal[suno_v_lalal[\"won\"] == \"lalal\"]\n",
    "lalal_wins.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "11",
   "metadata": {},
   "outputs": [],
   "source": [
    "suno_wins = suno_v_lalal[suno_v_lalal[\"won\"] == \"original\"]\n",
    "suno_wins.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "12",
   "metadata": {},
   "outputs": [],
   "source": [
    "for row_id in suno_wins[\"row_id\"]:\n",
    "    row_metadata = metadata[row_id]\n",
    "    source_a = row_metadata[\"source_a_fp\"]\n",
    "    source_b = row_metadata[\"source_b_fp\"]\n",
    "    if \"lalal\" in source_a:\n",
    "        suno_source = source_b\n",
    "        lalal_source = source_a\n",
    "    else:\n",
    "        suno_source = source_a\n",
    "        lalal_source = source_b\n",
    "\n",
    "    suno_audio = Audio.from_file(suno_source)\n",
    "    lalal_audio = Audio.from_file(lalal_source)\n",
    "\n",
    "    suno_audio.play()\n",
    "    lalal_audio.play()\n",
    "    clear_output()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13",
   "metadata": {},
   "outputs": [],
   "source": [
    "for row_id in lalal_wins[\"row_id\"]:\n",
    "    row_metadata = metadata[row_id]\n",
    "    source_a = row_metadata[\"source_a_fp\"]\n",
    "    source_b = row_metadata[\"source_b_fp\"]\n",
    "    if \"lalal\" in source_a:\n",
    "        suno_source = source_b\n",
    "        lalal_source = source_a\n",
    "    else:\n",
    "        suno_source = source_a\n",
    "        lalal_source = source_b\n",
    "\n",
    "    suno_audio = Audio.from_file(suno_source)\n",
    "    lalal_audio = Audio.from_file(lalal_source)\n",
    "\n",
    "    suno_audio.play()\n",
    "    lalal_audio.play()\n",
    "    clear_output()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14",
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "import librosa\n",
    "import numpy as np\n",
    "from pathlib import Path\n",
    "\n",
    "\n",
    "def calculate_loudness_lufs(file_path):\n",
    "    \"\"\"\n",
    "    Calculate loudness in LUFS (Loudness Units relative to Full Scale).\n",
    "    This is a perceptually relevant loudness measure.\n",
    "    \"\"\"\n",
    "    try:\n",
    "        y, sr = librosa.load(file_path)\n",
    "        # Calculate RMS energy and convert to LUFS approximation\n",
    "        rms = librosa.feature.rms(y=y)[0]\n",
    "        # Convert RMS to dB and approximate LUFS\n",
    "        loudness_db = 20 * np.log10(\n",
    "            np.mean(rms) + 1e-8\n",
    "        )  # Add small value to avoid log(0)\n",
    "        return loudness_db\n",
    "    except Exception as e:\n",
    "        print(f\"Error processing {file_path}: {e}\")\n",
    "        return None\n",
    "\n",
    "\n",
    "def calculate_loudness_rms(file_path):\n",
    "    \"\"\"\n",
    "    Alternative: Calculate RMS loudness in dB.\n",
    "    \"\"\"\n",
    "    try:\n",
    "        y, sr = librosa.load(file_path)\n",
    "        rms = np.sqrt(np.mean(y**2))\n",
    "        loudness_db = 20 * np.log10(rms + 1e-8)\n",
    "        return loudness_db\n",
    "    except Exception as e:\n",
    "        print(f\"Error processing {file_path}: {e}\")\n",
    "        return None\n",
    "\n",
    "\n",
    "def compare_pairwise_loudness(\n",
    "    model_a_files, model_b_files, loudness_func=calculate_loudness_lufs\n",
    "):\n",
    "    \"\"\"\n",
    "    Compare loudness between pairwise audio files from two models.\n",
    "\n",
    "    Args:\n",
    "        model_a_files: List of file paths for model A\n",
    "        model_b_files: List of file paths for model B (same order as model A)\n",
    "        loudness_func: Function to calculate loudness (default: calculate_loudness_lufs)\n",
    "\n",
    "    Returns:\n",
    "        Dictionary with comparison results\n",
    "    \"\"\"\n",
    "    if len(model_a_files) != len(model_b_files):\n",
    "        raise ValueError(\"Model A and Model B must have the same number of files\")\n",
    "\n",
    "    results = {\n",
    "        \"model_a_loudness\": [],\n",
    "        \"model_b_loudness\": [],\n",
    "        \"loudness_differences\": [],  # Model A - Model B\n",
    "        \"pair_names\": [],\n",
    "        \"valid_pairs\": [],\n",
    "    }\n",
    "\n",
    "    print(f\"Processing {len(model_a_files)} pairs...\")\n",
    "\n",
    "    for i, (file_a, file_b) in enumerate(zip(model_a_files, model_b_files)):\n",
    "        print(f\"Processing pair {i+1}/{len(model_a_files)}\")\n",
    "\n",
    "        loudness_a = loudness_func(file_a)\n",
    "        loudness_b = loudness_func(file_b)\n",
    "\n",
    "        if loudness_a is not None and loudness_b is not None:\n",
    "            difference = loudness_a - loudness_b\n",
    "\n",
    "            results[\"model_a_loudness\"].append(loudness_a)\n",
    "            results[\"model_b_loudness\"].append(loudness_b)\n",
    "            results[\"loudness_differences\"].append(difference)\n",
    "            results[\"pair_names\"].append(f\"Pair_{i+1}\")\n",
    "            results[\"valid_pairs\"].append(i)\n",
    "\n",
    "    print(f\"Successfully processed {len(results['valid_pairs'])} pairs\")\n",
    "    return results\n",
    "\n",
    "\n",
    "def plot_scatter_comparison(results, model_a_name=\"Model A\", model_b_name=\"Model B\"):\n",
    "    \"\"\"\n",
    "    Plot scatter comparison of Model A vs Model B loudness.\n",
    "    \"\"\"\n",
    "    plt.figure(figsize=(8, 6))\n",
    "\n",
    "    plt.scatter(\n",
    "        results[\"model_a_loudness\"],\n",
    "        results[\"model_b_loudness\"],\n",
    "        alpha=0.6,\n",
    "        s=50,\n",
    "        color=\"blue\",\n",
    "    )\n",
    "\n",
    "    # Add diagonal line for reference (equal loudness)\n",
    "    min_loud = min(min(results[\"model_a_loudness\"]), min(results[\"model_b_loudness\"]))\n",
    "    max_loud = max(max(results[\"model_a_loudness\"]), max(results[\"model_b_loudness\"]))\n",
    "    plt.plot(\n",
    "        [min_loud, max_loud],\n",
    "        [min_loud, max_loud],\n",
    "        \"r--\",\n",
    "        alpha=0.7,\n",
    "        label=\"Equal Loudness\",\n",
    "    )\n",
    "\n",
    "    plt.xlabel(f\"{model_a_name} Loudness (dB)\")\n",
    "    plt.ylabel(f\"{model_b_name} Loudness (dB)\")\n",
    "    plt.title(\"Pairwise Loudness Comparison\")\n",
    "    plt.grid(True, alpha=0.3)\n",
    "    plt.legend()\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "def plot_difference_histogram(results, model_a_name=\"Model A\", model_b_name=\"Model B\"):\n",
    "    \"\"\"\n",
    "    Plot histogram of loudness differences.\n",
    "    \"\"\"\n",
    "    plt.figure(figsize=(8, 6))\n",
    "\n",
    "    plt.hist(\n",
    "        results[\"loudness_differences\"],\n",
    "        bins=20,\n",
    "        alpha=0.7,\n",
    "        color=\"green\",\n",
    "        edgecolor=\"black\",\n",
    "        linewidth=0.5,\n",
    "    )\n",
    "    plt.axvline(0, color=\"red\", linestyle=\"--\", linewidth=2, label=\"No Difference\")\n",
    "    plt.axvline(\n",
    "        np.mean(results[\"loudness_differences\"]),\n",
    "        color=\"orange\",\n",
    "        linestyle=\"--\",\n",
    "        linewidth=2,\n",
    "        label=f'Mean: {np.mean(results[\"loudness_differences\"]):.2f} dB',\n",
    "    )\n",
    "\n",
    "    plt.xlabel(f\"Loudness Difference ({model_a_name} - {model_b_name}) [dB]\")\n",
    "    plt.ylabel(\"Frequency\")\n",
    "    plt.title(\"Distribution of Loudness Differences\")\n",
    "    plt.grid(True, alpha=0.3)\n",
    "    plt.legend()\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "def plot_pairwise_bars(results, model_a_name=\"Model A\", model_b_name=\"Model B\"):\n",
    "    \"\"\"\n",
    "    Plot bar chart of differences by pair.\n",
    "    \"\"\"\n",
    "    plt.figure(figsize=(12, 6))\n",
    "\n",
    "    pair_indices = range(len(results[\"loudness_differences\"]))\n",
    "    colors = [\"red\" if diff > 0 else \"blue\" for diff in results[\"loudness_differences\"]]\n",
    "\n",
    "    bars = plt.bar(\n",
    "        pair_indices, results[\"loudness_differences\"], color=colors, alpha=0.6\n",
    "    )\n",
    "    plt.axhline(0, color=\"black\", linestyle=\"-\", linewidth=1)\n",
    "\n",
    "    plt.xlabel(\"Audio Pair Index\")\n",
    "    plt.ylabel(\"Loudness Difference (dB)\")\n",
    "    plt.title(\n",
    "        f\"Loudness Differences by Pair\\n(Red: {model_a_name} Louder, Blue: {model_b_name} Louder)\"\n",
    "    )\n",
    "    plt.grid(True, alpha=0.3, axis=\"y\")\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "def plot_boxplot_comparison(results, model_a_name=\"Model A\", model_b_name=\"Model B\"):\n",
    "    \"\"\"\n",
    "    Plot box plot comparison of loudness distributions.\n",
    "    \"\"\"\n",
    "    plt.figure(figsize=(8, 6))\n",
    "\n",
    "    box_data = [results[\"model_a_loudness\"], results[\"model_b_loudness\"]]\n",
    "    box_labels = [model_a_name, model_b_name]\n",
    "\n",
    "    bp = plt.boxplot(box_data, labels=box_labels, patch_artist=True)\n",
    "    bp[\"boxes\"][0].set_facecolor(\"lightblue\")\n",
    "    bp[\"boxes\"][1].set_facecolor(\"lightcoral\")\n",
    "\n",
    "    plt.ylabel(\"Loudness (dB)\")\n",
    "    plt.title(\"Loudness Distribution Comparison\")\n",
    "    plt.grid(True, alpha=0.3)\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "def print_summary_statistics(results, model_a_name=\"Model A\", model_b_name=\"Model B\"):\n",
    "    \"\"\"\n",
    "    Print detailed summary statistics.\n",
    "    \"\"\"\n",
    "    print(\"\\n\" + \"=\" * 50)\n",
    "    print(\"LOUDNESS COMPARISON SUMMARY\")\n",
    "    print(\"=\" * 50)\n",
    "    print(f\"Number of pairs analyzed: {len(results['loudness_differences'])}\")\n",
    "    print(f\"\\n{model_a_name} Statistics:\")\n",
    "    print(f\"  Mean: {np.mean(results['model_a_loudness']):.2f} dB\")\n",
    "    print(f\"  Std:  {np.std(results['model_a_loudness']):.2f} dB\")\n",
    "    print(\n",
    "        f\"  Range: {np.min(results['model_a_loudness']):.2f} to {np.max(results['model_a_loudness']):.2f} dB\"\n",
    "    )\n",
    "\n",
    "    print(f\"\\n{model_b_name} Statistics:\")\n",
    "    print(f\"  Mean: {np.mean(results['model_b_loudness']):.2f} dB\")\n",
    "    print(f\"  Std:  {np.std(results['model_b_loudness']):.2f} dB\")\n",
    "    print(\n",
    "        f\"  Range: {np.min(results['model_b_loudness']):.2f} to {np.max(results['model_b_loudness']):.2f} dB\"\n",
    "    )\n",
    "\n",
    "    print(f\"\\nLoudness Differences ({model_a_name} - {model_b_name}):\")\n",
    "    print(f\"  Mean: {np.mean(results['loudness_differences']):.2f} dB\")\n",
    "    print(f\"  Std:  {np.std(results['loudness_differences']):.2f} dB\")\n",
    "    print(\n",
    "        f\"  Range: {np.min(results['loudness_differences']):.2f} to {np.max(results['loudness_differences']):.2f} dB\"\n",
    "    )\n",
    "\n",
    "    louder_a = sum(1 for diff in results[\"loudness_differences\"] if diff > 0)\n",
    "    louder_b = len(results[\"loudness_differences\"]) - louder_a\n",
    "    print(\n",
    "        f\"\\n{model_a_name} is louder in {louder_a}/{len(results['loudness_differences'])} pairs ({louder_a/len(results['loudness_differences'])*100:.1f}%)\"\n",
    "    )\n",
    "    print(\n",
    "        f\"{model_b_name} is louder in {louder_b}/{len(results['loudness_differences'])} pairs ({louder_b/len(results['loudness_differences'])*100:.1f}%)\"\n",
    "    )\n",
    "\n",
    "\n",
    "# Example usage:\n",
    "if __name__ == \"__main__\":\n",
    "    # Example file lists - make sure they're in corresponding order\n",
    "    model_a_files = [\n",
    "        \"path/to/model_a/audio_1.wav\",\n",
    "        \"path/to/model_a/audio_2.wav\",\n",
    "        \"path/to/model_a/audio_3.wav\",\n",
    "        # ... more files\n",
    "    ]\n",
    "\n",
    "    model_b_files = [\n",
    "        \"path/to/model_b/audio_1.wav\",  # Corresponding to model_a audio_1\n",
    "        \"path/to/model_b/audio_2.wav\",  # Corresponding to model_a audio_2\n",
    "        \"path/to/model_b/audio_3.wav\",  # Corresponding to model_a audio_3\n",
    "        # ... more files\n",
    "    ]\n",
    "\n",
    "# Compare loudness\n",
    "results = compare_pairwise_loudness(suno_files, lalal_files)\n",
    "\n",
    "# Create plots\n",
    "plot_scatter_comparison(results, \"Original\", \"Lalal\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "15",
   "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
}
