{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a0bd1142",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import json\n",
    "import os"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7c722038",
   "metadata": {},
   "outputs": [],
   "source": [
    "output_dir = \"/home/christian/code/christian/metadata/vox\"\n",
    "output_filename = \"artist_id_to_vox_paths_tr.json\"\n",
    "filepath = os.path.join(output_dir, output_filename)\n",
    "\n",
    "metas = json.load(open(filepath))\n",
    "print(len(metas))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "31d06b48",
   "metadata": {},
   "outputs": [],
   "source": [
    "metas[list(metas.keys())[0]]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "da9789d1",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "from typing import List, Tuple\n",
    "from suno_utils.audio import Audio\n",
    "\n",
    "\n",
    "def _fast_trim_mono(\n",
    "    x: np.ndarray,  # shape: (samples,), float32/64 in [-1, 1]\n",
    "    sr: int,  # sample rate (Hz)\n",
    "    thresh_db_rel: float = -35,  # keep where RMS > max_RMS + thresh (dB)\n",
    "    win_ms: float = 20.0,  # moving RMS window size (ms)\n",
    "    pad_ms: float = 20.0,  # pad around kept regions (ms)\n",
    "    min_keep_ms: float = 40.0,  # drop kept bits shorter than this (ms)\n",
    ") -> Tuple[np.ndarray, List[Tuple[int, int]]]:\n",
    "    \"\"\"\n",
    "    Ultra-fast silence trimmer for mono audio. No convolutions, all O(n).\n",
    "    Returns (trimmed_audio, kept_spans) with kept_spans in original sample indices.\n",
    "    \"\"\"\n",
    "    assert x.ndim == 1, \"Expected mono waveform of shape (samples,)\"\n",
    "    n = x.size\n",
    "    if n == 0:\n",
    "        return x[:0], []\n",
    "\n",
    "    # --- Moving RMS via cumulative sums (box filter), O(n) ---\n",
    "    # Compute moving average of power over a window, then sqrt.\n",
    "    win = max(1, int(round(sr * win_ms / 1000.0)))\n",
    "    if win > n:\n",
    "        win = n\n",
    "\n",
    "    # power and cumulative sum (use float64 for numeric safety)\n",
    "    sq = x.astype(np.float64) ** 2\n",
    "    csum = np.empty(n + 1, dtype=np.float64)\n",
    "    csum[0] = 0.0\n",
    "    np.cumsum(sq, out=csum[1:])  # csum[k] = sum_{i<k} sq[i]\n",
    "\n",
    "    # moving average (valid positions)\n",
    "    # ma_valid[t] = mean of sq[t : t+win]\n",
    "    ma_valid = (csum[win:] - csum[:-win]) / win  # length n - win + 1\n",
    "\n",
    "    # Center-align to original length by padding equally on both sides\n",
    "    left = win // 2\n",
    "    right = n - (ma_valid.size + left)\n",
    "    rms = np.sqrt(np.pad(ma_valid, (left, right), mode=\"edge\"))\n",
    "\n",
    "    # --- Threshold relative to max ---\n",
    "    eps = 1e-12\n",
    "    rel_db = 20.0 * np.log10(np.maximum(rms, eps) / (np.max(rms) + eps))\n",
    "    mask = rel_db > thresh_db_rel  # True = keep\n",
    "\n",
    "    # --- Turn mask into spans, expand by pad, merge, drop short ---\n",
    "    pad = max(0, int(round(sr * pad_ms / 1000.0)))\n",
    "    min_keep = max(1, int(round(sr * min_keep_ms / 1000.0)))\n",
    "\n",
    "    # Find rising/falling edges\n",
    "    m = mask.astype(np.int8)\n",
    "    edges = np.flatnonzero(np.diff(m, prepend=0, append=0))\n",
    "    # edges come in pairs [start0, end0, start1, end1, ...]\n",
    "    starts = edges[::2]\n",
    "    ends = edges[1::2]\n",
    "\n",
    "    if starts.size == 0:\n",
    "        return x[:0], []\n",
    "\n",
    "    # Expand by pad and clamp\n",
    "    starts = np.maximum(0, starts - pad)\n",
    "    ends = np.minimum(n, ends + pad)\n",
    "\n",
    "    # Merge overlaps and drop short spans\n",
    "    spans: List[Tuple[int, int]] = []\n",
    "    s_prev = int(starts[0])\n",
    "    e_prev = int(ends[0])\n",
    "    for s, e in zip(starts[1:], ends[1:]):\n",
    "        s = int(s)\n",
    "        e = int(e)\n",
    "        if s <= e_prev:  # overlap/adjacent -> merge\n",
    "            e_prev = max(e_prev, e)\n",
    "        else:\n",
    "            if (e_prev - s_prev) >= min_keep:\n",
    "                spans.append((s_prev, e_prev))\n",
    "            s_prev, e_prev = s, e\n",
    "    # last span\n",
    "    if (e_prev - s_prev) >= min_keep:\n",
    "        spans.append((s_prev, e_prev))\n",
    "\n",
    "    if not spans:\n",
    "        return x[:0], []\n",
    "\n",
    "    # --- Concatenate kept spans (one pass) ---\n",
    "    parts = [x[a:b] for (a, b) in spans]\n",
    "    y = np.concatenate(parts, axis=0).astype(x.dtype)\n",
    "    return y, spans\n",
    "\n",
    "def _get_segments(audio_np: np.ndarray, sample_rate: int, num_segments: int = 3, segment_duration_sec: float = 3.0):\n",
    "    \"\"\"\n",
    "    audio_np: np.ndarray of shape (samples,)\n",
    "    sample_rate: int\n",
    "    num_segments: int, number of random segments to crop and return\n",
    "    segment_duration_sec: float, length (seconds) of each segment to return\n",
    "    Returns:\n",
    "        segments: list of np.ndarray of shape (segment_samples,)\n",
    "    \"\"\"\n",
    "    total_samples = len(audio_np)\n",
    "    segment_samples = int(segment_duration_sec * sample_rate)\n",
    "    if total_samples < segment_samples:\n",
    "        # Pad if audio is too short\n",
    "        padded = np.pad(audio_np, (0, segment_samples - total_samples), mode='constant')\n",
    "        return [padded.copy() for _ in range(num_segments)]\n",
    "\n",
    "    segments = []\n",
    "    for _ in range(num_segments):\n",
    "        start_idx = np.random.randint(0, total_samples - segment_samples + 1)\n",
    "        seg = audio_np[start_idx:start_idx + segment_samples]\n",
    "        segments.append(seg.copy())\n",
    "    return segments"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "19de0e16",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "import random\n",
    "from torch.utils.data import Dataset\n",
    "from suno_utils.audio import Audio\n",
    "\n",
    "# collate function\n",
    "def artist_segment_collate_fn(batch):\n",
    "    \"\"\"\n",
    "    Collate function for batching artist segment samples.\n",
    "\n",
    "    Each element in `batch` is a tuple:\n",
    "        (artist_id: str, artist_index: int, segments: List[np.ndarray])\n",
    "\n",
    "    Returns:\n",
    "        flat_artist_ids: List[str]                  # len = batch_size * num_segments\n",
    "        flat_artist_indices: torch.LongTensor       # shape (batch_size * num_segments,)\n",
    "        flat_segments: torch.FloatTensor            # shape (batch_size * num_segments, segment_samples)\n",
    "    \"\"\"\n",
    "    flat_artist_ids = []\n",
    "    flat_artist_indices = []\n",
    "    flat_segments = []\n",
    "\n",
    "    for artist_id, artist_index, segments in batch:\n",
    "        # segments: list of np.ndarray (num_segments, segment_samples)\n",
    "        for seg in segments:\n",
    "            flat_artist_ids.append(artist_id)\n",
    "            flat_artist_indices.append(artist_index)\n",
    "            flat_segments.append(torch.tensor(seg, dtype=torch.float32))\n",
    "\n",
    "    flat_artist_indices = torch.tensor(flat_artist_indices, dtype=torch.long)\n",
    "    flat_segments = torch.stack(flat_segments, dim=0)  # (batch_size * num_segments, segment_samples)\n",
    "\n",
    "    return flat_artist_ids, flat_artist_indices, flat_segments\n",
    "\n",
    "\n",
    "\n",
    "class BasicIterableDataset(torch.utils.data.IterableDataset):\n",
    "    def __init__(self, metas, num_segments: int = 3, segment_duration_s: float = 3.0):\n",
    "        super(BasicIterableDataset, self).__init__()\n",
    "        self.metas = metas\n",
    "        self.num_segments = num_segments\n",
    "        self.segment_duration_s = segment_duration_s\n",
    "        \n",
    "        # Create artist_id to index mapping for classifier\n",
    "        self.artist_ids = list(metas.keys())\n",
    "        self.artist_id_to_index = {artist_id: idx for idx, artist_id in enumerate(self.artist_ids)}\n",
    "        self.num_artists = len(self.artist_ids)\n",
    "        \n",
    "        print(f\"Dataset initialized with {self.num_artists} artists\")\n",
    "\n",
    "    def __iter__(self):\n",
    "        while True:\n",
    "            # Randomly sample an artist\n",
    "            artist_id = random.choice(self.artist_ids)\n",
    "            artist_index = self.artist_id_to_index[artist_id]\n",
    "            \n",
    "            # Randomly sample a stem from this artist\n",
    "            stems_dicts = self.metas[artist_id]\n",
    "            stem_dict = random.choice(stems_dicts)\n",
    "\n",
    "            # load this audio file\n",
    "            audio = Audio.from_file(stem_dict[\"path\"], n_channels=1)\n",
    "\n",
    "            # trim silence\n",
    "            audio_trim, _ = _fast_trim_mono(audio.array_float, audio.sample_rate)\n",
    "\n",
    "            # get segments\n",
    "            segments = _get_segments(audio_trim, audio.sample_rate, self.num_segments, self.segment_duration_s)\n",
    "\n",
    "            # don't yield the segments separately\n",
    "            # we will return list with the artist ids and indices\n",
    "            # and then use a special collate function to merge them\n",
    "            yield (artist_id, artist_index, segments)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a83a72c2",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f50678e1",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Here is my collate function for segment batches\n",
    "\n",
    "# collate function\n",
    "\n",
    "# collate function\n",
    "def artist_segment_collate_fn(batch):\n",
    "    \"\"\"\n",
    "    Collate function for batching artist segment samples.\n",
    "\n",
    "    Each element in `batch` is a tuple:\n",
    "        (artist_id: str, artist_index: int, segments: List[np.ndarray])\n",
    "\n",
    "    Returns:\n",
    "        flat_artist_ids: List[str]                  # len = batch_size * num_segments\n",
    "        flat_artist_indices: torch.LongTensor       # shape (batch_size * num_segments,)\n",
    "        flat_segments: torch.FloatTensor            # shape (batch_size * num_segments, segment_samples)\n",
    "    \"\"\"\n",
    "    flat_artist_ids = []\n",
    "    flat_artist_indices = []\n",
    "    flat_segments = []\n",
    "\n",
    "    for artist_id, artist_index, segments in batch:\n",
    "        # segments: list of np.ndarray (num_segments, segment_samples)\n",
    "        for seg in segments:\n",
    "            flat_artist_ids.append(artist_id)\n",
    "            flat_artist_indices.append(artist_index)\n",
    "            flat_segments.append(torch.tensor(seg, dtype=torch.float32))\n",
    "\n",
    "    flat_artist_indices = torch.tensor(flat_artist_indices, dtype=torch.long)\n",
    "    flat_segments = torch.stack(flat_segments, dim=0)  # (batch_size * num_segments, segment_samples)\n",
    "\n",
    "    return flat_artist_ids, flat_artist_indices, flat_segments\n",
    "\n",
    "\n",
    "# test the collate function\n",
    "dataset = BasicIterableDataset(metas)\n",
    "samples = [next(iter(dataset)) for _ in range(4)]\n",
    "batch_artist_ids, batch_artist_indices, batch_segments = artist_segment_collate_fn(list(samples))\n",
    "\n",
    "print(f\"\\nBatch info:\")\n",
    "print(f\"  batch_artist_ids: {batch_artist_ids}\")\n",
    "print(f\"  batch_artist_indices: {batch_artist_indices}\")\n",
    "print(f\"  batch_segments shape: {batch_segments.shape}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4012eaf5",
   "metadata": {},
   "outputs": [],
   "source": [
    "# speaker_embed_arcface.py\n",
    "import math\n",
    "from dataclasses import dataclass\n",
    "from typing import Optional\n",
    "\n",
    "import torch\n",
    "import torch.nn as nn\n",
    "import torch.nn.functional as F\n",
    "\n",
    "\n",
    "# ----------------------------\n",
    "# TDNN building blocks\n",
    "# ----------------------------\n",
    "class TDNNBlock(nn.Module):\n",
    "    \"\"\"\n",
    "    1D time-dilated conv (x-vector style) with ReLU+BN.\n",
    "    Input:  (B, C_in, T)\n",
    "    Output: (B, C_out, T)\n",
    "    \"\"\"\n",
    "    def __init__(self, c_in: int, c_out: int, kernel: int = 5, dilation: int = 1):\n",
    "        super().__init__()\n",
    "        pad = dilation * (kernel // 2)\n",
    "        self.conv = nn.Conv1d(c_in, c_out, kernel_size=kernel, dilation=dilation, padding=pad, bias=False)\n",
    "        self.bn   = nn.BatchNorm1d(c_out)\n",
    "        self.act  = nn.ReLU(inplace=True)\n",
    "\n",
    "    def forward(self, x):\n",
    "        x = self.conv(x)\n",
    "        x = self.bn(x)\n",
    "        x = self.act(x)\n",
    "        return x\n",
    "\n",
    "\n",
    "class StatsPooling(nn.Module):\n",
    "    \"\"\"\n",
    "    Mean+Std pooling over time.\n",
    "    Input:  (B, C, T)\n",
    "    Output: (B, 2*C)\n",
    "    \"\"\"\n",
    "    def forward(self, x):\n",
    "        # x: (B, C, T)\n",
    "        mean = x.mean(dim=-1)\n",
    "        std  = x.std(dim=-1, unbiased=False)\n",
    "        return torch.cat([mean, std], dim=1)\n",
    "\n",
    "\n",
    "# ----------------------------\n",
    "# Speaker Encoder -> Embedding\n",
    "# ----------------------------\n",
    "@dataclass\n",
    "class EncoderConfig:\n",
    "    n_mels: int = 80\n",
    "    tdnn_channels: tuple = (256, 256, 256, 256, 256)  # 5 TDNN layers\n",
    "    tdnn_kernels:  tuple = (5,   5,   7,   1,   1)\n",
    "    tdnn_dilations:tuple = (1,   2,   3,   1,   1)\n",
    "    emb_hidden: int = 256   # penultimate projection before final embedding\n",
    "    emb_dim: int = 128      # final embedding dimension\n",
    "    dropout: float = 0.1\n",
    "\n",
    "\n",
    "class SpeakerEncoder(nn.Module):\n",
    "    \"\"\"\n",
    "    Input:  log-mel features (B, F, T); F = n_mels\n",
    "    Output: L2-normalized embedding (B, emb_dim)\n",
    "    \"\"\"\n",
    "    def __init__(self, cfg: EncoderConfig):\n",
    "        super().__init__()\n",
    "        self.cfg = cfg\n",
    "\n",
    "        C = [cfg.n_mels] + list(cfg.tdnn_channels)\n",
    "        self.tdnn = nn.Sequential(*[\n",
    "            TDNNBlock(C[i], C[i+1], kernel=cfg.tdnn_kernels[i], dilation=cfg.tdnn_dilations[i])\n",
    "            for i in range(len(cfg.tdnn_channels))\n",
    "        ])\n",
    "\n",
    "        self.pool = StatsPooling()                                # (B, 2*C_last)\n",
    "        pooled_dim = 2 * cfg.tdnn_channels[-1]\n",
    "\n",
    "        self.fc1 = nn.Linear(pooled_dim, cfg.emb_hidden, bias=False)\n",
    "        self.bn1 = nn.BatchNorm1d(cfg.emb_hidden)\n",
    "        self.drop= nn.Dropout(p=cfg.dropout)\n",
    "\n",
    "        self.fc2 = nn.Linear(cfg.emb_hidden, cfg.emb_dim, bias=True)\n",
    "\n",
    "        # Kaiming for convs; Xavier for linears is fine\n",
    "        self.apply(self._init_weights)\n",
    "\n",
    "    @staticmethod\n",
    "    def _init_weights(m):\n",
    "        if isinstance(m, nn.Conv1d):\n",
    "            nn.init.kaiming_normal_(m.weight, nonlinearity='relu')\n",
    "        elif isinstance(m, nn.Linear):\n",
    "            nn.init.xavier_uniform_(m.weight)\n",
    "            if m.bias is not None:\n",
    "                nn.init.zeros_(m.bias)\n",
    "\n",
    "    def forward(self, mels: torch.Tensor) -> torch.Tensor:\n",
    "        \"\"\"\n",
    "        mels: (B, F, T)\n",
    "        returns: L2-normalized embeddings (B, emb_dim)\n",
    "        \"\"\"\n",
    "        x = mels  # (B, F, T)\n",
    "        x = self.tdnn(x)               # (B, C, T)\n",
    "        x = self.pool(x)               # (B, 2C)\n",
    "        x = self.fc1(x)                # (B, H)\n",
    "        x = self.bn1(x)\n",
    "        x = F.relu(x, inplace=True)\n",
    "        x = self.drop(x)\n",
    "        x = self.fc2(x)                # (B, D)\n",
    "        # L2 normalize to put on the unit hypersphere\n",
    "        x = F.normalize(x, p=2, dim=-1)\n",
    "        return x\n",
    "\n",
    "    @torch.no_grad()\n",
    "    def embed(self, mels: torch.Tensor) -> torch.Tensor:\n",
    "        self.eval()\n",
    "        return self.forward(mels)\n",
    "\n",
    "\n",
    "# ----------------------------\n",
    "# ArcFace / AAM-Softmax Head\n",
    "# ----------------------------\n",
    "class ArcMarginProduct(nn.Module):\n",
    "    \"\"\"\n",
    "    Implements AAM-Softmax (ArcFace) logits on-the-fly.\n",
    "    - Weight matrix W is L2-normalized per row.\n",
    "    - Inputs are expected already L2-normalized.\n",
    "    logits = s * cos(theta + m) for the target class, s * cos(theta) otherwise.\n",
    "\n",
    "    Args:\n",
    "      in_features:  embedding dim D\n",
    "      num_classes:  number of speakers C\n",
    "      s:            scale (30 is common)\n",
    "      m:            angular margin (0.2~0.3 common)\n",
    "      easy_margin:  if True, use the easy-margin variant\n",
    "      ls_eps:       label smoothing epsilon (optional)\n",
    "    \"\"\"\n",
    "    def __init__(self,\n",
    "                 in_features: int,\n",
    "                 num_classes: int,\n",
    "                 s: float = 30.0,\n",
    "                 m: float = 0.2,\n",
    "                 easy_margin: bool = False,\n",
    "                 ls_eps: float = 0.0):\n",
    "        super().__init__()\n",
    "        self.in_features = in_features\n",
    "        self.num_classes = num_classes\n",
    "        self.s = s\n",
    "        self.m = m\n",
    "        self.easy_margin = easy_margin\n",
    "        self.ls_eps = ls_eps\n",
    "\n",
    "        self.weight = nn.Parameter(torch.empty(num_classes, in_features))\n",
    "        nn.init.xavier_uniform_(self.weight)\n",
    "\n",
    "        # Precompute margin constants\n",
    "        self.cos_m = math.cos(m)\n",
    "        self.sin_m = math.sin(m)\n",
    "        self.th = math.cos(math.pi - m)  # cos(pi - m)\n",
    "        self.mm = math.sin(math.pi - m) * m\n",
    "\n",
    "    def forward(self, emb: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:\n",
    "        \"\"\"\n",
    "        emb:    (B, D) L2-normalized\n",
    "        labels: (B,)   int64 in [0, C-1]\n",
    "        returns: scaled logits for CE loss, shape (B, C)\n",
    "        \"\"\"\n",
    "        # Normalize class weights\n",
    "        W = F.normalize(self.weight, p=2, dim=1)  # (C, D)\n",
    "\n",
    "        # Cosine similarity between emb and each class weight\n",
    "        # cos_theta: (B, C)\n",
    "        cos_theta = torch.matmul(emb, W.t()).clamp(-1.0, 1.0)\n",
    "\n",
    "        # Gather target cosine\n",
    "        idx = torch.arange(emb.size(0), device=emb.device)\n",
    "        cos_theta_y = cos_theta[idx, labels]  # (B,)\n",
    "\n",
    "        # Compute cos(theta + m) via trig identity\n",
    "        sin_theta_y = torch.sqrt(torch.clamp(1.0 - cos_theta_y * cos_theta_y, min=0.0))\n",
    "        cos_theta_m = cos_theta_y * self.cos_m - sin_theta_y * self.sin_m  # (B,)\n",
    "\n",
    "        if self.easy_margin:\n",
    "            # If easy margin: if cos(theta_y) > 0 use margin, else keep original cos\n",
    "            cond = (cos_theta_y > 0).to(cos_theta.dtype)\n",
    "            cos_theta_y_m = cond * cos_theta_m + (1 - cond) * cos_theta_y\n",
    "        else:\n",
    "            # Classic ArcFace margin decision\n",
    "            cond = (cos_theta_y > self.th).to(cos_theta.dtype)\n",
    "            cos_theta_y_m = cond * cos_theta_m + (1 - cond) * (cos_theta_y - self.mm)\n",
    "\n",
    "        # Replace target logit\n",
    "        logits = cos_theta.clone()\n",
    "        logits[idx, labels] = cos_theta_y_m\n",
    "\n",
    "        # Scale\n",
    "        logits = logits * self.s\n",
    "\n",
    "        # Optional label smoothing (applied in CE). We return logits; apply CE outside,\n",
    "        # but provide a helper to build smoothed targets if needed.\n",
    "        return logits\n",
    "\n",
    "\n",
    "# ----------------------------\n",
    "# Full Model = Encoder + ArcFace head\n",
    "# ----------------------------\n",
    "class SpeakerEmbeddingModel(nn.Module):\n",
    "    \"\"\"\n",
    "    Training:\n",
    "        logits = model(mels, labels) -> (B, C)  # pass to nn.CrossEntropyLoss\n",
    "    Inference:\n",
    "        emb = model.embed(mels) -> (B, D) L2-normalized\n",
    "    \"\"\"\n",
    "    def __init__(self, n_classes: int, cfg: Optional[EncoderConfig] = None,\n",
    "                 s: float = 30.0, m: float = 0.2, easy_margin: bool = False, ls_eps: float = 0.0):\n",
    "        super().__init__()\n",
    "        self.cfg = cfg or EncoderConfig()\n",
    "        self.encoder = SpeakerEncoder(self.cfg)\n",
    "        self.arcface = ArcMarginProduct(\n",
    "            in_features=self.cfg.emb_dim,\n",
    "            num_classes=n_classes,\n",
    "            s=s, m=m, easy_margin=easy_margin, ls_eps=ls_eps\n",
    "        )\n",
    "\n",
    "    def forward(self, mels: torch.Tensor, labels: Optional[torch.Tensor] = None):\n",
    "        \"\"\"\n",
    "        mels:   (B, F, T) float\n",
    "        labels: (B,) long, required for training with ArcFace\n",
    "        \"\"\"\n",
    "        emb = self.encoder(mels)  # (B, D), L2-normalized\n",
    "        if labels is None:\n",
    "            return emb\n",
    "        logits = self.arcface(emb, labels)  # (B, C)\n",
    "        return logits\n",
    "\n",
    "    @torch.no_grad()\n",
    "    def embed(self, mels: torch.Tensor) -> torch.Tensor:\n",
    "        return self.encoder.embed(mels)\n",
    "\n",
    "\n",
    "import torch\n",
    "import torchaudio as ta\n",
    "import torchaudio.functional as AF\n",
    "import torchaudio.transforms as AT\n",
    "\n",
    "def wav_to_logmels(\n",
    "    wav: torch.Tensor, sr: int,\n",
    "    target_sr: int = 16_000, n_mels: int = 80,\n",
    "    win_ms: float = 25, hop_ms: float = 10,\n",
    "    fmin: float = 20.0, fmax: float | None = None,\n",
    "    top_db: float = 80.0, cmvn: bool = True\n",
    ") -> torch.Tensor:\n",
    "    # mix to mono if (B, C, T)\n",
    "    if wav.dim() == 3:\n",
    "        wav = wav.mean(1)\n",
    "    # resample if needed\n",
    "    if sr != target_sr:\n",
    "        wav = AF.resample(wav, sr, target_sr)\n",
    "    n_fft = int(target_sr * win_ms / 1000)\n",
    "    win_length = n_fft\n",
    "    hop_length = int(target_sr * hop_ms / 1000)\n",
    "\n",
    "    # Ensure MelSpectrogram and its buffers are on the same device as input\n",
    "    device = wav.device\n",
    "\n",
    "    # Build MelSpectrogram on correct device & move module to input's device\n",
    "    mel_spect = AT.MelSpectrogram(\n",
    "        sample_rate=target_sr, n_fft=n_fft, win_length=win_length, hop_length=hop_length,\n",
    "        f_min=fmin, f_max=fmax or target_sr/2, n_mels=n_mels,\n",
    "        window_fn=lambda window_length: torch.hann_window(window_length, device=device),\n",
    "        power=2.0, mel_scale=\"slaney\", norm=\"slaney\"\n",
    "    )\n",
    "    mel_spect = mel_spect.to(device)\n",
    "    mel = mel_spect(wav)  # (B, n_mels, T)\n",
    "\n",
    "    # AmplitudeToDB -- ensure on same device as well\n",
    "    db_xfm = AT.AmplitudeToDB(stype=\"power\", top_db=top_db).to(device)\n",
    "    logmel = db_xfm(mel)\n",
    "\n",
    "    if cmvn:\n",
    "        mu = logmel.mean(dim=(-1, -2), keepdim=True)\n",
    "        sd = logmel.std(dim=(-1, -2), keepdim=True).clamp_min(1e-5)\n",
    "        logmel = (logmel - mu) / sd\n",
    "\n",
    "    return logmel\n",
    "\n",
    "# ----------------------------\n",
    "# Usage sketch\n",
    "# ----------------------------\n",
    "if __name__ == \"__main__\":\n",
    "    B, bins, T = 8, 80, 300        # batch=8, 80 mels, 300 frames (~3s if 100 fps)\n",
    "    C = 1000                    # number of speakers (classes)\n",
    "    x = torch.randn(B, bins, T)\n",
    "    y = torch.randint(0, C, (B,))\n",
    "\n",
    "    model = SpeakerEmbeddingModel(n_classes=C, cfg=EncoderConfig())\n",
    "    logits = model(x, y)        # (B, C)\n",
    "    loss = F.cross_entropy(logits, y)\n",
    "\n",
    "    loss.backward()\n",
    "    with torch.no_grad():\n",
    "        emb = model.embed(x)    # (B, D), L2-normalized\n",
    "        assert torch.allclose(emb.norm(dim=-1), torch.ones(B), atol=1e-4)\n",
    "        print(\"Embedding shape:\", emb.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6cbc28dc",
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm\n",
    "\n",
    "metas_subset = {k: v for k, v in metas.items() if k in list(metas.keys())[:10]}\n",
    "\n",
    "dataset = BasicIterableDataset(metas_subset, segment_duration_s=3.0, num_segments=6)\n",
    "dataloader = torch.utils.data.DataLoader(dataset, batch_size=4, collate_fn=artist_segment_collate_fn, num_workers=4, prefetch_factor=4)\n",
    "\n",
    "model = SpeakerEmbeddingModel(n_classes=len(dataset.artist_ids), cfg=EncoderConfig())\n",
    "model.cuda()\n",
    "# count number of parameters in millions\n",
    "num_params = sum(p.numel() for p in model.parameters()) / 1e6\n",
    "print(f\"Number of parameters: {num_params:.2f}M\")\n",
    "optimizer = torch.optim.AdamW(model.parameters(), lr=0.0001)\n",
    "# test the collate function\n",
    "\n",
    "pbar = tqdm(dataloader)\n",
    "for batch_idx, batch in enumerate(pbar):\n",
    "    batch_artist_ids, batch_artist_indices, batch_segments = batch\n",
    "\n",
    "    # move to gpu\n",
    "    batch_segments = batch_segments.cuda()\n",
    "    batch_artist_indices = batch_artist_indices.cuda()\n",
    "\n",
    "    with torch.no_grad():\n",
    "        batch_logmels = wav_to_logmels(batch_segments, 48000)\n",
    "\n",
    "    print(f\"{batch_logmels.std()}, {batch_logmels.mean()}, {batch_logmels.max()}, {batch_logmels.min()}\")\n",
    "\n",
    "    logits = model(batch_logmels, batch_artist_indices)\n",
    "    print(f\"{logits.std()}, {logits.mean()}, {logits.max()}, {logits.min()}\")\n",
    "    \n",
    "    loss = F.cross_entropy(logits, batch_artist_indices)\n",
    "    pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n",
    "    loss.backward()\n",
    "    optimizer.step()\n",
    "    optimizer.zero_grad()\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "480dc612",
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "# --------------------------------------------------------\n",
    "# 2. Basic cosine & EER helpers\n",
    "# --------------------------------------------------------\n",
    "def cosine(a, b):\n",
    "    return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b) + 1e-12))\n",
    "\n",
    "\n",
    "def compute_eer(scores, labels):\n",
    "    s, y = np.array(scores), np.array(labels)\n",
    "    order = np.argsort(s)\n",
    "    s, y = s[order], y[order]\n",
    "    P, N = max(1, y.sum()), max(1, (1 - y).sum())\n",
    "    tar_cum, imp_cum = np.cumsum(y), np.cumsum(1 - y)\n",
    "    FRR = tar_cum / P\n",
    "    FAR = 1.0 - imp_cum / N\n",
    "    diff = FAR - FRR\n",
    "    idx = np.where(np.sign(diff[:-1]) != np.sign(diff[1:]))[0]\n",
    "    if len(idx) == 0:\n",
    "        j = np.argmin(np.abs(diff))\n",
    "        return 0.5 * (FAR[j] + FRR[j]), s[j]\n",
    "    i = idx[0]\n",
    "    x0, x1, y0, y1 = s[i], s[i + 1], diff[i], diff[i + 1]\n",
    "    thr = x0 if y1 == y0 else x0 - y0 * (x1 - x0) / (y1 - y0)\n",
    "    w = 0 if x1 == x0 else (thr - x0) / (x1 - x0)\n",
    "    FAR_thr = FAR[i] + w * (FAR[i + 1] - FAR[i])\n",
    "    FRR_thr = FRR[i] + w * (FRR[i + 1] - FRR[i])\n",
    "    return float(0.5 * (FAR_thr + FRR_thr)), float(thr)\n",
    "\n",
    "\n",
    "# --------------------------------------------------------\n",
    "# 3. Single-file embedding helper\n",
    "# --------------------------------------------------------\n",
    "def embed_path(stem_dict, model, device, seg_s=3.0, rng=None, sr_target=48000):\n",
    "    \"\"\"\n",
    "    stem_dict: {\"path\": ..., \"duration_s\": ...}\n",
    "    \"\"\"\n",
    "    # load & trim\n",
    "    audio = Audio.from_file(stem_dict[\"path\"], n_channels=1)\n",
    "    audio_trim, _ = _fast_trim_mono(audio.array_float, audio.sample_rate)\n",
    "\n",
    "    wav = torch.tensor(audio_trim, dtype=torch.float32).unsqueeze(0)  # (1, T)\n",
    "    sr = audio.sample_rate\n",
    "\n",
    "    # resample\n",
    "    if sr != sr_target:\n",
    "        wav = AF.resample(wav, sr, sr_target)\n",
    "        sr = sr_target\n",
    "\n",
    "    # crop or pad to seg_s\n",
    "    seg_len = int(seg_s * sr)\n",
    "    T = wav.shape[-1]\n",
    "    if T < seg_len:\n",
    "        rep = math.ceil(seg_len / T)\n",
    "        wav = wav.repeat(1, rep)[:, :seg_len]\n",
    "    else:\n",
    "        start = rng.randrange(0, T - seg_len) if rng else (T - seg_len) // 2\n",
    "        wav = wav[:, start:start + seg_len]\n",
    "\n",
    "    # compute features + embed\n",
    "    mels = wav_to_logmels(wav.unsqueeze(0), sr)\n",
    "    with torch.no_grad():\n",
    "        emb = model.embed(mels.to(device)).cpu().numpy().squeeze()\n",
    "    emb /= np.linalg.norm(emb) + 1e-12\n",
    "    return emb\n",
    "\n",
    "# --------------------------------------------------------\n",
    "# 4. Main evaluation logic\n",
    "# --------------------------------------------------------\n",
    "def evaluate_eer_from_json(\n",
    "    json_path,\n",
    "    model,\n",
    "    num_speakers=200,\n",
    "    device=\"cuda\" if torch.cuda.is_available() else \"cpu\",\n",
    "    seed=0,\n",
    "):\n",
    "    rng = random.Random(seed)\n",
    "\n",
    "    with open(json_path) as f:\n",
    "        data = json.load(f)\n",
    "\n",
    "    speakers = [s for s, lst in data.items() if len(lst) >= 2]\n",
    "    #rng.shuffle(speakers)\n",
    "    speakers = speakers[:num_speakers]\n",
    "\n",
    "    print(f\"Using {len(speakers)} speakers with >=2 files\")\n",
    "\n",
    "    enroll_embs, test_embs = {}, {}\n",
    "    for spk in tqdm(speakers, desc=\"Embedding speakers\"):\n",
    "        enroll_embs[spk] = embed_path(data[spk][0], model, device, rng=rng)\n",
    "        test_embs[spk] = embed_path(data[spk][-1], model, device, rng=rng)\n",
    "\n",
    "    # build positive/negative pairs\n",
    "    scores, labels = [], []\n",
    "    for i, spk in enumerate(speakers):\n",
    "        # same speaker\n",
    "        scores.append(cosine(enroll_embs[spk], test_embs[spk]))\n",
    "        labels.append(1)\n",
    "        # random different speaker\n",
    "        j = rng.randrange(len(speakers))\n",
    "        while j == i:\n",
    "            j = rng.randrange(len(speakers))\n",
    "        other = speakers[j]\n",
    "        scores.append(cosine(enroll_embs[spk], test_embs[other]))\n",
    "        labels.append(0)\n",
    "\n",
    "    eer, thr = compute_eer(scores, labels)\n",
    "    print(f\"EER = {eer*100:.2f}%   threshold = {thr:.3f}   pairs = {len(scores)}\")\n",
    "    return eer, thr"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9e9418e0",
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "json_path = \"/home/christian/code/christian/metadata/vox/artist_id_to_vox_paths_val.json\"\n",
    "evaluate_eer_from_json(json_path, model, num_speakers=10)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b0e5a3af",
   "metadata": {},
   "outputs": [],
   "source": [
    "import sys\n",
    "sys.path.append(\"/home/christian/code/christian/vox_id/\")\n",
    "from train import BasicIterableDataset, artist_segment_collate_fn"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fd4312d2",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "\n",
    "metas_filepath = \"/home/christian/code/christian/metadata/vox/artist_id_to_vox_paths_val.json\"\n",
    "dataset = BasicIterableDataset(metas_filepath, segment_duration_s=3.0, num_segments=6, read_duration_s=30.0)\n",
    "dataloader = torch.utils.data.DataLoader(dataset, batch_size=32, collate_fn=artist_segment_collate_fn, num_workers=4, prefetch_factor=4)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fb5ce302",
   "metadata": {},
   "outputs": [],
   "source": [
    "from tqdm import tqdm\n",
    "for idx, batch in enumerate(tqdm(dataloader)):\n",
    "    artist_ids, artist_indices, segments = batch\n",
    "    #print(artist_indices, segments.shape)\n",
    "    \n",
    "    if idx > 100:\n",
    "        break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f341534d",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env_fa2",
   "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.15"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
