{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:51:34.766203Z",
     "start_time": "2024-04-20T21:51:33.143327Z"
    }
   },
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\"\n",
    "\n",
    "from suno_utils.fine.generation import FineConfig, Fine, load_model, generate\n",
    "\n",
    "import torch\n",
    "\n",
    "device = \"cuda\"  # examples: 'cpu', 'cuda', 'cuda:0', 'cuda:1', etc."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:51:36.606778Z",
     "start_time": "2024-04-20T21:51:34.768208Z"
    }
   },
   "outputs": [],
   "source": [
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-02_00-21-37/best_ckpt.pt\"  # 100M\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-03_01-37-50/best_ckpt.pt\"  # 300M\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-12_17-46-09/best_ckpt.pt\"  # 800M\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-14_21-01-29/best_ckpt.pt\"  # 300M 5 delay\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-20_03-07-30/best_ckpt.pt\"  # 300M 1 delay\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-21_00-49-07/best_ckpt.pt\"  # remove last 2\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-21_00-56-34/best_ckpt.pt\"  # 0.05 noise\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-21_22-46-12/best_ckpt.pt\"  # 0.1 noise\n",
    "# ckpt_path = \"/app/suno/checkpoints/2023-12-23_15-38-52/best_ckpt.pt\"  # 0.1 noise 1s ctx\n",
    "ckpt_path = \"/app/suno/checkpoints/7b_fine/best_ckpt.pt\"  # 7b\n",
    "\n",
    "fine_model = load_model(ckpt_path, device)\n",
    "config = fine_model.config\n",
    "config"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:51:36.996636Z",
     "start_time": "2024-04-20T21:51:36.608990Z"
    }
   },
   "outputs": [],
   "source": [
    "from suno_utils.models.dac.model.dac2 import DAC\n",
    "import typing as tp\n",
    "\n",
    "\n",
    "class CoarseCodec(torch.nn.Module):\n",
    "    def __init__(self, codec_path: str):\n",
    "        super().__init__()\n",
    "\n",
    "        sd = torch.load(codec_path, map_location=\"cpu\")\n",
    "        model = DAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "        model.load_state_dict(sd[\"state_dict\"])\n",
    "        model.eval()\n",
    "        self.model = model\n",
    "\n",
    "    @torch.inference_mode()\n",
    "    def forward(\n",
    "        self,\n",
    "        audios: tp.Union[torch.Tensor, tp.List[torch.Tensor], tp.Tuple[torch.Tensor]],\n",
    "        device: tp.Union[torch.device, str],\n",
    "    ) -> tp.Tuple[torch.Tensor, torch.Tensor]:\n",
    "        if isinstance(audios, list) or isinstance(audios, tuple):\n",
    "            audios = torch.cat(audios, dim=0)  # (B, C, T)\n",
    "\n",
    "        audios = audios.to(device)\n",
    "        if len(audios.shape) == 2:\n",
    "            audios = audios.unsqueeze(0)\n",
    "        encoded = self.model(audios, 48000)\n",
    "        return encoded[\"codes\"], encoded[\"z\"], encoded[\"audio\"]\n",
    "\n",
    "\n",
    "# coarse_codec = CoarseCodec(\"/home/victor/data/models/chirp_v2/dac_2c_25x8.pt\").cuda()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:51:38.695349Z",
     "start_time": "2024-04-20T21:51:36.999079Z"
    }
   },
   "outputs": [],
   "source": [
    "from suno_utils.models.dac.model.dac2_12 import DAC as DAC12\n",
    "\n",
    "\n",
    "class FineCodec(torch.nn.Module):\n",
    "    def __init__(self, codec_path: str):\n",
    "        super().__init__()\n",
    "\n",
    "        sd = torch.load(codec_path, map_location=\"cpu\")\n",
    "        model = DAC12(**sd[\"metadata\"][\"kwargs\"])\n",
    "        model.load_state_dict(sd[\"state_dict\"])\n",
    "        model.eval()\n",
    "        self.model = model\n",
    "\n",
    "    @torch.inference_mode()\n",
    "    def forward(\n",
    "        self,\n",
    "        audios: tp.Union[torch.Tensor, tp.List[torch.Tensor], tp.Tuple[torch.Tensor]],\n",
    "        device: tp.Union[torch.device, str],\n",
    "    ) -> tp.Tuple[torch.Tensor, torch.Tensor]:\n",
    "        if isinstance(audios, list) or isinstance(audios, tuple):\n",
    "            audios = torch.cat(audios, dim=0)  # (B, C, T)\n",
    "\n",
    "        audios = audios.to(device)\n",
    "        if len(audios.shape) == 2:\n",
    "            audios = audios.unsqueeze(0)\n",
    "        encoded = self.model(audios, 48000)\n",
    "        return encoded[\"codes\"], encoded[\"z\"], encoded[\"audio\"]\n",
    "\n",
    "    @torch.inference_mode()\n",
    "    def decode(self, codes):\n",
    "        # zero out padding for now\n",
    "        codes = torch.where(codes < self.model.quantizer.codebook_size, codes, 0)\n",
    "        latents = self.model.quantizer.from_codes(codes)[0]\n",
    "        return self.model.decode(latents)\n",
    "\n",
    "\n",
    "coarse_codec = FineCodec(\"/home/victor/data/models/chirp_v2/dac_2c_25x12.pt\").cuda()\n",
    "fine_codec = FineCodec(\"/home/victor/data/models/dac_100x16.pth\").cuda()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:51:38.830269Z",
     "start_time": "2024-04-20T21:51:38.697742Z"
    }
   },
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "import os\n",
    "import torch.nn.functional as F\n",
    "from tqdm import tqdm\n",
    "\n",
    "\n",
    "def make_coarse_from_clip(audio: Audio, remove_last=0):\n",
    "    arr = torch.from_numpy(audio.array_float).unsqueeze(0)\n",
    "    x, _, aud = coarse_codec(arr, \"cuda\")\n",
    "    return pad_coarse_codes(x, remove_last=remove_last)\n",
    "\n",
    "\n",
    "def pad_coarse_codes(x, remove_last=0):\n",
    "    if remove_last > 0:\n",
    "        # replace with padding\n",
    "        x[:, -remove_last:] = config.coarse_pad_token\n",
    "    # pad to length + 1 for inference token\n",
    "    x = F.pad(\n",
    "        x,\n",
    "        (0, config.coarse_samples - x.shape[-1] + 1),\n",
    "        \"constant\",\n",
    "        config.coarse_pad_token,\n",
    "    )\n",
    "\n",
    "    # add fine streams\n",
    "    x = F.pad(\n",
    "        x,\n",
    "        (0, 0, 0, config.fine_n_codebooks),\n",
    "        \"constant\",\n",
    "        config.fine_pad_token,\n",
    "    )\n",
    "    # set inference token\n",
    "    x[:, config.coarse_n_codebooks :, -1] = config.fine_infer_token\n",
    "    return x.squeeze(0)\n",
    "\n",
    "\n",
    "def make_coarse_batch_from_audio(audio: Audio, remove_last=0):\n",
    "    coarse_codes = []\n",
    "    ctx_s = config.coarse_samples // config.coarse_rate_hz\n",
    "    for i in range(0, int(audio.duration_s), ctx_s):\n",
    "        clip = audio.get_slice(i, i + ctx_s)\n",
    "        coarse_codes.append(make_coarse_from_clip(clip, remove_last=remove_last))\n",
    "    coarse_codes = torch.stack(coarse_codes, dim=0)\n",
    "    return coarse_codes\n",
    "\n",
    "\n",
    "def cycle_clip(audio: Audio, remove_last=0):\n",
    "    # split audio into 20s clips\n",
    "    coarse_codes = make_coarse_batch_from_audio(audio, remove_last=remove_last)\n",
    "    print(f\"batch size: {coarse_codes.shape[0]}\")\n",
    "    fine_codes = generate(\n",
    "        fine_model, coarse_codes, config.t_fine - 1, temperature=0.8, top_k=500\n",
    "    )\n",
    "\n",
    "    fine_codes = torch.transpose(fine_codes, 0, 1)\n",
    "    fine_codes = fine_codes.reshape(1, config.fine_n_codebooks, -1)\n",
    "\n",
    "    decoded = fine_codec.decode(fine_codes.cuda())\n",
    "\n",
    "    decoded_audio = Audio.from_array_float(decoded.squeeze(0).cpu().numpy(), 48000)\n",
    "\n",
    "    return decoded_audio"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:51:56.874588Z",
     "start_time": "2024-04-20T21:51:38.831996Z"
    }
   },
   "outputs": [],
   "source": [
    "# generate golden samples\n",
    "sample_dir = \"/home/victor/glockenspiel/descript-audio-codec/data/golden/\"\n",
    "# sample_dir = \"/home/victor/audio_samples/\"\n",
    "for sample_path in os.listdir(sample_dir):\n",
    "    print(f\"Sample: {sample_path}, original\")\n",
    "    audio = Audio.from_file(\n",
    "        os.path.join(\n",
    "            sample_dir,\n",
    "            sample_path,\n",
    "        ),\n",
    "        sample_rate=48000,\n",
    "        n_channels=2,\n",
    "    )\n",
    "    # audio = audio.get_slice(0, 4)\n",
    "    # audio.play()\n",
    "\n",
    "    print(f\"Sample: {sample_path}, reconstructed\")\n",
    "    decoded_audio = cycle_clip(audio, remove_last=0)\n",
    "    print(\"Decoded\")\n",
    "    decoded_audio.play()\n",
    "    stem = os.path.splitext(sample_path)[0]\n",
    "    decoded_audio.write_mp3(f\"audios/{stem}_decoded.mp3\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:52:10.469221Z",
     "start_time": "2024-04-20T21:51:56.876754Z"
    }
   },
   "outputs": [],
   "source": [
    "from suno_utils.utils.s3 import _download_s3_file\n",
    "from tempfile import NamedTemporaryFile\n",
    "import numpy as np\n",
    "\n",
    "# upsample a chirp\n",
    "\n",
    "\n",
    "def upsample_chirp(id: str, play_original=True, max_duration=120):\n",
    "    print(id)\n",
    "    with NamedTemporaryFile(suffix=\".npz\") as f:\n",
    "        _download_s3_file(\n",
    "            f\"s3://suno-data-uploads/studio/uploads/{id}.npz\",\n",
    "            f.name,\n",
    "        )\n",
    "        print(f\"keys: {np.load(f.name).files}\")\n",
    "        npf = np.load(f.name)\n",
    "        if \"v3.0_raw\" in npf:\n",
    "            array = npf[\"v3.0_raw\"]\n",
    "        else:\n",
    "            array = npf[\"v1_raw\"]\n",
    "        array = torch.from_numpy(array).long().to(\"cuda\").T\n",
    "        print(array.shape)\n",
    "\n",
    "    if play_original:\n",
    "        original_audio = Audio.from_s3(\n",
    "            f\"s3://suno-data-uploads/studio/uploads/{id}.mp3\"\n",
    "        )\n",
    "        original_audio.get_slice(0, max_duration).play(compress=False)\n",
    "    original_duration = original_audio.duration_s\n",
    "    # make coarse batch by splitting into 2s clips\n",
    "    coarse_codes = []\n",
    "    for i in range(0, array.size(-1), config.coarse_samples):\n",
    "        code = array[1:, i : i + config.coarse_samples]\n",
    "        coarse_codes.append(pad_coarse_codes(code.unsqueeze(0)))\n",
    "\n",
    "    coarse_codes = torch.stack(coarse_codes, dim=0)\n",
    "\n",
    "    print(coarse_codes.shape)\n",
    "    print(coarse_codes.max())\n",
    "\n",
    "    fine_codes = generate(\n",
    "        fine_model, coarse_codes, config.t_fine - 1, temperature=0.8, top_k=5\n",
    "    )\n",
    "    fine_codes = torch.transpose(fine_codes, 0, 1)\n",
    "    fine_codes = fine_codes.reshape(1, config.fine_n_codebooks, -1)\n",
    "\n",
    "    audio = fine_codec.decode(fine_codes.cuda())\n",
    "    audio = Audio.from_array_float(audio.squeeze(0).cpu().numpy(), 48000)\n",
    "    audio.play()\n",
    "    \n",
    "    arr = torch.from_numpy(audio.array_float).unsqueeze(0)\n",
    "    coarse_codes, _, aud = coarse_codec(arr, \"cuda\")\n",
    "    print(coarse_codes.shape)\n",
    "    recode_audio = coarse_codec.decode(coarse_codes)\n",
    "    print(\"recode audio\")\n",
    "    recode_audio = Audio.from_array_float(recode_audio.squeeze(0).cpu().numpy(), 48000)\n",
    "    recode_audio.play()\n",
    "    return audio.get_segment(to_s=original_duration)\n",
    "\n",
    "\n",
    "# v2\n",
    "# upsample_chirp(\"00022ad0-fbea-48b6-a8e6-353156322526\").play()\n",
    "# upsample_chirp(\"fa7482fa-0b35-4dfb-a882-3cb05c1a1386\").play()\n",
    "\n",
    "# v3\n",
    "# upsample_chirp(\"1acfc44f-44bb-4869-9432-42f220d5559a\").play()\n",
    "# upsample_chirp(\"2b24ac12-08f0-489b-92a0-666bb6a73c3b\").play()\n",
    "# upsample_chirp(\"7c997c6b-c458-4893-9124-6acb6ab95db6\").play()\n",
    "upsample_chirp(\"c6d5b0ed-2d04-4566-8b03-6fecac2ea353\").play(compress=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:52:21.162047Z",
     "start_time": "2024-04-20T21:52:10.470824Z"
    }
   },
   "outputs": [],
   "source": [
    "upsample_chirp(\"c873959a-5ea2-450a-b87f-1a3e6ac2ec2a\").play(compress=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "ExecuteTime": {
     "end_time": "2024-04-20T21:52:37.600153Z",
     "start_time": "2024-04-20T21:52:21.164124Z"
    }
   },
   "outputs": [],
   "source": [
    "upsample_chirp(\"a2ca0dd5-6c7c-45b7-90ce-c4d10f652786\").play(compress=False) #.to_wav(\"test.wav\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3 (ipykernel)",
   "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.14"
  },
  "toc": {
   "base_numbering": 1,
   "nav_menu": {},
   "number_sections": true,
   "sideBar": true,
   "skip_h1_title": false,
   "title_cell": "Table of Contents",
   "title_sidebar": "Contents",
   "toc_cell": false,
   "toc_position": {},
   "toc_section_display": true,
   "toc_window_display": false
  }
 },
 "nbformat": 4,
 "nbformat_minor": 4
}
