{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import random\n",
    "import matplotlib.pyplot as plt\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "seg_len = 750\n",
    "embed_dim = 128\n",
    "infill_ctx_vae_embeds = torch.randn(seg_len, embed_dim)\n",
    "empty_ctx_vae_embeds = torch.randn(seg_len, embed_dim)\n",
    "infill_ctx_mask = torch.ones(seg_len)\n",
    "empty_ctx_mask = torch.zeros(seg_len)\n",
    "\n",
    "pattern = 1\n",
    "seq_len = infill_ctx_vae_embeds.shape[1]\n",
    "\n",
    "# This will be the same as previous\n",
    "# In this implementation:\n",
    "# mask value of 1 (True) means keep this token (provide context)\n",
    "# mask value of 0 (False) means mask this token (model will infill)\n",
    "\n",
    "infill_ctx_vae_embeds = infill_ctx_vae_embeds.permute(1, 0)  # channels, seq_len\n",
    "infill_ctx_mask = torch.ones(infill_ctx_vae_embeds.shape[1]).bool()  # Start with all True (all context)\n",
    "\n",
    "empty_ctx_vae_embeds = empty_ctx_vae_embeds.permute(1, 0)  # channels, seq_len\n",
    "\n",
    "# Randomly choose one of four masking patterns:\n",
    "# 0. Left masked - keep right side\n",
    "# 1. Right masked - keep left side\n",
    "# 2. Middle masked - keep edges\n",
    "# 3. Edges masked - keep middle\n",
    "pattern = random.randint(0, 3)\n",
    "seq_len = infill_ctx_vae_embeds.shape[1]\n",
    "\n",
    "if pattern == 0:\n",
    "    # Left masked - keep right side\n",
    "    n = random.randint(1, seq_len - 1)\n",
    "    infill_ctx_mask[:n] = False  # Set left portion to False (masked)\n",
    "    infill_ctx_vae_embeds[..., :n] = empty_ctx_vae_embeds[..., :n]\n",
    "\n",
    "elif pattern == 1:\n",
    "    # Right masked - keep left side\n",
    "    n = random.randint(1, seq_len - 1)\n",
    "    infill_ctx_mask[-n:] = False  # Set right portion to False (masked)\n",
    "    infill_ctx_vae_embeds[..., -n:] = empty_ctx_vae_embeds[..., -n:]\n",
    "\n",
    "elif pattern == 2:\n",
    "    # Middle masked - keep edges\n",
    "    # Ensure there's at least one token on each edge\n",
    "    middle_start = random.randint(1, seq_len // 2)\n",
    "    middle_end = random.randint(middle_start + 1, seq_len - 1)\n",
    "    \n",
    "    infill_ctx_mask[middle_start:middle_end] = False  # Set middle to False (masked)\n",
    "    infill_ctx_vae_embeds[..., middle_start:middle_end] = empty_ctx_vae_embeds[..., middle_start:middle_end]\n",
    "\n",
    "else:  # pattern == 3\n",
    "    # Edges masked - keep middle\n",
    "    left_size = random.randint(1, seq_len // 3)  # Left masked region\n",
    "    right_start = random.randint(seq_len // 2, seq_len - 1)  # Start of right masked region\n",
    "    \n",
    "    # Mask left edge\n",
    "    infill_ctx_mask[:left_size] = False\n",
    "    infill_ctx_vae_embeds[..., :left_size] = empty_ctx_vae_embeds[..., :left_size]\n",
    "    \n",
    "    # Mask right edge\n",
    "    infill_ctx_mask[right_start:] = False\n",
    "    infill_ctx_vae_embeds[..., right_start:] = empty_ctx_vae_embeds[..., right_start:]\n",
    "\n",
    "plt.plot(infill_ctx_mask.squeeze().cpu().numpy())"
   ]
  },
  {
   "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
}
