{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import audiotools as at\n",
    "import tqdm\n",
    "\n",
    "from matplotlib import pyplot as plt\n",
    "import os\n",
    "from pathlib import Path\n",
    "from tqdm import tqdm\n",
    "\n",
    "eval_set = \"sample\"\n",
    "# eval_set = \"golden\"\n",
    "# goldens = f\"/home/victor/glockenspiel/descript-audio-codec/data/{eval_set}\"\n",
    "goldens = \"/app/suno/data/fad/reference_set_1k\"\n",
    "outputs = f\"{goldens}_outputs\"\n",
    "\n",
    "\n",
    "assert os.path.exists(goldens)\n",
    "assert os.path.exists(outputs)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from audiotools import AudioSignal\n",
    "\n",
    "# load outputs and match them to goldens\n",
    "output_files = at.util.find_audio(outputs)\n",
    "golden_files = at.util.find_audio(goldens)\n",
    "paired_signals = {}\n",
    "for f in tqdm(output_files):\n",
    "    for g in golden_files:\n",
    "        # check if f starts with g.filename\n",
    "        fp = Path(f)\n",
    "        gp = Path(g)\n",
    "        if fp.name.startswith(gp.stem):\n",
    "            # check same length\n",
    "            fs, gs = AudioSignal(f, duration=10), AudioSignal(g, duration=10)\n",
    "            # print(f, g)\n",
    "            length_dif = abs(fs.signal_length - gs.signal_length)\n",
    "            # if length_dif > 4800:\n",
    "            #     print(f\"Length difference of {length_dif} between {f} and {g}\")\n",
    "            #     continue\n",
    "\n",
    "            # truncate to the shortest\n",
    "            min_length = min(fs.signal_length, gs.signal_length)\n",
    "            fs = fs.truncate_samples(min_length)\n",
    "            gs = gs.truncate_samples(min_length)\n",
    "            assert fs.signal_length == gs.signal_length\n",
    "            paired_signals[fp.stem] = (fs, gs)\n",
    "            break\n",
    "print(f\"Paired {len(paired_signals)} outputs with goldens\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# from audiotools import metrics\n",
    "from audiotools import AudioSignal\n",
    "from dataclasses import dataclass\n",
    "\n",
    "import dac.nn.loss as losses\n",
    "\n",
    "\n",
    "@dataclass\n",
    "class State:\n",
    "    stft_loss: losses.MultiScaleSTFTLoss\n",
    "    mel_loss: losses.MelSpectrogramLoss\n",
    "    waveform_loss: losses.L1Loss\n",
    "    sisdr_loss: losses.SISDRLoss\n",
    "\n",
    "\n",
    "def get_metrics(signal: AudioSignal, recons: AudioSignal, state: State, sr=48000):\n",
    "    output = {}\n",
    "    x = signal.clone().resample(sr)\n",
    "    y = recons.clone().resample(sr)\n",
    "    output.update(\n",
    "        {\n",
    "            f\"mel\": state.mel_loss(x, y).item(),\n",
    "            f\"stft\": state.stft_loss(x, y).item(),\n",
    "            # f\"waveform\": state.waveform_loss(x, y).item(),\n",
    "            # f\"sisdr\": state.sisdr_loss(x, y).item(),\n",
    "            # f\"visqol-audio-{k}\": metrics.quality.visqol(x, y),\n",
    "            # f\"visqol-speech-{k}\": metrics.quality.visqol(x, y, \"speech\"),\n",
    "        }\n",
    "    )\n",
    "    output[\"path\"] = signal.path_to_file\n",
    "    output.update(signal.metadata)\n",
    "    return output\n",
    "\n",
    "\n",
    "state = State(\n",
    "    stft_loss=losses.MultiScaleSTFTLoss(),\n",
    "    mel_loss=losses.MelSpectrogramLoss(),\n",
    "    waveform_loss=losses.L1Loss(),\n",
    "    sisdr_loss=losses.SISDRLoss(),\n",
    ")\n",
    "\n",
    "metrics = []\n",
    "for k, (output, golden) in tqdm(paired_signals.items()):\n",
    "    m = get_metrics(output, golden, state)\n",
    "    m[\"model\"] = k.split(\"_model_\")[1]\n",
    "    m[\"golden\"] = k.split(\"_model_\")[0]\n",
    "    metrics.append(m)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import pandas as pd\n",
    "\n",
    "df = pd.DataFrame.from_dict(metrics)\n",
    "# df"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# aggregate by model name\n",
    "# group by model name, mean numeric columns\n",
    "grouped = df.groupby(\"model\")[\n",
    "    [\n",
    "        \"mel\",\n",
    "        \"stft\",\n",
    "        # \"waveform\",\n",
    "        # \"sisdr\",\n",
    "    ]\n",
    "]\n",
    "# check each group has the same number of samples\n",
    "assert len(grouped) == len(set(df[\"model\"]))\n",
    "samples_per_model = grouped.size()\n",
    "for k, v in samples_per_model.items():\n",
    "    assert (\n",
    "        v == samples_per_model[0]\n",
    "    ), f\"Model {k} has {v} samples, expected {samples_per_model[0]}\"\n",
    "\n",
    "grouped = grouped.mean()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# plot each metric\n",
    "for col in grouped.columns:\n",
    "    grouped[col].plot.bar()\n",
    "    plt.title(f\"{eval_set} {col} loss\")\n",
    "    plt.xticks(rotation=45)\n",
    "    y_min = grouped[col].min()\n",
    "    y_max = grouped[col].max()\n",
    "    y_range = (y_max - y_min) * 0.1\n",
    "    plt.ylim(y_min - y_range, y_max + y_range)\n",
    "\n",
    "    plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Find hard samples"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# only include 7B\n",
    "df_g = df[df[\"model\"] == \"7B_25x12\"]\n",
    "df_g = df_g.groupby(\"golden\")[[\"mel\", \"stft\"]].mean()\n",
    "# sort by stft\n",
    "df_g = df_g.sort_values(\"stft\")\n",
    "df_g.plot.hist(bins=20, alpha=0.5)\n",
    "\n",
    "# print top 10\n",
    "print(df_g.head(10))\n",
    "\n",
    "# print bottom 10\n",
    "print(df_g.tail(10))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "\n",
    "\n",
    "def play_audio_and_recons(path):\n",
    "    audio = Audio.from_file(path).get_segment(0, 10)\n",
    "    audio.play()\n",
    "\n",
    "    # also play a reconstruction from the 7B model\n",
    "    reconstruction = f\"{outputs}/{row[0]}_model_7B_25x12.mp3\"\n",
    "    assert os.path.exists(reconstruction)\n",
    "    audio = Audio.from_file(reconstruction).get_segment(0, 10)\n",
    "    audio.play()\n",
    "\n",
    "\n",
    "# play the worst\n",
    "for row in df_g.tail(5).iterrows():\n",
    "    print(row[0])\n",
    "    ground_truth = f\"{goldens}/{row[0]}.mp3\"\n",
    "    assert os.path.exists(ground_truth)\n",
    "    play_audio_and_recons(ground_truth)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# play the best\n",
    "for row in df_g.head(5).iterrows():\n",
    "    print(row[0])\n",
    "    ground_truth = f\"{goldens}/{row[0]}.mp3\"\n",
    "    assert os.path.exists(ground_truth)\n",
    "    play_audio_and_recons(ground_truth)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno",
   "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.12"
  },
  "orig_nbformat": 4
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
