{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import json\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "import labelbox as lb\n",
    "from datetime import datetime\n",
    "from uuid import uuid4\n",
    "\n",
    "from langdetect import detect\n",
    "from bs4 import BeautifulSoup\n",
    "from suno_utils.utils.s3 import check_s3_file_exists\n",
    "\n",
    "import tempfile\n",
    "from suno_utils.utils.s3 import upload_s3_files, download_s3_files\n",
    "\n",
    "import re\n",
    "\n",
    "pattern = r\"^([^_]+)_format_trimmed_([^_]+)$\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "inst_pairs = pd.read_csv(\n",
    "    \"/app2/suno/data/sara/musdb/audio_comparison2/vox_pairs_trimmed.txt\"\n",
    ")\n",
    "bass_pairs = pd.read_csv(\n",
    "    \"/app2/suno/data/sara/musdb/audio_comparison2/bass_pairs_trimmed.txt\"\n",
    ")\n",
    "drums_pairs = pd.read_csv(\n",
    "    \"/app2/suno/data/sara/musdb/audio_comparison2/drums_pairs_trimmed.txt\"\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "metadata": {},
   "outputs": [],
   "source": [
    "pairs = pd.concat([inst_pairs, bass_pairs, drums_pairs], ignore_index=True).sample(\n",
    "    frac=1\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Create a combined dataframe of all model->path mappings\n",
    "df_bass = pairs[pairs[\"instrument\"] == \"instrumental\"]\n",
    "model_path_pairs = pd.concat(\n",
    "    [\n",
    "        df_bass[[\"model_a\", \"source_a_fp\"]].rename(\n",
    "            columns={\"model_a\": \"model\", \"source_a_fp\": \"path\"}\n",
    "        ),\n",
    "        df_bass[[\"model_b\", \"source_b_fp\"]].rename(\n",
    "            columns={\"model_b\": \"model\", \"source_b_fp\": \"path\"}\n",
    "        ),\n",
    "    ]\n",
    ")\n",
    "\n",
    "# Group by model and get unique paths\n",
    "model_to_paths = (\n",
    "    model_path_pairs.groupby(\"model\")[\"path\"].apply(lambda x: list(set(x))).to_dict()\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "for model, sources in model_to_paths.items():\n",
    "    print(model, len(sources))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5",
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "import librosa\n",
    "import numpy as np\n",
    "from collections import defaultdict\n",
    "\n",
    "\n",
    "def calculate_spectral_centroid(file_path):\n",
    "    \"\"\"Calculate the spectral centroid for an audio file.\"\"\"\n",
    "    try:\n",
    "        y, sr = librosa.load(file_path)\n",
    "        spectral_centroid = librosa.feature.spectral_centroid(y=y, sr=sr)[0]\n",
    "        return np.mean(spectral_centroid)\n",
    "    except Exception as e:\n",
    "        print(f\"Error processing {file_path}: {e}\")\n",
    "        return None\n",
    "\n",
    "\n",
    "def plot_spectral_centroid_distributions(model_files_dict):\n",
    "    \"\"\"\n",
    "    Plot spectral centroid distributions for each model in separate graphs.\n",
    "\n",
    "    Args:\n",
    "        model_files_dict: Dictionary with model names as keys and lists of file paths as values\n",
    "    \"\"\"\n",
    "    # Calculate spectral centroids for each model\n",
    "    model_centroids = {}\n",
    "\n",
    "    for model_name, file_paths in model_files_dict.items():\n",
    "        centroids = []\n",
    "        print(f\"Processing {model_name}...\")\n",
    "\n",
    "        for file_path in file_paths:\n",
    "            centroid = calculate_spectral_centroid(file_path)\n",
    "            if centroid is not None:\n",
    "                centroids.append(centroid)\n",
    "\n",
    "        model_centroids[model_name] = centroids\n",
    "        print(f\"Processed {len(centroids)} files for {model_name}\")\n",
    "\n",
    "    # Create separate plots for each model\n",
    "    n_models = len(model_centroids)\n",
    "    fig, axes = plt.subplots(n_models, 1, figsize=(10, 4 * n_models))\n",
    "\n",
    "    # Handle case where there's only one model\n",
    "    if n_models == 1:\n",
    "        axes = [axes]\n",
    "\n",
    "    colors = [\"blue\", \"red\", \"green\", \"orange\", \"purple\", \"brown\", \"pink\", \"gray\"]\n",
    "\n",
    "    for idx, (model_name, centroids) in enumerate(model_centroids.items()):\n",
    "        ax = axes[idx]\n",
    "\n",
    "        if centroids:\n",
    "            # Plot histogram\n",
    "            ax.hist(\n",
    "                centroids,\n",
    "                bins=30,\n",
    "                alpha=0.7,\n",
    "                color=colors[idx % len(colors)],\n",
    "                edgecolor=\"black\",\n",
    "                linewidth=0.5,\n",
    "            )\n",
    "\n",
    "            ax.set_title(\n",
    "                f\"Spectral Centroid Distribution - {model_name}\",\n",
    "                fontsize=14,\n",
    "                fontweight=\"bold\",\n",
    "            )\n",
    "            ax.set_xlabel(\"Spectral Centroid (Hz)\")\n",
    "            ax.set_ylabel(\"Frequency\")\n",
    "            ax.grid(True, alpha=0.3)\n",
    "        else:\n",
    "            ax.text(\n",
    "                0.5,\n",
    "                0.5,\n",
    "                f\"No valid data for {model_name}\",\n",
    "                transform=ax.transAxes,\n",
    "                ha=\"center\",\n",
    "                va=\"center\",\n",
    "            )\n",
    "            ax.set_title(\n",
    "                f\"Spectral Centroid Distribution - {model_name}\",\n",
    "                fontsize=14,\n",
    "                fontweight=\"bold\",\n",
    "            )\n",
    "\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "def plot_spectral_centroid_distributions_density(model_files_dict):\n",
    "    \"\"\"\n",
    "    Alternative version using density plots (KDE) for separate graphs.\n",
    "\n",
    "    Args:\n",
    "        model_files_dict: Dictionary with model names as keys and lists of file paths as values\n",
    "    \"\"\"\n",
    "    from scipy import stats\n",
    "\n",
    "    # Calculate spectral centroids for each model\n",
    "    model_centroids = {}\n",
    "\n",
    "    for model_name, file_paths in model_files_dict.items():\n",
    "        centroids = []\n",
    "        print(f\"Processing {model_name}...\")\n",
    "\n",
    "        for file_path in file_paths:\n",
    "            centroid = calculate_spectral_centroid(file_path)\n",
    "            if centroid is not None:\n",
    "                centroids.append(centroid)\n",
    "\n",
    "        model_centroids[model_name] = centroids\n",
    "        print(f\"Processed {len(centroids)} files for {model_name}\")\n",
    "\n",
    "    # Create separate density plots for each model\n",
    "    n_models = len(model_centroids)\n",
    "    fig, axes = plt.subplots(n_models, 1, figsize=(10, 4 * n_models))\n",
    "\n",
    "    # Handle case where there's only one model\n",
    "    if n_models == 1:\n",
    "        axes = [axes]\n",
    "\n",
    "    colors = [\"blue\", \"red\", \"green\", \"orange\", \"purple\", \"brown\", \"pink\", \"gray\"]\n",
    "\n",
    "    for idx, (model_name, centroids) in enumerate(model_centroids.items()):\n",
    "        ax = axes[idx]\n",
    "\n",
    "        if centroids and len(centroids) > 1:\n",
    "            color = colors[idx % len(colors)]\n",
    "\n",
    "            # Create density plot\n",
    "            density = stats.gaussian_kde(centroids)\n",
    "            x_range = np.linspace(min(centroids), max(centroids), 200)\n",
    "            ax.fill_between(x_range, density(x_range), alpha=0.4, color=color)\n",
    "            ax.plot(x_range, density(x_range), color=color, linewidth=2)\n",
    "\n",
    "            ax.set_title(\n",
    "                f\"Spectral Centroid Density Distribution - {model_name}\",\n",
    "                fontsize=14,\n",
    "                fontweight=\"bold\",\n",
    "            )\n",
    "            ax.set_xlabel(\"Spectral Centroid (Hz)\")\n",
    "            ax.set_ylabel(\"Density\")\n",
    "            ax.grid(True, alpha=0.3)\n",
    "        else:\n",
    "            ax.text(\n",
    "                0.5,\n",
    "                0.5,\n",
    "                f\"No valid data for {model_name}\",\n",
    "                transform=ax.transAxes,\n",
    "                ha=\"center\",\n",
    "                va=\"center\",\n",
    "            )\n",
    "            ax.set_title(\n",
    "                f\"Spectral Centroid Density Distribution - {model_name}\",\n",
    "                fontsize=14,\n",
    "                fontweight=\"bold\",\n",
    "            )\n",
    "\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "plot_spectral_centroid_distributions(model_to_paths)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "import librosa\n",
    "import numpy as np\n",
    "from pathlib import Path\n",
    "\n",
    "\n",
    "def comprehensive_loudness_analysis(file_path):\n",
    "    \"\"\"Calculate multiple loudness metrics\"\"\"\n",
    "    try:\n",
    "        y, sr = librosa.load(file_path)\n",
    "\n",
    "        # RMS loudness\n",
    "        rms = librosa.feature.rms(y=y)[0]\n",
    "        rms_db = 20 * np.log10(np.mean(rms) + 1e-10)\n",
    "\n",
    "        # Peak amplitude\n",
    "        peak_db = 20 * np.log10(np.max(np.abs(y)) + 1e-10)\n",
    "\n",
    "        # Spectral centroid (brightness measure)\n",
    "        spectral_centroid = np.mean(librosa.feature.spectral_centroid(y=y, sr=sr))\n",
    "\n",
    "        return {\n",
    "            \"rms_db\": rms_db,\n",
    "            \"peak_db\": peak_db,\n",
    "            \"spectral_centroid\": spectral_centroid,\n",
    "        }\n",
    "    except Exception as e:\n",
    "        print(f\"Error processing {file_path}: {e}\")\n",
    "        return None\n",
    "\n",
    "\n",
    "# Calculate comprehensive metrics\n",
    "model_metrics = {}\n",
    "\n",
    "for model, paths in model_to_paths.items():\n",
    "    metrics_list = []\n",
    "    print(f\"Processing {model}...\")\n",
    "\n",
    "    for path in paths:\n",
    "        if Path(path).exists():\n",
    "            metrics = comprehensive_loudness_analysis(path)\n",
    "            if metrics is not None:\n",
    "                metrics_list.append(metrics)\n",
    "\n",
    "    model_metrics[model] = metrics_list\n",
    "\n",
    "# Generate comprehensive report\n",
    "print(\"\\n\" + \"=\" * 60)\n",
    "print(\"COMPREHENSIVE AUDIO ANALYSIS REPORT\")\n",
    "print(\"=\" * 60)\n",
    "\n",
    "for model, metrics_list in model_metrics.items():\n",
    "    if metrics_list:\n",
    "        # Extract values for each metric\n",
    "        rms_values = [m[\"rms_db\"] for m in metrics_list]\n",
    "        peak_values = [m[\"peak_db\"] for m in metrics_list]\n",
    "        centroid_values = [m[\"spectral_centroid\"] for m in metrics_list]\n",
    "\n",
    "        print(f\"\\n{model.upper()}:\")\n",
    "        print(f\"  Files processed: {len(metrics_list)}\")\n",
    "        print(f\"  RMS Loudness: {np.mean(rms_values):.2f} dB\")\n",
    "        print(f\"  Peak Loudness: {np.mean(peak_values):.2f} dB\")\n",
    "        print(f\"  Spectral Centroid: {np.mean(centroid_values):.0f} Hz\")\n",
    "    else:\n",
    "        print(f\"\\n{model.upper()}: No valid audio files found\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "metadata": {},
   "outputs": [],
   "source": [
    "preference_counts = score_column(\"won\", df)\n",
    "a_counts = score_column(\"source_a\", df)\n",
    "b_counts = score_column(\"source_b\", df)\n",
    "total_counts = a_counts + b_counts"
   ]
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
