{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import glob\n",
    "import json\n",
    "import torch\n",
    "import IPython\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "\n",
    "from suno_utils.utils.text import (\n",
    "    write_jsonl,\n",
    "    read_jsonl,\n",
    "    write_json,\n",
    "    read_json,\n",
    "    normalize_whitespace,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "root_dir = \"/home/christian/code/christian/outputs/quality_metas/audio_2ch_48khz_lg/train\"\n",
    "metas = glob.glob(os.path.join(root_dir, \"*.jsonl\"))\n",
    "\n",
    "metadata = {}\n",
    "\n",
    "for meta_filepath in metas:\n",
    "    subset_name = os.path.basename(meta_filepath).replace(\"_quality_metas.jsonl\", \"\")\n",
    "    subset_metas = read_jsonl(meta_filepath)\n",
    "    metadata[subset_name] = subset_metas\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# plot hitsogram of quality scores\n",
    "fig, axs = plt.subplots(2, 5, figsize=(20, 5), sharex=True, sharey=True)\n",
    "axs = np.reshape(axs, -1)\n",
    "bins = np.linspace(-5, 5, 50)\n",
    "for idx, (subset_name, subset_metas) in enumerate(metadata.items()):\n",
    "    scores = [float(meta[\"audio_quality\"][\"score\"]) for meta in subset_metas]\n",
    "    axs[idx].hist(scores, bins=bins, density=True)\n",
    "    axs[idx].set_title(subset_name)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# get audio for the worst quality in each subset\n",
    "subset_metas = metadata[\"genius_hq\"]\n",
    "\n",
    "# sort metas by the audio quality score\n",
    "subset_metas = sorted(\n",
    "    subset_metas, key=lambda x: float(x[\"audio_quality\"][\"score\"]), reverse=False\n",
    ")\n",
    "\n",
    "# get worst quality\n",
    "worst_metas = subset_metas[:10]\n",
    "for worst_meta in worst_metas:\n",
    "    print(worst_meta[\"audio_quality\"][\"score\"])\n",
    "    #IPython.display.display(IPython.display.Audio(worst_meta[\"audio_file\"][0]))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "for quality_threshold in [0.0, -0.5, -1, -2, -4, -6, -8]:\n",
    "    filtered_subset_metas = []\n",
    "    for metas in subset_metas:\n",
    "        if float(metas[\"audio_quality\"][\"score\"]) > quality_threshold:\n",
    "            filtered_subset_metas.append(metas)\n",
    "\n",
    "    keep_percent = len(filtered_subset_metas) / len(subset_metas) * 100\n",
    "    num_removed = len(subset_metas) - len(filtered_subset_metas)\n",
    "\n",
    "    print(f\"Quality threshold: {quality_threshold} keep: {len(filtered_subset_metas)} ({keep_percent:0.2f}%) remove: {num_removed} \")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "quality_threshold = 0.0\n",
    "subset_name = \"genius_hq\"\n",
    "root_dir = \"/home/christian/code/christian/outputs/quality_metas/audio_2ch_48khz_lg/\"\n",
    "\n",
    "for sub_dir in [\"train\", \"val\"]:\n",
    "    subset_metas = []\n",
    "\n",
    "    meta_filepath = os.path.join(root_dir, sub_dir, f\"{subset_name}_quality_metas.jsonl\")\n",
    "    subset_metas = read_jsonl(meta_filepath)\n",
    "\n",
    "    filtered_subset_metas = []\n",
    "    for metas in subset_metas:\n",
    "        if float(metas[\"audio_quality\"][\"score\"]) > quality_threshold:\n",
    "            filtered_subset_metas.append(metas)\n",
    "\n",
    "    keep_percent = len(filtered_subset_metas) / len(subset_metas) * 100\n",
    "    num_removed = len(subset_metas) - len(filtered_subset_metas)\n",
    "    print(f\"({sub_dir}) Quality threshold: {quality_threshold} keep: {len(filtered_subset_metas)} ({keep_percent:0.2f}%) remove: {num_removed} \")\n",
    "\n",
    "    write_jsonl(filtered_subset_metas, f\"/home/christian/code/christian/outputs/quality_metas/audio_2ch_48khz_lg/{sub_dir}/{subset_name}_filtered_0.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# listen to some tiktok covers data\n",
    "root_dir = \"/app/suno/christian/data/tiktok_covers_48khz/train\"\n",
    "filepaths = glob.glob(os.path.join(root_dir, \"*.wav\"))\n",
    "print(len(filepaths))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "# random filepath \n",
    "filepath = np.random.choice(filepaths)\n",
    "IPython.display.display(IPython.display.Audio(filepath, normalize=False))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load codec model\n",
    "import funcy\n",
    "import torchaudio\n",
    "from dac.model.dac4 import DAC\n",
    "from suno_utils.utils.s3 import read_from_s3\n",
    "\n",
    "\n",
    "device = \"cuda:0\"\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": [
    "# codec cycle random audio\n",
    "filepath = np.random.choice(filepaths)\n",
    "audio, sr = torchaudio.load(filepath)\n",
    "print(filepath)\n",
    "with torch.no_grad():\n",
    "    cycled_audio = model_100hz(audio.unsqueeze(0).to(device))[\"audio\"].squeeze(0).cpu()\n",
    "    peak = cycled_audio.abs().max()\n",
    "    if peak > 1:\n",
    "        cycled_audio /= peak\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(data=audio, rate=48_000, normalize=False))\n",
    "IPython.display.display(IPython.display.Audio(data=cycled_audio, rate=48_000, normalize=False))"
   ]
  },
  {
   "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
}
