{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import json\n",
    "import torch\n",
    "import IPython\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "from suno_utils.utils.text import read_jsonl, write_jsonl"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load genius metas \n",
    "genius_metas_filepath = \"/home/christian/code/christian/metadata/genius_hq_metas_audio_production.jsonl\"\n",
    "genius_metas = read_jsonl(genius_metas_filepath)\n",
    "print(len(genius_metas))\n",
    "\n",
    "# first assume everything is good\n",
    "filtered_ids = set([meta[\"id\"] for meta in genius_metas])\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load the train metas\n",
    "train_metas_filepath = \"/app/suno/data/diffusion_mix/vae_25hz_30s/metas_tr.jsonl\"\n",
    "train_metas = read_jsonl(train_metas_filepath)\n",
    "print(len(train_metas))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# try cutting on views\n",
    "main_metas_filepath = \"/home/christian/code/christian/metadata/genius_hq_metas.jsonl\"\n",
    "main_metas = read_jsonl(main_metas_filepath)\n",
    "print(len(main_metas))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "yt_metas_filepath = \"/home/christian/code/christian/metadata/youtube_music_metas.jsonl\"\n",
    "yt_metas = read_jsonl(yt_metas_filepath)\n",
    "print(len(yt_metas))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "main_metas_map = {meta[\"id\"]: meta for meta in main_metas}\n",
    "yt_metas_map = {meta[\"id\"]: meta for meta in yt_metas}\n",
    "\n",
    "# combine the two into one map\n",
    "combined_metas_map = {**main_metas_map, **yt_metas_map}\n",
    "print(len(combined_metas_map))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(yt_metas[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm\n",
    "count = 0\n",
    "valid_ids = set()\n",
    "for meta in tqdm(train_metas):\n",
    "    add_meta = False\n",
    "    if meta[\"id\"] in combined_metas_map:\n",
    "        # check if youtube views are available\n",
    "        youtube_views = combined_metas_map[meta[\"id\"]].get(\"youtube_views\", None)\n",
    "        genius_views = combined_metas_map[meta[\"id\"]].get(\"genius_views\", None)\n",
    "        view_count = combined_metas_map[meta[\"id\"]].get(\"view_count\", None)\n",
    "\n",
    "        # cut on these views\n",
    "        if youtube_views is not None:\n",
    "            exceed_youtube_views = youtube_views > 100_000\n",
    "        else:\n",
    "            exceed_youtube_views = True\n",
    "        if genius_views is not None:\n",
    "            exceed_genius_views = genius_views > 1_000\n",
    "        else:\n",
    "            exceed_genius_views = True\n",
    "        if view_count is not None:\n",
    "            exceed_view_count = view_count > 100_000\n",
    "        else:\n",
    "            exceed_view_count = True\n",
    "\n",
    "        if exceed_youtube_views and exceed_genius_views and exceed_view_count:\n",
    "            add_meta = True\n",
    "\n",
    "        if add_meta:\n",
    "            valid_ids.add(meta[\"id\"])\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(valid_ids))\n",
    "valid_hrs = len(valid_ids) * 30 / 3600\n",
    "print(f\"{valid_hrs:.2f} hours of data\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 90,
   "metadata": {},
   "outputs": [],
   "source": [
    "filtered_ids = valid_ids"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Stereo width"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# make a histogram of the stereo_width\n",
    "stereo_widths = [float(meta[\"features\"][\"stereo_width\"]) for meta in genius_metas]\n",
    "plt.hist(stereo_widths, bins=250)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# cut out everything below 0.1 and above 0.6\n",
    "min_stereo_width = 0.1\n",
    "max_stereo_width = 0.6  \n",
    "cut_stereo_widths = [width for width in stereo_widths if width > min_stereo_width and width < max_stereo_width]\n",
    "print(f\"{len(cut_stereo_widths)} remaining\")\n",
    "print(f\"removed {len(stereo_widths) - len(cut_stereo_widths)}\")\n",
    "# create a set of the ids that are within the range \n",
    "cut_stereo_widths_ids = set([meta[\"id\"] for meta in genius_metas if float(meta[\"features\"][\"stereo_width\"]) > min_stereo_width and float(meta[\"features\"][\"stereo_width\"]) < max_stereo_width])\n",
    "filtered_ids = filtered_ids.intersection(cut_stereo_widths_ids)\n",
    "print(len(filtered_ids))\n",
    "\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Spectral Centroid"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# make a histogram of the spectral_centroid\n",
    "spectral_centroids = [float(meta[\"features\"][\"spectral_centroid\"]) for meta in genius_metas]\n",
    "spectral_centroids = [centroid for centroid in spectral_centroids if np.isfinite(centroid)]\n",
    "plt.hist(spectral_centroids, bins=250)\n",
    "plt.show()\n",
    "\n",
    "# compute 2 sigma rule\n",
    "two_sigma = np.std(spectral_centroids) * 2\n",
    "min_spectral_centroid = np.mean(spectral_centroids) - two_sigma\n",
    "max_spectral_centroid = np.mean(spectral_centroids) + two_sigma\n",
    "print(min_spectral_centroid, max_spectral_centroid)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "min_spectral_centroid = 1000\n",
    "max_spectral_centroid = 5000\n",
    "cut_spectral_centroids = [centroid for centroid in spectral_centroids if centroid > min_spectral_centroid and centroid < max_spectral_centroid]\n",
    "print(f\"{len(cut_spectral_centroids)} remaining\")\n",
    "print(f\"removed {len(spectral_centroids) - len(cut_spectral_centroids)}\")\n",
    "# create a set of the ids that are within the range \n",
    "cut_spectral_centroids_ids = set([meta[\"id\"] for meta in genius_metas if float(meta[\"features\"][\"spectral_centroid\"]) > min_spectral_centroid and float(meta[\"features\"][\"spectral_centroid\"]) < max_spectral_centroid])\n",
    "filtered_ids = filtered_ids.intersection(cut_spectral_centroids_ids)\n",
    "print(len(filtered_ids))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Total Clips"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "total_clips = [int(meta[\"features\"][\"total_clips\"]) for meta in genius_metas]\n",
    "plt.hist(total_clips, bins=250)\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(max(total_clips), np.median(total_clips))\n",
    "max_total_clips = 100\n",
    "cut_total_clips = [clips for clips in total_clips if clips < max_total_clips]\n",
    "print(f\"{len(cut_total_clips)} remaining\")\n",
    "print(f\"removed {len(total_clips) - len(cut_total_clips)}\")\n",
    "# create a set of the ids that are within the range \n",
    "cut_total_clips_ids = set([meta[\"id\"] for meta in genius_metas if int(meta[\"features\"][\"total_clips\"]) < max_total_clips])\n",
    "filtered_ids = filtered_ids.intersection(cut_total_clips_ids)\n",
    "print(len(filtered_ids))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Loudness factor"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "loudness_factors = [float(meta[\"features\"][\"loudness_factor\"]) for meta in genius_metas]\n",
    "loudness_factors = [loudness for loudness in loudness_factors if np.isfinite(loudness)]\n",
    "plt.hist(loudness_factors, bins=250)\n",
    "\n",
    "# two sigma rule\n",
    "two_sigma = np.std(loudness_factors) * 1.5\n",
    "min_loudness_factor = np.mean(loudness_factors) - two_sigma\n",
    "max_loudness_factor = np.mean(loudness_factors) + two_sigma\n",
    "print(min_loudness_factor, max_loudness_factor)\n",
    "cut_loudness_factors = [loudness for loudness in loudness_factors if loudness > min_loudness_factor and loudness < max_loudness_factor]\n",
    "print(f\"removed {len(loudness_factors) - len(cut_loudness_factors)}\")\n",
    "\n",
    "plt.vlines(min_loudness_factor, 0, 100000, color=\"red\")\n",
    "plt.vlines(max_loudness_factor, 0, 100000, color=\"red\")\n",
    "\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# filter the ids\n",
    "cut_loudness_factors_ids = set([meta[\"id\"] for meta in genius_metas if float(meta[\"features\"][\"loudness_factor\"]) > min_loudness_factor and float(meta[\"features\"][\"loudness_factor\"]) < max_loudness_factor])\n",
    "filtered_ids = filtered_ids.intersection(cut_loudness_factors_ids)\n",
    "print(len(filtered_ids))\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Bass ratio\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "bass_ratios = [float(meta[\"features\"][\"bass_ratio\"]) for meta in genius_metas]\n",
    "bass_ratios = [ratio for ratio in bass_ratios if np.isfinite(ratio)]\n",
    "#plt.hist(bass_ratios, bins=250)\n",
    "#plt.show()\n",
    "\n",
    "# 2 sigma rule\n",
    "two_sigma = np.std(bass_ratios) * 0.1\n",
    "min_bass_ratio = np.mean(bass_ratios) - two_sigma\n",
    "max_bass_ratio = np.mean(bass_ratios) + two_sigma\n",
    "print(min_bass_ratio, max_bass_ratio)\n",
    "\n",
    "min_bass_ratio = -5\n",
    "max_bass_ratio = 5\n",
    "# cut out and replot the histogram\n",
    "cut_bass_ratios = [ratio for ratio in bass_ratios if ratio > min_bass_ratio and ratio < max_bass_ratio]\n",
    "print(f\"{len(cut_bass_ratios)} remaining\")\n",
    "print(f\"removed {len(bass_ratios) - len(cut_bass_ratios)}\")\n",
    "plt.hist(cut_bass_ratios, bins=250)\n",
    "plt.show()\n",
    "\n",
    "# filter the ids\n",
    "cut_bass_ratios_ids = set([meta[\"id\"] for meta in genius_metas if float(meta[\"features\"][\"bass_ratio\"]) > min_bass_ratio and float(meta[\"features\"][\"bass_ratio\"]) < max_bass_ratio])\n",
    "filtered_ids = filtered_ids.intersection(cut_bass_ratios_ids)\n",
    "print(len(filtered_ids))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "mid_ratios = [float(meta[\"features\"][\"mid_ratio\"]) for meta in genius_metas]\n",
    "mid_ratios = [ratio for ratio in mid_ratios if np.isfinite(ratio)]\n",
    "plt.hist(mid_ratios, bins=250)\n",
    "plt.show()\n",
    "\n",
    "min_mid_ratio = -1\n",
    "max_mid_ratio = 2\n",
    "cut_mid_ratios = [ratio for ratio in mid_ratios if ratio > min_mid_ratio and ratio < max_mid_ratio]\n",
    "print(f\"{len(cut_mid_ratios)} remaining\")\n",
    "print(f\"removed {len(mid_ratios) - len(cut_mid_ratios)}\")\n",
    "plt.hist(cut_mid_ratios, bins=250)\n",
    "plt.show()\n",
    "\n",
    "# filter the ids\n",
    "cut_mid_ratios_ids = set([meta[\"id\"] for meta in genius_metas if float(meta[\"features\"][\"mid_ratio\"]) > min_mid_ratio and float(meta[\"features\"][\"mid_ratio\"]) < max_mid_ratio])\n",
    "filtered_ids = filtered_ids.intersection(cut_mid_ratios_ids)\n",
    "print(len(filtered_ids))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "high_ratios = [float(meta[\"features\"][\"high_ratio\"]) for meta in genius_metas]\n",
    "high_ratios = [ratio for ratio in high_ratios if np.isfinite(ratio)]\n",
    "plt.hist(high_ratios, bins=250)\n",
    "plt.show()\n",
    "\n",
    "min_high_ratio = -5\n",
    "max_high_ratio = 5\n",
    "cut_high_ratios = [ratio for ratio in high_ratios if ratio > min_high_ratio and ratio < max_high_ratio]\n",
    "print(f\"{len(cut_high_ratios)} remaining\")\n",
    "print(f\"removed {len(high_ratios) - len(cut_high_ratios)}\")\n",
    "plt.hist(cut_high_ratios, bins=250)\n",
    "plt.show()\n",
    "\n",
    "# filter the ids\n",
    "cut_high_ratios_ids = set([meta[\"id\"] for meta in genius_metas if float(meta[\"features\"][\"high_ratio\"]) > min_high_ratio and float(meta[\"features\"][\"high_ratio\"]) < max_high_ratio])\n",
    "filtered_ids = filtered_ids.intersection(cut_high_ratios_ids)\n",
    "print(len(filtered_ids))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Final subset"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(len(filtered_ids))\n",
    "valid_hrs = len(filtered_ids) * 30 / 3600\n",
    "print(f\"{valid_hrs:.2f} hours of data\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(list(filtered_ids)[100])\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 103,
   "metadata": {},
   "outputs": [],
   "source": [
    "filtered_ids_metas = [{'id': id} for id in filtered_ids]\n",
    "write_jsonl(filtered_ids_metas, \"/app/suno/data/diffusion_mix/vae_25hz_30s/ft_ids.jsonl\")\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# check language distribution\n",
    "english_count = 0\n",
    "non_english_count = 0\n",
    "for meta in tqdm(filtered_ids_metas):\n",
    "    if combined_metas_map[meta[\"id\"]][\"lang\"] == \"en\":\n",
    "        english_count += 1\n",
    "    elif combined_metas_map[meta[\"id\"]][\"lang\"] == \"None\":\n",
    "        english_count += 1\n",
    "    else:\n",
    "        non_english_count += 1\n",
    "print(f\"English: {english_count}, Non-English: {non_english_count}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "combined_metas_map[list(filtered_ids)[100]]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 108,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torchaudio\n",
    "from suno_utils.utils.s3 import read_from_s3"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# listen to random example\n",
    "rand_id = list(filtered_ids)[np.random.randint(0, len(filtered_ids))]\n",
    "print(rand_id)\n",
    "rand_meta = combined_metas_map[rand_id]\n",
    "\n",
    "filepath = rand_meta.get(\"audio_filepath\", None)\n",
    "if filepath is None:\n",
    "    filepath = rand_meta.get(\"s3_filepath\", None)\n",
    "\n",
    "audio, sr = read_from_s3(filepath, read_f=torchaudio.load)\n",
    "start_s = 60.0\n",
    "end_s = start_s + 10.0\n",
    "IPython.display.Audio(audio[:, int(start_s*sr):int(end_s*sr)].numpy(), rate=sr)"
   ]
  },
  {
   "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
}
