{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import glob\n",
    "import torch\n",
    "import torchaudio\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "from tqdm import tqdm\n",
    "from sklearn.preprocessing import StandardScaler\n",
    "from sklearn.cluster import KMeans, MiniBatchKMeans\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Training (K-Means)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load some audios and compute melspectrograms\n",
    "audio_paths = [ \n",
    "    \"/home/christian/audio/reference-audio-wav/09 Sounds Like Hallelujah.wav\",\n",
    "\"/home/christian/audio/reference-audio-wav/02 Freddie Freeloader.wav\",\n",
    "\"/home/christian/audio/reference-audio-wav/01 No Son Of Mine.wav\",\n",
    " \"/home/christian/audio/reference-audio-wav/02 Dreams.wav\",\n",
    " \"/home/christian/audio/reference-audio-wav/01 J.S. Bach Suite No.1, S.1007, G major - I. Prelude.wav\",\n",
    "\"/home/christian/audio/reference-audio-wav/04 Fuckwithmeyouknowigotit.wav\",\n",
    "\"/home/christian/audio/reference-audio-wav/03 Always Be.wav\",\n",
    "]\n",
    "\n",
    "audio_paths = glob.glob(os.path.join(\"/app/suno/data/audio_2ch_12khz_lg/val/**/*.wav\"))\n",
    "\n",
    "SAMPLE_RATE = 12000\n",
    "N_MELS = 64\n",
    "BATCH_SIZE = 128\n",
    "N_SEC = 5\n",
    "\n",
    "melspec_encode = torchaudio.transforms.MelSpectrogram(\n",
    "    sample_rate=SAMPLE_RATE,\n",
    "    n_fft=256,            # Length of the FFT window\n",
    "    win_length=None,       # Window size\n",
    "    hop_length=128,        # Number of samples between successive frames\n",
    "    n_mels=N_MELS,            # Number of Mel bands\n",
    "    center=True,           # Whether the t-th frame is centered at t*hop_length\n",
    "    pad_mode='reflect',    # Padding mode\n",
    "    power=2.0              # Power of the norm\n",
    ")\n",
    "\n",
    "dataset = []\n",
    "audios_encode = []\n",
    "\n",
    "for audio_path in tqdm(audio_paths):\n",
    "    audio, sr = torchaudio.load(audio_path)\n",
    "\n",
    "\n",
    "    if sr != SAMPLE_RATE:\n",
    "        audio_encode = torchaudio.functional.resample(audio, sr, SAMPLE_RATE)\n",
    "    else:\n",
    "        audio_encode = audio\n",
    "\n",
    "    # crop to 60 sec\n",
    "    start_idx = 0\n",
    "    end_idx = start_idx + int(SAMPLE_RATE * N_SEC)\n",
    "    audio_encode = audio_encode[:, start_idx: end_idx]\n",
    "    \n",
    "    if audio_encode.shape[1] < SAMPLE_RATE * N_SEC:\n",
    "        continue\n",
    "\n",
    "    audios_encode.append(audio_encode)\n",
    "\n",
    "    if len(audios_encode) == BATCH_SIZE:\n",
    "        audios_encode = torch.cat(audios_encode, dim=0)\n",
    "        melspec = melspec_encode(audios_encode)\n",
    "        melspec_db = torchaudio.transforms.AmplitudeToDB()(melspec)\n",
    "        dataset.append(melspec_db.reshape(-1, N_MELS))\n",
    "        audios_encode = []\n",
    "\n",
    "X = torch.cat(dataset, dim=0)\n",
    "print(X.shape)\n",
    "        \n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "scaler = StandardScaler()\n",
    "X_scaled = scaler.fit_transform(X)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Define the number of clusters\n",
    "n_clusters = 32768\n",
    "\n",
    "# Initialize the KMeans model\n",
    "kmeans = MiniBatchKMeans(batch_size=32, n_clusters=n_clusters, random_state=0, verbose=1, max_iter=1000)\n",
    "\n",
    "# Fit the model to the data\n",
    "kmeans.fit(X_scaled)\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Predict"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "test_audio_path = \"/home/christian/audio/reference-audio-wav/08 Get Lucky.wav\"\n",
    "audio, sr = torchaudio.load(test_audio_path)\n",
    "\n",
    "if sr != SAMPLE_RATE:\n",
    "    audio_encode = torchaudio.functional.resample(audio, sr, SAMPLE_RATE)\n",
    "else:\n",
    "    audio_encode = audio\n",
    "\n",
    "# crop to 60 sec\n",
    "start_idx = 0\n",
    "end_idx = start_idx + int(SAMPLE_RATE * 5)\n",
    "audio_encode = audio_encode[:, start_idx: end_idx]\n",
    "\n",
    "# Compute the Mel spectrogram\n",
    "melspec = melspec_encode(audio_encode)\n",
    "\n",
    "# Convert to decibels\n",
    "melspec_db = torchaudio.transforms.AmplitudeToDB()(melspec)\n",
    "print(melspec_db.shape)\n",
    "melspec_db = melspec_db[0] # look just at the left channel\n",
    "\n",
    "# encode with learned clusters \n",
    "melspec_codes = kmeans.predict(melspec_db.T)\n",
    "recon_melspec_db = kmeans.cluster_centers_[melspec_codes].T\n",
    "print(recon_melspec_db.shape, melspec_db.shape)\n",
    "\n",
    "# plot\n",
    "fig, axs = plt.subplots(nrows=2, ncols=1)\n",
    "axs[0].imshow(melspec_db.numpy(), aspect='auto', origin='lower')\n",
    "axs[1].imshow(recon_melspec_db, aspect='auto', origin='lower')\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Plot Semantic"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# MERT\n",
    "from suno_utils.tasks.mert_25 import (\n",
    "    preload_models as preload_semantic_models,\n",
    "    encode as semantic_encode,\n",
    ")\n",
    "from suno_utils.utils.s3 import read_from_s3\n",
    "\n",
    "\n",
    "_ = preload_semantic_models(\n",
    "    checkpoint_filepath=\"s3://suno-data/georg/models/semantic/mert_25.pt\",\n",
    "    centroids_filepath=\"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\",\n",
    "    device=\"cuda\",\n",
    ")\n",
    "\n",
    "centroids = read_from_s3(\"s3://suno-data/georg/models/semantic/mert_25_2x4k.npy\", read_f=np.load)[0]\n",
    "embedding = torch.nn.Embedding(centroids.shape[0], centroids.shape[1])\n",
    "embedding.weight.data.copy_(torch.tensor(centroids, dtype=torch.float32))\n",
    "embedding.weight.requires_grad = False\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "test_audio_path = \"/home/christian/audio/reference-audio-wav/08 Get Lucky.wav\"\n",
    "audio, sr = torchaudio.load(test_audio_path)\n",
    "\n",
    "if sr != 24000:\n",
    "    audio_encode = torchaudio.functional.resample(audio, sr, 24000)\n",
    "else:\n",
    "    audio_encode = audio\n",
    "\n",
    "# crop to 60 sec\n",
    "start_idx = 0\n",
    "end_idx = start_idx + int(SAMPLE_RATE * 5)\n",
    "audio_encode = audio_encode[:, start_idx: end_idx].mean(dim=0, keepdim=True)\n",
    "\n",
    "# Compute the Mel spectrogram\n",
    "semantic_codes = torch.from_numpy(semantic_encode(audio_encode)).long()[:,0]\n",
    "semantic_latents = embedding(semantic_codes).detach().cpu().numpy()\n",
    "print(semantic_latents.shape)\n",
    "\n",
    "# plot\n",
    "fig, axs = plt.subplots(nrows=2, ncols=1)\n",
    "axs[1].imshow(semantic_latents, aspect='auto', origin='lower')\n",
    "#axs[1].imshow(recon_melspec_db, aspect='auto', origin='lower')"
   ]
  },
  {
   "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
}
