{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "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": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load ear model\n",
    "from suno_utils.tasks.ear import load_model\n",
    "ear_model = load_model(\"s3://suno-data/christian/checkpoints/ear/ear_v2_s3080.pt\", compile=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# temp let me copy some files from s3 to local \n",
    "\n",
    "\n",
    "work_items = read_jsonl(\n",
    "    \"/home/christian/code/christian/metadata/discogs_subset_sampled_metas.jsonl\"\n",
    ")\n",
    "print(len(work_items))\n",
    "# shuffle the work items\n",
    "random.shuffle(work_items)\n",
    "\n",
    "work_items = work_items[:100]\n",
    "\n",
    "\n",
    "out_dir = \"/home/christian/audio/discogs_subset_sampled_metas\"\n",
    "os.makedirs(out_dir, exist_ok=True)\n",
    "\n",
    "def download_file(w):\n",
    "    filepath = w[\"s3_filepath\"]\n",
    "    ext = filepath.split(\".\")[-1]\n",
    "    meta_id = w[\"id\"]\n",
    "    out_filepath = os.path.join(out_dir, f\"{meta_id}.{ext}\")\n",
    "    if os.path.exists(out_filepath):\n",
    "        return\n",
    "    os.system(f\"aws s3 cp {filepath} {out_filepath}\")\n",
    "\n",
    "# Use joblib for parallel downloads\n",
    "Parallel(n_jobs=8, backend=\"threading\")(\n",
    "    delayed(download_file)(w) for w in tqdm(work_items)\n",
    ")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# start with a main directory of audio files \n",
    "#audio_dir = \"/home/christian/audio/pencil-helmet-window\"\n",
    "#audio_dir = \"/home/christian/audio/farm-speaker-desk\"\n",
    "#audio_dir = \"/home/christian/audio/walking-flowers\"\n",
    "audio_dir = \"/home/christian/audio/auk-clips-up-u-2\"\n",
    "\n",
    "filepaths = glob.glob(os.path.join(audio_dir, \"**\", \"*.mp3\"))\n",
    "print(len(filepaths))\n",
    "\n",
    "use_wav_reference = False\n",
    "# ideally we also have a reference distribution from the training data \n",
    "if use_wav_reference:\n",
    "    reference_dir = \"/home/christian/audio/reference-audio-wav\"\n",
    "    reference_filepaths = glob.glob(os.path.join(reference_dir, \"*.wav\"))\n",
    "    print(len(reference_filepaths))\n",
    "else:\n",
    "    reference_dir  = \"/home/christian/audio/discogs_subset_sampled_metas\"\n",
    "    reference_filepaths = glob.glob(os.path.join(reference_dir, \"*.webm\"))\n",
    "    codec_reference_filepaths = glob.glob(os.path.join(reference_dir, \"*.mp3\"))\n",
    "\n",
    "    print(len(reference_filepaths), len(codec_reference_filepaths))\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import pandas as pd\n",
    "\n",
    "# load existing results \n",
    "df = pd.read_csv('outputs/audio_analysis_results_5.csv')\n",
    "\n",
    "# now filter the filepaths to exclude the ones that are already in the dataframe\n",
    "filepaths = [filepath for filepath in filepaths if filepath not in df['filepath'].values]\n",
    "\n",
    "print(len(filepaths))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 24,
   "metadata": {},
   "outputs": [],
   "source": [
    "import math\n",
    "import numpy as np\n",
    "import scipy.signal as signal\n",
    "\n",
    "def biquad_highshelf(fs, f0, dB_gain, Q=0.707):\n",
    "    A = 10**(dB_gain / 40)\n",
    "    w0 = 2 * math.pi * f0 / fs\n",
    "    alpha = math.sin(w0) / (2 * Q)\n",
    "    cos_w0 = math.cos(w0)\n",
    "\n",
    "    b0 =    A*( (A+1) + (A-1)*cos_w0 + 2*math.sqrt(A)*alpha )\n",
    "    b1 = -2*A*( (A-1) + (A+1)*cos_w0 )\n",
    "    b2 =    A*( (A+1) + (A-1)*cos_w0 - 2*math.sqrt(A)*alpha )\n",
    "    a0 =        (A+1) - (A-1)*cos_w0 + 2*math.sqrt(A)*alpha\n",
    "    a1 =  2*( (A-1) - (A+1)*cos_w0 )\n",
    "    a2 =        (A+1) - (A-1)*cos_w0 - 2*math.sqrt(A)*alpha\n",
    "\n",
    "    return [b0/a0, b1/a0, b2/a0, a1/a0, a2/a0]\n",
    "\n",
    "def apply_highshelf_filter(audio, fs, f0, dB_gain, Q=0.707):\n",
    "    \"\"\"\n",
    "    Apply a biquad high-shelf filter to the input audio.\n",
    "\n",
    "    Parameters:\n",
    "    - audio: 1D NumPy array of audio samples\n",
    "    - fs: Sampling rate in Hz\n",
    "    - f0: Cutoff frequency in Hz\n",
    "    - dB_gain: Gain in decibels (positive for boost, negative for cut)\n",
    "    - Q: Quality factor (default 0.707)\n",
    "\n",
    "    Returns:\n",
    "    - Filtered audio as a NumPy array\n",
    "    \"\"\"\n",
    "    b0, b1, b2, a1, a2 = biquad_highshelf(fs, f0, dB_gain, Q)\n",
    "    b = [b0, b1, b2]\n",
    "    a = [1.0, a1, a2]\n",
    "    return signal.lfilter(b, a, audio)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import IPython\n",
    "\n",
    "\n",
    "# we can create new audio files with some post-processing \n",
    "source_model = \"v45_2b_step_2_600_000\"\n",
    "# find all files that have this model name \n",
    "source_filepaths = glob.glob(os.path.join(audio_dir, \"**\", f\"*{source_model}.mp3\"))\n",
    "print(len(source_filepaths))\n",
    "\n",
    "# now we can create new audio files with some post-processing \n",
    "new_model_name = \"v45_2b_step_2_600_000_hshelf_6db\"\n",
    "\n",
    "def process_file(filepath, source_model, new_model_name):\n",
    "    # load the audio file\n",
    "    waveform, sr = torchaudio.load(filepath)\n",
    "    # apply the highshelf filter\n",
    "    filtered_waveform = apply_highshelf_filter(waveform.numpy(), sr, 8000, 6)\n",
    "    filtered_waveform = torch.from_numpy(filtered_waveform)\n",
    "    # save the filtered audio file\n",
    "    torchaudio.save(filepath.replace(source_model, new_model_name), filtered_waveform, sr)\n",
    "    return filepath\n",
    "\n",
    "# Use joblib to parallelize the processing\n",
    "results = Parallel(n_jobs=-1)(\n",
    "    delayed(process_file)(filepath, source_model, new_model_name) \n",
    "    for filepath in tqdm(source_filepaths)\n",
    ")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# apply codec\n",
    "# load codec\n",
    "from suno_utils.tasks.dac_vae_fixed_25hz import (\n",
    "    preload_models as preload_codec_models,\n",
    "    decode as codec_decode,\n",
    "    encode as codec_encode,\n",
    "    decode_stream_to_full_audio,\n",
    ")\n",
    "\n",
    "codec_filepath = \"s3://suno-data/minz/models/dac_vae_tuned_25hz.pth\"\n",
    "_ = preload_codec_models(codec_filepath)\n",
    "\n",
    "source_model = \"reference\"\n",
    "new_model_name = \"reference_codec_cycle\"\n",
    "\n",
    "source_filepaths = reference_filepaths\n",
    "\n",
    "for filepath in tqdm(source_filepaths):\n",
    "    # load the audio file\n",
    "    waveform, sr = torchaudio.load(filepath)\n",
    "    # crop waveform to 4min\n",
    "    waveform = waveform[:, :int(4 * 60 * sr)]\n",
    "    # encode the audio file\n",
    "    encoded_waveform = codec_encode(waveform, sr)   \n",
    "    # apply the codec\n",
    "    decoded_waveform = codec_decode(encoded_waveform, sr)\n",
    "\n",
    "    # save the decoded audio file\n",
    "    out_filepath = filepath.replace(source_model, new_model_name).replace(\".webm\", \".mp3\")\n",
    "    decoded_waveform.write_hq_mp3(out_filepath)\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": 13,
   "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_average_stereo_spectrum(waveform, sr):\n",
    "    assert waveform.dim() == 2 and waveform.shape[0] == 2\n",
    "    # split into left and right channels\n",
    "    left = waveform[0]\n",
    "    right = waveform[1]\n",
    "\n",
    "    # compute mid and side channels\n",
    "    mid = (left + right) / 2\n",
    "    side = (left - right) / 2\n",
    "    \n",
    "    # calculate spectrum for mid and side channels\n",
    "    spectrum_mid = calculate_average_spectrum_db(mid, sr)\n",
    "    spectrum_side = calculate_average_spectrum_db(side, sr)\n",
    "    \n",
    "    return spectrum_mid, spectrum_side\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",
    "def get_shimmerscore_from_file(filepath):\n",
    "    # crop waveforn to max of 4min\n",
    "    shimmer_score = shimmerscore(filepath)\n",
    "    return shimmer_score\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "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",
    "waveform, sr = torchaudio.load(\"/home/christian/audio/reference-audio-wav/02 Dreams.wav\")\n",
    "shimmer_score = get_shimmerscore(waveform, sr)\n",
    "print(shimmer_score)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "metadata": {},
   "outputs": [],
   "source": [
    "# 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",
    "    # average stereo spectrum\n",
    "    average_stereo_spectrum_mid, average_stereo_spectrum_side = calculate_average_stereo_spectrum(audio, sr)\n",
    "\n",
    "    # spectrum evolution\n",
    "    spectrum_first, spectrum_last = calculate_spectrum_evolution(audio, sr)\n",
    "\n",
    "    # shimmer score\n",
    "    shimmer_score = get_shimmerscore_from_file(filepath)\n",
    "\n",
    "    # check for metadata file\n",
    "    metadata = None\n",
    "    metadata_filepath = filepath.replace(\".mp3\", \"__metadata.json\")\n",
    "    if os.path.exists(metadata_filepath):\n",
    "        try:\n",
    "            with open(metadata_filepath, \"r\") as f:\n",
    "                metadata = json.load(f)\n",
    "        except Exception as e:\n",
    "            metadata = None\n",
    "\n",
    "    if metadata is not None:\n",
    "        ear_score = metadata[\"ear_score\"]\n",
    "        hoot_cer = metadata[\"hoot_cer\"]\n",
    "    else:\n",
    "        ear_score = None\n",
    "        hoot_cer = None\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",
    "        \"average_stereo_spectrum_mid\": average_stereo_spectrum_mid,\n",
    "        \"average_stereo_spectrum_side\": average_stereo_spectrum_side,\n",
    "        \"shimmer_score\": shimmer_score,\n",
    "        \"ear_score\": ear_score,\n",
    "        \"hoot_cer\": hoot_cer,\n",
    "    }\n",
    "    \n",
    "    \n",
    "\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 15,
   "metadata": {},
   "outputs": [],
   "source": [
    "results = {}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "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"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Compute EAR score separately (not in parallel to avoid memory issues)\n",
    "print(\"Computing EAR scores...\")\n",
    "device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",
    "\n",
    "def calculate_ear_score(audio_tensor, sr, device):\n",
    "    # Resample to 16kHz if needed\n",
    "    if sr != 48000:\n",
    "        audio_tensor = torchaudio.functional.resample(audio_tensor, sr, 48000)\n",
    "    # Move to device\n",
    "    audio_tensor = audio_tensor.to(device)\n",
    "    \n",
    "    # Get EAR score\n",
    "    with torch.no_grad():\n",
    "        score = ear_model.get_score(audio_tensor, sample_rate=48000)\n",
    "    \n",
    "    return score\n",
    "\n",
    "# only compute ear score for subset of models\n",
    "selected_models = ['v45_2b_step_2_600_000', \"8n_25hz_v45_ft_ear_5e5_150k_\", 'reference']  # <-- replace with the actual model names\n",
    "\n",
    "\n",
    "ear_results = {}\n",
    "for filepath in tqdm(filepaths):\n",
    "    # Extract model name from filepath\n",
    "    model_name = None\n",
    "    for selected_model in selected_models:\n",
    "        if selected_model in filepath:\n",
    "            model_name = selected_model\n",
    "            break\n",
    "    \n",
    "    # Skip if the filepath doesn't contain any of the selected models\n",
    "    if model_name is None:\n",
    "        continue\n",
    "    \n",
    "    # Create a result dictionary with model and filepath information\n",
    "    result = {\n",
    "        \"model\": model_name,\n",
    "        \"filepath\": filepath,\n",
    "        \"clip_id\": os.path.basename(filepath)\n",
    "    }\n",
    "    try:\n",
    "        # Load audio using torchaudio\n",
    "        waveform, sr = torchaudio.load(filepath)\n",
    "        \n",
    "        # Calculate EAR score\n",
    "        ear_score = calculate_ear_score(waveform, sr, device=device)\n",
    "        result[\"ear_score\"] = ear_score\n",
    "        # Clear GPU memory if using CUDA\n",
    "        #if device.type == \"cuda\":\n",
    "        torch.cuda.empty_cache()\n",
    "    except Exception as e:\n",
    "        print(f\"Error processing {filepath}: {e}\")\n",
    "        result[\"ear_score\"] = None\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# now process the reference audio in parallel\n",
    "reference_results = Parallel(n_jobs=-1)(\n",
    "    delayed(lambda filepath: {\n",
    "        \"model\": \"reference\",\n",
    "        \"clip_id\": os.path.basename(filepath),\n",
    "        \"filepath\": filepath,\n",
    "        **analyze_audio(filepath)\n",
    "    })(filepath) for filepath in tqdm(reference_filepaths)\n",
    ")\n",
    "\n",
    "# Add reference results to the main results dictionary\n",
    "results[\"reference\"] = reference_results\n",
    "\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# store the results into a pandas dataframe\n",
    "# each row is a clip, columns are model, id, lufs_db, stereo_width, clipped_samples, etc...\n",
    "import pandas as pd\n",
    "import json\n",
    "import numpy as np\n",
    "# Create a list to hold all data\n",
    "all_data = []\n",
    "\n",
    "# Flatten the nested dictionary structure\n",
    "for model_name, clips in results.items():\n",
    "    for clip in clips:\n",
    "        # Create a base row with the model name\n",
    "        row = {'model': model_name}\n",
    "        \n",
    "        # Dynamically add all keys from the clip dictionary\n",
    "        # This will work even if we change the structure of the results dict\n",
    "        for key, value in clip.items():\n",
    "            # Convert numpy arrays to lists for JSON serialization\n",
    "            if isinstance(value, np.ndarray):\n",
    "                row[key] = value.tolist()\n",
    "            else:\n",
    "                row[key] = value\n",
    "            \n",
    "        all_data.append(row)\n",
    "\n",
    "# Create the DataFrame\n",
    "df = pd.DataFrame(all_data)\n",
    "\n",
    "# Display the first few rows and column information\n",
    "print(\"DataFrame shape:\", df.shape)\n",
    "print(\"Columns:\", df.columns.tolist())\n",
    "display(df.head())\n",
    "\n",
    "# Try to load the previous results\n",
    "try:\n",
    "    # Load the previous results\n",
    "    old_df = pd.read_json('outputs/audio_analysis_results_5.json', orient='records')\n",
    "    \n",
    "    # Merge the old and new results\n",
    "    df = pd.concat([old_df, df])\n",
    "    \n",
    "    print(f\"Successfully merged with previous results. New shape: {df.shape}\")\n",
    "except Exception as e:\n",
    "    print(f\"Could not load previous results: {e}\")\n",
    "    print(\"Continuing with only new results.\")\n",
    "\n",
    "# Save the results using JSON to preserve array data types\n",
    "df.to_json('outputs/audio_analysis_results_5.json', orient='records')\n",
    "\n",
    "# Also save a CSV version for compatibility, but note that arrays will be converted to strings\n",
    "df.to_csv('outputs/audio_analysis_results_5.csv', index=False)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# list of unique model names\n",
    "unique_models = df['model'].unique()\n",
    "print(unique_models)\n",
    "\n",
    "# v2_infill_v1_t18_1E6_beta100_n32_bt4_3k_acc4_fix"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "# Specify the models to plot\n",
    "selected_models = [\"v45_2b_step_2mil_ft_8k_infill_apr21_t1_18_cs0\", \"v2_infill_d3_t10_1E6_beta100_n16_bt4_acc4_3k_12k\", \"16n_25hz_v45_infill_shared_flow_resume_1_5m\"]\n",
    "#selected_models = [\"reference\", \"diff_v2_2b_2mil_ft_infill_20250421_v2\", \"v45_2b_step_2mil_ft_8k_infill_apr21_t1_18_cs0\", \"v2_infill_d4_2_5E6_beta100_n16_bt2_acc4_4k_16k\", \"\"] #\"v2_infill_d3_t10_1E6_beta100_n16_bt4_acc4_3k_12k\", 'reference']  # <-- replace with the actual model names\n",
    "#selected_models = [\"16n_25hz_v45_infill_shared_flow_resume_750k\", \"16n_25hz_v45_infill_shared_flow_resume_1_25m\", \"16n_25hz_v45_infill_shared_flow_resume_1_5m\", \"16n_25hz_v45_infill_shared_flow_4b_750k\", \"16n_25hz_v45_infill_shared_flow_4b_1m\",\"16n_25hz_v45_infill_shared_flow_4b_1_25m\"]\n",
    "\n",
    "colors = plt.cm.tab10.colors\n",
    "\n",
    "# Subset the dataframe\n",
    "df_subset = df[df['model'].isin(selected_models)]\n",
    "\n",
    "# Define metrics to plot\n",
    "metrics = [\n",
    "    {'name': 'lufs_db', 'title': 'Loudness (LUFS)', 'filename': 'lufs_distribution.png', 'figsize': (6, 8)},\n",
    "    {'name': 'stereo_width', 'title': 'Stereo Width (0 = mono, 1 = wide)', 'filename': 'stereo_width_distribution.png', 'figsize': (6, 6)},\n",
    "    {'name': 'shimmer_score', 'title': 'Shimmer Score', 'filename': 'shimmer_score_distribution.png', 'figsize': (6, 6)},\n",
    "    {'name': 'lufs_db_factor', 'title': 'Loudness Factor', 'filename': 'lufs_db_factor_distribution.png', 'figsize': (6, 6)},\n",
    "    {'name': 'clipped_samples', 'title': 'Clipped Samples', 'filename': 'clipped_samples_distribution.png', 'figsize': (6, 6)},\n",
    "    {'name': 'ear_score', 'title': 'Ear Score', 'filename': 'ear_score_distribution.png', 'figsize': (6, 6)},\n",
    "    {'name': 'hoot_cer', 'title': 'Hoot CER', 'filename': 'hoot_cer_distribution.png', 'figsize': (6, 6)},\n",
    "]\n",
    "\n",
    "# Loop through each metric and create plots\n",
    "for metric in metrics:\n",
    "    # Calculate min and max for bins\n",
    "    metric_min = df_subset[metric['name']].min()\n",
    "    metric_max = df_subset[metric['name']].max()\n",
    "    metric_bins = np.linspace(metric_min, metric_max, 50)\n",
    "    \n",
    "    # Create figure and axes\n",
    "    fig, axes = plt.subplots(len(selected_models), 1, figsize=metric['figsize'], sharex=True)\n",
    "    \n",
    "    # Plot histograms for each model\n",
    "    for i, model in enumerate(selected_models):\n",
    "        model_data = df_subset[df_subset['model'] == model]\n",
    "        stats_string = f\"mean: {model_data[metric['name']].mean():.2f}, std: {model_data[metric['name']].std():.2f}\"\n",
    "        \n",
    "        axes[i].hist(model_data[metric['name']], bins=metric_bins, alpha=0.7, edgecolor='black',\n",
    "                 color=colors[i % len(colors)], label=stats_string)\n",
    "        axes[i].set_ylabel('Count', fontsize=10)\n",
    "        axes[i].grid(True, linestyle='--', alpha=0.5)\n",
    "        axes[i].set_title(model, fontsize=10)\n",
    "        axes[i].legend()\n",
    "    \n",
    "    # Set labels and title\n",
    "    axes[-1].set_xlabel(metric['title'], fontsize=12)\n",
    "    fig.suptitle(f\"{metric['title']} Distribution by Model\", fontsize=14)\n",
    "    \n",
    "    # Adjust layout and save\n",
    "    plt.tight_layout()\n",
    "    plt.subplots_adjust(top=0.9)\n",
    "    plt.savefig(f\"plots/{metric['filename']}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "np.array(df[\"average_spectrum_db\"].iloc[10_000])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "# ---- Config ----\n",
    "#selected_models = ['dit_v3_dpo_t10_3k_5e6_b100_t25', 'v1_t6_5E6_beta100_n16_bt4', 'v45_2b_step_2_600_000', 'v1_t6_5E6_beta100_n8_bt4_30k']\n",
    "#selected_models = ['v1_t10_5E6_beta100_n8_bt4_9k', 'v45_2b_step_2_600_000', 'reference', \"v1_t6_5E6_beta100_n16_bt4\"]  # <-- replace with the actual model names\n",
    "#selected_models = [\"reference\", \"\"]\n",
    "#\n",
    "#selected_models = [\"v45_2b_step_2_600_000\", \"v1_t12_5E6_beta100_n16_bt4_9k\", \"v1_t12_5E6_beta100_n16_bt4_3k\", \"v1_t6_5E6_beta100_n16_bt4\", 'reference']  # <-- replace with the actual model names\n",
    "\n",
    "\n",
    "\n",
    "sr = 48000       # replace with your actual sample rate\n",
    "n_fft = 16384      # replace with your actual FFT size\n",
    "df_subset = df[df['model'].isin(selected_models)]\n",
    "colors = plt.cm.tab10.colors\n",
    "normalize_at_1khz = False  # Flag to normalize spectra at 1kHz\n",
    "\n",
    "# ---- Frequency Axis ----\n",
    "first_spec = df['average_spectrum_db'].iloc[0]\n",
    "n_bins = np.array(first_spec).shape[1]\n",
    "freqs = np.linspace(0, sr / 2, n_bins)\n",
    "\n",
    "\n",
    "plt.figure(figsize=(7, 4))\n",
    "\n",
    "# Find the index closest to 1kHz for normalization\n",
    "if normalize_at_1khz:\n",
    "    idx_1khz = np.argmin(np.abs(freqs - 1000))\n",
    "\n",
    "for i, model in enumerate(selected_models):\n",
    "    model_data = df_subset[df_subset['model'] == model]\n",
    "\n",
    "    # Average across time axis first, then average over all clips\n",
    "    spectra = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db']])\n",
    "    mean_spectrum = spectra.mean(axis=0)\n",
    "    \n",
    "    # Normalize at 1kHz if flag is set\n",
    "    if normalize_at_1khz:\n",
    "        normalization_value = mean_spectrum[idx_1khz]\n",
    "        mean_spectrum = mean_spectrum - normalization_value\n",
    "\n",
    "    plt.plot(freqs, mean_spectrum, label=model, color=colors[i % len(colors)])\n",
    "\n",
    "plt.xscale('log')\n",
    "plt.xlabel('Frequency (Hz)', fontsize=12)\n",
    "plt.ylabel('Average Magnitude (dB)', fontsize=12)\n",
    "plt.title('Average Spectrum per Model' + (' (normalized at 1kHz)' if normalize_at_1khz else ''), fontsize=14)\n",
    "plt.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "plt.ylim(-40, 48) if not normalize_at_1khz else plt.ylim(-60, 48)\n",
    "plt.xlim(20, 24000)\n",
    "plt.legend()\n",
    "plt.tight_layout()\n",
    "#plt.show()\n",
    "plt.savefig('plots/average_spectrum_per_model.png')\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "# ---- Config ----\n",
    "#selected_models = ['dit_v3_dpo_t10_3k_5e6_b100_t25', 'v1_t6_5E6_beta100_n16_bt4', 'v45_2b_step_2_600_000', 'v1_t6_5E6_beta100_n8_bt4_30k']\n",
    "#selected_models = ['v1_t10_5E6_beta100_n8_bt4_9k', 'v45_2b_step_2_600_000', 'reference', \"v1_t6_5E6_beta100_n16_bt4\"]  # <-- replace with the actual model names\n",
    "#selected_models = [\"reference\", \"\"]\n",
    "#\n",
    "#selected_models = [\"v45_2b_step_2_600_000\", \"v1_t12_5E6_beta100_n16_bt4_9k\", \"v1_t12_5E6_beta100_n16_bt4_3k\", \"v1_t6_5E6_beta100_n16_bt4\", 'reference']  # <-- replace with the actual model names\n",
    "\n",
    "\n",
    "\n",
    "sr = 48000       # replace with your actual sample rate\n",
    "n_fft = 16384      # replace with your actual FFT size\n",
    "df_subset = df[df['model'].isin(selected_models)]\n",
    "colors = plt.cm.tab10.colors\n",
    "normalize_at_1khz = False  # Flag to normalize spectra at 1kHz\n",
    "\n",
    "# ---- Frequency Axis ----\n",
    "first_spec_mid = df['average_stereo_spectrum_mid'].iloc[0]\n",
    "n_bins = len(first_spec_mid)\n",
    "freqs = np.linspace(0, sr / 2, n_bins)\n",
    "\"\"\n",
    "# Create a figure with two subplots (mid and side)\n",
    "fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(7, 6), sharex=True)\n",
    "\n",
    "# Find the index closest to 1kHz for normalization\n",
    "if normalize_at_1khz:\n",
    "    idx_1khz = np.argmin(np.abs(freqs - 1000))\n",
    "\n",
    "# Plot Mid Spectrum\n",
    "for i, model in enumerate(selected_models):\n",
    "    model_data = df_subset[df_subset['model'] == model]\n",
    "    \n",
    "    # Average over all clips for mid channel\n",
    "    mid_spectra = np.array([s for s in model_data['average_stereo_spectrum_mid']])\n",
    "    mean_mid_spectrum = np.mean(mid_spectra, axis=0)\n",
    "    \n",
    "    # Normalize at 1kHz if flag is set\n",
    "    if normalize_at_1khz:\n",
    "        normalization_value = mean_mid_spectrum[idx_1khz]\n",
    "        mean_mid_spectrum = mean_mid_spectrum - normalization_value\n",
    "    \n",
    "    ax1.plot(freqs, mean_mid_spectrum, label=model, color=colors[i % len(colors)])\n",
    "\n",
    "# Plot Side Spectrum\n",
    "for i, model in enumerate(selected_models):\n",
    "    model_data = df_subset[df_subset['model'] == model]\n",
    "    \n",
    "    # Average over all clips for side channel\n",
    "    side_spectra = np.array([s for s in model_data['average_stereo_spectrum_side']])\n",
    "    mean_side_spectrum = np.mean(side_spectra, axis=0)\n",
    "    \n",
    "    # Normalize at 1kHz if flag is set\n",
    "    if normalize_at_1khz:\n",
    "        normalization_value = mean_side_spectrum[idx_1khz]\n",
    "        mean_side_spectrum = mean_side_spectrum - normalization_value\n",
    "    \n",
    "    ax2.plot(freqs, mean_side_spectrum, label=model, color=colors[i % len(colors)])\n",
    "\n",
    "# Configure Mid plot\n",
    "ax1.set_xscale('log')\n",
    "ax1.set_ylabel('Mid Channel (dB)', fontsize=12)\n",
    "ax1.set_title('Average Mid Spectrum per Model' + (' (normalized at 1kHz)' if normalize_at_1khz else ''), fontsize=14)\n",
    "ax1.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "ax1.set_ylim(-60, 48) if normalize_at_1khz else ax1.set_ylim(-40, 48)\n",
    "ax1.legend()\n",
    "\n",
    "# Configure Side plot\n",
    "ax2.set_xscale('log')\n",
    "ax2.set_xlabel('Frequency (Hz)', fontsize=12)\n",
    "ax2.set_ylabel('Side Channel (dB)', fontsize=12)\n",
    "ax2.set_title('Average Side Spectrum per Model' + (' (normalized at 1kHz)' if normalize_at_1khz else ''), fontsize=14)\n",
    "ax2.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "ax2.set_ylim(-60, 20) if normalize_at_1khz else ax2.set_ylim(-40, 48)\n",
    "ax2.set_xlim(20, 24000)\n",
    "#ax2.legend()\n",
    "\n",
    "plt.tight_layout()\n",
    "#plt.show()\n",
    "plt.savefig('plots/average_mid_side_spectrum_per_model.png')\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# plot all the spectra for one model\n",
    "#selected_models = [\"reference\", \"]\n",
    "\n",
    "\n",
    "sr = 48000       # replace with your actual sample rate\n",
    "n_fft = 16384      # replace with your actual FFT size\n",
    "normalize_at_1khz = False  # Flag to normalize spectra at 1kHz\n",
    "\n",
    "# ---- Frequency Axis ----\n",
    "first_spec = df['average_spectrum_db'].iloc[0]\n",
    "n_bins = np.array(first_spec).shape[1]\n",
    "freqs = np.linspace(0, sr / 2, n_bins)\n",
    "\n",
    "fig, axs = plt.subplots(figsize=(7,6), nrows=len(selected_models), ncols=1, sharex=True)\n",
    "\n",
    "for i, model in enumerate(selected_models):\n",
    "    for j, spectra in enumerate(df_subset[df_subset['model'] == model]['average_spectrum_db']):\n",
    "        axs[i].plot(freqs, np.mean(spectra, axis=0), label=f\"Clip {i}\", color=\"tab:blue\", alpha=0.33, linewidth=0.5)\n",
    "\n",
    "    axs[i].set_xscale('log')\n",
    "    #axs[i].set_xlabel('Frequency (Hz)', fontsize=12)\n",
    "    axs[i].set_ylabel('Average Magnitude (dB)', fontsize=12)\n",
    "    axs[i].set_title(model, fontsize=12)\n",
    "    axs[i].grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "    axs[i].set_xlim(20, 20000)\n",
    "    axs[i].set_ylim(-40, 40)\n",
    "plt.tight_layout()\n",
    "plt.savefig('plots/all_spectra_per_model.png')\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "# ---- Config ----\n",
    "#selected_models = ['dit_v3_dpo_t10_3k_5e6_b100_t25', 'v1_t6_5E6_beta100_n16_bt4', 'v45_2b_step_2_600_000', 'v1_t6_5E6_beta100_n8_bt4_30k']\n",
    "#selected_models = ['v1_t10_5E6_beta100_n8_bt4_9k', 'v45_2b_step_2_600_000', \"v1_t6_5E6_beta100_n16_bt4\"]  # <-- replace with the actual model names\n",
    "#selected_models = [\"v45_2b_step_2_600_000\", \"dit_v3_dpo_t10_3k_5e6_b100_t25\", \"v1_t6_5E6_beta100_n16_bt4\", \"v45_2b_shared_ctx_s3784\"]  # Models to compare against reference\n",
    "\n",
    "sr = 48000       # replace with your actual sample rate\n",
    "n_fft = 16384      # replace with your actual FFT size\n",
    "df_subset = df[df['model'].isin(selected_models + ['reference'])]  # Include reference for comparison\n",
    "colors = plt.cm.tab10.colors\n",
    "normalize_at_1khz = True  # Flag to normalize spectra at 1kHz\n",
    "\n",
    "# ---- Frequency Axis ----\n",
    "first_spec = df['average_spectrum_db'].iloc[0]\n",
    "n_bins = np.array(first_spec).shape[1]\n",
    "freqs = np.linspace(0, sr / 2, n_bins)\n",
    "\n",
    "plt.figure(figsize=(7, 4))\n",
    "\n",
    "# Get reference model data first\n",
    "reference_data = df_subset[df_subset['model'] == 'reference']\n",
    "reference_spectra = np.array([np.mean(s, axis=0) for s in reference_data['average_spectrum_db']])\n",
    "reference_mean_spectrum = reference_spectra.mean(axis=0)\n",
    "\n",
    "# Find the index closest to 1kHz for normalization\n",
    "if normalize_at_1khz:\n",
    "    idx_1khz = np.argmin(np.abs(freqs - 1000))\n",
    "    reference_normalization_value = reference_mean_spectrum[idx_1khz]\n",
    "    reference_mean_spectrum = reference_mean_spectrum - reference_normalization_value\n",
    "\n",
    "# Plot difference between each model and reference\n",
    "for i, model in enumerate(selected_models):\n",
    "    model_data = df_subset[df_subset['model'] == model]\n",
    "\n",
    "    # Average across time axis first, then average over all clips\n",
    "    spectra = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db']])\n",
    "    mean_spectrum = spectra.mean(axis=0)\n",
    "    \n",
    "    # Normalize at 1kHz if flag is set\n",
    "    if normalize_at_1khz:\n",
    "        normalization_value = mean_spectrum[idx_1khz]\n",
    "        mean_spectrum = mean_spectrum - normalization_value\n",
    "\n",
    "    # Calculate difference with reference\n",
    "    difference_spectrum = mean_spectrum - reference_mean_spectrum\n",
    "    \n",
    "    plt.plot(freqs, difference_spectrum, label=f\"{model}\", color=colors[i % len(colors)])\n",
    "\n",
    "# Add a black line at y=0\n",
    "plt.axhline(y=0, color='black', linestyle='-', linewidth=1)\n",
    "\n",
    "plt.xscale('log')\n",
    "plt.xlabel('Frequency (Hz)', fontsize=12)\n",
    "plt.ylabel('Difference in Magnitude (dB)', fontsize=12)\n",
    "plt.title('Spectral Difference vs Reference' + (' (normalized at 1kHz)' if normalize_at_1khz else ''), fontsize=14)\n",
    "plt.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "plt.ylim(-28, 28)  # Adjusted for difference values\n",
    "plt.xlim(20, 20000)  # Adjusted for frequency range\n",
    "plt.legend()\n",
    "plt.tight_layout()\n",
    "plt.savefig('plots/spectral_difference_vs_reference.png')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "# ---- Config ----\n",
    "#selected_models = ['dit_v3_dpo_t10_3k_5e6_b100_t25', 'v1_t6_5E6_beta100_n16_bt4', 'v45_2b_step_2_600_000', 'v1_t6_5E6_beta100_n8_bt4_30k']\n",
    "#selected_models = ['v1_t6_5E6_beta100_n8_bt4_30k', 'v45_2b_step_2_600_000']\n",
    "#selected_models = ['v45_2b_step_2_600_000', 'reference', \"v45_2b_shared_ctx_s3784\"]\n",
    "\n",
    "sr = 48000       # replace with your actual sample rate\n",
    "n_fft = 16384      # replace with your actual FFT size\n",
    "df_subset = df[df['model'].isin(selected_models)]\n",
    "colors = plt.cm.tab10.colors\n",
    "\n",
    "# ---- Frequency Axis ----\n",
    "first_spec = df['average_spectrum_db_last'].iloc[0]\n",
    "n_bins = np.array(first_spec).shape[1]\n",
    "freqs = np.linspace(0, sr / 2, n_bins)\n",
    "\n",
    "# Create a single plot for all models\n",
    "plt.figure(figsize=(6, 4))\n",
    "\n",
    "for i, model in enumerate(selected_models):\n",
    "    model_data = df_subset[df_subset['model'] == model]\n",
    "    \n",
    "    # Process last 30s spectra\n",
    "    spectra_last = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db_last']])\n",
    "    mean_spectrum_last = spectra_last.mean(axis=0)\n",
    "    \n",
    "    # Process first 30s spectra\n",
    "    spectra_first = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db_first']])\n",
    "    mean_spectrum_first = spectra_first.mean(axis=0)\n",
    "\n",
    "    # Calculate delta between last and first\n",
    "    spectra_delta = mean_spectrum_last - mean_spectrum_first\n",
    "\n",
    "    # Plot only the delta for each model\n",
    "    plt.plot(freqs, spectra_delta, label=f\"{model}\", color=colors[i % len(colors)])\n",
    "    \n",
    "plt.xscale('log')\n",
    "plt.ylabel('Delta Magnitude (dB)', fontsize=12)\n",
    "plt.xlabel('Frequency (Hz)', fontsize=12)\n",
    "plt.title('Spectral Change (Last 30s - First 30s)', fontsize=14)\n",
    "plt.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "plt.ylim(-10, 10)  # Adjusted for delta values\n",
    "plt.xlim(20, 20000)\n",
    "plt.legend()\n",
    "plt.tight_layout()\n",
    "plt.savefig('outputs/spectral_change_last_30s_first_30s.png')\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "# ---- Config ----\n",
    "#selected_models = ['dit_v3_dpo_t10_3k_5e6_b100_t25', 'v1_t6_5E6_beta100_n16_bt4', 'v45_2b_step_2_600_000', 'v1_t6_5E6_beta100_n8_bt4_30k']\n",
    "#selected_models = ['v1_t6_5E6_beta100_n8_bt4_30k', 'v45_2b_step_2_600_000']\n",
    "#selected_models = ['v1_t10_5E6_beta100_n8_bt4_9k', 'v45_2b_step_2_600_000', 'reference' , \"v45_2b_shared_ctx_s3784\"]\n",
    "\n",
    "sr = 48000       # replace with your actual sample rate\n",
    "n_fft = 16384      # replace with your actual FFT size\n",
    "df_subset = df[df['model'].isin(selected_models)]\n",
    "colors = plt.cm.tab10.colors\n",
    "\n",
    "# ---- Frequency Axis ----\n",
    "first_spec = df['average_spectrum_db_last'].iloc[0]\n",
    "n_bins = np.array(first_spec).shape[1]\n",
    "freqs = np.linspace(0, sr / 2, n_bins)\n",
    "\n",
    "# Create a subplot for each model \n",
    "fig, axes = plt.subplots(len(selected_models), 1, figsize=(6, (2.5*len(selected_models))), sharex=True)\n",
    "\n",
    "for i, model in enumerate(selected_models):\n",
    "    model_data = df_subset[df_subset['model'] == model]\n",
    "    ax = axes[i] if len(selected_models) > 1 else axes\n",
    "\n",
    "    # Process last 30s spectra\n",
    "    spectra_last = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db_last']])\n",
    "    mean_spectrum_last = spectra_last.mean(axis=0)\n",
    "    \n",
    "    # Process first 30s spectra\n",
    "    spectra_first = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db_first']])\n",
    "    mean_spectrum_first = spectra_first.mean(axis=0)\n",
    "\n",
    "    spectra_delta = mean_spectrum_last - mean_spectrum_first\n",
    "\n",
    "    # Plot both spectra on the same subplot\n",
    "    ax.plot(freqs, mean_spectrum_last, label=\"Last 30s\", linestyle='--', color=colors[0])\n",
    "    ax.plot(freqs, mean_spectrum_first, label=\"First 30s\", color=colors[1])\n",
    "    ax.plot(freqs, spectra_delta, label=\"Delta\", color=colors[2])\n",
    "    \n",
    "    ax.set_xscale('log')\n",
    "    ax.set_ylabel('Average Magnitude (dB)', fontsize=12)\n",
    "    ax.set_title(f'Model: {model}', fontsize=14)\n",
    "    ax.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "    ax.set_ylim(-30, 30)\n",
    "    ax.set_xlim(20, 20000)\n",
    "    ax.legend()\n",
    "\n",
    "# Set common x-label\n",
    "plt.xlabel('Frequency (Hz)', fontsize=12)\n",
    "plt.tight_layout()\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 17,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\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_decay(audio_obj_first: Audio, audio_obj_last: Audio, n_fft=16384, hop_length=8192):\n",
    "    spectrum_first = calculate_average_spectrum_db(audio_obj_first.array_float, n_fft, hop_length)\n",
    "    spectrum_last = calculate_average_spectrum_db(audio_obj_last.array_float, n_fft, hop_length)\n",
    "    difference = spectrum_last - spectrum_first\n",
    "    decay = np.sum(np.abs(difference))\n",
    "    return decay, difference\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load audio\n",
    "from suno_utils.audio import Audio  \n",
    "#audio, sr = torchaudio.load(\"/home/christian/audio/reference-audio-wav/02 Dreams.wav\")\n",
    "#audio = Audio.from_file(\"/home/christian/audio/pencil-helmet-window/f0e56a18-96ba-41b3-b0ae-77fd8cd52f09/f0e56a18-96ba-41b3-b0ae-77fd8cd52f09_v2_infill_v1_t6_5E6_beta1000_n16_bt4_2k_repro_main_step10_nctx0_3.mp3\")\n",
    "#audio = Audio.from_file(\"/home/christian/audio/pencil-helmet-window/febc3fdf-cdb2-4b58-b5ff-39d58246ae5d/febc3fdf-cdb2-4b58-b5ff-39d58246ae5d_16n_v45_infill_ear_sft_3e5_t6_100k.mp3\") #sft6 \n",
    "audio = Audio.from_file(\"/home/christian/audio/pencil-helmet-window/febc3fdf-cdb2-4b58-b5ff-39d58246ae5d/febc3fdf-cdb2-4b58-b5ff-39d58246ae5d_v2_infill_v1_t17_5E6_beta1000_n16_bt4_9k_repro_main_3k_step10_nctx0_5.mp3\")\n",
    "audio_first = audio.get_slice(0, 30)\n",
    "audio_last = audio.get_slice(90, 120)\n",
    "\n",
    "print(audio_first.array_float.shape)\n",
    "print(audio_last.array_float.shape)\n",
    "# calculate decay\n",
    "decay, difference = calculate_decay(audio_first, audio_last)\n",
    "print(decay)\n",
    "\n",
    "plt.plot(difference)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "from suno_utils.audio import Audio\n",
    "\n",
    "audio = Audio.from_file(\"/home/christian/audio/pencil-helmet-window/febc3fdf-cdb2-4b58-b5ff-39d58246ae5d/febc3fdf-cdb2-4b58-b5ff-39d58246ae5d_v2_infill_v1_t17_5E6_beta1000_n16_bt4_9k_repro_main_3k_step10_nctx0_5.mp3\")\n",
    "audio_first = audio.get_slice(0, 30)\n",
    "audio_last = audio.get_slice(90, 120)\n",
    "\n",
    "n_fft = 16384      # replace with your actual FFT size\n",
    "hop_length = n_fft // 2  # Standard hop length\n",
    "\n",
    "spectrum_first = calculate_average_spectrum_db(audio_first.array_float, n_fft, hop_length)\n",
    "spectrum_last = calculate_average_spectrum_db(audio_last.array_float, n_fft, hop_length)\n",
    "\n",
    "colors = plt.cm.tab10.colors\n",
    "\n",
    "# ---- Frequency Axis ----\n",
    "first_spec = df['average_spectrum_db_last'].iloc[0]\n",
    "n_bins = np.array(first_spec).shape[1]\n",
    "freqs = np.linspace(0, sr / 2, n_bins)\n",
    "\n",
    "# Create a single plot for all models\n",
    "plt.figure(figsize=(6, 4))\n",
    "\n",
    "for i, model in enumerate(selected_models):\n",
    "    model_data = df_subset[df_subset['model'] == model]\n",
    "    \n",
    "    # Process last 30s spectra\n",
    "    spectra_last = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db_last']])\n",
    "    mean_spectrum_last = spectra_last.mean(axis=0)\n",
    "    \n",
    "    # Process first 30s spectra\n",
    "    spectra_first = np.array([np.mean(s, axis=0) for s in model_data['average_spectrum_db_first']])\n",
    "    mean_spectrum_first = spectra_first.mean(axis=0)\n",
    "\n",
    "    # Calculate delta between last and first\n",
    "    spectra_delta = mean_spectrum_last - mean_spectrum_first\n",
    "\n",
    "    # Plot only the delta for each model\n",
    "    plt.plot(freqs, spectra_delta, label=f\"{model}\", color=colors[i % len(colors)])\n",
    "    \n",
    "plt.xscale('log')\n",
    "plt.ylabel('Delta Magnitude (dB)', fontsize=12)\n",
    "plt.xlabel('Frequency (Hz)', fontsize=12)\n",
    "plt.title('Spectral Change (Last 30s - First 30s)', fontsize=14)\n",
    "plt.grid(True, which='both', linestyle='--', alpha=0.5)\n",
    "plt.ylim(-24, 24)  # Adjusted for delta values\n",
    "plt.xlim(1000, 20000)\n",
    "plt.legend()\n",
    "plt.tight_layout()\n",
    "plt.savefig('outputs/spectral_change_last_30s_first_30s.png')\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
}
