{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 10,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"4\"\n",
    "import torch\n",
    "import torch.nn as nn\n",
    "import math\n",
    "\n",
    "def mdct(x):\n",
    "    N = x.shape[-1]\n",
    "    n = torch.arange(N, device=x.device)\n",
    "    k = torch.arange(N // 2, device=x.device)\n",
    "\n",
    "    arg = (\n",
    "        (math.pi / (2 * N))\n",
    "        * ((2 * n + 1 + N // 2).view(-1, 1))\n",
    "        * ((2 * k + 1).view(1, -1))\n",
    "    )\n",
    "    mdct_matrix = torch.cos(arg) * (2.0 / N) ** 0.5\n",
    "\n",
    "    return torch.matmul(x, mdct_matrix)\n",
    "\n",
    "\n",
    "def imdct(X):\n",
    "    half_N = X.shape[-1]\n",
    "    N = half_N * 2\n",
    "    n = torch.arange(N, device=X.device)\n",
    "    k = torch.arange(half_N, device=X.device)\n",
    "\n",
    "    arg = (\n",
    "        (math.pi / (2 * N))\n",
    "        * ((2 * n + 1 + N // 2).view(-1, 1))\n",
    "        * ((2 * k + 1).view(1, -1))\n",
    "    )\n",
    "    imdct_matrix = torch.cos(arg) * (2.0 / N) ** 0.5\n",
    "\n",
    "    return torch.matmul(X, imdct_matrix.T) * 2.0\n",
    "\n",
    "\n",
    "def audio_to_mdct_frames(audio, frame_size=1920, midside=False):\n",
    "    batch, channels, samples = audio.shape\n",
    "    hop_size = frame_size // 2\n",
    "\n",
    "    if midside:\n",
    "        mid = (audio[:, 0, :] + audio[:, 1, :]) / 2.0\n",
    "        side = (audio[:, 0, :] - audio[:, 1, :]) / 2.0\n",
    "        audio = torch.stack([mid, side], dim=1)\n",
    "\n",
    "    n_frames = (samples - frame_size) // hop_size + 1\n",
    "\n",
    "    frames = []\n",
    "    for i in range(n_frames):\n",
    "        start = i * hop_size\n",
    "        frame = audio[:, :, start : start + frame_size]\n",
    "        if frame.shape[-1] == frame_size:\n",
    "            frames.append(frame)\n",
    "\n",
    "    frames = torch.stack(frames, dim=2)\n",
    "\n",
    "    window = torch.sin(\n",
    "        torch.pi / frame_size * (torch.arange(frame_size, device=audio.device) + 0.5)\n",
    "    )\n",
    "    frames = frames * window.view(1, 1, 1, -1)\n",
    "\n",
    "    shape = frames.shape\n",
    "    frames_reshaped = frames.reshape(-1, frame_size)\n",
    "    mdct_coeffs = mdct(frames_reshaped)\n",
    "\n",
    "    return mdct_coeffs.reshape(shape[0], shape[1], shape[2], -1)\n",
    "\n",
    "\n",
    "def mdct_frames_to_audio(mdct_coeffs, frame_size=1920, midside=False):\n",
    "    batch, channels, n_frames, half_frame_size = mdct_coeffs.shape\n",
    "    hop_size = frame_size // 2\n",
    "\n",
    "    # Reshape and apply IMDCT\n",
    "    shape = mdct_coeffs.shape\n",
    "    coeffs_reshaped = mdct_coeffs.reshape(-1, half_frame_size)\n",
    "    frames = imdct(coeffs_reshaped)\n",
    "    frames = frames.reshape(shape[0], shape[1], shape[2], -1)\n",
    "\n",
    "    # Apply window\n",
    "    window = torch.sin(\n",
    "        torch.pi\n",
    "        / frame_size\n",
    "        * (torch.arange(frame_size, device=mdct_coeffs.device) + 0.5)\n",
    "    )\n",
    "    frames = frames * window.view(1, 1, 1, -1)\n",
    "\n",
    "    # Calculate total samples and create output shape\n",
    "    total_samples = (n_frames - 1) * hop_size + frame_size\n",
    "\n",
    "    # Reshape frames to prepare for folding\n",
    "    frames = frames.permute(0, 1, 3, 2)  # [batch, channels, frame_size, n_frames]\n",
    "    frames = frames.reshape(batch * channels, frame_size, n_frames)\n",
    "\n",
    "    # Use fold operation to overlap-add frames\n",
    "    output = torch.nn.functional.fold(\n",
    "        frames,\n",
    "        output_size=(1, total_samples),\n",
    "        kernel_size=(1, frame_size),\n",
    "        stride=(1, hop_size),\n",
    "    )\n",
    "\n",
    "    # Reshape output to expected dimensions\n",
    "    output = output.view(batch, channels, total_samples)\n",
    "\n",
    "    if midside:\n",
    "        mid = output[:, 0, :]\n",
    "        side = output[:, 1, :]\n",
    "        left = mid + side\n",
    "        right = mid - side\n",
    "        output = torch.stack([left, right], dim=1)\n",
    "\n",
    "    return output\n",
    "\n",
    "class WaveformMDCTVAE(nn.Module):\n",
    "    def __init__(\n",
    "        self,\n",
    "        frame_size: int = 1920,\n",
    "        latent_dim: int = 256,\n",
    "        hidden_dims: list = None,\n",
    "        dropout: float = 0.1,\n",
    "        midside: bool = False,\n",
    "    ):\n",
    "        super().__init__()\n",
    "\n",
    "        self.frame_size = frame_size\n",
    "        self.n_coeffs = frame_size // 2\n",
    "        self.latent_dim = latent_dim\n",
    "        self.midside = midside\n",
    "\n",
    "        if hidden_dims is None:\n",
    "            hidden_dims = [512, 256]\n",
    "\n",
    "        # Encoder layers\n",
    "        modules = []\n",
    "        input_dim = 2 * self.n_coeffs\n",
    "\n",
    "        for h_dim in hidden_dims:\n",
    "            modules.append(\n",
    "                nn.Sequential(\n",
    "                    nn.Linear(input_dim, h_dim),\n",
    "                    # nn.LayerNorm(h_dim),\n",
    "                    nn.LeakyReLU(),\n",
    "                    nn.Dropout(dropout),\n",
    "                    nn.Linear(h_dim, h_dim),\n",
    "                )\n",
    "            )\n",
    "            input_dim = h_dim\n",
    "\n",
    "        self.encoder = nn.Sequential(*modules)\n",
    "        self.fc_mu = nn.Linear(hidden_dims[-1], latent_dim)\n",
    "        self.fc_var = nn.Linear(hidden_dims[-1], latent_dim)\n",
    "\n",
    "        # Decoder layers\n",
    "        modules = []\n",
    "        hidden_dims.reverse()\n",
    "\n",
    "        self.decoder_input = nn.Sequential(\n",
    "            nn.Linear(latent_dim, hidden_dims[0]),\n",
    "            nn.LayerNorm(hidden_dims[0]),\n",
    "            nn.LeakyReLU(),\n",
    "            nn.Dropout(dropout),\n",
    "        )\n",
    "\n",
    "        for i in range(len(hidden_dims) - 1):\n",
    "            modules.append(\n",
    "                nn.Sequential(\n",
    "                    nn.Linear(hidden_dims[i], hidden_dims[i + 1]),\n",
    "                    # nn.LayerNorm(hidden_dims[i + 1]),\n",
    "                    nn.LeakyReLU(),\n",
    "                    nn.Dropout(dropout),\n",
    "                    nn.Linear(hidden_dims[i + 1], hidden_dims[i + 1]),\n",
    "                )\n",
    "            )\n",
    "\n",
    "        self.decoder = nn.Sequential(*modules)\n",
    "        self.final_layer = nn.Linear(hidden_dims[-1], 2 * self.n_coeffs)\n",
    "\n",
    "    def _encode(self, mdct_frames: torch.Tensor) -> list[torch.Tensor]:\n",
    "        batch_size, _, n_frames, _ = mdct_frames.shape\n",
    "        ch1 = mdct_frames[:, 0, :, :]\n",
    "        ch2 = mdct_frames[:, 1, :, :]\n",
    "\n",
    "        x = torch.cat((ch1, ch2), dim=-1)\n",
    "        result = self.encoder(x)\n",
    "        mu = self.fc_mu(result)\n",
    "        log_var = self.fc_var(result)\n",
    "\n",
    "        return [mu, log_var]\n",
    "\n",
    "    def _decode(self, z: torch.Tensor) -> torch.Tensor:\n",
    "        batch_size, n_frames, _ = z.shape\n",
    "\n",
    "        result = self.decoder_input(z)\n",
    "        result = self.decoder(result)\n",
    "        result = self.final_layer(result)\n",
    "        ch1 = result[..., : self.n_coeffs]\n",
    "        ch2 = result[..., self.n_coeffs :]\n",
    "        result = torch.stack((ch1, ch2), dim=1)\n",
    "\n",
    "        return result\n",
    "\n",
    "    def reparameterize(self, mu: torch.Tensor, log_var: torch.Tensor) -> torch.Tensor:\n",
    "        if self.training:\n",
    "            std = torch.exp(0.5 * log_var)\n",
    "            eps = torch.randn_like(std)\n",
    "            return eps * std + mu\n",
    "        else:\n",
    "            return mu\n",
    "\n",
    "    def forward(\n",
    "        self, waveform: torch.Tensor\n",
    "    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n",
    "        \"\"\"\n",
    "        Forward pass handling both waveform conversion and VAE operations.\n",
    "\n",
    "        Args:\n",
    "            waveform: Input audio tensor of shape (batch_size, channels, samples)\n",
    "\n",
    "        Returns:\n",
    "            tuple containing:\n",
    "            - reconstructed waveform\n",
    "            - reconstructed MDCT coefficients\n",
    "            - original MDCT coefficients (for loss computation)\n",
    "            - mu\n",
    "            - log_var\n",
    "        \"\"\"\n",
    "        # Convert input waveform to MDCT frames\n",
    "        mdct_frames = audio_to_mdct_frames(\n",
    "            waveform, frame_size=self.frame_size, midside=self.midside\n",
    "        )  # (batch, channels, n_frames, n_coeffs)\n",
    "\n",
    "        # Encode and decode\n",
    "        mu, log_var = self._encode(mdct_frames)\n",
    "        z = self.reparameterize(mu, log_var)\n",
    "        mdct_recon = self._decode(z)\n",
    "\n",
    "        # Convert back to waveform\n",
    "        waveform_recon = mdct_frames_to_audio(\n",
    "            mdct_recon, frame_size=self.frame_size, midside=self.midside\n",
    "        )\n",
    "\n",
    "        return waveform_recon, mdct_recon, mdct_frames, mu, log_var\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "model_filepath = \"/app/suno/christian/checkpoints/mdct-codec/2024-12-11_16-31-48_s4310/last_ckpt.pt\"\n",
    "\n",
    "ckpt = torch.load(model_filepath)\n",
    "print(ckpt[\"run_config\"][\"model\"])\n",
    "model = WaveformMDCTVAE(**ckpt[\"run_config\"][\"model\"])\n",
    "state_dict = ckpt[\"model\"]\n",
    "#new_state_dict = {}\n",
    "#for key, value in state_dict.items():\n",
    "#    new_key = key.replace(\"_orig_mod.module.\", \"\") \n",
    "#    new_state_dict[new_key] = value\n",
    "model.load_state_dict(state_dict)\n",
    "model.eval()\n",
    "model.cuda()\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import IPython\n",
    "import torchaudio\n",
    "test_filepath = \"/home/christian/audio/50_genre_songs/Miles Davis - Freddie Freeloader (Official Audio).mp3\"\n",
    "waveform, sr = torchaudio.load(test_filepath)\n",
    "\n",
    "if sr != 48000:\n",
    "    waveform = torchaudio.transforms.Resample(sr, 48000)(waveform)\n",
    "\n",
    "waveform = waveform[:, :48000*30].unsqueeze(0)\n",
    "\n",
    "with torch.no_grad():\n",
    "    waveform_recon, mdct_recon, mdct_frames, mu, log_var = model(waveform.cuda())\n",
    "\n",
    "    print(mdct_frames[0, 0, :, :])\n",
    "    print(mdct_recon[0, 0, :, :])\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(waveform[0].cpu().numpy(), rate=48000))\n",
    "IPython.display.display(IPython.display.Audio(waveform_recon[0].cpu().numpy(), rate=48000))\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.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
