{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import json\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"3\"\n",
    "import glob\n",
    "import torch\n",
    "import random\n",
    "import torchaudio\n",
    "import numpy as np\n",
    "import soundfile as sf\n",
    "import pyloudnorm as pyln\n",
    "import torchaudio.functional as F\n",
    "\n",
    "from joblib import Parallel, delayed\n",
    "from tqdm import tqdm\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "from suno_utils.tasks.shimmerscore import shimmerscore"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [],
   "source": [
    "def calculate_stereo_width(waveform):\n",
    "    # Split into left and right channels\n",
    "    left = waveform[0]\n",
    "    right = waveform[1]\n",
    "  \n",
    "    # Compute mid/side representation\n",
    "    mid = (left + right) / 2\n",
    "    side = (left - right) / 2\n",
    "    \n",
    "    # Compute RMS energy of mid and side channels\n",
    "    mid_energy = torch.sqrt(torch.mean(mid ** 2))\n",
    "    side_energy = torch.sqrt(torch.mean(side ** 2))\n",
    "    \n",
    "    # Compute stereo width based on mid/side ratio\n",
    "    # Normalize to range 0-1 using sigmoid-like function\n",
    "    width_ratio = (side_energy / (mid_energy + 1e-8)).item()\n",
    "    stereo_width = 2 * (1 / (1 + np.exp(-width_ratio)) - 0.5)\n",
    "\n",
    "    return stereo_width\n",
    "\n",
    "def calculate_loudness(waveform, sr):\n",
    "    meter = pyln.Meter(sr)\n",
    "    lufs_db = meter.integrated_loudness(waveform)\n",
    "    return lufs_db\n",
    "\n",
    "def calculate_loudness_factor(waveform, sr):\n",
    "    meter = pyln.Meter(sr)\n",
    "    normalized_waveform = waveform / np.clip(np.max(np.abs(waveform)), 1e-10, None)\n",
    "    lufs_db = meter.integrated_loudness(normalized_waveform)\n",
    "    return lufs_db\n",
    "\n",
    "def calculate_possible_clipped_samples(waveform):\n",
    "    # Convert to numpy if it's a torch tensor\n",
    "    if isinstance(waveform, torch.Tensor):\n",
    "        waveform_np = waveform.numpy()\n",
    "    else:\n",
    "        waveform_np = waveform\n",
    "    \n",
    "    # Count samples that are at or above the clipping threshold\n",
    "    return np.sum(np.abs(waveform_np) >= 1.0).item() if isinstance(np.sum(np.abs(waveform_np) >= 1.0), torch.Tensor) else np.sum(np.abs(waveform_np) >= 1.0)\n",
    "\n",
    "def calculate_average_spectrum_db(waveform, n_fft=16384, hop_length=8192):\n",
    "    # Keep as torch tensor or convert to torch tensor if it's numpy\n",
    "    if not isinstance(waveform, torch.Tensor):\n",
    "        waveform = torch.from_numpy(waveform)\n",
    "\n",
    "    # Calculate the average spectrum using STFT for efficiency\n",
    "    n_fft = 2048  # Choose an appropriate FFT size\n",
    "    hop_length = n_fft // 4  # Standard hop length\n",
    "    \n",
    "    # Compute STFT using torch\n",
    "    if waveform.dim() > 1:\n",
    "        # For stereo, compute STFT for each channel\n",
    "        stft_results = []\n",
    "        for channel in range(waveform.shape[0]):\n",
    "            stft = torch.stft(\n",
    "                waveform[channel], \n",
    "                n_fft=n_fft, \n",
    "                hop_length=hop_length, \n",
    "                window=torch.hann_window(n_fft), \n",
    "                return_complex=True\n",
    "            )\n",
    "            # Get magnitude\n",
    "            stft_magnitude = torch.abs(stft)\n",
    "            stft_results.append(stft_magnitude)\n",
    "        \n",
    "        # Average across time frames for each channel\n",
    "        magnitude_spectrum = torch.stack([torch.mean(stft, dim=1) for stft in stft_results])\n",
    "    else:\n",
    "        # For mono\n",
    "        stft = torch.stft(\n",
    "            waveform, \n",
    "            n_fft=n_fft, \n",
    "            hop_length=hop_length, \n",
    "            window=torch.hann_window(n_fft), \n",
    "            return_complex=True\n",
    "        )\n",
    "        # Get magnitude\n",
    "        stft_magnitude = torch.abs(stft)\n",
    "        magnitude_spectrum = torch.mean(stft_magnitude, dim=1)\n",
    "    \n",
    "    # Convert to dB scale\n",
    "    spectrum_db = 20 * torch.log10(magnitude_spectrum + 1e-10)  # Adding small value to avoid log(0)\n",
    "    \n",
    "    # Convert to numpy for consistency with the rest of the code\n",
    "    return spectrum_db.numpy()\n",
    "    \n",
    "def calculate_spectrum_evolution(waveform, sr, n_fft=16384, hop_length=8192):\n",
    "    # compute spectrum for first 30s\n",
    "    waveform_first = waveform[:, :30 * sr]\n",
    "    spectrum_first = calculate_average_spectrum_db(waveform_first, n_fft, hop_length)\n",
    "\n",
    "    # compute spectrum for last 30s\n",
    "    waveform_last = waveform[:, -30 * sr:]\n",
    "    spectrum_last = calculate_average_spectrum_db(waveform_last, n_fft, hop_length)\n",
    "\n",
    "    return spectrum_first, spectrum_last\n",
    "\n",
    "def get_ear_score(waveform, sr):\n",
    "    # crop waveforn to max of 4min\n",
    "    waveform = waveform[:, :int(4 * 60 * sr)]\n",
    "    with torch.no_grad():\n",
    "        ear_score = ear_model.get_score(waveform, sample_rate=sr)\n",
    "    return ear_score\n",
    "\n",
    "def get_shimmerscore(waveform, sr):\n",
    "    # crop waveforn to max of 4min\n",
    "    waveform = waveform[:, :int(4 * 60 * sr)]\n",
    "    shimmer_score = shimmerscore(waveform, sr)\n",
    "    return shimmer_score\n",
    "\n",
    "# so let's measure some things \n",
    "# loudness, average spectrum\n",
    "# peak values\n",
    "# ear scores \n",
    "\n",
    "\n",
    "def analyze_audio(filepath):\n",
    "    audio, sr = torchaudio.load(filepath)\n",
    "\n",
    "    if sr != 48000:\n",
    "        audio = torchaudio.functional.resample(audio, sr, 48000)\n",
    "\n",
    "    # calculate loudness\n",
    "    lufs_db = calculate_loudness(audio.permute(1, 0).numpy(), sr)\n",
    "    lufs_db_factor = calculate_loudness_factor(audio.permute(1, 0).numpy(), sr)\n",
    "\n",
    "    # calculate stereo width\n",
    "    stereo_width = calculate_stereo_width(audio)\n",
    "\n",
    "    # clipped samples\n",
    "    clipped_samples = calculate_possible_clipped_samples(audio)\n",
    "\n",
    "    # average spectrum\n",
    "    average_spectrum_db = calculate_average_spectrum_db(audio)\n",
    "\n",
    "    # spectrum evolution\n",
    "    spectrum_first, spectrum_last = calculate_spectrum_evolution(audio, sr)\n",
    "\n",
    "    # shimmer score\n",
    "    shimmer_score = get_shimmerscore(audio, sr)\n",
    "\n",
    "    return {\n",
    "        \"lufs_db\": lufs_db,\n",
    "        \"lufs_db_factor\": lufs_db_factor,\n",
    "        \"stereo_width\": stereo_width,\n",
    "        \"clipped_samples\": clipped_samples,\n",
    "        \"average_spectrum_db\": average_spectrum_db,\n",
    "        \"average_spectrum_db_first\": spectrum_first,\n",
    "        \"average_spectrum_db_last\": spectrum_last,\n",
    "        \"shimmer_score\": shimmer_score,\n",
    "    }\n",
    "        "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "LOCAL_DATA_DIR = \"/app/suno/data/slakh2100_flac_redux/\"\n",
    "S3_DATA_DIR = \"s3://suno-data/datasets/slakh2100_flac_redux\"\n",
    "\n",
    "mix_filepaths = glob.glob(os.path.join(LOCAL_DATA_DIR, \"**/*/\"))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "results = {}\n",
    "\n",
    "from tqdm import tqdm\n",
    "from joblib import Parallel, delayed\n",
    "\n",
    "def process_file(filepath):\n",
    "    clip_id = os.path.basename(filepath).split(\"_\")[0]\n",
    "    model_name = \"_\".join(os.path.basename(filepath).split(\"_\")[1:]).replace(\".mp3\", \"\")\n",
    "    \n",
    "    stats = analyze_audio(filepath)\n",
    "    \n",
    "    return {\n",
    "        \"model_name\": model_name,\n",
    "        \"clip_id\": clip_id,\n",
    "        \"filepath\": filepath,\n",
    "        **stats\n",
    "    }\n",
    "\n",
    "# Limit to first 21 files for testing\n",
    "limited_filepaths = filepaths\n",
    "\n",
    "# Process files in parallel\n",
    "processed_results = Parallel(n_jobs=-1)(\n",
    "    delayed(process_file)(filepath) for filepath in tqdm(limited_filepaths)\n",
    ")\n",
    "\n",
    "# Organize results by model\n",
    "for result in processed_results:\n",
    "    model_name = result.pop(\"model_name\")\n",
    "    if model_name not in results:\n",
    "        results[model_name] = []\n",
    "    results[model_name].append(result)\n",
    "\n"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_diff",
   "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.12.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
