{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import torch\n",
    "import funcy\n",
    "import IPython\n",
    "import numpy as np\n",
    "\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"6\"\n",
    "\n",
    "from suno_utils.utils.text import (    \n",
    "    write_jsonl,\n",
    "    read_jsonl,\n",
    "    write_json,\n",
    "    read_json,\n",
    "    normalize_whitespace,\n",
    ")\n",
    "\n",
    "from dac.model.dac4 import DAC\n",
    "from suno_utils.utils.s3 import read_from_s3"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load VAE\n",
    "device = \"cuda:0\"\n",
    "#device = \"cpu\"\n",
    "# checkpoint_filepath = \"/app/suno/christian/checkpoints/dac/100hz_vae_peaq_kl_0.005/best/dac/weights.pth\"\n",
    "checkpoint_filepath = \"s3://suno-data/christian/100hz_vae_peaq_kl_0.005.pth\"\n",
    "load_f = funcy.partial(torch.load, map_location=\"cpu\")\n",
    "\n",
    "if checkpoint_filepath.startswith(\"s3://\"):\n",
    "    sd = read_from_s3(checkpoint_filepath, read_f=load_f)\n",
    "else:\n",
    "    sd = load_f(checkpoint_filepath)\n",
    "\n",
    "sd[\"metadata\"][\"kwargs\"] = {\n",
    "    k: v\n",
    "    for k, v in sd[\"metadata\"][\"kwargs\"].items()\n",
    "    if k in DAC.__init__.__code__.co_varnames\n",
    "}\n",
    "model_100hz = DAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "model_100hz.load_state_dict(sd[\"state_dict\"])\n",
    "model_100hz.eval()\n",
    "model_100hz.to(device)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load VAE\n",
    "device = \"cuda:0\"\n",
    "#device = \"cpu\"\n",
    "# checkpoint_filepath = \"/app/suno/christian/checkpoints/dac/100hz_vae_peaq_kl_0.005/best/dac/weights.pth\"\n",
    "checkpoint_filepath = \"s3://suno-data/christian/25hz_vae_peaq_kl_0.005.pth\"\n",
    "load_f = funcy.partial(torch.load, map_location=\"cpu\")\n",
    "\n",
    "if checkpoint_filepath.startswith(\"s3://\"):\n",
    "    sd = read_from_s3(checkpoint_filepath, read_f=load_f)\n",
    "else:\n",
    "    sd = load_f(checkpoint_filepath)\n",
    "\n",
    "sd[\"metadata\"][\"kwargs\"] = {\n",
    "    k: v\n",
    "    for k, v in sd[\"metadata\"][\"kwargs\"].items()\n",
    "    if k in DAC.__init__.__code__.co_varnames\n",
    "}\n",
    "model_25hz = DAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "model_25hz.load_state_dict(sd[\"state_dict\"])\n",
    "model_25hz.eval()\n",
    "model_25hz.to(device)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load VAE fixed\n",
    "device = \"cuda:0\"\n",
    "#device = \"cpu\"\n",
    "\n",
    "\n",
    "from suno_utils.models.dac.model.dac_vae_peaq import DAC\n",
    "\n",
    "\n",
    "#checkpoint_filepath = \"s3://suno-data/minz/models/dac_vae_fixed_25hz.pth\"\n",
    "checkpoint_filepath = \"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\"\n",
    "load_f = funcy.partial(torch.load, map_location=\"cpu\")\n",
    "\n",
    "if checkpoint_filepath.startswith(\"s3://\"):\n",
    "    sd = read_from_s3(checkpoint_filepath, read_f=load_f)\n",
    "else:\n",
    "    sd = load_f(checkpoint_filepath)\n",
    "\n",
    "sd[\"metadata\"][\"kwargs\"] = {\n",
    "    k: v\n",
    "    for k, v in sd[\"metadata\"][\"kwargs\"].items()\n",
    "    if k in DAC.__init__.__code__.co_varnames\n",
    "}\n",
    "model_25hz = DAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "model_25hz.load_state_dict(sd[\"state_dict\"])\n",
    "model_25hz.eval()\n",
    "model_25hz.to(device)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.dac_vae_fixed_25hz import decode, preload_models\n",
    "_ = preload_models(checkpoint_filepath=\"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "device = \"cpu\"\n",
    "\n",
    "checkpoint_filepath = \"s3://suno-data/georg/models/codec/dac_2c_25x12.pt\"\n",
    "load_f = funcy.partial(torch.load, map_location=\"cpu\")\n",
    "\n",
    "if checkpoint_filepath.startswith(\"s3://\"):\n",
    "    sd = read_from_s3(checkpoint_filepath, read_f=load_f)\n",
    "else:\n",
    "    sd = load_f(checkpoint_filepath)\n",
    "\n",
    "sd[\"metadata\"][\"kwargs\"] = {\n",
    "    k: v\n",
    "    for k, v in sd[\"metadata\"][\"kwargs\"].items()\n",
    "    if k in DAC.__init__.__code__.co_varnames\n",
    "}\n",
    "model_25hz_codec = DAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "model_25hz_codec.load_state_dict(sd[\"state_dict\"])\n",
    "model_25hz_codec.eval()\n",
    "model_25hz_codec.to(device)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas = read_jsonl(\"/app/suno/data/diffusion_mix/metadata_6min/metas.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas_phonemes = read_jsonl(\"/app/suno/data/diffusion_mix/metadata_6min/metas_phonemized.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas_phonemes[20000]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "lyric_metas = [meta for meta in metas_phonemes if \"text\" in meta]\n",
    "phoneme_metas = [meta for meta in metas_phonemes if meta[\"phonemized_text\"] is not None]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(lyric_metas), len(phoneme_metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(metas), len(metas_phonemes))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(metas_phonemes[1000][\"text\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "#base_dir = \"/app/suno/data/diffusion_mix/dac_vae_fixed_25hz\"\n",
    "base_dir = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz\"\n",
    "base_dir = \"/app2/suno/data/christian/outputs/v3-bootstrap-data-t4/memmaps/syn_sft_t2\"\n",
    "#base_dir = \"/app/suno/data/diffusion_v5/v0\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_val.jsonl\", progress=True)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "idx = 999\n",
    "meta = metas[idx]\n",
    "print(meta[\"tags\"])\n",
    "audio = decode(vae_data[idx])\n",
    "audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "from sklearn.decomposition import PCA\n",
    "import matplotlib.pyplot as plt\n",
    "from tqdm import tqdm\n",
    "\n",
    "# assume vae_latents is your (3363, 750, 128) array\n",
    "X = vae_data.reshape(-1, 128)  # flatten to (3363 * 750, 128)\n",
    "\n",
    "errors = []\n",
    "components_range = range(1, 5)\n",
    "\n",
    "for n in tqdm(components_range):\n",
    "    pca = PCA(n_components=n)\n",
    "    X_pca = pca.fit_transform(X)\n",
    "    X_recon = pca.inverse_transform(X_pca)\n",
    "    error = np.mean((X - X_recon) ** 2)\n",
    "    errors.append(error)\n",
    "\n",
    "plt.plot(components_range, errors)\n",
    "plt.xlabel('Number of PCA Components')\n",
    "plt.ylabel('Reconstruction MSE')\n",
    "plt.title('PCA Reconstruction Error')\n",
    "plt.grid()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "pca = PCA(n_components=128)\n",
    "pca.fit(X)\n",
    "explained = pca.explained_variance_ratio_\n",
    "cumulative = np.cumsum(explained)\n",
    "\n",
    "plt.plot(range(1, 129), cumulative)\n",
    "plt.xlabel('Number of PCA Components')\n",
    "plt.ylabel('Cumulative Explained Variance')\n",
    "plt.title('PCA Explained Variance')\n",
    "plt.grid()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Fit PCA on full data\n",
    "from sklearn.preprocessing import StandardScaler\n",
    "\n",
    "X = vae_data.reshape(-1, 128)\n",
    "scaler = StandardScaler()\n",
    "X_scaled = scaler.fit_transform(X)\n",
    "\n",
    "pca = PCA(n_components=128)\n",
    "pca.fit(X_scaled)\n",
    "\n",
    "# Find number of components to retain 95% variance\n",
    "cumulative = np.cumsum(pca.explained_variance_ratio_)\n",
    "n_components_95 = np.searchsorted(cumulative, 0.95) + 1\n",
    "print(n_components_95)\n",
    "\n",
    "# Fit PCA with just enough components\n",
    "pca_95 = PCA(n_components=n_components_95)\n",
    "X_pca = pca_95.fit_transform(X)\n",
    "\n",
    "# Project and reconstruct one example\n",
    "example = vae_data[100]  # shape (750, 128)\n",
    "example_scaled = scaler.transform(example)  # standardize\n",
    "\n",
    "example_pca = pca_95.transform(example_scaled)\n",
    "example_recon_scaled = pca_95.inverse_transform(example_pca)\n",
    "\n",
    "example_recon = scaler.inverse_transform(example_recon_scaled)  # back to original scale\n",
    "\n",
    "audio = decode(example)\n",
    "audio.play()\n",
    "\n",
    "audio = decode(example_recon)\n",
    "audio.play()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "\n",
    "base_dir = \"/app/suno/data/splice_samples_30s/dac_vae_tuned_25hz/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_tr.jsonl\", progress=True)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_tr.bin\"\n",
    "#semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "#semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "#semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)\n",
    "\n",
    "print(vae_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# collect all tags and count occurrences\n",
    "tag_counts = {}\n",
    "for meta in tqdm(metas):\n",
    "    for tag in meta[\"tags\"]:\n",
    "        if tag in tag_counts:\n",
    "            tag_counts[tag] += 1\n",
    "        else:\n",
    "            tag_counts[tag] = 1\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "tag_counts\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "idx = 5400\n",
    "print(metas[idx][\"tags\"])\n",
    "tags = \", \".join(metas[idx][\"tags\"])\n",
    "print(tags)\n",
    "#audio = decode(vae_data[idx])\n",
    "#audio.play()\n",
    "\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# duration histogram\n",
    "import matplotlib.pyplot as plt\n",
    "plt.hist([meta[\"original_duration_s\"] for meta in metas], bins=100)\n",
    "plt.show()\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# find the max of vae_data in each example\n",
    "num = vae_data.shape[0]\n",
    "new_max_vals = []\n",
    "for i in range(num):\n",
    "    max_val = np.max(np.abs(vae_data[i]))\n",
    "    new_max_vals.append(max_val)\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "base_dir = \"/app/suno/data/diffusion_mix/vae_25hz_30s/\"\n",
    "#base_dir = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz\"\n",
    "#base_dir = \"/app/suno/data/diffusion_v5/v0\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_val.jsonl\", progress=True)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "old_vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "np.max(vae_data[i])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# find the max of vae_data in each example\n",
    "num = vae_data.shape[0]\n",
    "old_max_vals = []\n",
    "for i in range(num):\n",
    "    max_val = np.max(np.abs(vae_data[i]))\n",
    "    old_max_vals.append(max_val)\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "base_dir = \"/app/suno/data/diffusion_mix/vae_100hz_30s/\"\n",
    "#base_dir = \"/app/suno/data/diffusion_mix/dac_vae_tuned_25hz\"\n",
    "#base_dir = \"/app/suno/data/diffusion_v5/v0\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_val.jsonl\", progress=True)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data_100hz = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "max_vals_100hz = []\n",
    "for i in range(vae_data_100hz.shape[0]):\n",
    "    max_val = np.max(np.abs(vae_data_100hz[i]))\n",
    "    max_vals_100hz.append(max_val)\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "plt.hist(old_max_vals, bins=np.linspace(0, 20, 250))\n",
    "plt.hist(new_max_vals, bins=np.linspace(0, 20, 250))\n",
    "plt.hist(max_vals_100hz, bins=np.linspace(0, 20, 250))\n",
    "plt.show()\n",
    "\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "index = 70005\n",
    "meta = metas[index]\n",
    "print(meta)\n",
    "text_aligned = meta.get(\"text_aligned\")\n",
    "print(text_aligned)\n",
    "\n",
    "audio = decode(vae_data[index])\n",
    "audio.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "base_dir = \"/app/suno/data/diffusion_mix/vae_25hz_30s\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_val.jsonl\", progress=False)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM).astype(np.float32)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS).astype(np.int16)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def fold_tensor(x: torch.Tensor, patch_size: int) -> torch.Tensor:\n",
    "    \"\"\"Fold first dimension into last dimension with given patch size.\"\"\"\n",
    "    batch, dim = x.shape\n",
    "    assert batch % patch_size == 0, f\"First dimension {batch} must be divisible by patch_size {patch_size}\"\n",
    "    \n",
    "    # Ensure contiguous memory layout before reshaping\n",
    "    x = x.contiguous()\n",
    "    # Reshape to (new_batch, patch_size, dim)\n",
    "    x = x.reshape(-1, patch_size, dim)\n",
    "    # Flatten patch dimension into feature dimension\n",
    "    return x.reshape(batch // patch_size, patch_size * dim)\n",
    "\n",
    "def unfold_tensor(x: torch.Tensor, original_dim: int) -> torch.Tensor:\n",
    "    \"\"\"Inverse operation of fold_tensor.\"\"\"\n",
    "    batch, dim = x.shape\n",
    "    patch_size = dim // original_dim\n",
    "    \n",
    "    # Ensure contiguous memory layout before reshaping\n",
    "    x = x.contiguous()\n",
    "    # Reshape back to (batch, patch_size, original_dim)\n",
    "    x = x.reshape(batch, patch_size, original_dim)\n",
    "    # Flatten first two dimensions\n",
    "    return x.reshape(batch * patch_size, original_dim)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 3000\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "base_dir = \"/app/suno/data/diffusion_mix/vae_100hz_30s\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_val.jsonl\", progress=False)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM).astype(np.float32)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS).astype(np.int16)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "rand_idx = np.random.randint(0, len(metas))\n",
    "print(rand_idx)\n",
    "print(metas[rand_idx])\n",
    "vae_seq = vae_data[rand_idx]\n",
    "vae_seq = torch.from_numpy(vae_seq.copy()).to(device).float()\n",
    "\n",
    "# vae_seq is (1, 3000, 128)\n",
    "vae_seq_folded = fold_tensor(vae_seq, 8)\n",
    "print(vae_seq_folded.shape)\n",
    "vae_seq_unfolded = unfold_tensor(vae_seq_folded, VAE_DIM)\n",
    "print(vae_seq_unfolded.shape)\n",
    "\n",
    "\n",
    "audio = model_100hz.decode(vae_seq_unfolded.unsqueeze(0).permute(0, 2, 1))[0].detach().cpu()         \n",
    "audio /= audio.abs().max().clamp(1e-8)\n",
    "print(audio.mean())\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(audio.numpy(), rate=48000))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# measure the mean and std for the latents\n",
    "print(\"mean\", np.mean(vae_data), \"std\", np.std(vae_data))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(0.47132844 * 2.0)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "memmap_metas_text = [meta for meta in metas if \"text\" in meta]\n",
    "memmap_metas_phonemes = [meta for meta in metas if meta[\"phonemized_text\"] is not None]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(memmap_metas_text), len(memmap_metas_phonemes), len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# measure mean and median of phoneme lengths\n",
    "phoneme_lengths = [len(meta[\"phonemized_text\"]) for meta in memmap_metas_phonemes]\n",
    "print(np.mean(phoneme_lengths), np.median(phoneme_lengths), np.max(phoneme_lengths))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "# make histogram of  n_vae_tokens \n",
    "vae_token_lengths = [meta[\"n_vae_tokens\"] for meta in metas]\n",
    "print(np.mean(vae_token_lengths), np.median(vae_token_lengths), np.max(vae_token_lengths), np.min(vae_token_lengths))\n",
    "plt.hist(vae_token_lengths, bins=50)\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# make a histogram of phoneme lengths\n",
    "import matplotlib.pyplot as plt\n",
    "plt.hist(phoneme_lengths, bins=50)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "rand_idx = np.random.randint(0, len(metas))\n",
    "print(rand_idx)\n",
    "print(metas[rand_idx])\n",
    "\n",
    "for key, val in metas[rand_idx].items():\n",
    "    print(f\"{key}: {val}\")\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "vae_seq = vae_data[rand_idx]\n",
    "print(vae_seq.shape)\n",
    "# crop out vae pad tokens before decode\n",
    "vae_seq = vae_seq[:metas[rand_idx][\"n_vae_tokens\"]]\n",
    "print(vae_seq.shape)\n",
    "vae_seq = torch.from_numpy(vae_seq.copy()).to(\"cpu\").unsqueeze(0).long()\n",
    "# vae_seq is (1, 3000, 128)\n",
    "\n",
    "with torch.no_grad():\n",
    "    #audio = model_25hz_codec.decode(vae_seq.permute(0, 2, 1))[0].detach().cpu()  \n",
    "    z_q, _, _ = model_25hz_codec.quantizer.from_codes(vae_seq.permute(0, 2, 1))     \n",
    "    audio = model_25hz_codec.decode(z_q)[0].detach().cpu()\n",
    "    audio /= audio.abs().max().clamp(1e-8)\n",
    "    print(audio.mean())\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(audio.numpy(), rate=48000))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "vae_seq_float16 = vae_seq.half()\n",
    "vag_seq_float_16 = vae_seq_float16.float()\n",
    "\n",
    "with torch.no_grad():\n",
    "    audio = model_100hz.decode(vag_seq_float_16.permute(0, 2, 1))[0].detach().cpu()         \n",
    "    audio /= audio.abs().max().clamp(1e-8)\n",
    "    print(audio.mean())\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(audio.numpy(), rate=44100))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "VAE_T_MEMMAP = 1000\n",
    "VAE_DIM = 128\n",
    "\n",
    "data_dir = \"/app/suno/christian/data/suno_diffusion_genius_hq_lyrics/\"\n",
    "vae_memmap_filepath = os.path.join(data_dir, \"vae_train.bin\")\n",
    "# load output memmap\n",
    "vae_data = np.memmap(os.path.join(vae_memmap_filepath), dtype=np.float32, mode=\"r\")\n",
    "vae_data = vae_data.reshape(-1, VAE_DIM, VAE_T_MEMMAP)\n",
    "\n",
    "metas = read_jsonl(\"/app/suno/christian/data/suno_diffusion_genius_hq_lyrics/train_metas.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "rand_idx = np.random.randint(0, len(metas))\n",
    "print(rand_idx)\n",
    "print(metas[rand_idx])\n",
    "vae_seq = vae_data[rand_idx]\n",
    "vae_seq = torch.from_numpy(vae_seq.copy()).to(device).unsqueeze(0).float()\n",
    "\n",
    "audio = model_100hz.decode(vae_seq)[0].detach().cpu()         \n",
    "audio /= audio.abs().max().clamp(1e-8)\n",
    "print(audio.mean())\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(audio.numpy(), rate=44100))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "N_TOKENS_MEMMAP = 6016\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "COARSE_N_CODEBOOKS = 12\n",
    "COARSE_PAD_TOKEN = 2048\n",
    "\n",
    "# load metadata\n",
    "data_dir = \"/app/suno/data/chirp_v4_ft/base_v2\"\n",
    "\n",
    "metas_filename = os.path.join(data_dir, \"metas_tr.jsonl\")\n",
    "memmap_filepath = os.path.join(data_dir, \"data_tr.bin\")\n",
    "new_metas = read_jsonl(metas_filename)\n",
    "new_info = read_json(os.path.join(data_dir, \"info_tr.json\"))\n",
    "\n",
    "# load output memmap\n",
    "out_mm = np.memmap(os.path.join(memmap_filepath), dtype=np.uint16, mode=\"r\")\n",
    "out_mm = out_mm.reshape(-1, N_TOKENS_MEMMAP, SEMANTIC_N_CODEBOOKS + COARSE_N_CODEBOOKS)\n",
    "\n",
    "print(out_mm.shape, len(new_metas))\n",
    "print(new_info.keys())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# randomly listen to some stuff\n",
    "import random\n",
    "from suno_utils.tasks.dac_2c_12cb import preload_models as preload_codec_models\n",
    "from suno_utils.tasks.dac_2c_12cb import (\n",
    "    encode as codec_encode,\n",
    "    decode as codec_decode,\n",
    "    EMBEDDING_RATE as CODEC_EMBEDDING_RATE,\n",
    ")\n",
    "\n",
    "_ = preload_codec_models(\"/app/suno/data/chirp_v4/models/dac_2c_25x12.pt\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "idx = random.choice(new_info[\"genius_hq_lyrics\"][\"idx_list\"])\n",
    "# idx_key = random.choice(list(test_info.keys()))\n",
    "# print(idx_key)\n",
    "# idx = random.choice(test_info[idx_key][\"idx_list\"])\n",
    "assert \"original_duration_s\" in new_metas[idx]\n",
    "print(idx)\n",
    "print(\"dataset:\", new_metas[idx].get(\"dataset\"))\n",
    "print(\"tags:\", new_metas[idx].get(\"tags\"))\n",
    "print(\"quality\", new_metas[idx].get(\"audio_quality\")[\"score\"])\n",
    "arr = out_mm[idx, 1:].copy().astype(np.int16)[:, 1:]\n",
    "pad_idx_arr = np.where(arr == COARSE_PAD_TOKEN)[0]\n",
    "if len(pad_idx_arr) > 0:\n",
    "    arr = arr[: pad_idx_arr[0], :]\n",
    "a = codec_decode(arr)\n",
    "a.play(compress=False)\n",
    "print(\"text:\", new_metas[idx].get(\"text\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "base_dir = \"/app/suno/data/diffusion_mix/vae_100hz_30s\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_context_aligned_val.jsonl\", progress=False)\n",
    "#vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "#semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for idx, meta in tqdm(enumerate(metas)):\n",
    "    start_s = meta[\"start_s\"]\n",
    "    end_s = meta[\"end_s\"]\n",
    "    print(start_s, end_s)\n",
    "    # check if the start_s is 0\n",
    "    if float(start_s) != 0.0:\n",
    "        prev_meta_idx = idx - 1\n",
    "        if prev_meta_idx >= 0:\n",
    "            prev_meta = metas[prev_meta_idx]\n",
    "            if prev_meta[\"end_s\"] == meta[\"start_s\"]:\n",
    "                print(\"found a match\", idx, prev_meta_idx)\n",
    "            "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas[1]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# we are going to loop over metas and create pointers to all previous blocks with the same id\n",
    "# this will allow us to use the same memmap but then index for previous context \n",
    "# to do this we will add a new field to each meta called \"prev_context_id\"\n",
    "# this will be a list of indices of the previous blocks with the same id\n",
    "# note: this assumes that the metas are in order and that the id is the same for all blocks with the same id\n",
    "\n",
    "new_metas = []\n",
    "\n",
    "for idx, meta in tqdm(enumerate(metas)):\n",
    "    new_meta = meta.copy()\n",
    "    meta_id = meta[\"id\"]\n",
    "    # check start_s\n",
    "    start_s = meta[\"start_s\"]\n",
    "    # this is the first block so we don't have any context\n",
    "    new_meta[\"prev_context_id\"] = []\n",
    "\n",
    "    if float(start_s) != 0:\n",
    "        # look at previous indicies from the current index and find the first index where start_s is 0\n",
    "        for i in range(idx - 1, -1, -1):\n",
    "            # check if the id is the same\n",
    "            if metas[i][\"id\"] == meta_id:\n",
    "                new_meta[\"prev_context_id\"].append(i)\n",
    "            else:\n",
    "                break\n",
    "    new_metas.append(new_meta)\n",
    "\n",
    "print(len(new_metas), len(metas))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# save to new metas file\n",
    "write_jsonl(new_metas, f\"{base_dir}/metas_context_tr.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# now we will need to add alignments back into the metas \n",
    "alignment_metas = read_from_s3(\n",
    "    \"s3://suno-data/datasets/metadata/alignments/genius_alignments_v11.jsonl\",\n",
    "    read_f=read_jsonl,\n",
    ")\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for idx, item in enumerate(alignment_metas[0][1]):\n",
    "    duration_s = item[\"end_s\"] - item[\"start_s\"]\n",
    "    print(idx*30, (idx+1)*30, item[\"start_s\"], item[\"end_s\"], f\"{duration_s:.2f}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "from suno_utils.audio import Audio\n",
    "# download npz and mp3 from s3\n",
    "npz_data = read_from_s3(\"s3://suno-data/christian/data/upsample_v4_t_5_20241018/25hz_20241031_v1/14034bf0-e25e-4597-8aec-4034f68262c3/14034bf0-e25e-4597-8aec-4034f68262c3-11722.npz\", read_f=np.load)\n",
    "mp3 = Audio.from_s3(\"s3://suno-data/christian/data/upsample_v4_t_5_20241018/25hz_20241031_v1/14034bf0-e25e-4597-8aec-4034f68262c3/14034bf0-e25e-4597-8aec-4034f68262c3-11722.mp3\", n_channels=2)\n",
    "\n",
    "upsampled_latents = npz_data[\"upsampled_latents\"]\n",
    "semantic_codes = npz_data[\"semantic_codes\"]\n",
    "\n",
    "print(upsampled_latents.shape, semantic_codes.shape)\n",
    "\n",
    "# decode the latents\n",
    "with torch.no_grad():\n",
    "    audio = model_25hz.decode(torch.from_numpy(upsampled_latents).to(\"cpu\").float().permute(0, 2, 1))[0].detach().cpu()\n",
    "    audio /= audio.abs().max().clamp(1e-8)\n",
    "    print(audio.mean())\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(audio.numpy(), rate=48000))\n",
    "IPython.display.display(IPython.display.Audio(mp3.array_float, rate=48000))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.suno_codec_vae import (  \n",
    "    preload_models as preload_sac_vae_models,\n",
    "    encode as sac_vae_encode,\n",
    "    decode as sac_vae_decode,\n",
    ")\n",
    "\n",
    "_ = preload_sac_vae_models(\"s3://suno-data/minz/models/codec/vae_37epoch.ckpt\")\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.models.dac.model.sac_vae import SAC\n",
    "\n",
    "device = \"cuda:0\"\n",
    "#checkpoint_filepath = \"s3://suno-data/minz/models/sac_vae_25hz.pth\"\n",
    "checkpoint_filepath = \"s3://suno-data/minz/models/codec/vae_37epoch.ckpt\"\n",
    "load_f = funcy.partial(torch.load, map_location=\"cpu\")\n",
    "\n",
    "if checkpoint_filepath.startswith(\"s3://\"):\n",
    "    sd = read_from_s3(checkpoint_filepath, read_f=load_f)\n",
    "else:\n",
    "    sd = load_f(checkpoint_filepath)\n",
    "\n",
    "sd[\"metadata\"][\"kwargs\"] = {\n",
    "    k: v\n",
    "    for k, v in sd[\"metadata\"][\"kwargs\"].items()\n",
    "    if k in SAC.__init__.__code__.co_varnames\n",
    "}\n",
    "model_25hz_vae = SAC(**sd[\"metadata\"][\"kwargs\"])\n",
    "model_25hz_vae.load_state_dict(sd[\"state_dict\"])\n",
    "model_25hz_vae.eval()\n",
    "model_25hz_vae.to(device)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 128\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "base_dir = \"/app/suno/data/diffusion_mix/sac_vae_25hz_30s\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_tr.jsonl\", progress=False)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_tr.bin\"\n",
    "semantic_memmap_filepath = f\"{base_dir}/data_semantic_tr.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)#.astype(np.float32)\n",
    "semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)#.astype(np.int16)\n",
    "\n",
    "print(vae_data.shape, semantic_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "rand_idx = 150\n",
    "print(metas[rand_idx])\n",
    "vae_seq = vae_data[rand_idx]\n",
    "vae_seq = torch.from_numpy(vae_seq.copy()).to(\"cuda\").float()\n",
    "print(vae_seq.shape)\n",
    "\n",
    "# vae_seq is (1, 750, 128)\n",
    "audio = codec_decode(vae_seq)\n",
    "audio.play(compress=False)\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "base_dir = \"s3://suno-data/datasets/bundles/v3/diffusion_mix_fix/sac_vae_25hz/\"\n",
    "npz_filename = \"part_999.npz\"\n",
    "\n",
    "npz_data = read_from_s3(os.path.join(base_dir, npz_filename), read_f=np.load)\n",
    "print(npz_data.keys())\n",
    "latents = npz_data[\"3aa154f4-e94b-43ef-9a34-5ffc41174830\"]\n",
    "print(latents.shape)\n",
    "\n",
    "latents = torch.from_numpy(latents.copy())\n",
    "latents = latents[:750,:].to(device).float()\n",
    "print(latents.shape)\n",
    "# vae_seq is (1, 750, 128)\n",
    "#audio = sac_vae_decode(vae_seq)\n",
    "with torch.no_grad():\n",
    "    audio = model_25hz.decode(latents.unsqueeze(0).permute(0, 2, 1))[0].detach().cpu()\n",
    "    audio /= audio.abs().max().clamp(1e-8)\n",
    "    print(audio.mean())\n",
    "\n",
    "#audio = sac_vae_decode(latents)\n",
    "#audio.play(compress=False)\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "IPython.display.display(IPython.display.Audio(audio.numpy(), rate=48000))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.tasks.suno_codec_vae import (   \n",
    "    preload_models,\n",
    "    encode as codec_encode,\n",
    "    decode as codec_decode,\n",
    ")   \n",
    "\n",
    "_ = preload_models(\"s3://suno-data/minz/models/codec/vae_37epoch.ckpt\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "VAE_DIM = 768\n",
    "VAE_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "SEMANTIC_N_CODEBOOKS = 1\n",
    "SEMANTIC_N_MEMMAP_TOKENS = 750\n",
    "\n",
    "#base_dir = \"/app/suno/data/diffusion_mix/dac_vae_fixed_25hz\"\n",
    "base_dir = \"/app/suno/data/diffusion_mix/semantic_cont_25hz_30s\"\n",
    "#base_dir = \"/app/suno/data/diffusion_v5/v0\"\n",
    "#base_dir = \"/mnt/localdisk/cjs_shards/\"\n",
    "metas = read_jsonl(f\"{base_dir}/metas_val.jsonl\", progress=True)\n",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "#semantic_memmap_filepath = f\"{base_dir}/data_semantic_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "#semantic_memmap = np.memmap(semantic_memmap_filepath, dtype=np.uint16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "#semantic_data = semantic_memmap.reshape(-1, SEMANTIC_N_MEMMAP_TOKENS, SEMANTIC_N_CODEBOOKS)\n",
    "\n",
    "print(vae_data.shape, len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "metas[3]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# compute the std of the latents\n",
    "vad_data_array = vae_data.astype(np.float32)\n",
    "print(vad_data_array.shape)\n",
    "# compute the std of the latents\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "latent_std = np.std(vad_data_array)\n",
    "print(latent_std.shape)\n",
    "print(latent_std)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "1/0.28950977"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "0.28950977*3.45"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_diff",
   "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.12.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
