{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import numpy as np\n",
    "import glob\n",
    "import json\n",
    "import matplotlib.pyplot as plt"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "filepaths = glob.glob(\"/app/suno/data/diff_sft/v45_2b_step_2_600_000/pos_interesting_clips_up_u_1_20241201_full_v2/*_vae.npz\")\n",
    "print(len(filepaths))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "filepaths[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "idx = 2000\n",
    "vae_latents = np.load(filepaths[idx])[\"arr_0\"] * 0.4\n",
    "print(vae_latents.shape)\n",
    "\n",
    "noise = np.random.randn(vae_latents.shape[0], vae_latents.shape[1])\n",
    "\n",
    "hist_values = vae_latents.flatten()\n",
    "hist_noise_values = noise.flatten()\n",
    "\n",
    "vae_latents_std = np.std(vae_latents)\n",
    "vae_latents_mean = np.mean(vae_latents)\n",
    "vae_latents_max = np.max(vae_latents)\n",
    "vae_latents_min = np.min(vae_latents)\n",
    "print(f\"vae_latents_std: {vae_latents_std}\")\n",
    "print(f\"vae_latents_mean: {vae_latents_mean}\")\n",
    "print(f\"vae_latents_max: {vae_latents_max}\")\n",
    "print(f\"vae_latents_min: {vae_latents_min}\")\n",
    "\n",
    "noise_min = np.min(noise)\n",
    "noise_max = np.max(noise)\n",
    "noise_mean = np.mean(noise)\n",
    "noise_std = np.std(noise)\n",
    "#print(f\"noise_min: {noise_min}\")\n",
    "#print(f\"noise_max: {noise_max}\")\n",
    "#print(f\"noise_mean: {noise_mean}\")\n",
    "#print(f\"noise_std: {noise_std}\")\n",
    "\n",
    "plt.hist(hist_values, bins=100, alpha=0.5)\n",
    "plt.hist(hist_noise_values, bins=100, alpha=0.5)\n",
    "plt.show()\n",
    "\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/dac_vae_fixed_25hz\"\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 = 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)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "idx = 2500\n",
    "vae_latents = vae_data[idx].astype(np.float32) * 0.4\n",
    "print(vae_latents.shape)\n",
    "\n",
    "noise = np.random.randn(vae_latents.shape[0], vae_latents.shape[1])\n",
    "\n",
    "hist_values = vae_latents.flatten()\n",
    "hist_noise_values = noise.flatten()\n",
    "\n",
    "vae_latents_std = np.std(vae_latents)\n",
    "vae_latents_mean = np.mean(vae_latents)\n",
    "vae_latents_max = np.max(vae_latents)\n",
    "vae_latents_min = np.min(vae_latents)\n",
    "print(f\"vae_latents_std: {vae_latents_std}\")\n",
    "print(f\"vae_latents_mean: {vae_latents_mean}\")\n",
    "print(f\"vae_latents_max: {vae_latents_max}\")\n",
    "print(f\"vae_latents_min: {vae_latents_min}\")\n",
    "\n",
    "noise_min = np.min(noise)\n",
    "noise_max = np.max(noise)\n",
    "noise_mean = np.mean(noise)\n",
    "noise_std = np.std(noise)\n",
    "#print(f\"noise_min: {noise_min}\")\n",
    "#print(f\"noise_max: {noise_max}\")\n",
    "#print(f\"noise_mean: {noise_mean}\")\n",
    "#print(f\"noise_std: {noise_std}\")\n",
    "\n",
    "plt.hist(hist_values, bins=100, alpha=0.5)\n",
    "#plt.hist(hist_noise_values, bins=100, alpha=0.5)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "gen_max_vals = []\n",
    "real_max_vals = []\n",
    "from tqdm import tqdm\n",
    "\n",
    "for n in tqdm(range(3000)):\n",
    "    gen_vae_latents = np.load(filepaths[n])[\"arr_0\"][:750] * 0.4\n",
    "\n",
    "    real_vae_latents = vae_data[n].astype(np.float32) * 0.4\n",
    "    gen_max_val = np.max(np.abs(gen_vae_latents))\n",
    "    real_max_val = np.max(np.abs(real_vae_latents))\n",
    "    gen_max_vals.append(gen_max_val)\n",
    "    real_max_vals.append(real_max_val)\n",
    "\n",
    "\n",
    "plt.hist(gen_max_vals, bins=100, alpha=0.5, label=\"gen\")\n",
    "plt.hist(real_max_vals, bins=100, alpha=0.5, label=\"real\")\n",
    "plt.legend()\n",
    "plt.title(\"Max absolute value of VAE latents (dac_vae_tuned_25hz)\")\n",
    "plt.xlabel(\"Max absolute value (with scaling factor 0.4)\")\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "first_chunk_max_vals = []\n",
    "second_chunk_max_vals = []\n",
    "third_chunk_max_vals = []\n",
    "\n",
    "for n in tqdm(range(3000)):\n",
    "    vae_latents = np.load(filepaths[n])[\"arr_0\"] * 0.4\n",
    "\n",
    "    if vae_latents.shape[0] < 2250:\n",
    "        continue\n",
    "\n",
    "    first_vae_latents = vae_latents[:750]\n",
    "    second_vae_latents = vae_latents[750:1500]\n",
    "    third_vae_latents = vae_latents[1500:2250]\n",
    "    first_chunk_max_val = np.max(np.abs(first_vae_latents))\n",
    "    second_chunk_max_val = np.max(np.abs(second_vae_latents))\n",
    "    third_chunk_max_val = np.max(np.abs(third_vae_latents))\n",
    "    first_chunk_max_vals.append(first_chunk_max_val)\n",
    "    second_chunk_max_vals.append(second_chunk_max_val)\n",
    "    third_chunk_max_vals.append(third_chunk_max_val)\n",
    "\n",
    "plt.hist(first_chunk_max_vals, bins=100, alpha=0.5, label=\"first chunk\")\n",
    "plt.hist(second_chunk_max_vals, bins=100, alpha=0.5, label=\"second chunk\")\n",
    "plt.hist(third_chunk_max_vals, bins=100, alpha=0.5, label=\"third chunk\")\n",
    "plt.legend()\n",
    "plt.title(\"Max absolute value of VAE latents (dac_vae_tuned_25hz)\")\n",
    "plt.xlabel(\"Max absolute value (with scaling factor 0.4)\")\n",
    "plt.show()\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",
    "vae_memmap_filepath = f\"{base_dir}/data_vae_val.bin\"\n",
    "\n",
    "# load memmaps\n",
    "vae_memmap = np.memmap(vae_memmap_filepath, dtype=np.float16, mode=\"r\")\n",
    "\n",
    "# reshape memmaps\n",
    "old_vae_data = vae_memmap.reshape(-1, VAE_N_MEMMAP_TOKENS, VAE_DIM)\n",
    "\n",
    "print(old_vae_data.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "dac_tuned_max_vals = []\n",
    "old_dac_max_vals = []\n",
    "\n",
    "for n in tqdm(range(3000)):\n",
    "    dac_tuned_vae_latents = vae_data[n].astype(np.float32) * 0.4\n",
    "    old_dac_vae_latents = old_vae_data[n].astype(np.float32) * 2.5\n",
    "\n",
    "    dac_tuned_max_val = np.max(np.abs(dac_tuned_vae_latents))\n",
    "    old_dac_max_val = np.max(np.abs(old_dac_vae_latents))\n",
    "    dac_tuned_max_vals.append(dac_tuned_max_val)\n",
    "    old_dac_max_vals.append(old_dac_max_val)\n",
    "\n",
    "\n",
    "plt.hist(dac_tuned_max_vals, bins=100, alpha=0.5, label=\"dac_tuned\")\n",
    "plt.hist(old_dac_max_vals, bins=100, alpha=0.5, label=\"old_dac\")\n",
    "plt.legend()\n",
    "plt.title(\"Max absolute value of VAE latents (val data)\")\n",
    "plt.xlabel(\"Max absolute value (with scaling factor applied)\")\n",
    "plt.show()\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
}
