{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import glob\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"1\"\n",
    "import glob\n",
    "import torch\n",
    "import faiss\n",
    "import IPython\n",
    "import numpy as np\n",
    "import torchaudio\n",
    "import itertools\n",
    "\n",
    "from tqdm import tqdm\n",
    "from dac.model.dac2 import DAC\n",
    "from dac.model.discriminator2 import Discriminator as Discriminator_import\n",
    "from dac.nn import loss as loss_import\n",
    "from dac.utils.accelerator import Accelerator\n",
    "from dac.utils import load_model\n",
    "\n",
    "import matplotlib.pyplot as plt\n",
    "from sklearn.preprocessing import StandardScaler\n",
    "from sklearn.cluster import KMeans, MiniBatchKMeans\n",
    "\n",
    "from suno_utils.utils.s3 import read_from_s3"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Run K-means to quantize"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load data from memmaps\n",
    "SAMPLE_RATE = 48_000\n",
    "VAE_DIM = 128\n",
    "VAE_T_MEMMAP = 1000\n",
    "MERT_DIM = 1024\n",
    "MERT_T_MEMMAP = 250\n",
    "USE_VAL = False\n",
    "NUM_FRAMES = int(10 * SAMPLE_RATE)\n",
    "BATCH_SIZE = 32\n",
    "MAX_FILES_PER_SUBSET = 50_0000\n",
    "\n",
    "#root_dir = \"/app/suno/christian/data/ursing_vocal_48khz/embed\"\n",
    "root_dir = \"/app/suno/christian/data/suno_diffusion/\"\n",
    "\n",
    "if USE_VAL:\n",
    "    in_vae_mm = np.memmap(\n",
    "        f\"{root_dir}/vae_val.bin\", \n",
    "        dtype=np.float32, \n",
    "        mode=\"r\", \n",
    "    )\n",
    "\n",
    "    in_mert_mm = np.memmap(\n",
    "        f\"{root_dir}/mert_val.bin\", \n",
    "        dtype=np.float32, \n",
    "        mode=\"r\", \n",
    "    )\n",
    "else:\n",
    "    in_vae_mm = np.memmap(\n",
    "        f\"{root_dir}/vae_train.bin\", \n",
    "        dtype=np.float32, \n",
    "        mode=\"r\", \n",
    "    )\n",
    "\n",
    "    in_mert_mm = np.memmap(\n",
    "        f\"{root_dir}/mert_train.bin\", \n",
    "        dtype=np.float32, \n",
    "        mode=\"r\", \n",
    "    )\n",
    "\n",
    "# reshape to (n_examples, VAE_DIM, T_MEMMAP)\n",
    "vae_data = in_vae_mm.reshape(-1, VAE_DIM, VAE_T_MEMMAP)\n",
    "# reshape to (n_examples, VAE_DIM * T_MEMMAP)\n",
    "vae_data = vae_data.reshape(-1, VAE_DIM)\n",
    "print(vae_data.shape)\n",
    "\n",
    "# reshape to (n_examples, T_MEMMAP)\n",
    "mert_data = in_mert_mm.reshape(-1, MERT_DIM, MERT_T_MEMMAP)\n",
    "# reshape to (n_examples, MERT_DIM * T_MEMMAP)\n",
    "mert_data = mert_data.reshape(-1, MERT_DIM)\n",
    "print(mert_data.shape)\n",
    "\n",
    "# randomly sample examples\n",
    "print(\"Sampling...\")\n",
    "sample_size = 5000000\n",
    "#if vae_data.shape[0] > sample_size:\n",
    "#    vae_data_subset = vae_data[np.random.choice(vae_data.shape[0], sample_size, replace=False)]\n",
    "#    print(\"vae_data_subset\", vae_data_subset.shape)\n",
    "#else:\n",
    "vae_data_subset = vae_data\n",
    "\n",
    "if mert_data.shape[0] > sample_size:\n",
    "    mert_data_subset = mert_data[np.random.choice(mert_data.shape[0], sample_size, replace=False)]\n",
    "    print(\"mert_data_subset\", mert_data_subset.shape)\n",
    "else:\n",
    "    mert_data_subset = mert_data\n",
    "\n",
    "# this randomly sampling took 34 minutes"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# check normalization of subset data\n",
    "print(\"Checking normalization...\")\n",
    "#scaler = StandardScaler()\n",
    "#scaler.fit(mert_data_subset)\n",
    "#print(\"mean\", scaler.mean_)\n",
    "#print(\"var\", scaler.var_)\n",
    "mert_mean = np.mean(mert_data_subset, axis=0)\n",
    "mert_std = np.std(mert_data_subset, axis=0)\n",
    "print(\"mean\", mert_mean, mert_mean.shape)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# k-means"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "os.environ[\"OMP_NUM_THREADS\"] = \"1\" # use fewer threads for faiss"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "centroid_sizes = [2048, 4096, 8192, 16384, 32768, 65536]\n",
    "errors = []\n",
    "\n",
    "for ncentroids in centroid_sizes:\n",
    "    print(f\"Running KMeans for {ncentroids} centroids\")\n",
    "    # run clustering for MERT\n",
    "    niter = 100\n",
    "    nredo = 3\n",
    "    verbose = True\n",
    "    d = mert_data.shape[-1] # embed dim\n",
    "\n",
    "    mert_kmeans = faiss.Kmeans(d, ncentroids, niter=niter, nredo=nredo, verbose=verbose, gpu=True)\n",
    "    mert_kmeans.train(mert_data_subset)\n",
    "\n",
    "    # evaluate by quantizing the data and then reconstructing\n",
    "    mert_codes = mert_kmeans.index.search(mert_data_subset[:1000, :], 1)[1]\n",
    "    mert_reconstructed = mert_kmeans.centroids[mert_codes.flatten()]\n",
    "    mert_error = np.mean((mert_data_subset[:1000] - mert_reconstructed)**2)\n",
    "    errors.append(mert_error)\n",
    "    print(\"Reconstruction error\", mert_error)\n",
    "\n",
    "# Reconstruction error 0.4775746\n",
    "# Reconstruction error 0.4622394"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "final_results = [result[-1] for result in results]\n",
    "print(final_results)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, axs = plt.subplots(nrows=len(results), sharex=True, figsize=(5, 10))\n",
    "\n",
    "for idx, result in enumerate(results):\n",
    "    # split result into 10 item chunks\n",
    "    sub_results = [result[i:i+10] for i in range(0, len(result), 10)]\n",
    "    for sub_result in sub_results:\n",
    "        axs[idx].plot(sub_result)\n",
    "    #axs[idx].set_title(f\"{centroid_sizes[idx]} centroids\")\n",
    "    #axs[idx].set_yscale(\"log\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, ax = plt.subplots()\n",
    "t = np.arange(len(centroid_sizes))\n",
    "ax.plot(t, errors)\n",
    "ax.set_xlabel(\"Number of centroids\")\n",
    "ax.set_ylabel(\"Error\")\n",
    "ax.set_xticks(np.arange(len(centroid_sizes)))\n",
    "ax.set_xticklabels(centroid_sizes)\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# save centroids to npy and store on s3\n",
    "np.save(f\"/app/suno/christian/data/vae/ursing_vae_100hz_32768.npy\", vae_kmeans.centroids)\n",
    "os.system(\"aws s3 cp /app/suno/christian/data/vae/ursing_vae_100hz_32768.npy s3://suno-data/christian/ursing_vae_100hz_32768.npy\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# MERT k-means"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# run clustering for MERT\n",
    "ncentroids = 32768\n",
    "niter = 200\n",
    "nredo = 10\n",
    "verbose = True\n",
    "SAMPLE_RATE = 48_000\n",
    "d = mert_data_subset.shape[-1] # embed dim\n",
    "\n",
    "mert_kmeans = faiss.Kmeans(d, ncentroids, niter=niter, nredo=nredo, verbose=verbose, gpu=True)\n",
    "mert_kmeans.train(mert_data_subset)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# save centroids to npy and store on s3\n",
    "np.save(f\"/app/suno/christian/data/suno_diffusion/mert_25hz_{ncentroids}.npy\", mert_kmeans.centroids)\n",
    "os.system(f\"aws s3 cp /app/suno/christian/data/suno_diffusion/mert_25hz_{ncentroids}.npy s3://suno-data/christian/mert_100hz_{ncentroids}.npy\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# 4. Quantize VAE embeds with learned centroids"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "os.environ[\"OMP_NUM_THREADS\"] = \"32\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# read centroids from S3\n",
    "centroid_path = \"s3://suno-data/christian/ursing_vae_100hz_32768.npy\"\n",
    "centroids = read_from_s3(centroid_path, read_f=np.load)\n",
    "print(centroids.shape)\n",
    "\n",
    "d = 128\n",
    "\n",
    "# create faiss kmeans index using centroids\n",
    "index_cpu = faiss.IndexFlatL2(d)\n",
    "index_cpu.add(centroids)\n",
    "\n",
    "# mvoe to gpu \n",
    "res = faiss.StandardGpuResources()  # Use a single GPU\n",
    "gpu_index = faiss.index_cpu_to_gpu(res, 0, index_cpu)  # 0 is the GPU ID\n",
    "\n",
    "# load memmap of embeddings\n",
    "SAMPLE_RATE = 48_000\n",
    "VAE_DIM = 128\n",
    "T_MEMMAP = 1000\n",
    "USE_VAL = False\n",
    "NUM_FRAMES = int(10 * SAMPLE_RATE)\n",
    "BATCH_SIZE = 1000\n",
    "\n",
    "if USE_VAL:\n",
    "    subset_name = \"val\"\n",
    "else:\n",
    "    subset_name = \"train\"\n",
    "\n",
    "in_mm = np.memmap(\n",
    "    f\"/app/suno/christian/data/ursing_vocal_48khz/embed/vae_{subset_name}.bin\", \n",
    "    dtype=np.float32, \n",
    "    mode=\"r\", \n",
    ")\n",
    "in_mm = in_mm.reshape(-1, VAE_DIM, T_MEMMAP)\n",
    "print(in_mm.shape)\n",
    "\n",
    "num_examples = in_mm.shape[0]\n",
    "\n",
    "# create a new memmap to store codes for the entire dataset\n",
    "out_mm = np.memmap(\n",
    "    f\"/app/suno/christian/data/ursing_vocal_48khz/embed/vae_discrete_{subset_name}.bin\", \n",
    "    dtype=np.uint16, \n",
    "    mode=\"w+\", \n",
    "    shape=(num_examples * T_MEMMAP)\n",
    ")\n",
    "\n",
    "# quantize in_mm in batches \n",
    "for idx in tqdm(range(0, num_examples, BATCH_SIZE)):\n",
    "    batch = in_mm[idx:idx+BATCH_SIZE, ...].reshape(-1, VAE_DIM)\n",
    "    closest_cluster = gpu_index.search(batch, 1)[1]\n",
    "    out_mm[idx*T_MEMMAP:(idx+BATCH_SIZE)*T_MEMMAP] = closest_cluster.flatten()\n",
    "    \n",
    "\n",
    "# quantize with kmeans clusters\n",
    "# get the closest cluster for each sample\n",
    "#closest_cluster = kmeans.index.search(z.numpy(), 1)[1]\n",
    "#print(closest_cluster)\n",
    "# get sequence of cluster centroids\n",
    "#cluster_centroids = kmeans.centroids[closest_cluster]\n",
    "#z_q = torch.from_numpy(cluster_centroids).squeeze(1)\n",
    "out_mm.flush()\n",
    "del out_mm\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "x = np.random.rand(1, 128)\n",
    "# find closest cluster in index\n",
    "closest_cluster = kmeans.index.search(x, 1)[1]\n",
    "print(closest_cluster)"
   ]
  },
  {
   "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.14"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
