{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"7\"\n",
    "\n",
    "import suno_utils.tasks.audio_features.beat_this_downbeat\n",
    "from suno_utils.audio import Audio\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "import logging\n",
    "logging.basicConfig(level=logging.INFO)\n",
    "\n",
    "extractor = suno_utils.tasks.audio_features.beat_this_downbeat.BeatThisDownbeatExtractor(device=\"cuda\", model_path=\"s3://suno-data/m4burns/beat_this_rc_12l.pt\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "# try these\n",
    "\n",
    "uuid = \"8dc2c1f3-213d-4503-84e2-65216085cc9a\"\n",
    "#uuid = \"5658e5c0-1252-495f-a14f-f92d978aeeba\"\n",
    "#uuid = \"039e43d6-07c7-48a9-af12-165e775a1fd7\"\n",
    "#uuid = \"55a2d2f8-8317-4b28-b468-ad2a95f60f21\"\n",
    "\n",
    "# downbeat placement could still use some improvement\n",
    "# but we only need beat 1 of the first full bar to be correct for studio\n",
    "\n",
    "audio = Audio.from_s3(f\"s3://suno-data-uploads/studio/uploads/{uuid}.mp3\")\n",
    "audio_mono = Audio.convert(audio, n_channels=1, sample_rate=audio.sample_rate, byte_width=audio.byte_width)\n",
    "\n",
    "out = extractor.extract(audio_mono)\n",
    "beat_info = np.array(out[\"downbeats\"])\n",
    "beats = beat_info[:, 0]\n",
    "downbeats = beat_info[beat_info[:, 1] == 1][:, 0]\n",
    "\n",
    "median_ibd = np.median(np.diff(beats))\n",
    "max_ibd = np.max(np.diff(beats))\n",
    "min_ibd = np.min(np.diff(beats))\n",
    "bpm_estimate = 60 / median_ibd.item()\n",
    "print(\"bpm estimate: \", bpm_estimate)\n",
    "print(\"max ibd: \", max_ibd)\n",
    "print(\"min ibd: \", min_ibd)\n",
    "\n",
    "def overlay_metronome(audio: Audio, beat_times: np.ndarray, downbeat_times: np.ndarray):\n",
    "    import librosa\n",
    "    downbeat_times_set = set(downbeat_times.tolist())\n",
    "    audio_beat = librosa.clicks(times = np.array([b for b in beat_times.tolist() if b not in downbeat_times_set]), sr=audio.sample_rate, click_freq=1000, length=audio.array_float.shape[-1])\n",
    "    audio_downbeat = librosa.clicks(times = downbeat_times, sr=audio.sample_rate, click_freq=1500, length=audio.array_float.shape[-1])\n",
    "    if audio.array_float.ndim > 1:\n",
    "        audio_beat = audio_beat[None, :]\n",
    "        audio_downbeat = audio_downbeat[None, :]\n",
    "    return Audio.from_array_float(audio.array_float + audio_beat.reshape(1, -1) + audio_downbeat.reshape(1, -1), audio.sample_rate, max_allowed_val=10.0)\n",
    "\n",
    "def beats_to_tempo(beats: np.ndarray):\n",
    "    return 60 / np.diff(beats, axis=-1)\n",
    "\n",
    "def get_tempo_cv(beats: np.ndarray):\n",
    "    tempo = beats_to_tempo(beats)\n",
    "    return np.std(tempo, axis=-1) / np.mean(tempo, axis=-1)\n",
    "\n",
    "T_PREFIX = 10\n",
    "\n",
    "beats = np.array(out[\"raw_downbeats\"])\n",
    "beats_refined = np.array(out[\"downbeats\"])\n",
    "onsets = extractor.get_onsets(audio_mono, as_timeseries=True)\n",
    "\n",
    "plt.figure(figsize=(10, 3))\n",
    "plt.title(\"Onset Signal with Original and Refined Beats\")\n",
    "plt.xlabel(\"Frame (at 200 Hz)\")\n",
    "plt.ylabel(\"Onset Strength\")\n",
    "plt.plot(onsets[:T_PREFIX*200])\n",
    "beats_at_200hz = np.round(beats[beats[:, 0] < T_PREFIX, 0] * 200).astype(int)\n",
    "beats_refined_at_200hz = np.round(beats_refined[beats_refined[:, 0] < T_PREFIX, 0] * 200).astype(int)\n",
    "plt.scatter(beats_at_200hz, onsets[beats_at_200hz], marker='o', color='red')\n",
    "plt.scatter(beats_refined_at_200hz, onsets[beats_refined_at_200hz], marker='x', color='blue')\n",
    "plt.legend(['Onset', 'Original Beats', 'Refined Beats'], loc='upper right')\n",
    "plt.show()\n",
    "\n",
    "print(\"original cv: \", get_tempo_cv(beats[:, 0]))\n",
    "print(\"original sum onset energy: \", np.sum(onsets[beats_at_200hz]))\n",
    "print(\"refined cv: \", get_tempo_cv(beats_refined[:, 0]))\n",
    "print(\"refined sum onset energy: \", np.sum(onsets[beats_refined_at_200hz]))\n",
    "\n",
    "plt.figure(figsize=(10, 3))\n",
    "plt.title(\"Beat-to-Beat Tempo: Original vs Refined\")\n",
    "plt.xlabel(\"Beat Index\")\n",
    "plt.ylabel(\"Tempo (BPM)\")\n",
    "plt.plot(beats_to_tempo(beats[:, 0]), label=\"Original\")\n",
    "plt.plot(beats_to_tempo(beats_refined[:, 0]), label=\"Refined\")\n",
    "plt.legend()\n",
    "plt.show()\n",
    "\n",
    "print(\"original\")\n",
    "overlay_metronome(audio, beats[beats[:,1] != 1, 0], beats[beats[:,1] == 1, 0]).play()\n",
    "\n",
    "print(\"refined\")\n",
    "overlay_metronome(audio, beats_refined[beats_refined[:, 1] != 1, 0], beats_refined[beats_refined[:, 1] == 1, 0]).play()"
   ]
  }
 ],
 "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
}
