{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\"\n",
    "import torch\n",
    "import torchaudio\n",
    "import numpy as np\n",
    "from torch import nn, einsum\n",
    "\n",
    "from suno_utils.utils.s3 import read_from_s3\n",
    "from suno_utils.tasks.dac_2c_12cb import DAC\n",
    "from suno_utils.models.musicfm.modeling_MusicFM import MusicFM_MERTLong\n",
    "from suno_utils.models.dac.nn.quantize_2 import ResidualVectorQuantize\n",
    "\n",
    "from suno_utils.tasks.dac_peaq100 import (\n",
    "    load_model as load_vae_model,\n",
    "    encode as vae_encode,\n",
    "    decode as vae_decode,\n",
    ")\n",
    "\n",
    "\n",
    "# MERT\n",
    "from suno_utils.tasks.mert_25 import (\n",
    "    preload_models as preload_semantic_models,\n",
    "    encode as semantic_encode,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load MERT semantic\n",
    "_ = preload_semantic_models(\n",
    "    checkpoint_filepath=\"s3://suno-data/georg/models/semantic/mert_25.pt\",\n",
    "    centroids_filepath=\"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\",\n",
    "    device=\"cuda\",\n",
    ")\n",
    "centroid_path = \"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\"\n",
    "\n",
    "centroids = read_from_s3(centroid_path, read_f=np.load)[0]\n",
    "embedding = torch.nn.Embedding(centroids.shape[0], centroids.shape[1])\n",
    "embedding.weight.data.copy_(torch.tensor(centroids, dtype=torch.float32))\n",
    "embedding.weight.requires_grad = False\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load some random audio \n",
    "audio, sr = torchaudio.load(\"/home/christian/audio/bad-audio/bill-evans-intro.wav\")\n",
    "\n",
    "# semantic encode to seq of codes\n",
    "codes = semantic_encode([audio.mean(dim=0, keepdim=True)])\n",
    "codes = torch.tensor(codes).long()\n",
    "print(codes.shape)\n",
    "\n",
    "# decode using pretrained embedding\n",
    "embeds = embedding(codes)\n",
    "\n",
    "print(\"embed\", embeds.shape, embeds.min(), embeds.max())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def load_rvq(codec_path: str, codec_input_dim: int, codec_n_codebooks: int, codec_codebook_size: int, codec_codebook_dim: int, codec_quantizer_dropout: float):\n",
    "    sd = read_from_s3(codec_path, read_f=torch.load)\n",
    "    model = ResidualVectorQuantize(\n",
    "        input_dim=codec_input_dim,\n",
    "        n_codebooks=codec_n_codebooks,\n",
    "        codebook_size=codec_codebook_size,\n",
    "        codebook_dim=codec_codebook_dim,\n",
    "        quantizer_dropout=codec_quantizer_dropout,\n",
    "    )\n",
    "\n",
    "    model.load_state_dict(\n",
    "        {k[10:]: v for k, v in sd[\"state_dict\"].items() if k.startswith(\"quantize\")}\n",
    "    )\n",
    "    model.eval()\n",
    "\n",
    "    for param_name, param in model.named_parameters():\n",
    "        param.requires_grad = False\n",
    "\n",
    "    return model\n",
    "\n",
    "def decode_vq(rvq, codes, n_quantizers):\n",
    "    z_q = 0\n",
    "    for i, quantizer in enumerate(rvq.quantizers[:n_quantizers]):\n",
    "        _z_q = quantizer.embed_code(codes[:, :, i]).transpose(1, 2)\n",
    "        _z_q = quantizer.out_proj(_z_q)\n",
    "        z_q += _z_q.transpose(1, 2)\n",
    "    return z_q"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "class ScaledSinusoidalEmbedding(nn.Module):\n",
    "    def __init__(self, dim, theta=10000):\n",
    "        super().__init__()\n",
    "        assert (dim % 2) == 0, \"dimension must be divisible by 2\"\n",
    "        self.scale = nn.Parameter(torch.ones(1) * dim**-0.5)\n",
    "\n",
    "        half_dim = dim // 2\n",
    "        freq_seq = torch.arange(half_dim).float() / half_dim\n",
    "        inv_freq = theta**-freq_seq\n",
    "        self.register_buffer(\"inv_freq\", inv_freq, persistent=False)\n",
    "\n",
    "    def forward(self, x, pos=None, seq_start_pos=None):\n",
    "        seq_len, device = x.shape[1], x.device\n",
    "\n",
    "        if pos is None:\n",
    "            pos = torch.arange(seq_len, device=device)\n",
    "\n",
    "        if seq_start_pos is not None:\n",
    "            pos = pos - seq_start_pos[..., None]\n",
    "\n",
    "        emb = einsum(\"i, j -> i j\", pos, self.inv_freq)\n",
    "        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n",
    "        return emb * self.scale\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "pos_embedding = ScaledSinusoidalEmbedding(128)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "model = load_vae_model(\n",
    "    checkpoint_filepath=\"s3://suno-data/minz/models/dac_100hz_12cb.pth\"\n",
    ")\n",
    "\n",
    "codec_input_dim = 128 \n",
    "codec_n_codebooks = 12\n",
    "codec_codebook_size = 32768\n",
    "codec_codebook_dim = 8\n",
    "codec_quantizer_dropout = 0.0\n",
    "\n",
    "rvq = load_rvq(\"s3://suno-data/minz/models/dac_100hz_12cb.pth\", codec_input_dim, codec_n_codebooks, codec_codebook_size, codec_codebook_dim, codec_quantizer_dropout)\n",
    "\n",
    "input_memmap_path = \"/app/suno/christian/data/enhance/input_tr.bin\"\n",
    "corrupt_memmap_path = \"/app/suno/christian/data/enhance/corrupt_tr.bin\"\n",
    "\n",
    "t_memmap = 1000\n",
    "n_codebooks = 12    \n",
    "embed_dim = 128\n",
    "\n",
    "# load the input and corrupt memmaps\n",
    "input_data = np.memmap(input_memmap_path, dtype=np.uint16, mode=\"r\")\n",
    "input_data = input_data.reshape(-1, t_memmap, n_codebooks)\n",
    "\n",
    "corrupt_data = np.memmap(corrupt_memmap_path, dtype=np.uint16, mode=\"r\")\n",
    "corrupt_data = corrupt_data.reshape(-1, t_memmap, n_codebooks)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for rand_index in range(1):\n",
    "\n",
    "    input_ex = torch.from_numpy(input_data[rand_index, ...].copy()).long()\n",
    "    corrupt_ex = torch.from_numpy(corrupt_data[rand_index, ...].copy()).long()\n",
    "    print(\"input_ex\", input_ex.shape)\n",
    "\n",
    "    # decode vq to continuous embed\n",
    "    input_zq = decode_vq(rvq, input_ex.unsqueeze(0), 12)\n",
    "    corrupt_zq = decode_vq(rvq, corrupt_ex.unsqueeze(0), 12)\n",
    "    #input_zq /= 20.0\n",
    "    print(\"input_zq\", input_zq.shape, input_zq.min(), input_zq.max())\n",
    "\n",
    "    with torch.no_grad():\n",
    "        pos_embed = pos_embedding(input_zq)\n",
    "        print(\"pos_embed\", pos_embed.shape, pos_embed.min(), pos_embed.max())\n",
    "\n",
    "\n",
    "    embed = input_zq + pos_embed\n",
    "    print(\"embed\", embed.min(), embed.max())\n",
    "\n",
    "    print(\"zq\", input_zq.shape)\n",
    "\n",
    "# decode embed back to audio\n",
    "#input_audio = model.decode(input_zq.permute(0, 2, 1))[0].detach().cpu()\n",
    "#corrupt_audio = model.decode(corrupt_zq.permute(0, 2, 1))[0].detach().cpu()\n",
    "#print(\"audio\", input_audio.shape)\n",
    "\n"
   ]
  },
  {
   "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.14"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
