{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "os.environ['PATH'] += \":/home/m4burns/\"\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"7\"\n",
    "\n",
    "import suno_utils.tasks.audio_features.beat_this_downbeat\n",
    "from suno_utils.tasks.audio_features.downbeats_data_prep.augment import stretch_audio\n",
    "from suno_utils.audio import Audio\n",
    "\n",
    "import logging\n",
    "logging.basicConfig(level=logging.WARN)\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": [
    "def extract_beats(audio: np.ndarray) -> np.ndarray:\n",
    "    audio_mono = Audio.convert(audio, n_channels=1, sample_rate=audio.sample_rate, byte_width=audio.byte_width)\n",
    "    out = extractor.extract(audio_mono)\n",
    "    beats_refined = np.array(out[\"downbeats\"])\n",
    "    return beats_refined[:, 0]\n",
    "\n",
    "def beats_to_tempo(beats: np.ndarray) -> float:\n",
    "    return 60 / np.diff(beats, axis=-1)\n",
    "\n",
    "def est_target_tempo(beats: np.ndarray) -> int:\n",
    "    tempos = 60 / np.diff(beats, axis=-1)\n",
    "    return round(np.median(tempos))\n",
    "\n",
    "def get_target_beats(beat_times: np.ndarray, target_tempo: int) -> list[float]:\n",
    "    target_beat_time = 60. / target_tempo\n",
    "    print(f\"Beat length: {target_beat_time}\")\n",
    "    target_beats = []\n",
    "\n",
    "    for idx, beat in enumerate(beat_times):\n",
    "        if idx == 0: # first beat\n",
    "            next_beat = beat\n",
    "        else:\n",
    "            last_beat = target_beats[idx - 1]\n",
    "            next_beat = last_beat + target_beat_time\n",
    "        target_beats.append(next_beat)\n",
    "\n",
    "    return target_beats\n",
    "\n",
    "def get_warp_markers(beat_times: np.ndarray, target_tempo) -> list[(float, float)]:\n",
    "    target_beats = get_target_beats(beat_times, target_tempo)\n",
    "    return list(zip(beat_times.tolist(), target_beats))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "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\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "audio = Audio.from_s3(f\"s3://suno-data-uploads/studio/uploads/{uuid}.mp3\")\n",
    "audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "audio = Audio.from_s3(f\"s3://suno-data-uploads/studio/uploads/{uuid}.mp3\")\n",
    "beat_times = extract_beats(audio)\n",
    "target_tempo = est_target_tempo(beat_times)\n",
    "print(f\"Normalizing to {target_tempo} bpm\")\n",
    "warp_markers = get_warp_markers(beat_times, target_tempo)\n",
    "out_audio = stretch_audio(audio, warp_markers)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5",
   "metadata": {},
   "outputs": [],
   "source": [
    "out_audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.figure(figsize=(10, 3))\n",
    "plt.title(\"Beat-to-Beat Tempo\")\n",
    "plt.xlabel(\"Beat Index\")\n",
    "plt.ylabel(\"Tempo (BPM)\")\n",
    "plt.plot(beats_to_tempo(beat_times), label=\"Original\")\n",
    "plt.plot(beats_to_tempo(np.asarray([wm[1] for wm in warp_markers])), label=\"Normalized\")\n",
    "plt.legend()\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "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
}
