{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"6\"\n",
    "\n",
    "import numpy as np\n",
    "import torch\n",
    "\n",
    "from suno_utils.utils.s3 import read_from_s3\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.diffusion.generation import DiffusionGenerationConfig, generate, preload_models, generate_chunk\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "preload_models()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\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",
    "#gen_id = \"99bffa17-7e59-47b4-a048-5528cbda05d5\" # sister\n",
    "#gen_id = \"081d73c4-7805-4212-9c80-8db1137ca3c4\" # friends\n",
    "#gen_id = \"562f762d-6ced-4080-9af1-910ee3d0a5dc\" # something real\n",
    "#gen_id = \"23c15c62-494d-422d-8a60-8b0454044322\" # rubber duck\n",
    "#gen_id = \"4b140a9e-964b-422c-85b5-5861ad1a9d38\" # once\n",
    "#gen_id = \"7b214347-fa38-4e9b-96f4-f7ec65adea45\" # rock n roll\n",
    "gen_id = \"a5e2198a-f352-4abb-9a24-7f81b143ded3\" # stone\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",
    "\n",
    "print(s3_filepath)\n",
    "data = read_from_s3(s3_filepath, read_f=np.load)\n",
    "\n",
    "audio = Audio.from_s3(mp3_filepath, n_channels=2).get_slice(0, 60.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",
    "else:\n",
    "    raise ValueError(\"No codes found\")\n",
    "\n",
    "text_data = read_from_s3(f\"s3://suno-data-uploads/studio/uploads/{gen_id}_hoot.json\")\n",
    "aligned_lyrics = json.loads(text_data)\n",
    "print(aligned_lyrics)\n",
    "audio.normalize_volume().play()\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[750:1500]\n",
    "print(semantic_codes.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# get vae npz from another song to use as history\n",
    "gen_id = \"707299a3-7668-484f-8479-421e516c6916\" # another song\n",
    "gen_id = \"ca293380-933b-4358-b3a4-e0460d415d02\"\n",
    "\n",
    "s3_filepath = f\"s3://suno-data-uploads/studio/uploads/{gen_id}_vae.npz\"\n",
    "mp3_filepath = f\"s3://suno-data-uploads/studio/uploads/{gen_id}.mp3\"\n",
    "\n",
    "audio = Audio.from_s3(mp3_filepath, n_channels=2)\n",
    "# get the last 30 seconds of audio \n",
    "audio = audio.get_slice(-30.0, None)\n",
    "audio.normalize_volume().play()\n",
    "\n",
    "\n",
    "print(s3_filepath)\n",
    "data = read_from_s3(s3_filepath, read_f=np.load)\n",
    "vae_latents = torch.from_numpy(data[\"vae_latents\"]).float().cuda()\n",
    "print(vae_latents.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#history_latents = torch.randn(1, semantic_codes.shape[0], 128).cuda() * 1000000.0\n",
    "#history_latents = torch.randn(1, semantic_codes.shape[0], 128).cuda()\n",
    "\n",
    "history_latents = vae_latents[-semantic_codes.shape[0]:]\n",
    "print(history_latents.shape)\n",
    "\n",
    "for use_history in [True, False]:\n",
    "    print(f\"use_history: {use_history}\")\n",
    "    generation_config = DiffusionGenerationConfig(\n",
    "        audio=semantic_codes,\n",
    "        history_latents=history_latents if use_history else None,\n",
    "        lyrics=lyrics,\n",
    "        ctx_cfg_coef=5.0,\n",
    "        text_cfg_coef=2.0,\n",
    "        seed=420\n",
    "    )\n",
    "\n",
    "    audio = generate(generation_config)\n",
    "    audio.normalize_volume().play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env",
   "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.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
