{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "os.environ['CUDA_VISIBLE_DEVICES'] = '5'\n",
    "\n",
    "import numpy as np\n",
    "from suno_utils.audio import Audio\n",
    "from dataclasses import dataclass\n",
    "from pathlib import Path"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "window_size = 1024\n",
    "hop_length = 128\n",
    "sr = 22050\n",
    "high_cutoff = 9950\n",
    "low_cutoff = 1000\n",
    "\n",
    "\n",
    "def spectrogram(signal):\n",
    "    window = np.hanning(window_size)\n",
    "    low_bin = int(low_cutoff * window_size / sr)\n",
    "    high_bin = int(high_cutoff * window_size / sr)\n",
    "\n",
    "    hops = signal.shape[0] // hop_length\n",
    "    real_spec = np.zeros((hops, high_bin - low_bin + 1))\n",
    "    for i in range(hops):\n",
    "        start = i * hop_length\n",
    "        if start + window_size > signal.shape[0]:\n",
    "            break\n",
    "        real_spec[i, :] = 10 * np.log10(\n",
    "            np.abs(np.fft.fft(signal[start : start + window_size] * window)[low_bin : high_bin + 1])\n",
    "            + 1e-6\n",
    "        )\n",
    "    return real_spec"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "@dataclass\n",
    "class Clip:\n",
    "    audio: Audio\n",
    "    shimmers: np.ndarray   # seconds\n",
    "    spectrogram: np.ndarray\n",
    "\n",
    "    @classmethod\n",
    "    def from_negative_audio(cls, path: Path):\n",
    "        audio = Audio.from_file(str(path), sample_rate=sr, n_channels=1)\n",
    "        return cls(\n",
    "            audio,\n",
    "            np.array([]),\n",
    "            spectrogram(audio.array_float)\n",
    "        )\n",
    "\n",
    "    @classmethod\n",
    "    def from_csv_path(cls, path: Path):\n",
    "        with open(path, 'r') as f:\n",
    "            lines = f.readlines()\n",
    "        audio = None\n",
    "        for suffix in ['wav', 'opus']:\n",
    "            try:\n",
    "                audio = Audio.from_file(str(path.with_suffix(f'.{suffix}')), sample_rate=sr, n_channels=1)\n",
    "                break\n",
    "            except FileNotFoundError:\n",
    "                pass\n",
    "        if audio is None:\n",
    "            raise FileNotFoundError(f\"No audio file found for {path}\")\n",
    "        shimmers = np.array([float(line.split(',')[0].strip()) for line in lines])\n",
    "        return cls(\n",
    "            audio,\n",
    "            shimmers,\n",
    "            spectrogram(audio.array_float)\n",
    "        )\n",
    "\n",
    "NON_SUNO_DIR = Path('/home/m4burns/flac_out')\n",
    "ANNOTATIONS_DIR = Path('/home/m4burns/glockenspiel/suno_utils/notebooks/shimmerscore/shimmer_annotations')\n",
    "ANNOTATIONS_DIR_V2 = ANNOTATIONS_DIR / 'v2'\n",
    "\n",
    "clips = []\n",
    "\n",
    "# clips += [\n",
    "#     Clip.from_csv_path(path)\n",
    "#     for path in ANNOTATIONS_DIR.glob('*.csv')\n",
    "# ]\n",
    "\n",
    "clips += [\n",
    "    Clip.from_csv_path(path)\n",
    "    for path in ANNOTATIONS_DIR_V2.glob('*.csv')\n",
    "]\n",
    "\n",
    "clips += [\n",
    "    Clip.from_negative_audio(path)\n",
    "    for path in list(NON_SUNO_DIR.glob('*.flac'))[:5]\n",
    "]\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(f\"Loaded {sum([len(clip.shimmers) for clip in clips])} annotations from {len(clips)} files\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "defect_len_ms = 120\n",
    "defect_len_hops = int(defect_len_ms * sr / 1000 / hop_length)\n",
    "\n",
    "label_jitter_ms = 10\n",
    "label_jitter_hops = int(label_jitter_ms * sr / 1000 / hop_length)\n",
    "\n",
    "defect_len_hops"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import torch.nn as nn\n",
    "import torch.nn.functional as F\n",
    "import torch.optim as optim\n",
    "from torch.utils.data import Dataset, DataLoader, random_split\n",
    "import numpy as np\n",
    "from collections import Counter\n",
    "\n",
    "def add_spectral_noise(log_spec, snr_range):\n",
    "    \"\"\"\n",
    "    Add noise to log-magnitude spectra\n",
    "    \"\"\"\n",
    "    \n",
    "    # Convert from log back to linear magnitude\n",
    "    mag_spec = torch.pow(10.0, log_spec / 10.0)\n",
    "    \n",
    "    # Calculate noise magnitude directly from SNR\n",
    "    snr_db = torch.empty(1).uniform_(snr_range[0], snr_range[1])\n",
    "    signal_power = torch.mean(mag_spec)\n",
    "    snr_linear = torch.pow(10.0, snr_db / 10.0)\n",
    "    noise_magnitude = signal_power / snr_linear\n",
    "    \n",
    "    # Generate complex Gaussian noise\n",
    "    noise_real = torch.randn(mag_spec.shape)\n",
    "    noise_imag = torch.randn(mag_spec.shape)\n",
    "    noise_complex = (noise_real + 1j * noise_imag) / torch.sqrt(torch.tensor(2.0))\n",
    "    \n",
    "    # Scale noise to desired magnitude\n",
    "    noise_scaled = noise_complex * noise_magnitude\n",
    "    \n",
    "    # Add noise in complex domain (using random phase for original signal)\n",
    "    random_phase = torch.empty(mag_spec.shape).uniform_(0, 2*torch.pi)\n",
    "    noisy_complex = mag_spec * torch.exp(1j * random_phase) + noise_scaled\n",
    "    \n",
    "    # Convert back to log magnitude WITHOUT adding epsilon\n",
    "    noisy_mag = torch.abs(noisy_complex)\n",
    "    noisy_log_spec = 10 * torch.log10(noisy_mag)\n",
    "    \n",
    "    return noisy_log_spec\n",
    "\n",
    "class SpectrogramDataset(Dataset):\n",
    "    def __init__(self, spectrograms, labels):\n",
    "        \"\"\"\n",
    "        Dataset for spectrogram defect detection\n",
    "        \n",
    "        Args:\n",
    "            spectrograms (np.ndarray): Array of spectrograms, shape (N, C, T)\n",
    "            labels (np.ndarray): Binary labels, shape (N, T)\n",
    "        \"\"\"\n",
    "        self.spectrograms = torch.FloatTensor(spectrograms)\n",
    "        self.labels = torch.FloatTensor(labels)\n",
    "        self.is_augmenting = True\n",
    "        \n",
    "    def __len__(self):\n",
    "        return len(self.spectrograms)\n",
    "\n",
    "    def get_class_weights(self):\n",
    "        \"\"\"Calculate class weights based on label distribution\"\"\"\n",
    "        # Flatten labels across all time steps\n",
    "        flat_labels = self.labels.numpy().flatten()\n",
    "        counts = Counter(flat_labels)\n",
    "        total = len(flat_labels)\n",
    "        \n",
    "        # Compute inverse frequency weights\n",
    "        weights = {\n",
    "            label: total / (len(counts) * count)\n",
    "            for label, count in counts.items()\n",
    "        }\n",
    "        \n",
    "        return torch.FloatTensor([weights[0], weights[1]])\n",
    "\n",
    "    def augment(self, spectrogram, label):\n",
    "        \"\"\"Apply data augmentation to a spectrogram\"\"\"\n",
    "        # Convert to tensor if not already\n",
    "        spec = torch.as_tensor(spectrogram)\n",
    "\n",
    "        # Random gain between -15 and 3 dB\n",
    "        gain = torch.randint(-15, 3, (1,))\n",
    "        spec = spec + gain\n",
    "\n",
    "        # Random zoom and crop frequency axis of spectrogram\n",
    "        # TODO needs scale invariance\n",
    "        # if torch.rand(1) < 0.1:\n",
    "        #     zoom_factor = 1 + torch.rand(1) * 0.15\n",
    "        #     zoomed_spec_len = int(spec.shape[0] * zoom_factor)\n",
    "        #     if zoomed_spec_len > spec.shape[0]:\n",
    "        #         start_bin = torch.randint(0, int(spec.shape[0] * zoom_factor - spec.shape[0]), (1,))\n",
    "        #         spec = torch.nn.functional.interpolate(spec.T.unsqueeze(0), size=zoomed_spec_len, mode='linear', align_corners=False).squeeze(0).T[start_bin:start_bin+spec.shape[0]]\n",
    "        \n",
    "        # Add noise\n",
    "        if torch.rand(1) < 0.5:\n",
    "            spec = add_spectral_noise(spec, snr_range=(5, 30))\n",
    "        \n",
    "        # # Channel masking\n",
    "        # if torch.rand(1) < 0.5:\n",
    "        #     num_channels_to_mask = torch.randint(1, spec.shape[1]//4, (1,))\n",
    "        #     channel_indices = torch.randperm(spec.shape[1])[:num_channels_to_mask]\n",
    "        #     spec[:, channel_indices] = -80\n",
    "        \n",
    "        # # Time masking\n",
    "        # if torch.rand(1) < 0.5:\n",
    "        #     num_steps_to_mask = torch.randint(1, spec.shape[0]//4, (1,))\n",
    "        #     time_indices = torch.randperm(spec.shape[0])[:num_steps_to_mask]\n",
    "        #     spec[time_indices, :] = -80\n",
    "\n",
    "        # Randomly jitter labels\n",
    "        if torch.rand(1) < 0.5:\n",
    "            shimmer_times = torch.where(label == 1)[0]\n",
    "            jitter = torch.randint(-label_jitter_hops, label_jitter_hops + 1, (shimmer_times.shape[0],))\n",
    "            shimmer_times += jitter\n",
    "            label = torch.zeros_like(label)\n",
    "            label[torch.clamp(shimmer_times, 0, label.shape[0] - 1)] = 1\n",
    "        \n",
    "        return spec, label\n",
    "    \n",
    "    def __getitem__(self, idx):\n",
    "        spec = self.spectrograms[idx]\n",
    "        label = self.labels[idx]\n",
    "\n",
    "        if self.is_augmenting and torch.rand(1) < 0.8:\n",
    "            spec, label = self.augment(spec, label)\n",
    "            \n",
    "        return spec, label\n",
    "\n",
    "class SmoothFocalLoss(nn.Module):\n",
    "    \"\"\"Focal Loss with gaussian smoothed targets for misaligned labels\"\"\"\n",
    "    def __init__(self, device, alpha=None, gamma=2, window_size=41, sigma=None):\n",
    "        super().__init__()\n",
    "        self.alpha = alpha  # Class weights\n",
    "        self.gamma = gamma  # Focusing parameter\n",
    "        self.window_size = window_size\n",
    "        \n",
    "        # Create gaussian window - do this once in init\n",
    "        if sigma is None:\n",
    "            sigma = window_size / 6.0  # Default sigma based on window size\n",
    "        \n",
    "        # Create gaussian window\n",
    "        x = torch.linspace(-window_size//2, window_size//2, window_size)\n",
    "        gaussian = torch.exp(-(x ** 2) / (2 * sigma ** 2)).to(device)\n",
    "        self.register_buffer('window', gaussian / gaussian.sum())\n",
    "        \n",
    "    def forward(self, inputs, targets):\n",
    "        # Smooth the targets with gaussian window\n",
    "        smoothed_targets = F.conv1d(\n",
    "            targets.unsqueeze(1).float(),\n",
    "            self.window.view(1, 1, -1),\n",
    "            padding=self.window_size//2\n",
    "        ).squeeze(1)\n",
    "        \n",
    "        # Calculate focal loss with smoothed targets\n",
    "        bce_loss = F.binary_cross_entropy_with_logits(\n",
    "            inputs, \n",
    "            smoothed_targets,\n",
    "            reduction='none',\n",
    "            pos_weight=self.alpha[1] if self.alpha is not None else None\n",
    "        )\n",
    "        \n",
    "        pt = torch.exp(-bce_loss)\n",
    "        focal_loss = ((1 - pt) ** self.gamma) * bce_loss\n",
    "        \n",
    "        return focal_loss.mean()\n",
    "\n",
    "class AudioDefectCNN(nn.Module):\n",
    "    def __init__(self, in_channels, kernel_size=23, n_filters=64):\n",
    "        \"\"\"\n",
    "        CNN for detecting defects in audio spectrograms\n",
    "        \n",
    "        Args:\n",
    "            in_channels (int): Number of input channels in spectrogram\n",
    "            kernel_size (int): Size of convolutional kernel (default 22 for defect length)\n",
    "            n_filters (int): Number of convolutional filters\n",
    "        \"\"\"\n",
    "        super().__init__()\n",
    "        \n",
    "        # First conv layer with kernel size matching defect length\n",
    "        self.conv1 = nn.Conv1d(\n",
    "            in_channels=in_channels,\n",
    "            out_channels=n_filters,\n",
    "            kernel_size=kernel_size,\n",
    "            stride=1,  # Same length output\n",
    "            padding=kernel_size // 2  # Same padding\n",
    "        )\n",
    "        self.bn1 = nn.BatchNorm1d(n_filters)\n",
    "        \n",
    "        # Second conv layer for feature extraction\n",
    "        self.conv2 = nn.Conv1d(\n",
    "            in_channels=n_filters,\n",
    "            out_channels=n_filters * 2,\n",
    "            kernel_size=3,\n",
    "            padding=1\n",
    "        )\n",
    "        self.bn2 = nn.BatchNorm1d(n_filters * 2)\n",
    "\n",
    "        # Final conv layer to produce logits\n",
    "        self.conv3 = nn.Conv1d(\n",
    "            in_channels=n_filters * 2,\n",
    "            out_channels=1,  # Binary classification\n",
    "            kernel_size=1\n",
    "        )\n",
    "        \n",
    "        self.dropout = nn.Dropout(0.2)\n",
    "        \n",
    "    def forward(self, x):\n",
    "        \"\"\"\n",
    "        Forward pass\n",
    "        \n",
    "        Args:\n",
    "            x (torch.Tensor): Input log spectrogram of shape (batch_size, channels, time)\n",
    "            \n",
    "        Returns:\n",
    "            torch.Tensor: Logits for defect classification\n",
    "        \"\"\"\n",
    "        input_length = x.size(-1)\n",
    "\n",
    "        # First conv with ReLU activation and batch norm\n",
    "        x = self.conv1(x)\n",
    "        x = self.bn1(x)\n",
    "        x = F.relu(x)\n",
    "        assert x.size(-1) == input_length, f\"Length mismatch after conv1: got {x.size(-1)}, expected {input_length}\"\n",
    "\n",
    "        x = self.dropout(x)\n",
    "        \n",
    "        # Second conv with ReLU activation and batch norm\n",
    "        x = self.conv2(x)\n",
    "        x = self.bn2(x)\n",
    "        x = F.relu(x)\n",
    "\n",
    "        x = self.dropout(x)\n",
    "        \n",
    "        # Final conv to produce logits\n",
    "        x = self.conv3(x)\n",
    "        \n",
    "        # Remove channel dimension\n",
    "        return x.squeeze(1)\n",
    "\n",
    "class DefectDetector:\n",
    "    def __init__(self, model, class_weights=None, focal_gamma=2, device='cuda' if torch.cuda.is_available() else 'cpu'):\n",
    "        \"\"\"\n",
    "        Wrapper class for training and inference\n",
    "        \n",
    "        Args:\n",
    "            model (nn.Module): The neural network model\n",
    "            device (str): Device to run on ('cuda' or 'cpu')\n",
    "        \"\"\"\n",
    "        self.model = model.to(device)\n",
    "        self.device = device\n",
    "        # Use Focal Loss with class weights if provided\n",
    "        if class_weights is not None:\n",
    "            self.criterion = SmoothFocalLoss(device=device, alpha=class_weights.to(device) / 4, gamma=focal_gamma, window_size=21)\n",
    "        else:\n",
    "            self.criterion = nn.BCEWithLogitsLoss()\n",
    "        self.optimizer = optim.Adam(model.parameters(), lr=0.00001)\n",
    "        \n",
    "    def train_epoch(self, train_loader):\n",
    "        \"\"\"\n",
    "        Train for one epoch\n",
    "        \n",
    "        Args:\n",
    "            train_loader (DataLoader): Training data loader\n",
    "            \n",
    "        Returns:\n",
    "            float: Average loss for the epoch\n",
    "        \"\"\"\n",
    "        self.model.train()\n",
    "        total_loss = 0\n",
    "        \n",
    "        for batch_idx, (spectrogram, labels) in enumerate(train_loader):\n",
    "            spectrogram = spectrogram.to(self.device)\n",
    "            labels = labels.to(self.device)\n",
    "            \n",
    "            self.optimizer.zero_grad()\n",
    "            logits = self.model(spectrogram)\n",
    "            loss = self.criterion(logits, labels)\n",
    "            \n",
    "            loss.backward()\n",
    "            self.optimizer.step()\n",
    "            \n",
    "            total_loss += loss.item()\n",
    "            \n",
    "        return total_loss / len(train_loader)\n",
    "    \n",
    "    def evaluate(self, test_loader, cheat=False):\n",
    "        \"\"\"\n",
    "        Evaluate model on test set\n",
    "        \n",
    "        Args:\n",
    "            test_loader (DataLoader): Test data loader\n",
    "            \n",
    "        Returns:\n",
    "            tuple: (average loss, f1 score)\n",
    "        \"\"\"\n",
    "        self.model.eval()\n",
    "        total_loss = 0\n",
    "        all_predictions = []\n",
    "        all_labels = []\n",
    "        \n",
    "        with torch.no_grad():\n",
    "            for spectrogram, labels in test_loader:\n",
    "                spectrogram = spectrogram.to(self.device)\n",
    "                labels = labels.to(self.device)\n",
    "                \n",
    "                logits = self.model(spectrogram)\n",
    "                loss = self.criterion(logits, labels)\n",
    "                \n",
    "                predictions = (torch.sigmoid(logits) > 0.5).float()\n",
    "                if cheat:\n",
    "                    predictions = labels.clone().float()\n",
    "                all_predictions.append(predictions.cpu())\n",
    "                all_labels.append(labels.cpu())\n",
    "                total_loss += loss.item()\n",
    "        \n",
    "        # Concatenate all batches\n",
    "        all_predictions = torch.cat(all_predictions, dim=0)\n",
    "        all_labels = torch.cat(all_labels, dim=0)\n",
    "        \n",
    "        # Calculate F1 score with tolerance window\n",
    "        f1_score = self._calculate_f1_with_tolerance(all_predictions, all_labels)\n",
    "        \n",
    "        return total_loss / len(test_loader), f1_score\n",
    "    \n",
    "    def _calculate_f1_with_tolerance(self, predictions, labels, tolerance_ms=100):\n",
    "        \"\"\"\n",
    "        Calculate F1 score with tolerance window of +/- 100ms\n",
    "        \n",
    "        Args:\n",
    "            predictions (torch.Tensor): Model predictions\n",
    "            labels (torch.Tensor): Ground truth labels\n",
    "            tolerance_ms (int): Tolerance window in milliseconds\n",
    "            \n",
    "        Returns:\n",
    "            float: F1 score\n",
    "        \"\"\"\n",
    "        # Convert tolerance to number of frames (assuming 20ms per frame)\n",
    "        tolerance_frames = int(tolerance_ms * sr / 1000 / hop_length)\n",
    "        \n",
    "        tp = 0\n",
    "        fp = 0\n",
    "        fn = 0\n",
    "        \n",
    "        for i in range(len(predictions)):\n",
    "            pred_positives = torch.where(predictions[i] == 1)[0]\n",
    "            true_positives = torch.where(labels[i] == 1)[0]\n",
    "            \n",
    "            # For each predicted positive, check if it's within tolerance of a true positive\n",
    "            matched_true_pos = set()\n",
    "            \n",
    "            matched = False\n",
    "            for pred_pos in pred_positives:\n",
    "                for true_pos in true_positives:\n",
    "                    if abs(pred_pos - true_pos) <= tolerance_frames:\n",
    "                        matched = True\n",
    "                        if int(true_pos) not in matched_true_pos:\n",
    "                            tp += 1\n",
    "                            matched_true_pos.add(int(true_pos))\n",
    "                            break\n",
    "                if not matched:\n",
    "                    fp += 1\n",
    "            # Count false negatives (true positives that weren't matched)\n",
    "            fn += len(true_positives) - len(matched_true_pos)\n",
    "        \n",
    "        # Calculate precision, recall and F1\n",
    "        precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n",
    "        recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n",
    "        f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0\n",
    "        \n",
    "        return f1\n",
    "    \n",
    "    def predict(self, spectrogram):\n",
    "        \"\"\"\n",
    "        Make predictions on new data\n",
    "        \n",
    "        Args:\n",
    "            spectrogram (torch.Tensor): Input spectrogram of shape (batch_size, channels, time)\n",
    "            \n",
    "        Returns:\n",
    "            torch.Tensor: Binary predictions for each time step\n",
    "        \"\"\"\n",
    "        self.model.eval()\n",
    "        with torch.no_grad():\n",
    "            spectrogram = spectrogram.to(self.device)\n",
    "            logits = self.model(spectrogram)\n",
    "            predictions = torch.sigmoid(logits) > 0.5\n",
    "        return predictions.cpu()\n",
    "\n",
    "def prepare_data(spectrograms, labels, batch_size=8, train_ratio=0.8):\n",
    "    \"\"\"\n",
    "    Prepare train and test DataLoaders\n",
    "    \n",
    "    Args:\n",
    "        spectrograms (np.ndarray): Array of spectrograms\n",
    "        labels (np.ndarray): Binary labels\n",
    "        batch_size (int): Batch size for training\n",
    "        train_ratio (float): Ratio of data to use for training\n",
    "        \n",
    "    Returns:\n",
    "        tuple: (train_loader, test_loader)\n",
    "    \"\"\"\n",
    "    dataset = SpectrogramDataset(spectrograms, labels)\n",
    "    \n",
    "    # Calculate split sizes\n",
    "    train_size = int(train_ratio * len(dataset))\n",
    "    test_size = len(dataset) - train_size\n",
    "    \n",
    "    # Split dataset\n",
    "    train_dataset, test_dataset = random_split(\n",
    "        dataset, \n",
    "        [train_size, test_size],\n",
    "        generator=torch.Generator().manual_seed(42)\n",
    "    )\n",
    "    \n",
    "    # Create data loaders\n",
    "    train_loader = DataLoader(\n",
    "        train_dataset,\n",
    "        batch_size=batch_size,\n",
    "        shuffle=True,\n",
    "        num_workers=4\n",
    "    )\n",
    "    \n",
    "    test_loader = DataLoader(\n",
    "        test_dataset,\n",
    "        batch_size=batch_size,\n",
    "        shuffle=False,\n",
    "        num_workers=4\n",
    "    )\n",
    "    \n",
    "    return train_loader, test_loader, dataset\n",
    "\n",
    "in_channels = clips[0].spectrogram.shape[1]\n",
    "sequence_length = 1000\n",
    "\n",
    "# cut clips to sequence length\n",
    "spectrograms = []\n",
    "labels = []\n",
    "for clip in clips:\n",
    "    for i in range(0, clip.spectrogram.shape[0] - sequence_length + 1, sequence_length):\n",
    "        time_lhs = i * hop_length / sr\n",
    "        time_rhs = time_lhs + sequence_length * hop_length / sr\n",
    "        this_spec = clip.spectrogram[i:i+sequence_length]\n",
    "        this_labels = np.zeros(this_spec.shape[0])\n",
    "        shimmers_in_this_seg = clip.shimmers[np.logical_and(clip.shimmers >= time_lhs, clip.shimmers <= time_rhs)]\n",
    "        this_labels[np.floor((shimmers_in_this_seg - time_lhs) * sr / hop_length).astype(int)] = 1\n",
    "        spectrograms.append(this_spec)\n",
    "        labels.append(this_labels)\n",
    "\n",
    "\n",
    "spectrograms = np.array(spectrograms).transpose(0, 2, 1)\n",
    "labels = np.array(labels)\n",
    "\n",
    "train_loader, test_loader, dataset = prepare_data(spectrograms, labels)\n",
    "\n",
    "def train():\n",
    "    # Plot a few examples of spectrograms with their labels\n",
    "    # import matplotlib.pyplot as plt\n",
    "\n",
    "    # Find spectrograms with non-empty labels\n",
    "    # non_empty_indices = [i for i, label in enumerate(labels) if np.any(label == 1)]\n",
    "    \n",
    "    # Plot 3 examples\n",
    "    # fig, axes = plt.subplots(3, 1, figsize=(15, 12))\n",
    "    # for i, ax in enumerate(axes):\n",
    "    #     if i < len(non_empty_indices):\n",
    "    #         idx = non_empty_indices[i]\n",
    "    #         spec = spectrograms[idx]\n",
    "    #         label = labels[idx]\n",
    "            \n",
    "    #         # Plot spectrogram\n",
    "    #         im = ax.imshow(spec.T, aspect='auto', origin='lower')\n",
    "            \n",
    "    #         # Add vertical lines for label positions\n",
    "    #         label_positions = np.where(label == 1)[0]\n",
    "    #         for pos in label_positions:\n",
    "    #             ax.axvline(x=pos, color='red', linestyle='-', alpha=0.5)\n",
    "            \n",
    "    #         ax.set_title(f'Spectrogram {idx} with {len(label_positions)} shimmer points')\n",
    "    #         plt.colorbar(im, ax=ax)\n",
    "    \n",
    "    # plt.tight_layout()\n",
    "    # plt.show()\n",
    "\n",
    "    # Create model and trainer\n",
    "    model = AudioDefectCNN(in_channels=in_channels)\n",
    "    detector = DefectDetector(model, class_weights=dataset.get_class_weights())\n",
    "\n",
    "    # Display initial F1 score\n",
    "    dataset.is_augmenting = False\n",
    "    test_loss, f1_score = detector.evaluate(test_loader)\n",
    "    print(f\"Initial F1 Score: {f1_score:.4f}\")\n",
    "\n",
    "    _, cheat_f1_score = detector.evaluate(test_loader, cheat=True)\n",
    "    print(f\"Initial F1 Score (cheat): {cheat_f1_score:.4f}\")\n",
    "\n",
    "    # Training loop\n",
    "    n_epochs = 3000\n",
    "    best_f1_score = 0\n",
    "    for epoch in range(n_epochs):\n",
    "        dataset.is_augmenting = True\n",
    "        train_loss = detector.train_epoch(train_loader)\n",
    "        dataset.is_augmenting = False\n",
    "        test_loss, test_f1_score = detector.evaluate(test_loader)\n",
    "\n",
    "        if test_f1_score > best_f1_score:\n",
    "            best_f1_score = test_f1_score\n",
    "            print(f\"New best F1 score: {best_f1_score:.4f}\")\n",
    "            torch.save(model.state_dict(), \"best_model.pt\")\n",
    "        \n",
    "        print(f\"Epoch {epoch+1}/{n_epochs} | Train Loss: {train_loss:.4f} | Test Loss: {test_loss:.4f} | F1 Score: {test_f1_score:.4f} | Best F1: {best_f1_score:.4f}\")\n",
    "\n",
    "    return model, train_loader, test_loader, dataset\n",
    "\n",
    "model, train_loader, test_loader, dataset = train()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "# Get predictions on test set\n",
    "model = AudioDefectCNN(in_channels=clips[0].spectrogram.shape[1])\n",
    "model.load_state_dict(torch.load(\"best_model.pt\", weights_only=True, map_location=torch.device('cpu')))\n",
    "model.eval()\n",
    "with torch.no_grad():\n",
    "    # Select random samples from test set\n",
    "    dataset.is_augmenting = False\n",
    "    num_samples = 6\n",
    "    rand_indices = np.random.randint(0, len(test_loader.dataset), num_samples)\n",
    "    samples = [test_loader.dataset[i] for i in rand_indices]\n",
    "    \n",
    "    # Separate spectrograms and labels\n",
    "    spectrograms = torch.stack([s[0] for s in samples])\n",
    "    labels = torch.stack([s[1] for s in samples])\n",
    "    \n",
    "    # Get model predictions\n",
    "    predictions = model(spectrograms)\n",
    "    #predictions = (torch.sigmoid(predictions) > 0.5).float()\n",
    "\n",
    "# Create visualization\n",
    "fig, axes = plt.subplots(num_samples, 1, figsize=(15, 5*num_samples))\n",
    "if num_samples == 1:\n",
    "    axes = [axes]\n",
    "\n",
    "for i, ax in enumerate(axes):\n",
    "    # Plot spectrogram\n",
    "    spec = spectrograms[i].numpy()\n",
    "    im = ax.imshow(spec, aspect='auto', origin='lower')\n",
    "    \n",
    "    # Plot ground truth labels in red\n",
    "    label_positions = torch.where(labels[i] == 1)[0]\n",
    "    for pos in label_positions.tolist():\n",
    "        ax.axvline(x=pos, color='red', linestyle='-', alpha=0.5, label='Ground Truth' if pos == label_positions[0] else None)\n",
    "    \n",
    "    # Plot predictions as a continuous blue curve\n",
    "    pred_curve = torch.sigmoid(predictions[i]).numpy()\n",
    "    ax.plot(np.arange(spec.shape[1]), spec.shape[0] * pred_curve, color='blue', linestyle='--', alpha=0.5, label='Prediction')\n",
    "    \n",
    "    # Annotate points where prediction crosses 0.5 threshold\n",
    "    threshold_crossings = np.where((pred_curve[:-1] < 0.5) & (pred_curve[1:] >= 0.5))[0]\n",
    "    for pos in threshold_crossings:\n",
    "        ax.scatter(pos+1, spec.shape[0] * 0.5, color='cyan', marker='^', s=100, \n",
    "                  label='Threshold (0.5)' if pos == threshold_crossings[0] else None)\n",
    "    \n",
    "    ax.set_title(f'Sample {i+1}')\n",
    "    ax.legend()\n",
    "    plt.colorbar(im, ax=ax)\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm\n",
    "\n",
    "def shimmer_score(audio):\n",
    "    if isinstance(audio, str):\n",
    "        audio = Audio.from_file(audio, sample_rate=sr, n_channels=1)\n",
    "    assert isinstance(audio, Audio)\n",
    "    assert audio.n_channels == 1\n",
    "    assert audio.sample_rate == sr\n",
    "    spec = spectrogram(audio.array_float)\n",
    "\n",
    "    spec_torch = torch.from_numpy(spec.T).unsqueeze(0).to(dtype=torch.float32)\n",
    "    with torch.no_grad():\n",
    "        preds = model(spec_torch)\n",
    "        probs = torch.sigmoid(preds).squeeze()\n",
    "\n",
    "    # Find rising edges across 0.5 threshold\n",
    "    prev_below = probs[:-1] < 0.5\n",
    "    next_above = probs[1:] >= 0.5\n",
    "    rising_edge_indices = torch.where(prev_below & next_above)[0] + 1\n",
    "\n",
    "    seconds_per_frame = hop_length / sr\n",
    "    rising_edge_times = rising_edge_indices.numpy() * seconds_per_frame\n",
    "\n",
    "    return rising_edge_times, torch.sum(probs > 0.5).item() / audio.duration_s\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from concurrent.futures import ThreadPoolExecutor, as_completed\n",
    "\n",
    "def shimmer_score_chunks(path):\n",
    "    audio = Audio.from_file(str(path), sample_rate=sr, n_channels=1)\n",
    "\n",
    "    for pos in range(0, int(audio.duration_s - 10), 10):\n",
    "        chunk = audio.get_segment(pos, min(pos + 10, audio.duration_s))\n",
    "        if chunk.sample_rate != sr:\n",
    "            continue # what?\n",
    "        yield shimmer_score(chunk)[1]\n",
    "\n",
    "def evaluate_shimmer_scores(file_paths):\n",
    "    \"\"\"Evaluate shimmer scores for a list of audio files in parallel.\"\"\"\n",
    "    with ThreadPoolExecutor(max_workers=20) as executor:\n",
    "        futures = [executor.submit(shimmer_score_chunks, p) for p in file_paths]\n",
    "        scores = []\n",
    "        for future in tqdm(as_completed(futures), total=len(file_paths)):\n",
    "            scores.append(list(future.result()))\n",
    "    return scores\n",
    "\n",
    "# # some non-suno audio\n",
    "non_suno_paths = list(Path('/home/m4burns/flac_out').glob('*.flac'))[:50]\n",
    "non_suno_scores = evaluate_shimmer_scores(non_suno_paths)\n",
    "\n",
    "# training set\n",
    "train_paths = list(Path('shimmer_annotations').glob('*.wav')) + list(Path('shimmer_annotations/v2').glob('*.opus'))\n",
    "train_scores = evaluate_shimmer_scores(train_paths)\n",
    "\n",
    "# training set (only new files)\n",
    "train_paths_only_new = [p for p in train_paths if p.name.count('-') == 4]\n",
    "train_scores_only_new = evaluate_shimmer_scores(train_paths_only_new)\n",
    "\n",
    "# old shimmery data\n",
    "old_paths = list(Path('../shimmery').glob('*.wav'))\n",
    "old_scores = evaluate_shimmer_scores(old_paths)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print(np.max(non_suno_scores))\n",
    "Audio.from_file(str(non_suno_paths[np.argmax(non_suno_scores)]), n_channels=2, sample_rate=48000).play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#shimmer_score(\"shimmer_annotations/v2/22bf283f-982b-40c3-aecc-3878c5583062.opus\")\n",
    "#shimmer_score(\"neg_test.opus\")\n",
    "\n",
    "#[np.mean(x) for x in train_scores_only_new]\n",
    "\n",
    "shimmer_score(Audio.from_s3(\"s3://suno-data-uploads/studio/uploads/ceb33a1e-b9e1-4ecc-8c50-7185fb5b5995.opus\").convert(n_channels=1, sample_rate=sr, byte_width=2))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def flatten(xss):\n",
    "    return [x for xs in xss for x in xs]\n",
    "\n",
    "plt.figure(figsize=(10, 6))\n",
    "plt.hist([flatten(non_suno_scores), flatten(train_scores), flatten(train_scores_only_new), flatten(old_scores)], \n",
    "         label=['Non-Suno', 'Training Set (All)', 'Training Set (Only diff v2)', 'Diff v1 pre-fix'],\n",
    "         bins=20, alpha=0.7)\n",
    "plt.xlabel('Shimmer Score')\n",
    "plt.ylabel('Count')\n",
    "plt.title('Distribution of Shimmer Scores Across Different Sets')\n",
    "plt.legend()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "state_dict = {\n",
    "    \"in_channels\": clips[0].spectrogram.shape[1],\n",
    "    \"kernel_size\": model.conv1.kernel_size[0],\n",
    "    \"n_filters\": model.conv1.out_channels,\n",
    "    \"window_size\": window_size,\n",
    "    \"hop_length\": hop_length,\n",
    "    \"sample_rate\": sr,\n",
    "    \"high_cutoff\": high_cutoff,\n",
    "    \"low_cutoff\": low_cutoff,\n",
    "    \"state_dict\": model.state_dict(),\n",
    "}\n",
    "\n",
    "torch.save(state_dict, \"shimmer_cnn_2025-04-25.pt\")"
   ]
  }
 ],
 "metadata": {
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
