{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "import os\n",
    "import IPython\n",
    "import torchaudio\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "from suno_utils.audio import Audio\n",
    "from suno_utils.utils.s3 import read_from_s3"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "spectral_descriptors = {\n",
    "    \"very_dull\": {\n",
    "        \"range\": (0, 500),\n",
    "        \"descriptors\": [\"subby\", \"muddy\", \"sub-frequency dominant\", \"fundamental-focused\", \"subsonic-rich\", \"very dull\"]\n",
    "    },\n",
    "    \"dull\": {\n",
    "        \"range\": (500, 2000),\n",
    "        \"descriptors\": [\"warm\", \"low-mid dominant\", \"rolled-off\", \"dull\"]\n",
    "    },\n",
    "    \"bright\": {\n",
    "        \"range\": (2000, 3500),\n",
    "        \"descriptors\": [\"present\", \"presence-boosted\", \"thin\", \"bright\"]\n",
    "    },\n",
    "    \"very_bright\": {\n",
    "        \"range\": (3500, 24000),\n",
    "        \"descriptors\": [\"sibilant\", \"high-frequency dominant\", \"treble-emphasized\", \"piercing\", \"very bright\"]\n",
    "    }\n",
    "}\n",
    "\n",
    "stereo_descriptors = {\n",
    "    \"mono\": {\n",
    "        \"range\": (0.0, 0.15),\n",
    "        \"descriptors\": [\"mono\", \"centered\", \"point-source\", \"phase-coherent\", \"mid-signal\", \"mono-like\"]\n",
    "    },\n",
    "    \"normal-stereo\": {\n",
    "        \"range\": (0.15, 0.3),\n",
    "        \"descriptors\": [\"balanced stereo\", \"natural\", \"stereo-coherent\", \"standard-width\", \"focused-field\"]\n",
    "    },\n",
    "    \"wide\": {\n",
    "        \"range\": (0.3, 1.0),\n",
    "        \"descriptors\": [\"wide\", \"decorrelated\", \"extreme-width\", \"stereoized\", \"outside-speakers\"]\n",
    "    }\n",
    "}\n",
    "\n",
    "loudness_descriptors = {\n",
    "    \"quiet\": {\n",
    "        \"range\": (-80, -16),\n",
    "        \"descriptors\": [\"soft\", \"low-level\", \"attenuated\", \"low-amplitude\", \"faint\", \"quiet\"]\n",
    "    },\n",
    "    \"normalized\": {\n",
    "        \"range\": (-16, -10),\n",
    "        \"descriptors\": [\"normalized\", \"balanced loudness\", \"standard loudness\"]\n",
    "    },\n",
    "    \"loud\": {\n",
    "        \"range\": (-10, 0),\n",
    "        \"descriptors\": [\"hot\", \"maximized\", \"peak-limited\", \"brick-wall-limited\", \"compressed\", \"loud\"]\n",
    "    }\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_tags_from_audio_production_features(features):\n",
    "    tags = []\n",
    "    categories = []\n",
    "    \n",
    "    descriptor_map = {\n",
    "        \"spectral_centroid\": spectral_descriptors,\n",
    "        \"stereo_width\": stereo_descriptors, \n",
    "        \"loudness_factor\": loudness_descriptors\n",
    "    }\n",
    "    \n",
    "    for name, value in features.items():\n",
    "        if name not in descriptor_map:\n",
    "            continue\n",
    "            \n",
    "        for category, info in descriptor_map[name].items():\n",
    "            if info[\"range\"][0] <= float(value) <= info[\"range\"][1]:\n",
    "                categories.append(category)\n",
    "                tags.extend(random.sample(info[\"descriptors\"], 2))\n",
    "                break\n",
    "                \n",
    "    return list(set(tags)), categories"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 158,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test the tag function \n",
    "rand_meta = random.choice(audio_production_metas)\n",
    "print(rand_meta[\"features\"])\n",
    "source_meta = base_metas_map[rand_meta[\"id\"]]\n",
    "tags = get_tags_from_audio_production_features(rand_meta[\"features\"])\n",
    "print(tags)\n",
    "# get the audio file\n",
    "audio_path = source_meta[\"audio_filepath\"]\n",
    "audio, sr = read_from_s3(audio_path, read_f=torchaudio.load)\n",
    "# first 10 seconds\n",
    "audio = audio[:, sr*30:sr*40]\n",
    "IPython.display.display(IPython.display.Audio(audio, rate=sr))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# lets count how many we have in each subset for each feature\n",
    "for feature_key in [\"spectral_centroid\", \"stereo_width\", \"loudness_factor\"]:\n",
    "    print(feature_key)\n",
    "    if feature_key == \"spectral_centroid\":\n",
    "        for tag, descriptor in spectral_descriptors.items():\n",
    "            print(tag, len([meta for meta in audio_production_metas if descriptor[\"range\"][0] <= float(meta[\"features\"][feature_key]) <= descriptor[\"range\"][1]]))\n",
    "    elif feature_key == \"stereo_width\":\n",
    "        for tag, descriptor in stereo_descriptors.items():\n",
    "            print(tag, len([meta for meta in audio_production_metas if descriptor[\"range\"][0] <= float(meta[\"features\"][feature_key]) <= descriptor[\"range\"][1]]))\n",
    "    elif feature_key == \"loudness_factor\":\n",
    "        for tag, descriptor in loudness_descriptors.items():\n",
    "            print(tag, len([meta for meta in audio_production_metas if descriptor[\"range\"][0] <= float(meta[\"features\"][feature_key]) <= descriptor[\"range\"][1]]))\n",
    "\n",
    "    print()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Enable retina (high-DPI) display\n",
    "%config InlineBackend.figure_format = 'retina'\n",
    "\n",
    "# Set the default figure size and DPI\n",
    "plt.rcParams['figure.figsize'] = [8, 6]\n",
    "plt.rcParams['figure.dpi'] = 100\n",
    "plt.rcParams['savefig.dpi'] = 300\n",
    "\n",
    "# Improve font quality\n",
    "plt.rcParams['pdf.fonttype'] = 42\n",
    "plt.rcParams['ps.fonttype'] = 42\n",
    "\n",
    "# Use high-quality rendering\n",
    "plt.rcParams['image.interpolation'] = 'nearest'\n",
    "plt.rcParams['image.resample'] = True"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "audio_production_metas = read_jsonl(\n",
    "    \"/home/christian/code/christian/metadata/genius_hq_metas_audio_production.jsonl\"\n",
    ")\n",
    "print(f\"Loaded {len(audio_production_metas)} audio production metas\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 16,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "from sklearn.preprocessing import StandardScaler\n",
    "from sklearn.ensemble import IsolationForest\n",
    "\n",
    "def classify_audio_quality(features, contamination=0.1):\n",
    "    \"\"\"\n",
    "    Classifies audio quality based on outliers in feature dimensions\n",
    "    Returns: quality scores (-1: low quality, 1: high quality)\n",
    "    \"\"\"\n",
    "    scaler = StandardScaler()\n",
    "    scaled_features = scaler.fit_transform(features)\n",
    "    \n",
    "    # Get outlier scores per feature dimension\n",
    "    quality_scores = np.zeros((features.shape[0], features.shape[1]))\n",
    "    for i in range(features.shape[1]):\n",
    "        detector = IsolationForest(contamination=contamination, random_state=42, verbose=2)\n",
    "        quality_scores[:, i] = detector.fit_predict(scaled_features[:, i].reshape(-1, 1))\n",
    "    \n",
    "    # Average across features\n",
    "    return np.mean(quality_scores, axis=1)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "# Collect values and track IDs\n",
    "feature_values = {}\n",
    "ids = []\n",
    "valid_indices = []  # Track which samples were valid\n",
    "\n",
    "for idx, ap_meta in enumerate(tqdm(audio_production_metas)):\n",
    "    features = ap_meta[\"features\"]\n",
    "    valid_sample = True\n",
    "    \n",
    "    # First check if this sample has valid values for all features\n",
    "    temp_values = {}\n",
    "    for feature_key in features.keys():\n",
    "        if feature_key not in [\"id\", \"spectral_character\"]:\n",
    "            try:\n",
    "                value = float(features[feature_key])\n",
    "                if not (np.isnan(value) or np.isinf(value)):\n",
    "                    temp_values[feature_key] = value\n",
    "                else:\n",
    "                    valid_sample = False\n",
    "                    break\n",
    "            except (ValueError, TypeError):\n",
    "                valid_sample = False\n",
    "                break\n",
    "    \n",
    "    # If sample is valid, add its values and ID\n",
    "    if valid_sample and temp_values:\n",
    "        for feature_key, value in temp_values.items():\n",
    "            if feature_key not in feature_values:\n",
    "                feature_values[feature_key] = []\n",
    "            feature_values[feature_key].append(value)\n",
    "        ids.append(ap_meta[\"id\"])\n",
    "        valid_indices.append(idx)\n",
    "\n",
    "print(f\"Number of features: {len(feature_values)}\")\n",
    "feature_values = np.array(list(feature_values.values()))\n",
    "print(f\"Feature array shape: {feature_values.shape}\")\n",
    "print(f\"Number of valid samples: {len(ids)}\")\n",
    "\n",
    "# feature_values.T will give shape (n_samples, n_features)\n",
    "# ids contains corresponding IDs in same order\n",
    "# valid_indices contains original indices in audio_production_metas"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# clean up the feature values by removing infs and nans but keep the shape\n",
    "quality_scores = classify_audio_quality(feature_values.T)\n",
    "print(quality_scores)\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.hist(quality_scores, bins=100)\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 48,
   "metadata": {},
   "outputs": [],
   "source": [
    "ranges = np.arange(-1.0, 1.25, 0.25)\n",
    "samples = {}\n",
    "\n",
    "for i in range(len(ranges) - 1):\n",
    "   range_min, range_max = ranges[i], ranges[i+1]\n",
    "   mask = (quality_scores >= range_min) & (quality_scores < range_max)\n",
    "   indices = np.where(mask)[0]\n",
    "   \n",
    "   if len(indices) > 0:\n",
    "       sample_indices = np.random.choice(indices, min(3, len(indices)), replace=False)\n",
    "       samples[f\"{range_min:.2f} to {range_max:.2f}\"] = [\n",
    "           {\n",
    "               'score': quality_scores[idx],\n",
    "               'id': ids[idx],\n",
    "               'meta': audio_production_metas[valid_indices[idx]]\n",
    "           }\n",
    "           for idx in sample_indices\n",
    "       ]\n",
    "\n",
    "for range_name, examples in samples.items():\n",
    "   print(f\"\\nRange {range_name}:\")\n",
    "   for ex in examples:\n",
    "       print(f\"  Score: {ex['score']:.3f}, ID: {ex['id']}\")\n",
    "       # listen to the audio by getting the audio_filepath \n",
    "       audio_path = base_metas_map[ex[\"id\"]][\"audio_filepath\"]\n",
    "       audio, sr = read_from_s3(audio_path, read_f=torchaudio.load)\n",
    "       audio = audio[:, :sr*30]\n",
    "       IPython.display.display(IPython.display.Audio(audio, rate=sr))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# get the id of the lowest quality scores\n",
    "lowest_quality_ids = np.argsort(quality_scores)[:3]\n",
    "for idx in lowest_quality_ids:\n",
    "    print(quality_scores[idx])\n",
    "    meta_id = audio_production_metas[idx][\"id\"]\n",
    "    print(base_metas_map[meta_id][\"audio_filepath\"])\n",
    "    audio, sr = read_from_s3(base_metas_map[meta_id][\"audio_filepath\"], read_f=torchaudio.load)\n",
    "    audio = audio[:, :sr*30]\n",
    "    IPython.display.display(IPython.display.Audio(audio, rate=sr))\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "audio_production_metas = read_jsonl(\n",
    "    \"/home/christian/code/christian/metadata/youtube_music_metas_audio_production.jsonl\"\n",
    ")\n",
    "print(f\"Loaded {len(audio_production_metas)} audio production metas\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "base_metas = read_jsonl(\n",
    "    \"/home/christian/code/christian/metadata/genius_hq_metas.jsonl\"\n",
    ")\n",
    "print(f\"Loaded {len(base_metas)} base metas\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "metadata": {},
   "outputs": [],
   "source": [
    "base_metas_map = {meta[\"id\"]: meta for meta in base_metas}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(audio_production_metas[0][\"features\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import random\n",
    "# lets listen to random audios with certain spectral centroid\n",
    "spectral_centroid_ranges = [\n",
    "    (0, 250),\n",
    "    (250, 2000),\n",
    "    (2000, 5000),\n",
    "    (5000, 22050),\n",
    "]\n",
    "\n",
    "for sc_range in spectral_centroid_ranges:\n",
    "    print(sc_range)\n",
    "    in_range_metas = []\n",
    "    # find and audio production meta with spectral centroid in the range\n",
    "    for meta in audio_production_metas:\n",
    "        sc = float(meta[\"features\"][\"spectral_centroid\"])\n",
    "        if sc_range[0] <= sc <= sc_range[1]:\n",
    "            in_range_metas.append(meta)\n",
    "\n",
    "    # select three random metas from the in range metas\n",
    "    in_range_metas = random.sample(in_range_metas, 3)\n",
    "    for meta in in_range_metas:\n",
    "        print(meta[\"features\"][\"spectral_centroid\"])\n",
    "        audio_path = base_metas_map[meta[\"id\"]][\"audio_filepath\"]\n",
    "        print(audio_path)\n",
    "        audio, sr = read_from_s3(audio_path, read_f=torchaudio.load)\n",
    "        IPython.display.display(IPython.display.Audio(audio, rate=sr))\n",
    "        break\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# clipped samples histogram\n",
    "spectral_centroids = [float(meta[\"features\"][\"spectral_centroid\"]) for meta in audio_production_metas]\n",
    "spectral_centroids = [sc for sc in spectral_centroids if np.isfinite(sc)]\n",
    "\n",
    "\n",
    "# add top 5% and bottom 5% to the histogram\n",
    "plt.hist(spectral_centroids, bins=1000)\n",
    "#plt.axvline(np.percentile(spectral_centroids, 95), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.axvline(np.percentile(spectral_centroids, 5), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.yscale('log')\n",
    "plt.show()\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# clipped samples histogram\n",
    "spectral_centroids = [float(meta[\"features\"][\"loudness_factor\"]) for meta in audio_production_metas]\n",
    "spectral_centroids = [sc for sc in spectral_centroids if np.isfinite(sc)]\n",
    "\n",
    "\n",
    "# add top 5% and bottom 5% to the histogram\n",
    "plt.hist(spectral_centroids, bins=1000)\n",
    "#plt.axvline(np.percentile(spectral_centroids, 95), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.axvline(np.percentile(spectral_centroids, 5), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.yscale('log')\n",
    "plt.show()\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# clipped samples histogram\n",
    "spectral_centroids = [float(meta[\"features\"][\"stereo_width\"]) for meta in audio_production_metas]\n",
    "spectral_centroids = [sc for sc in spectral_centroids if np.isfinite(sc)]\n",
    "\n",
    "\n",
    "# add top 5% and bottom 5% to the histogram\n",
    "plt.hist(spectral_centroids, bins=100)\n",
    "#plt.axvline(np.percentile(spectral_centroids, 95), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.axvline(np.percentile(spectral_centroids, 5), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.yscale('log')\n",
    "plt.show()\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "metadata": {},
   "outputs": [],
   "source": [
    "def create_float_bins(values, bin_width=0.01):\n",
    "    \"\"\"\n",
    "    Creates bins for float values rounded to 2 decimal places.\n",
    "    Returns bin edges that ensure all values are properly captured.\n",
    "    \n",
    "    Args:\n",
    "        values (array-like): Float values rounded to 2 decimal places\n",
    "        bin_width (float): Width of each bin (default 0.01)\n",
    "        \n",
    "    Returns:\n",
    "        list: Bin edges that cover all values\n",
    "    \"\"\"\n",
    "    min_val = min(values)\n",
    "    max_val = max(values)\n",
    "    \n",
    "    # Adjust min/max to ensure we capture all values\n",
    "    min_edge = (min_val // bin_width) * bin_width\n",
    "    max_edge = ((max_val // bin_width) + 1) * bin_width\n",
    "    \n",
    "    # Generate bin edges\n",
    "    bins = []\n",
    "    current = min_edge\n",
    "    while current <= max_edge:\n",
    "        bins.append(round(current, 2))\n",
    "        current += bin_width\n",
    "        \n",
    "    return bins\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# clipped samples histogram\n",
    "stereo_width = [float(meta[\"features\"][\"stereo_width\"]) for meta in audio_production_metas]\n",
    "stereo_width = [sw for sw in stereo_width if np.isfinite(sw)]\n",
    "\n",
    "bins = create_float_bins(stereo_width)\n",
    "\n",
    "# add top 5% and bottom 5% to the histogram\n",
    "plt.hist(stereo_width, bins=bins)\n",
    "#plt.axvline(np.percentile(stereo_width, 95), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.axvline(np.percentile(spectral_centroids, 5), color='k', linestyle='dashed', linewidth=2)\n",
    "#plt.yscale('log')\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# calc the top and bottom 5%\n",
    "top_5_percent = np.percentile(spectral_centroids, 95)\n",
    "bottom_5_percent = np.percentile(spectral_centroids, 5)\n",
    "print(bottom_5_percent, top_5_percent)\n",
    "\n",
    "# count how many samples are in the top and bottom 5%\n",
    "top_5_percent_count = len([sc for sc in spectral_centroids if sc > top_5_percent])\n",
    "bottom_5_percent_count = len([sc for sc in spectral_centroids if sc < bottom_5_percent])\n",
    "print(top_5_percent_count, bottom_5_percent_count)\n",
    "\n",
    "print(min(spectral_centroids), max(spectral_centroids), np.median(spectral_centroids))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# clipped samples histogram\n",
    "stereo_width = [float(meta[\"features\"][\"stereo_width\"]) for meta in audio_production_metas]\n",
    "stereo_width = [sw for sw in stereo_width if np.isfinite(sw)]\n",
    "\n",
    "# add top 5% and bottom 5% to the histogram\n",
    "plt.hist(stereo_width, bins=250)\n",
    "plt.axvline(np.percentile(stereo_width, 95), color='k', linestyle='dashed', linewidth=2)\n",
    "plt.axvline(np.percentile(stereo_width, 5), color='k', linestyle='dashed', linewidth=2)\n",
    "plt.show()\n",
    "\n",
    "print(min(stereo_width), max(stereo_width), np.median(stereo_width))\n",
    "print(np.percentile(stereo_width, 95), np.percentile(stereo_width, 5))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# loudness factor histogram\n",
    "loudness_factor = [float(meta[\"features\"][\"loudness_factor\"]) for meta in audio_production_metas]\n",
    "loudness_factor = [lf for lf in loudness_factor if np.isfinite(lf)]\n",
    "print(min(loudness_factor), max(loudness_factor), np.median(loudness_factor))\n",
    "plt.hist(loudness_factor, bins=250)\n",
    "plt.yscale('log')\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# listen to audio for the top 5 loudness factors\n",
    "sorted_metas = sorted(audio_production_metas, key=lambda x: float(x[\"features\"][\"loudness_factor\"]), reverse=False)\n",
    "top_lf = sorted_metas[:10]\n",
    "for meta in top_lf:\n",
    "    print(meta[\"features\"][\"loudness_factor\"])\n",
    "    audio_path = base_metas_map[meta[\"id\"]][\"audio_filepath\"]\n",
    "    print(audio_path)\n",
    "    audio, sr = read_from_s3(audio_path, read_f=torchaudio.load)\n",
    "    print(audio.min(), audio.max())\n",
    "    IPython.display.display(IPython.display.Audio(audio, rate=sr))\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# xy scatter plot of stereo width and spectral centroid\n",
    "stereo_width = [float(meta[\"features\"][\"stereo_width\"]) for meta in audio_production_metas]\n",
    "spectral_centroids = [float(meta[\"features\"][\"spectral_centroid\"]) for meta in audio_production_metas]\n",
    "\n",
    "plt.scatter(stereo_width, spectral_centroids, s=1, alpha=0.1)\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import pandas as pd\n",
    "import plotly.express as px\n",
    "\n",
    "import plotly.io as pio\n",
    "pio.renderers.default = 'notebook'\n",
    "\n",
    "\n",
    "# Sample data from the user's code\n",
    "stereo_width = [float(meta[\"features\"][\"stereo_width\"]) for meta in audio_production_metas]\n",
    "spectral_centroids = [float(meta[\"features\"][\"spectral_centroid\"]) for meta in audio_production_metas]\n",
    "loudness_factor = [float(meta[\"features\"][\"loudness_factor\"]) for meta in audio_production_metas]\n",
    "# if -inf found in loudness factor, replace with -80\n",
    "loudness_factor = [lf if lf != float('-inf') else -80.0 for lf in loudness_factor]\n",
    "# if nan found in spectral centroid, replace with 0\n",
    "spectral_centroids = [sc if np.isfinite(sc) else 0.0 for sc in spectral_centroids]\n",
    "\n",
    "\n",
    "# Normalize the features using z-score\n",
    "def z_score_normalize(data):\n",
    "    mean = np.mean(data)\n",
    "    std = np.std(data)\n",
    "    print(mean, std)\n",
    "    data = [(x - mean) / std for x in data]\n",
    "    # repalce infs and nans\n",
    "    data = [x if np.isfinite(x) else 0.0 for x in data]\n",
    "    return data\n",
    "\n",
    "stereo_width_norm = z_score_normalize(stereo_width)\n",
    "spectral_centroids_norm = z_score_normalize(spectral_centroids)\n",
    "loudness_factor_norm = z_score_normalize(loudness_factor)\n",
    "\n",
    "# Prepare the data as a DataFrame for Plotly\n",
    "df = pd.DataFrame({\n",
    "    \"Stereo Width (z-score)\": stereo_width_norm,\n",
    "    \"Spectral Centroid (z-score)\": spectral_centroids_norm,\n",
    "    \"Loudness Factor (z-score)\": loudness_factor_norm\n",
    "})\n",
    "\n",
    "if True:\n",
    "    # Create the interactive 3D scatter plot\n",
    "    fig = px.scatter_3d(\n",
    "        df,\n",
    "        x=\"Stereo Width (z-score)\",\n",
    "        y=\"Spectral Centroid (z-score)\",\n",
    "        z=\"Loudness Factor (z-score)\",\n",
    "        opacity=0.7,\n",
    "        title=\"3D Scatter Plot of Audio Features\",\n",
    "        labels={\n",
    "            \"Stereo Width (z-score)\": \"Stereo Width (z-score)\",\n",
    "            \"Spectral Centroid (z-score)\": \"Spectral Centroid (z-score)\",\n",
    "            \"Loudness Factor (z-score)\": \"Loudness Factor (z-score)\"\n",
    "        }\n",
    "    )\n",
    "\n",
    "    # Show the plot\n",
    "    fig.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# listen to audio for the top 5 spectral centroids\n",
    "sorted_metas = sorted(audio_production_metas, key=lambda x: float(x[\"features\"][\"spectral_centroid\"]), reverse=True)\n",
    "top_sc = sorted_metas[:10]\n",
    "for meta in top_sc:\n",
    "    print(meta[\"features\"][\"spectral_centroid\"])\n",
    "    audio_path = base_metas_map[meta[\"id\"]][\"audio_filepath\"]\n",
    "    print(audio_path)\n",
    "    audio, sr = read_from_s3(audio_path, read_f=torchaudio.load)\n",
    "    IPython.display.display(IPython.display.Audio(audio, rate=sr))\n",
    "    \n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# bass ratio histogram\n",
    "bass_ratio = [float(meta[\"features\"][\"bass_ratio\"]) for meta in audio_production_metas]\n",
    "bass_ratio = [br for br in bass_ratio if np.isfinite(br)]\n",
    "print(min(bass_ratio), max(bass_ratio), np.median(bass_ratio))\n",
    "# add top 5% and bottom 5% to the histogram\n",
    "plt.hist(bass_ratio, bins=100)\n",
    "plt.axvline(np.percentile(bass_ratio, 95), color='k', linestyle='dashed', linewidth=2)\n",
    "plt.axvline(np.percentile(bass_ratio, 5), color='k', linestyle='dashed', linewidth=2)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 39,
   "metadata": {},
   "outputs": [],
   "source": [
    "# listen to audio for the top 5 bass ratios\n",
    "sorted_metas = sorted(audio_production_metas, key=lambda x: float(x[\"features\"][\"bass_ratio\"]), reverse=True)\n",
    "top_br = sorted_metas[:10]\n",
    "for meta in top_br:\n",
    "    print(meta[\"features\"][\"bass_ratio\"])\n",
    "    audio_path = base_metas_map[meta[\"id\"]][\"audio_filepath\"]\n",
    "    print(audio_path)\n",
    "    audio, sr = read_from_s3(audio_path, read_f=torchaudio.load)\n",
    "    IPython.display.display(IPython.display.Audio(audio, rate=sr))\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# top five most clipped samples\n",
    "clipped_samples = [float(meta[\"features\"][\"total_clips\"]) for meta in audio_production_metas]\n",
    "clipped_samples = [cs for cs in clipped_samples if np.isfinite(cs)]\n",
    "clipped_samples = [cs for cs in clipped_samples if cs > 2]\n",
    "print(min(clipped_samples), max(clipped_samples), np.median(clipped_samples))\n",
    "# add top 5% and bottom 5% to the histogram\n",
    "plt.hist(clipped_samples, bins=100)\n",
    "plt.show()"
   ]
  },
  {
   "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
}
