{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# get semantic codes from a song\n",
    "import json\n",
    "import torch\n",
    "import os\n",
    "from suno_utils.audio import Audio\n",
    "import numpy as np\n",
    "from suno_utils.utils.s3 import read_from_s3\n",
    "\n",
    "\n",
    "# new codec\n",
    "#gen_id = \"85be20bd-a388-4872-8a0d-77e3385e9bde\"\n",
    "gen_id = \"549cfa66-65cf-4bee-b264-389c0c029b45\"\n",
    "#gen_id = \"72adda71-4382-4cfa-b21f-150fb1809044\"\n",
    "#gen_id = \"7be341c0-b9af-48da-856b-136f409bba01\"\n",
    "#gen_id = \"1305dd9b-1625-4859-84ed-d40ee56b4d86\"\n",
    "\n",
    "# greek\n",
    "#gen_id = \"579dafcc-dc96-4993-9572-f904d4c37f3d\"\n",
    "\n",
    "# staging clips\n",
    "#gen_id = \"e65d270a-61f1-462f-9c38-8544bfc51a33\"\n",
    "\n",
    "\n",
    "s3_filepath = f\"s3://suno-data-uploads/studio/uploads/{gen_id}.npz\"\n",
    "mp3_filepath = f\"s3://suno-data-uploads/studio/uploads/{gen_id}.mp3\"\n",
    "vae_filepath = f\"s3://suno-data-uploads/studio/uploads/{gen_id}_vae.npz\"\n",
    "\n",
    "\n",
    "print(s3_filepath)\n",
    "data = read_from_s3(s3_filepath, read_f=np.load)\n",
    "\n",
    "vae_data = read_from_s3(vae_filepath, read_f=np.load)\n",
    "print(vae_data[\"vae_latents\"].shape)\n",
    "audio = Audio.from_s3(mp3_filepath, n_channels=2)#.get_slice(0, 120.01)\n",
    "\n",
    "if \"v3.0_raw\" in data:\n",
    "    codes = data[\"v3.0_raw\"]\n",
    "elif \"v3.5_raw\" in data:\n",
    "    codes = data[\"v3.5_raw\"]\n",
    "elif \"v4.0_raw\" in data:\n",
    "    codes = data[\"v4.0_raw\"]\n",
    "elif \"v5.0_raw\" in data:\n",
    "    codes = data[\"v5.0_raw\"]\n",
    "else:\n",
    "    raise ValueError(\"No codes found\")\n",
    "\n",
    "\n",
    "#aws s3 cp s3://suno-data-uploads/studio/uploads/2e0eec9e-4f86-46c5-9a43-508829fb9d0b_hoot.json text_data.json\n",
    "\n",
    "#text_data = read_from_s3(f\"s3://suno-data-uploads/studio/uploads/{gen_id}_hoot.json\")\n",
    "os.system(f\"aws s3 cp s3://suno-data-uploads/studio/uploads/{gen_id}_hoot.json text_data.json\")\n",
    "text_data = open(\"text_data.json\", \"r\", encoding=\"utf-8\").read()\n",
    "aligned_lyrics = json.loads(text_data)\n",
    "print(aligned_lyrics)\n",
    "audio.normalize_volume().play()\n",
    "\n",
    "tags = \"pop\"\n",
    "\n",
    "lyrics = \"\"\n",
    "for elem in aligned_lyrics:\n",
    "    if \"word\" in elem:\n",
    "        lyrics += elem[\"word\"]\n",
    "\n",
    "semantic_codes = torch.from_numpy(codes[:, 0]).long()#.cuda()\n",
    "#semantic_codes = semantic_codes[:3000]\n",
    "print(semantic_codes.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from madmom.features import DBNDownBeatTrackingProcessor\n",
    "from suno_utils.tasks.audio_features.downbeat import DownbeatExtractor\n",
    "import torch.nn.functional as F\n",
    "import torch\n",
    "\n",
    "downbeat_extractor = DownbeatExtractor()\n",
    "\n",
    "downbeat_processor = DBNDownBeatTrackingProcessor(\n",
    "    fps=100, num_threads=2, beats_per_bar=[3, 4]\n",
    ")\n",
    "\n",
    "with torch.no_grad():\n",
    "    y_hat = (\n",
    "        F.softmax(\n",
    "            downbeat_extractor.mert_to_downbeat_model(\n",
    "                torch.tensor(semantic_codes, dtype=torch.int64)\n",
    "                .reshape(1, -1, 1)\n",
    "                .to(downbeat_extractor.device)\n",
    "            ),\n",
    "            dim=-1,\n",
    "        )\n",
    "        .cpu()\n",
    "        .numpy()\n",
    "        .squeeze()\n",
    "    )\n",
    "    #print(y_hat.shape)\n",
    "    downbeats = downbeat_processor.process(y_hat[:, :2])\n",
    "    print(downbeats)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env2",
   "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": 2
}
