{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "os.environ['CUDA_VISIBLE_DEVICES'] = \"2\"\n",
    "\n",
    "import math\n",
    "import torch\n",
    "import einsum\n",
    "import numpy as np\n",
    "from torch import nn, einsum\n",
    "import torch.optim as optim\n",
    "from tqdm import tqdm\n",
    "from torch.nn.functional import mse_loss\n",
    "\n",
    "from suno_utils.utils.text import read_jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Define the noise schedule and sampling loop\n",
    "def get_alphas_sigmas(t):\n",
    "    \"\"\"Returns the scaling factors for the clean image (alpha) and for the\n",
    "    noise (sigma), given a timestep.\"\"\"\n",
    "    return torch.cos(t * math.pi / 2), torch.sin(t * math.pi / 2)\n",
    "\n",
    "def alpha_sigma_to_t(alpha, sigma):\n",
    "    \"\"\"Returns a timestep, given the scaling factors for the clean image and for\n",
    "    the noise.\"\"\"\n",
    "    return torch.atan2(sigma, alpha) / math.pi * 2\n",
    "\n",
    "def t_to_alpha_sigma(t):\n",
    "    \"\"\"Returns the scaling factors for the clean image and for the noise, given\n",
    "    a timestep.\"\"\"\n",
    "    return torch.cos(t * math.pi / 2), torch.sin(t * math.pi / 2)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "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",
    "\n",
    "# Define a causal Transformer-based autoregressive model\n",
    "class SimpleARModel(nn.Module):\n",
    "    def __init__(self, input_dim, hidden_dim, output_dim, n_heads, n_layers):\n",
    "        super(SimpleARModel, self).__init__()\n",
    "        self.embedding = nn.Linear(input_dim, hidden_dim)\n",
    "        encoder_layer = nn.TransformerEncoderLayer(d_model=hidden_dim, nhead=n_heads)\n",
    "        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=n_layers)\n",
    "        self.linear = nn.Linear(hidden_dim, output_dim)\n",
    "        self.pos_embedding = ScaledSinusoidalEmbedding(hidden_dim)\n",
    "\n",
    "    def forward(self, x):\n",
    "        x = self.embedding(x) + self.pos_embedding(x)\n",
    "        seq_len = x.size(1)\n",
    "        mask = nn.Transformer.generate_square_subsequent_mask(seq_len).to(x.device)\n",
    "        x = self.transformer(x, mask=mask, is_causal=True)\n",
    "        return self.linear(x)\n",
    "\n",
    "\n",
    "\n",
    "# Define the diffusion model (small MLP)\n",
    "class DiffusionMLP(nn.Module):\n",
    "    def __init__(self, input_dim, hidden_dim):\n",
    "        super(DiffusionMLP, self).__init__()\n",
    "        self.mlp = nn.Sequential(\n",
    "            nn.Linear(input_dim, hidden_dim),\n",
    "            nn.ReLU(),\n",
    "            nn.Linear(hidden_dim, input_dim)\n",
    "        )\n",
    "\n",
    "    def forward(self, x, z): \n",
    "        # consider also condition on time step of diffusion process\n",
    "        return self.mlp(x + z)\n",
    "\n",
    "# Define the diffusion loss function\n",
    "def diffusion_loss(z: torch.Tensor, x: torch.Tensor, diffusion_model: nn.Module):\n",
    "    # Draw uniformly distributed continuous timesteps\n",
    "    #t = self.rng.draw(reals.shape[0])[:, 0].to(self.device)\n",
    "    t = torch.rand(x.size(0), device=x.device)\n",
    "\n",
    "    # Replace 1% of t with ones to ensure training on terminal SNR\n",
    "    t = torch.where(torch.rand_like(t) < 0.01, torch.ones_like(t), t)\n",
    "\n",
    "    # Calculate the noise schedule parameters for those timesteps\n",
    "    alphas, sigmas = get_alphas_sigmas(t)\n",
    "\n",
    "    # combine noise with inputs\n",
    "    noise = torch.randn_like(x)\n",
    "    noised_inputs = x * alphas + noise * sigmas\n",
    "    targets = noise * alphas - x * sigmas\n",
    "\n",
    "    # Calculate the predicted noise\n",
    "    noise_pred = diffusion_model(noised_inputs, z)\n",
    "\n",
    "    return mse_loss(noise_pred, targets)\n",
    "\n",
    "# Example dataset\n",
    "class SimpleDataset(torch.utils.data.Dataset):\n",
    "    def __init__(self, size, seq_length, feature_dim):\n",
    "        self.data = torch.randn(size, seq_length, feature_dim)\n",
    "\n",
    "    def __len__(self):\n",
    "        return len(self.data)\n",
    "\n",
    "    def __getitem__(self, idx):\n",
    "        return self.data[idx]\n",
    "    \n",
    "\n",
    "class VAEMemmapDataset(torch.utils.data.Dataset):\n",
    "    def __init__(\n",
    "        self,\n",
    "        vae_memmap_path: str,\n",
    "        infer_embed: torch.Tensor = None,\n",
    "        vae_metas_path: str = None,\n",
    "        vae_dim: int = 128,\n",
    "        n_tokens_memmap: int = 1000,\n",
    "    ):\n",
    "        \"\"\"For use in training unconditional diffusion model.\n",
    "\n",
    "        When a metas file is provided, the metadata is loaded and returned with the data.\n",
    "        This can be used for lyric conditioning, etc.\n",
    "\n",
    "        \"\"\"\n",
    "        super().__init__()\n",
    "        self.vae_metas_path = vae_metas_path\n",
    "        self.infer_embed = infer_embed\n",
    "        self.vae_dim = vae_dim\n",
    "        self.n_tokens_memmap = n_tokens_memmap\n",
    "\n",
    "        vae_data = np.memmap(vae_memmap_path, dtype=np.float32, mode=\"r\")\n",
    "        vae_data = vae_data.reshape(-1, vae_dim, n_tokens_memmap)\n",
    "        self.vae_data = vae_data\n",
    "        print(f\"Found {vae_data.shape[0]} examples.\")\n",
    "\n",
    "        if vae_metas_path is not None:\n",
    "            self.metas = read_jsonl(vae_metas_path)\n",
    "            assert len(self.metas) == self.vae_data.shape[0]\n",
    "            print(\"Loaded metadata for\", len(self.metas), \"examples.\")\n",
    "        else:\n",
    "            self.metas = None\n",
    "\n",
    "    def __len__(self):\n",
    "        return self.vae_data.shape[0]\n",
    "\n",
    "    def __getitem__(self, idx: int):\n",
    "        info = {}\n",
    "        vae_embeds = torch.from_numpy(self.vae_data[idx, ...].copy()).float()\n",
    "        info[\"idx\"] = idx\n",
    "        info[\"seconds_start\"] = 0\n",
    "        info[\"seconds_total\"] = 10.0\n",
    "\n",
    "        if self.metas is not None:\n",
    "            info[\"lyrics\"] = self.metas[idx][\"lyrics\"]\n",
    "\n",
    "        # add infer embed to the first item in the sequence\n",
    "        if self.infer_embed is not None:\n",
    "            vae_embeds = torch.cat([self.infer_embed.unsqueeze(1), vae_embeds], dim=-1)\n",
    "\n",
    "        return (vae_embeds.permute(1, 0), info)\n",
    "\n",
    "# Training loop\n",
    "def train(model, diffusion_model, dataloader, optimizer, diffusion_optimizer, epochs):\n",
    "    model.train()\n",
    "    diffusion_model.train()\n",
    "    for epoch in range(epochs):\n",
    "        pbar = tqdm(dataloader)\n",
    "        for batch in pbar:\n",
    "            vae_embeds, metadata = batch\n",
    "            optimizer.zero_grad()\n",
    "            diffusion_optimizer.zero_grad()\n",
    "\n",
    "            # only use the first 10 tokens\n",
    "            vae_embeds = vae_embeds[:10, :]\n",
    "            vae_embeds = vae_embeds.cuda()\n",
    "\n",
    "            z = model(vae_embeds)  # Autoregressive model's output\n",
    "            loss = diffusion_loss(z, vae_embeds, diffusion_model)\n",
    "\n",
    "            loss.backward()\n",
    "            optimizer.step()\n",
    "            diffusion_optimizer.step()\n",
    "\n",
    "            pbar.set_description(f\"Epoch {epoch + 1}/{epochs}, Loss: {loss.item()}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "# Parameters\n",
    "vae_dim = 128 # this is the dimension of the VAE embeddings\n",
    "input_dim = vae_dim\n",
    "hidden_dim = 768\n",
    "output_dim = vae_dim\n",
    "n_heads = 4\n",
    "n_layers = 8\n",
    "batch_size = 128\n",
    "epochs = 10\n",
    "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",
    "\n",
    "infer_embed = torch.zeros(vae_dim)\n",
    "\n",
    "# Instantiate models, dataset, and dataloader\n",
    "model = SimpleARModel(input_dim, hidden_dim, output_dim, n_heads, n_layers)\n",
    "diffusion_model = DiffusionMLP(input_dim, hidden_dim)\n",
    "\n",
    "# print number of parameters in model and diffusion_model\n",
    "print(\"AR Model has\", sum(p.numel() for p in model.parameters()) / 1e6, \"M parameters\")\n",
    "print(\"Diffusion model has\", sum(p.numel() for p in diffusion_model.parameters()) / 1e6, \"M parameters\")\n",
    "\n",
    "dataset = VAEMemmapDataset(\n",
    "    \"/app/suno/christian/data/suno_diffusion_tiktok_covers_lyrics/vae_val.bin\", \n",
    "    infer_embed=infer_embed,\n",
    "    vae_metas_path=\"/app/suno/christian/data/suno_diffusion_tiktok_covers_lyrics/val_metas.jsonl\",\n",
    ")\n",
    "dataloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, num_workers=32)\n",
    "\n",
    "# Optimizers\n",
    "optimizer = optim.Adam(model.parameters(), lr=1e-4)\n",
    "diffusion_optimizer = optim.Adam(diffusion_model.parameters(), lr=1e-4)\n",
    "\n",
    "# move to GPU\n",
    "model = model.cuda()\n",
    "diffusion_model = diffusion_model.cuda()\n",
    "\n",
    "# Train the model\n",
    "train(model, diffusion_model, dataloader, optimizer, diffusion_optimizer, epochs)\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Inference"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# inference\n",
    "def reverse_diffusion(diffusion_model, z, num_steps=100):\n",
    "    \"\"\"Run the reverse diffusion process starting from a noisy embedding.\"\"\"\n",
    "    x_t = torch.randn_like(z)\n",
    "    for t in range(num_steps, 0, -1):\n",
    "        x_t = diffusion_step(diffusion_model, x_t, z, t)\n",
    "    return x_t\n",
    "\n",
    "def diffusion_step(diffusion_model, x_t, z, t):\n",
    "    \"\"\"Single step of the reverse diffusion process.\"\"\"\n",
    "    alpha_t = 0.5  # Example noise schedule\n",
    "    noise_pred = diffusion_model(x_t, z)\n",
    "    x_t = (x_t - noise_pred) / torch.sqrt(alpha_t)\n",
    "    return x_t\n",
    "\n",
    "def generate_sequence(model, diffusion_model, start_token, seq_len, num_diffusion_steps=100, device='cpu'):\n",
    "    model.eval()\n",
    "    diffusion_model.eval()\n",
    "    generated_sequence = [start_token]\n",
    "    \n",
    "    for _ in range(seq_len - 1):\n",
    "        current_sequence = torch.tensor(generated_sequence, dtype=torch.float32).unsqueeze(0).to(device)\n",
    "        \n",
    "        with torch.no_grad():\n",
    "            # Get the conditioning vector z from the AR model\n",
    "            z = model(current_sequence)\n",
    "            \n",
    "            # Run the reverse diffusion process to generate the next token embedding\n",
    "            next_token_embedding = reverse_diffusion(diffusion_model, z[:, -1, :], num_steps=num_diffusion_steps)\n",
    "        \n",
    "        generated_sequence.append(next_token_embedding.squeeze(0).tolist())\n",
    "    \n",
    "    return generated_sequence\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "generated_sequence = generate_sequence(model, diffusion_model, start_token, seq_len, device=device)\n"
   ]
  }
 ],
 "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
}
