{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/home/christian/miniconda3/envs/suno_env/lib/python3.10/site-packages/torch/utils/_pytree.py:185: FutureWarning: optree is installed but the version is too old to support PyTorch Dynamo in C++ pytree. C++ pytree support is disabled. Please consider upgrading optree using `python3 -m pip install --upgrade 'optree>=0.13.0'`.\n",
      "  warnings.warn(\n"
     ]
    }
   ],
   "source": [
    "import os\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\"\n",
    "import math\n",
    "import torch\n",
    "import numpy as np\n",
    "from tqdm import tqdm\n",
    "from typing import List, Dict, Union\n",
    "from sklearn.metrics import f1_score, roc_auc_score, average_precision_score"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "import sys\n",
    "sys.path.insert(0, \"/home/christian/code/christian/scripts\")\n",
    "\n",
    "from train_ear import AudioQualityModel, create_label_encoder, CorruptAudioDataset"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {},
   "outputs": [],
   "source": [
    "def validate(model, val_loader, label_to_idx: Dict[str, int]):\n",
    "    \"\"\"\n",
    "    Validate model with metrics including per-class AP scores mapped to corruption labels.\n",
    "    \n",
    "    Args:\n",
    "        model: The PyTorch model\n",
    "        val_loader: Validation data loader\n",
    "        label_to_idx: Dictionary mapping corruption labels to indices\n",
    "    \"\"\"\n",
    "    model.eval()\n",
    "    total_loss = 0\n",
    "    predictions = []\n",
    "\n",
    "    # Create reverse mapping from idx to label\n",
    "    idx_to_label = {v: k for k, v in label_to_idx.items()}\n",
    "\n",
    "    with torch.no_grad():\n",
    "        for batch in tqdm(val_loader):\n",
    "            audio, labels = batch\n",
    "            audio, labels = audio.cuda(), labels.cuda()\n",
    "            \n",
    "            scores = model(audio)\n",
    "            loss = torch.nn.functional.binary_cross_entropy_with_logits(scores, labels)\n",
    "            total_loss += loss.item()\n",
    "            \n",
    "            predictions.append((\n",
    "                torch.sigmoid(scores) > 0.5,\n",
    "                torch.sigmoid(scores),\n",
    "                labels\n",
    "            ))\n",
    "\n",
    "    # Concatenate all batches and move to CPU\n",
    "    preds, scores, labels = [torch.cat(x, dim=0).cpu().numpy() for x in zip(*predictions)]\n",
    "    \n",
    "    # Basic metrics\n",
    "    metrics = {\n",
    "        \"loss\": total_loss / len(val_loader),\n",
    "        \"accuracy\": (preds == labels).all(axis=1).mean(),\n",
    "        \"macro_f1\": f1_score(labels, preds, average=\"macro\", zero_division=0),\n",
    "        \"micro_f1\": f1_score(labels, preds, average=\"micro\", zero_division=0)\n",
    "    }\n",
    "    \n",
    "    # Overall and per-class AP scores\n",
    "    try:\n",
    "        metrics[\"ap\"] = average_precision_score(labels.ravel(), scores.ravel())\n",
    "        \n",
    "        # Calculate per-class AP with corruption labels\n",
    "        class_ap_dict = {}\n",
    "        ap_scores = []\n",
    "        \n",
    "        for i in range(labels.shape[1]):\n",
    "            try:\n",
    "                ap = average_precision_score(labels[:, i], scores[:, i])\n",
    "            except ValueError:\n",
    "                ap = 0.0\n",
    "            ap_scores.append(ap)\n",
    "            \n",
    "            # Get corruption label\n",
    "            corruption_label = idx_to_label[i]\n",
    "            class_ap_dict[corruption_label] = ap\n",
    "        \n",
    "        metrics[\"class_ap\"] = class_ap_dict\n",
    "        metrics[\"mean_class_ap\"] = np.mean(ap_scores)\n",
    "        \n",
    "        # Sort and store top/bottom performing corruptions\n",
    "        sorted_classes = sorted(class_ap_dict.items(), key=lambda x: x[1], reverse=True)\n",
    "        metrics[\"top_classes\"] = dict(sorted_classes[:100])\n",
    "        metrics[\"bottom_classes\"] = dict(sorted_classes[-100:])\n",
    "        \n",
    "    except ValueError:\n",
    "        metrics.update({\n",
    "            \"ap\": 0.0,\n",
    "            \"class_ap\": {label: 0.0 for label in label_to_idx.keys()},\n",
    "            \"mean_class_ap\": 0.0,\n",
    "            \"top_classes\": {},\n",
    "            \"bottom_classes\": {}\n",
    "        })\n",
    "    \n",
    "    return metrics\n",
    "\n",
    "def print_validation_metrics(metrics):\n",
    "    \"\"\"Pretty print the validation metrics with corruption labels.\"\"\"\n",
    "    print(f\"\\nOverall Metrics:\")\n",
    "    print(f\"Loss: {metrics['loss']:.4f}\")\n",
    "    print(f\"Accuracy: {metrics['accuracy']:.4f}\")\n",
    "    print(f\"Macro F1: {metrics['macro_f1']:.4f}\")\n",
    "    print(f\"Micro F1: {metrics['micro_f1']:.4f}\")\n",
    "    print(f\"Overall AP: {metrics['ap']:.4f}\")\n",
    "    print(f\"Mean Class AP: {metrics['mean_class_ap']:.4f}\")\n",
    "    \n",
    "    print(\"\\nTop 5 Best Detected Corruptions:\")\n",
    "    for corruption, ap in metrics['top_classes'].items():\n",
    "        print(f\"{corruption}: {ap:.4f}\")\n",
    "    \n",
    "    print(\"\\nTop 5 Worst Detected Corruptions:\")\n",
    "    for corruption, ap in metrics['bottom_classes'].items():\n",
    "        print(f\"{corruption}: {ap:.4f}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Load model"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "{'channel_imbalance': {'params': {'imbalance': [-1, -0.9, -0.8, -0.75, -0.7, -0.6, -0.5, -0.4, -0.3, 0.3, 0.4, 0.5, 0.6, 0.7, 0.75, 0.8, 0.9, 1]}}, 'lowpass': {'params': {'cutoff_hz': [500, 700, 1000, 1250, 1500, 1750, 2000, 2250, 2500, 2750, 3000, 3250, 3500, 3750, 4000, 4250, 4500, 4750, 5000, 5250, 5500, 5750, 6000, 6250, 6500, 6750, 7000, 7250, 7500, 7750, 8000]}}, 'highpass': {'params': {'cutoff_hz': [60, 70, 80, 90, 100, 110, 120, 130, 140, 150, 160, 170, 180, 190, 200, 250, 300, 350, 400, 450, 500, 550, 600, 650, 700, 750, 800, 1000, 1200, 1400, 1600, 1800, 2000, 2500, 3000, 4000]}}, 'tanh_distortion': {'params': {'gain_db': [4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24]}}, 'clipping_distortion': {'params': {'gain_db': [4, 6, 8, 10, 12, 14, 16, 18, 20]}}, 'hum': {'params': {'amplitude': [0.05, 0.075, 0.1, 0.125, 0.15, 0.175, 0.2, 0.25, 0.3, 0.5], 'freq': [50, 52, 54, 56, 58, 60, 62, 64, 66, 68]}}, 'comb_filter': {'params': {'delay_ms': [0.1, 0.8, 3.2, 12.8, 25.6], 'gain_db': [3, 6, 12]}}, 'add_clicks': {'params': {'density': [1e-05, 0.0001, 0.001, 0.01]}}, 'audio_codec': {'params': {'n_passes': [1, 2], 'bit_rate': [8000, 12000, 16000, 24000, 32000, 48000, 64000, 96000, 128000]}}, 'noise': {'params': {'gain_db': [-48, -42, -36, -30, -24, -18, -12, -6], 'noise_type': ['white', 'pink']}}, 'white_noise_burst': {'params': {'p_burst': [0.01, 0.1, 0.3]}}, 'stereo_width': {'params': {'width': [0.0, 4.0]}}, 'spectral_mask': {'params': {'threshold': [-24, -12, -6], 'ratio': [0.1]}}, 'freq_boost': {'params': {'gain_db': [-24, -22, -20, -18, -16, -14, -12, 12, 14, 16, 18, 20, 22, 24], 'freq_hz': [60, 80, 100, 120, 240, 480, 960, 1000, 2000, 3000, 4000, 5000, 6000, 7000, 8000, 9000, 10000, 12000]}}, 'bandpass': {'params': {'bandwidth': [0.4, 1.6, 6.4, 12.8], 'central_freq': [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000, 1100, 1200, 1300, 1400, 1500, 1600, 1700, 1800, 1900, 2000]}}}\n",
      "598\n"
     ]
    }
   ],
   "source": [
    "#model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2024-12-20_11-26-16_s1551/last_ckpt.pt\"\n",
    "#model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2025-01-06_14-26-49_s7229/last_ckpt.pt\" # ft\n",
    "#model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2025-01-07_22-12-58_s6801/last_ckpt.pt\" # compare\n",
    "#model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2025-01-08_18-40-22_s2240/last_ckpt.pt\" # compare2\n",
    "#model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2025-01-22_19-03-56_s5918/last_ckpt.pt\" # compare3\n",
    "#model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2025-02-14_17-09-21_s8910/last_ckpt.pt\"\n",
    "#model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2025-02-16_12-49-07_s1047/last_ckpt.pt\"\n",
    "model_filepath = \"/app/suno/christian/checkpoints/ear-v2/2025-02-17_19-46-45_s4563/last_ckpt.pt\"\n",
    "\n",
    "ckpt = torch.load(model_filepath)\n",
    "model = AudioQualityModel(**ckpt[\"run_config\"][\"model\"])\n",
    "state_dict = ckpt[\"model\"]\n",
    "new_state_dict = {}\n",
    "for key, value in state_dict.items():\n",
    "    new_key = key.replace(\"module.\", \"\")\n",
    "    new_state_dict[new_key] = value\n",
    "model.load_state_dict(new_state_dict)\n",
    "model.eval()\n",
    "model.cuda()\n",
    "\n",
    "# also load corruptions config\n",
    "corruptions = ckpt[\"corruptions\"]\n",
    "print(corruptions)\n",
    "label_encoder = create_label_encoder(corruptions)\n",
    "print(len(label_encoder))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "metadata": {},
   "outputs": [],
   "source": [
    "def compute_quality_score(predictions_dict):\n",
    "    weights = {\n",
    "        'noise': 1.5,\n",
    "        'clipping_distortion': 1.5,\n",
    "        'white_noise_burst': 1.3,\n",
    "        'audio_codec': 1.5,\n",
    "        'tanh_distortion': 1.5,\n",
    "        'bandpass': 1.2,\n",
    "        'lowpass': 1.5,\n",
    "        'highpass': 1.5,\n",
    "        'hum': 1.1,\n",
    "        'reverb': 0.9,\n",
    "        'comb_filter': 0.9,\n",
    "        'wow_flutter': 0.8,\n",
    "        'ring_modulation': 0.8,\n",
    "        'channel_imbalance': 0.7,\n",
    "        'stereo_width': 1.5,\n",
    "        'phase_randomize': 0.6,\n",
    "        'spectral_mask': 1.5\n",
    "    }\n",
    "    \n",
    "    severity_scales = {\n",
    "        'lowpass': lambda params: 1 + (8000 - float(params['cutoff_hz'])) / 8000,\n",
    "        'highpass': lambda params: 1 + float(params['cutoff_hz']) / 4000,\n",
    "        'audio_codec': lambda params: 1 + (128000 - float(params['bit_rate'])) / 128000,\n",
    "        'clipping_distortion': lambda params: 1 + float(params['gain_db']) / 20,\n",
    "        'tanh_distortion': lambda params: 1 + float(params['gain_db']) / 20,\n",
    "        'hum': lambda params: 1 + float(params['amplitude']) / 0.5,\n",
    "        'wow_flutter': lambda params: 1 + float(params['depth']) / 0.005\n",
    "    }\n",
    "    \n",
    "    def extract_params(corruption_str):\n",
    "        if ':' not in corruption_str:\n",
    "            return {}\n",
    "        param_str = corruption_str.split(':', 1)[1]\n",
    "        params = {}\n",
    "        for param in param_str.split(','):\n",
    "            if '=' in param:\n",
    "                key, value = param.split('=')\n",
    "                params[key.strip()] = value.strip()\n",
    "        return params\n",
    "    \n",
    "    def transform_prediction(p, severity_mult=1.0):\n",
    "        if p < 1e-6: return 0\n",
    "        return -10 * (p ** 0.3) * severity_mult\n",
    "    \n",
    "    score = 100\n",
    "    \n",
    "    for corruption_type, weight in weights.items():\n",
    "        relevant_items = [(k, v) for k, v in predictions_dict.items() if k.startswith(corruption_type)]\n",
    "        if relevant_items:\n",
    "            try:\n",
    "                max_pred = max(pred for _, pred in relevant_items if pred is not None and pred == pred)\n",
    "                max_pred_key = next(k for k, v in relevant_items if v == max_pred)\n",
    "                \n",
    "                # Apply severity scaling if available\n",
    "                severity_mult = 1.0\n",
    "                if corruption_type in severity_scales:\n",
    "                    params = extract_params(max_pred_key)\n",
    "                    try:\n",
    "                        severity_mult = severity_scales[corruption_type](params)\n",
    "                    except (KeyError, ValueError):\n",
    "                        pass\n",
    "                \n",
    "                penalty = transform_prediction(max_pred, severity_mult) * weight\n",
    "                score += penalty\n",
    "            except ValueError:\n",
    "                continue\n",
    "    \n",
    "    return max(0, min(100, score))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Run evaluation"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# create dataset for evaluation\n",
    "val_filepaths = \"/app/suno/data/audio_2ch_48khz_lg/ear_val.csv\"\n",
    "val_dataset = CorruptAudioDataset(\n",
    "    val_filepaths,\n",
    "    label_encoder,\n",
    "    corruptions,\n",
    "    sample_rate=48000,\n",
    "    num_workers=1,\n",
    "    chunk_size_s=5.0,\n",
    "    buffer_size=1000,\n",
    "    max_corruptions=1,\n",
    "    max_chunks_per_file=1,\n",
    "    no_corruption_probability=0.01,\n",
    ")\n",
    "\n",
    "val_loader = torch.utils.data.DataLoader(\n",
    "    val_dataset,\n",
    "    batch_size=32,\n",
    "    num_workers=4,\n",
    "    persistent_workers=True,  # this is necessary for the buffer to work\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "val_metrics = validate(model, val_loader, label_encoder)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print_validation_metrics(val_metrics)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Preference evaluation"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "94\n"
     ]
    }
   ],
   "source": [
    "import json\n",
    "import torchaudio\n",
    "#base_dir = \"/app/suno/christian/data/dpo_diffusion_test_set\"\n",
    "# load the json file\n",
    "#with open(os.path.join(base_dir, \"comparison_results.json\"), \"r\") as f:\n",
    "#    results = json.load(f)\n",
    "\n",
    "#print(len(results))\n",
    "\n",
    "\n",
    "\n",
    "results = []\n",
    "with open(\"/home/christian/code/christian/results.jsonl\", \"r\") as f:\n",
    "    for line in f:\n",
    "        results.append(json.loads(line))\n",
    "print(len(results))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": 12,
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|██████████| 94/94 [00:06<00:00, 15.58it/s]\n"
     ]
    }
   ],
   "source": [
    "# copy audio files to a new folder\n",
    "source_dir = \"/app/suno/christian/data/dpo_diffusion_test_set_1k\"\n",
    "base_dir = \"/home/christian/audio/ear-bench\"\n",
    "output_dirname = \"dpo_diffusion_test_set\"\n",
    "output_dir = os.path.join(base_dir, output_dirname)\n",
    "os.makedirs(output_dir, exist_ok=True)\n",
    "\n",
    "import shutil\n",
    "\n",
    "pbar = tqdm(results)\n",
    "for idx, result in enumerate(pbar):\n",
    "    # get the result\n",
    "    true_label = result[\"selected\"]\n",
    "    # get the audio files\n",
    "    audio_a = os.path.join(source_dir, result[\"request_id\"], result[\"audio_a_id\"] + \".mp3\")\n",
    "    audio_b = os.path.join(source_dir, result[\"request_id\"], result[\"audio_b_id\"] + \".mp3\")\n",
    "\n",
    "    # copy the audio files to the output folder\n",
    "    output_a = os.path.join(output_dir, f\"{idx:03d}_a_input_pref={True if true_label == 'a' else False}.mp3\")\n",
    "    output_b = os.path.join(output_dir, f\"{idx:03d}_b_input_pref={True if true_label == 'b' else False}.mp3\")\n",
    "    shutil.copy(audio_a, output_a)\n",
    "    shutil.copy(audio_b, output_b)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "metadata": {},
   "outputs": [],
   "source": [
    "def apply_highpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0):\n",
    "    return torchaudio.functional.highpass_biquad(audio, sample_rate, cutoff_hz)\n",
    "\n",
    "def apply_lowpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0):\n",
    "    return torchaudio.functional.lowpass_biquad(audio, sample_rate, cutoff_hz)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|██████████| 94/94 [01:01<00:00,  1.52it/s, accuracy=67, score_delta=0.0875]  "
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "0.6702127659574468\n",
      "Final performance: 63/94 p_val 0.0013\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "\n"
     ]
    }
   ],
   "source": [
    "predictions = []\n",
    "score_deltas = []\n",
    "from scipy import stats\n",
    "\n",
    "pbar = tqdm(results)\n",
    "\n",
    "idx_to_label = {idx: label for label, idx in label_encoder.items()}\n",
    "audio_a = None\n",
    "\n",
    "for result in pbar:\n",
    "    if audio_a is not None:\n",
    "        del audio_a\n",
    "        del audio_b\n",
    "    # get the audio files\n",
    "    audio_a = os.path.join(base_dir, result[\"request_id\"], result[\"audio_a_id\"] + \".mp3\")\n",
    "    audio_b = os.path.join(base_dir, result[\"request_id\"], result[\"audio_b_id\"] + \".mp3\")\n",
    "    # load the audio files\n",
    "    audio_a, sr_a = torchaudio.load(audio_a)\n",
    "    audio_b, sr_b = torchaudio.load(audio_b)\n",
    "\n",
    "    # crop to 30s of audio\n",
    "    start_s = 0\n",
    "    end_s = 120\n",
    "    audio_a = audio_a[:, start_s * sr_a:end_s * sr_a]\n",
    "    audio_b = audio_b[:, start_s * sr_b:end_s * sr_b]\n",
    "\n",
    "    # peak normalize\n",
    "    #audio_a /= audio_a.abs().max()\n",
    "    #audio_b /= audio_b.abs().max()\n",
    "\n",
    "    # get the result\n",
    "    true_label = result[\"selected\"]\n",
    "\n",
    "    # get score from model for each audio file\n",
    "    audio_a = audio_a.cuda()\n",
    "    audio_b = audio_b.cuda()\n",
    "    with torch.no_grad():\n",
    "\n",
    "        # let's kist alwaus corrupt b as a test\n",
    "        #true_label = \"a\"\n",
    "        #chunks_b = apply_highpass(chunks_b.clone(), sr_b, 100)\n",
    "\n",
    "        #score_compare = model(chunks_a, chunks_b)\n",
    "        #score_compare = torch.sigmoid(score_compare).cpu().mean().item()\n",
    "        score_a = model.get_score(audio_a)\n",
    "        score_b = model.get_score(audio_b)\n",
    "\n",
    "        #label_probs_a = {idx_to_label[i]: prob for i, prob in enumerate(score_a)}\n",
    "        #label_probs_b = {idx_to_label[i]: prob for i, prob in enumerate(score_b)}\n",
    "        \n",
    "        #for k, v in label_probs_a.items():\n",
    "        #     if v > 0.3:\n",
    "        #         print(k, v)\n",
    "        # get the quality score\n",
    "        #quality_score_a = compute_quality_score(label_probs_a)\n",
    "        #quality_score_b = compute_quality_score(label_probs_b)\n",
    "\n",
    "    # get the models choice\n",
    "    #model_choice = \"A\" if quality_score_a > quality_score_b else \"B\"\n",
    "    model_choice = \"a\" if score_a < score_b else \"b\"\n",
    "    score_delta = score_a - score_b\n",
    "    score_deltas.append(score_delta.item())\n",
    "    # check if the model choice is correct\n",
    "    correct = model_choice == true_label\n",
    "    predictions.append(correct)\n",
    "    pbar.set_postfix({\"accuracy\": np.mean(predictions) * 100, \"score_delta\": np.mean(np.abs(score_deltas))})\n",
    "print(np.mean(predictions))\n",
    "\n",
    "# Perform binomial test\n",
    "p_value = stats.binomtest(sum(predictions), len(predictions), 0.5).pvalue\n",
    "print(f\"Final performance: {sum(predictions)}/{len(predictions)} p_val {p_value:0.4f}\")\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# compute quality scores for 100 random files\n",
    "import random\n",
    "import glob\n",
    "\n",
    "# parallelize this\n",
    "from tqdm import tqdm\n",
    "import torchaudio\n",
    "import glob\n",
    "import numpy as np\n",
    "import torch\n",
    "from torch.utils.data import Dataset, DataLoader\n",
    "\n",
    "class AudioDataset(Dataset):\n",
    "    def __init__(self, filepaths, start_s=15.0, duration_s=30.0, target_sr=48000):\n",
    "        self.filepaths = filepaths\n",
    "        self.start_s = start_s\n",
    "        self.duration_s = duration_s\n",
    "        self.target_sr = target_sr\n",
    "        \n",
    "    def __len__(self):\n",
    "        return len(self.filepaths)\n",
    "    \n",
    "    def __getitem__(self, idx):\n",
    "        filepath = self.filepaths[idx]\n",
    "        try:\n",
    "            audio, sr = torchaudio.load(filepath)\n",
    "            if sr != self.target_sr:\n",
    "                audio = torchaudio.functional.resample(audio, sr, self.target_sr)\n",
    "                \n",
    "            # Crop to specified duration\n",
    "            start_idx = int(self.start_s * self.target_sr)\n",
    "            end_idx = int((self.start_s + self.duration_s) * self.target_sr)\n",
    "            audio = audio[:, start_idx:end_idx]\n",
    "\n",
    "            if audio.shape[1] < self.target_sr * self.duration_s:\n",
    "                return None\n",
    "\n",
    "            # break into 5s chunks\n",
    "            chunk_samples = 5 * self.target_sr\n",
    "            chunks = torch.split(audio, chunk_samples, dim=1)\n",
    "            chunks = [chunk / chunk.abs().max().clamp(1e-6) for chunk in chunks]\n",
    "            chunks = torch.stack(chunks)\n",
    "            \n",
    "            return {\n",
    "                'audio': chunks,\n",
    "                'filepath': filepath\n",
    "            }\n",
    "        except Exception as e:\n",
    "            print(f\"Error loading {filepath}: {str(e)}\")\n",
    "            return None\n",
    "\n",
    "def collate_fn(batch):\n",
    "    # Filter out None values from failed loads\n",
    "    batch = [b for b in batch if b is not None]\n",
    "    if not batch:\n",
    "        return None\n",
    "    \n",
    "    return {\n",
    "        'audio': torch.stack([item['audio'] for item in batch]),\n",
    "        'filepath': [item['filepath'] for item in batch]\n",
    "    }\n",
    "\n",
    "idx_to_label = {idx: label for label, idx in label_encoder.items()}\n",
    "\n",
    "# Setup\n",
    "root_dir = \"/app/suno/christian/data/dpo_diffusion_test_set_10k\"\n",
    "# get all the files recursively\n",
    "filepaths = glob.glob(os.path.join(root_dir, \"**\", \"*.mp3\"), recursive=True)\n",
    "filepaths = filepaths[:1000]\n",
    "print(len(filepaths))\n",
    "# Create dataset and dataloader\n",
    "dataset = AudioDataset(filepaths)\n",
    "dataloader = DataLoader(\n",
    "    dataset,\n",
    "    batch_size=8,  # Adjust based on your GPU memory\n",
    "    num_workers=16,  # Adjust based on your CPU cores\n",
    "    collate_fn=collate_fn,\n",
    "    shuffle=False,\n",
    "    pin_memory=True\n",
    ")\n",
    "\n",
    "results = []\n",
    "model = model.cuda()\n",
    "model.eval()\n",
    "\n",
    "with torch.no_grad():\n",
    "    for batch in tqdm(dataloader):\n",
    "        if batch is None:\n",
    "            continue\n",
    "            \n",
    "        # Move batch to GPU\n",
    "        audio = batch['audio'].cuda()\n",
    "        filepaths = batch['filepath']\n",
    "\n",
    "        # get the score for each chunk in parallel\n",
    "        audio = audio.view(audio.shape[0] * audio.shape[1], 2, -1)\n",
    "        \n",
    "        # Get predictions\n",
    "        preds = model(audio)\n",
    "        probs = torch.sigmoid(preds).cpu()\n",
    "\n",
    "        # move probs back to the original shape\n",
    "        probs = probs.reshape(audio.shape[0], -1, probs.shape[-1])\n",
    "\n",
    "        # average the probs\n",
    "        probs = probs.mean(dim=1).cpu().numpy()\n",
    "\n",
    "        # Process each item in the batch\n",
    "        for i in range(len(filepaths)):\n",
    "            label_probs = {idx_to_label[j]: prob for j, prob in enumerate(probs[i])}\n",
    "            \n",
    "            results.append({\n",
    "                \"filepath\": filepaths[i],\n",
    "                \"raw_preds\": probs[i],\n",
    "                \"label_predictions\": label_probs,\n",
    "                \"audio\": audio[i].cpu().numpy(),\n",
    "            })"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "scores = [compute_quality_score(result[\"label_predictions\"]) for result in results]\n",
    "print(len(scores))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "plt.hist(scores, bins=100)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# sort the results by score\n",
    "import IPython\n",
    "import IPython.display\n",
    "results = sorted(results, key=lambda x: compute_quality_score(x[\"label_predictions\"]), reverse=True)\n",
    "scores = [compute_quality_score(result[\"label_predictions\"]) for result in results]\n",
    "filepaths = [result[\"filepath\"] for result in results]\n",
    "\n",
    "\n",
    "# print the top 5 scores and lowest 5 scores\n",
    "#for i in range(5):\n",
    "#    print(scores[i])\n",
    "#    IPython.display.display(IPython.display.Audio(results[i][\"audio\"], rate=48000))\n",
    "    \n",
    "print()\n",
    "for i in range(len(scores) - 10, len(scores)):\n",
    "    print(filepaths[i])\n",
    "    for k, v in results[i][\"label_predictions\"].items():\n",
    "        if v > 0.5:\n",
    "            print(k, v)\n",
    "    print(scores[i])\n",
    "    IPython.display.display(IPython.display.Audio(results[i][\"audio\"], rate=48000))\n",
    "print()\n",
    "for i in range(10):\n",
    "    print(filepaths[i])\n",
    "    corruption = results[i][\"label_predictions\"]\n",
    "    for k, v in corruption.items():\n",
    "        if v > 0.5:\n",
    "            print(k, v)\n",
    "    print(scores[i])\n",
    "    IPython.display.display(IPython.display.Audio(results[i][\"audio\"], rate=48000))\n",
    "\n"
   ]
  },
  {
   "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
}
