{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":129543,"databundleVersionId":15525987,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"bff6be96","cell_type":"markdown","source":"# 🐆 Jaguar Re-Identification — EVA-02 Training Pipeline\n\n> **Goal:** Build a deep learning pipeline for individual jaguar identification (Re-ID)\n> using the EVA-02 Vision Transformer architecture.\n\n---\n\n## 📋 Table of Contents\n\n| #  | Section | Description |\n|----|---------|-------------|\n| 1  | Environment Setup & Imports | Importing required libraries |\n| 2  | Configuration | Managing all hyperparameters with a dataclass |\n| 3  | Utility Functions | Seed, visualization, and metric computation |\n| 4  | Loss Functions & Layers | ArcFace, Triplet Loss, GeM Pooling |\n| 5  | Model Architecture | EVA-02 based Re-ID model |\n| 6  | Dataset & Augmentation | Data loading, augmentation, and PK sampler |\n| 7  | Training Engine | Mixed-precision training loop |\n| 8  | Main Pipeline | End-to-end training flow |\n\n---\n\n**Model:** `eva02_large_patch14_448` (ImageNet-22K → ImageNet-1K fine-tuned)  \n**Techniques:** ArcFace + Batch Hard Triplet Loss + GeM Pooling + Gradient Accumulation  \n**Framework:** PyTorch + timm + Albumentations\n","metadata":{}},{"id":"5eaaa5ad","cell_type":"markdown","source":"## 1. Environment Setup & Imports\n\nAll project dependencies are imported here, organized into logical groups:\n- **Standard library:** `gc`, `logging`, `random`, `os`, `math`, `warnings`\n- **Scientific computing:** `numpy`, `pandas`, `sklearn`\n- **Deep learning:** `torch`, `timm`\n- **Visualization:** `matplotlib`, `seaborn`, `plotly`\n- **Data augmentation:** `albumentations`\n","metadata":{}},{"id":"d3e19f3e","cell_type":"code","source":"# =============================================================================\n# Standard Python libraries\n# =============================================================================\nimport gc                          # Memory management — cleanup after training\nimport logging                     # Structured log output\nimport math                        # Mathematical operations (cosine schedule etc.)\nimport os                          # OS-level operations\nimport random                      # Randomness control (seeding)\nimport warnings                    # Suppress unnecessary warnings\nfrom collections import defaultdict  # Label → index mapping\nfrom dataclasses import dataclass, field  # Configuration class\nfrom pathlib import Path           # Platform-independent file paths\nfrom typing import Any, Dict, List, Optional, Tuple, Union  # Type hints\n\n# =============================================================================\n# Scientific computing & machine learning\n# =============================================================================\nimport numpy as np                 # Numerical computations\nimport pandas as pd                # Tabular data (CSV reading, DataFrame)\nfrom sklearn.decomposition import PCA         # Embedding visualization (3D)\nfrom sklearn.model_selection import train_test_split  # Train/val split\n\n# =============================================================================\n# Deep learning — PyTorch ecosystem\n# =============================================================================\nimport torch                       # Core deep learning framework\nimport torch.nn as nn              # Neural network layers\nimport torch.nn.functional as F    # Functional API (normalize, avg_pool etc.)\nimport timm                        # Pretrained Vision Transformer models\nfrom torch.optim import AdamW      # AdamW optimizer (corrected weight decay)\nfrom torch.optim.lr_scheduler import CosineAnnealingLR  # Cosine LR scheduler\nfrom torch.utils.data import DataLoader, Dataset         # Data loading infrastructure\nfrom torch.utils.data.sampler import Sampler             # Custom batch sampler\n\n# =============================================================================\n# Data augmentation\n# =============================================================================\nimport albumentations as A                   # Image augmentation library\nfrom albumentations.pytorch import ToTensorV2  # NumPy → PyTorch tensor conversion\n\n# =============================================================================\n# Visualization\n# =============================================================================\nimport matplotlib.pyplot as plt     # Training curve plots\nimport seaborn as sns               # Statistical plot theme\nimport plotly.express as px         # 3D interactive embedding plot\n\n# =============================================================================\n# Image processing\n# =============================================================================\nfrom PIL import Image               # Reading image files\nfrom tqdm.auto import tqdm          # Progress bar (training loop)\n\n# =============================================================================\n# Suppress warnings & configure logger\n# =============================================================================\nwarnings.filterwarnings(\"ignore\")   # Suppress unnecessary warnings\n\n# Logging configuration — includes timestamp, module name, and severity level\nlogging.basicConfig(\n    level=logging.INFO,\n    format='%(asctime)s — %(name)s — %(levelname)s — %(message)s'\n)\nlogger = logging.getLogger(__name__)\n\nlogger.info(\"All libraries loaded successfully.\")\nlogger.info(f\"PyTorch version: {torch.__version__}\")\nlogger.info(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    logger.info(f\"GPU: {torch.cuda.get_device_name(0)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:00.186097Z","iopub.execute_input":"2026-03-27T08:30:00.186337Z","iopub.status.idle":"2026-03-27T08:30:18.971084Z","shell.execute_reply.started":"2026-03-27T08:30:00.186313Z","shell.execute_reply":"2026-03-27T08:30:18.970523Z"}},"outputs":[],"execution_count":null},{"id":"932e1046","cell_type":"markdown","source":"## 2. Configuration\n\nAll hyperparameters and file paths are centralized in a single `dataclass`.  \nBenefits of this approach:\n- **Single source of truth:** Parameter changes happen in one place\n- **Type safety:** Wrong parameter types are caught at definition\n- **Readability:** IDE autocomplete support\n\n### Key Parameters\n| Parameter | Value | Description |\n|-----------|-------|-------------|\n| `img_size` | 448 | Optimal resolution for EVA-02 Large |\n| `patch_size` | 14 | ViT patch size (448/14 = 32×32 grid) |\n| `accum_steps` | 8 | Gradient accumulation → Effective batch = 32 |\n| `arcface_m` | 0.50 | ArcFace angular margin (inter-class separation) |\n","metadata":{}},{"id":"ec811efe","cell_type":"code","source":"@dataclass\nclass Config:\n    \"\"\"\n    Configuration class holding all settings for the training pipeline.\n    \n    This class centralizes model architecture, training hyperparameters,\n    file paths, and ArcFace parameters in a single structure.\n    \"\"\"\n    \n    # -------------------------------------------------------------------------\n    # Model Settings\n    # -------------------------------------------------------------------------\n    # EVA-02 Large: 304M parameters, 448px resolution, ImageNet-22K pretrained\n    model_name: str = \"eva02_large_patch14_448.mim_m38m_ft_in22k_in1k\"\n    seed: int = 42  # Fixed seed for reproducibility\n    \n    # -------------------------------------------------------------------------\n    # File Paths (Kaggle environment by default)\n    # -------------------------------------------------------------------------\n    root_dir: Path = field(\n        default_factory=lambda: Path(\"/kaggle/input/competitions/round-2-jaguar-reidentification-challenge\")\n    )\n    output_dir: Path = field(default_factory=lambda: Path(\".\"))\n    \n    # -------------------------------------------------------------------------\n    # EVA-02 Architecture Specifics\n    # -------------------------------------------------------------------------\n    img_size: int = 448    # Input image resolution (square)\n    patch_size: int = 14   # Pixel size of each patch\n    \n    # -------------------------------------------------------------------------\n    # Training Hyperparameters\n    # -------------------------------------------------------------------------\n    epochs: int = 20                # Total number of training epochs\n    batch_size: int = 4             # Mini-batch size (adjusted for GPU memory)\n    accum_steps: int = 8            # Gradient accumulation steps → effective batch = 4 × 8 = 32\n    lr: float = 1e-4                # Initial learning rate\n    weight_decay: float = 0.05      # AdamW weight decay coefficient\n    \n    # Number of workers: auto-adjusted based on CPU core count (max 4)\n    num_workers: int = field(default_factory=lambda: min(4, os.cpu_count() or 1))\n    \n    # -------------------------------------------------------------------------\n    # ArcFace Hyperparameters\n    # -------------------------------------------------------------------------\n    arcface_s: float = 30.0   # Scaling factor — controls logit magnitude\n    arcface_m: float = 0.50   # Angular margin — enforces minimum inter-class angular separation\n    \n    # -------------------------------------------------------------------------\n    # Dataset Info (populated at runtime)\n    # -------------------------------------------------------------------------\n    num_classes: int = 0  # Number of unique jaguar identities (computed from CSV)\n    \n    # -------------------------------------------------------------------------\n    # Computed Properties\n    # -------------------------------------------------------------------------\n    @property\n    def grid_size(self) -> int:\n        \"\"\"Spatial grid size: 448 / 14 = 32 (ViT output grid).\"\"\"\n        return self.img_size // self.patch_size\n    \n    @property\n    def train_dir(self) -> Path:\n        \"\"\"Directory containing training images.\"\"\"\n        if not self.root_dir.exists():\n            return Path(\"train\")\n        return self.root_dir / \"train\"\n    \n    @property\n    def csv_path(self) -> Path:\n        \"\"\"Path to the training labels CSV file.\"\"\"\n        if not self.root_dir.exists():\n            return Path(\"train.csv\")\n        return self.root_dir / \"train.csv\"\n\n    @property\n    def device(self) -> torch.device:\n        \"\"\"Compute device: CUDA if GPU available, otherwise CPU.\"\"\"\n        return torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\ndef seed_everything(seed: int) -> None:\n    \"\"\"\n    Fixes all sources of randomness to guarantee experiment reproducibility.\n    \n    Affected components:\n        - Python random module\n        - NumPy random number generator\n        - PyTorch CPU & GPU random seeds\n        - cuDNN deterministic mode\n    \n    Args:\n        seed: The seed value to set.\n    \"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    \n    # Deterministic mode: results are exactly reproducible (slight performance cost)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    \n    logger.info(f\"Random seed set to {seed}.\")\n\n\n# Create configuration and display summary\nconfig = Config()\nseed_everything(config.seed)\n\nlogger.info(f\"Model: {config.model_name}\")\nlogger.info(f\"Resolution: {config.img_size}x{config.img_size} | Grid: {config.grid_size}x{config.grid_size}\")\nlogger.info(f\"Training: {config.epochs} epochs | Effective batch: {config.batch_size * config.accum_steps}\")\nlogger.info(f\"Device: {config.device}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:18.972389Z","iopub.execute_input":"2026-03-27T08:30:18.972826Z","iopub.status.idle":"2026-03-27T08:30:18.996147Z","shell.execute_reply.started":"2026-03-27T08:30:18.972801Z","shell.execute_reply":"2026-03-27T08:30:18.995453Z"}},"outputs":[],"execution_count":null},{"id":"c541e9a5","cell_type":"markdown","source":"## 3. Visualization & Evaluation Metrics\n\nThis section contains three core functions:\n\n1. **`plot_training_curves`** — Plots training loss, validation mAP, and learning rate curves\n2. **`visualize_3d_embeddings`** — Creates a 3D interactive PCA visualization of the embedding space\n3. **`compute_mAP`** — Computes standard Re-ID metrics: Mean Average Precision and Rank-1 accuracy\n\n### Re-ID Metrics Explained\n- **mAP (Mean Average Precision):** Average precision of correct matches across all queries\n- **Rank-1:** Percentage of queries where the top-1 retrieved result is correct\n","metadata":{}},{"id":"1071d25c","cell_type":"code","source":"def plot_training_curves(\n    history: Dict[str, List[float]], \n    save_path: Union[str, Path] = \"training_log.png\"\n) -> None:\n    \"\"\"\n    Generates a visual summary of the training process.\n    \n    Left panel: Training loss (red) + Validation mAP (green, secondary axis)\n    Right panel: Learning rate schedule (logarithmic scale)\n    \n    Args:\n        history: Dictionary containing 'loss', 'map', and 'lr' arrays.\n        save_path: File path to save the plot.\n    \"\"\"\n    sns.set_theme(style=\"whitegrid\")\n    fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\n    epochs_range = range(1, len(history['loss']) + 1)\n    \n    # ----- Left Panel: Loss & mAP -----\n    sns.lineplot(\n        x=epochs_range, y=history['loss'], \n        ax=axes[0], color=\"#FF5733\", marker=\"o\", label=\"Train Loss\"\n    )\n    axes[0].set_title(\"Training Loss & Validation mAP\", fontsize=13, fontweight=\"bold\")\n    axes[0].set_xlabel(\"Epoch\")\n    axes[0].set_ylabel(\"Loss\")\n\n    # Overlay validation mAP on secondary axis\n    if history.get('map'):\n        ax2 = axes[0].twinx()\n        sns.lineplot(\n            x=epochs_range, y=history['map'], \n            ax=ax2, color=\"green\", marker=\"x\", linestyle=\"--\", label=\"Val mAP\"\n        )\n        ax2.set_ylabel(\"mAP\", color=\"green\")\n        ax2.legend(loc=\"upper right\")\n\n    axes[0].legend(loc=\"upper left\")\n\n    # ----- Right Panel: Learning Rate -----\n    sns.lineplot(\n        x=epochs_range, y=history['lr'], \n        ax=axes[1], color=\"#337AFF\", linestyle=\"--\"\n    )\n    axes[1].set_title(\"Learning Rate Schedule\", fontsize=13, fontweight=\"bold\")\n    axes[1].set_xlabel(\"Epoch\")\n    axes[1].set_yscale(\"log\")  # Log scale — to visualize large LR differences\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    plt.close()\n    logger.info(f\"Training curves saved to {save_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:18.996963Z","iopub.execute_input":"2026-03-27T08:30:18.997218Z","iopub.status.idle":"2026-03-27T08:30:19.169209Z","shell.execute_reply.started":"2026-03-27T08:30:18.997190Z","shell.execute_reply":"2026-03-27T08:30:19.168343Z"}},"outputs":[],"execution_count":null},{"id":"1a3f5f41","cell_type":"code","source":"def visualize_3d_embeddings(\n    model: nn.Module, \n    loader: DataLoader, \n    device: torch.device, \n    output_dir: Path, \n    max_points: int = 1000\n) -> None:\n    \"\"\"\n    Creates a 3D PCA visualization of the model's embedding space.\n    \n    This visualization shows how well the model separates different jaguars.\n    Ideally, points belonging to the same identity should cluster together\n    while different identities remain well-separated.\n    \n    Args:\n        model: Trained Re-ID model.\n        loader: Evaluation data loader.\n        device: Compute device (CPU/GPU).\n        output_dir: Directory to save the HTML visualization.\n        max_points: Maximum number of points to plot (for performance).\n    \"\"\"\n    logger.info(\"Preparing 3D Embedding Visualization (PCA)...\")\n    model.eval()\n    embeddings, labels = [], []\n\n    # Collect embeddings in inference mode (no gradient computation)\n    with torch.no_grad():\n        for imgs, lbls in loader:\n            imgs = imgs.to(device)\n            feats = model(imgs, labels=None)  # labels=None → returns embedding only\n            embeddings.append(feats.cpu().numpy())\n            labels.extend(lbls.numpy())\n            if len(labels) >= max_points:\n                break\n\n    if not embeddings:\n        logger.warning(\"No embeddings generated for visualization.\")\n        return\n\n    # Concatenate all embeddings and reduce to 3D with PCA\n    X = np.vstack(embeddings)[:max_points]\n    y = np.array(labels)[:max_points]\n\n    pca = PCA(n_components=3)\n    X_3d = pca.fit_transform(X)\n\n    # Create interactive 3D scatter plot with Plotly\n    df_plot = pd.DataFrame(X_3d, columns=['x', 'y', 'z'])\n    df_plot['Jaguar_ID'] = [f\"ID_{i}\" for i in y]\n\n    fig = px.scatter_3d(\n        df_plot, x='x', y='y', z='z', color='Jaguar_ID',\n        title='EVA-02 Embedding Space (3D PCA Projection)',\n        size_max=5, opacity=0.7\n    )\n    fig.update_traces(marker=dict(size=4))\n    \n    save_path = output_dir / \"3d_embeddings_eva02.html\"\n    fig.write_html(str(save_path))\n    logger.info(f\"3D embedding visualization saved to {save_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.170916Z","iopub.execute_input":"2026-03-27T08:30:19.171228Z","iopub.status.idle":"2026-03-27T08:30:19.187968Z","shell.execute_reply.started":"2026-03-27T08:30:19.171199Z","shell.execute_reply":"2026-03-27T08:30:19.187281Z"}},"outputs":[],"execution_count":null},{"id":"c04ca753","cell_type":"code","source":"def compute_mAP(\n    features: np.ndarray, \n    labels: np.ndarray\n) -> Tuple[float, float]:\n    \"\"\"\n    Computes Mean Average Precision (mAP) and Rank-1 accuracy.\n    \n    Algorithm:\n    1. Build a cosine similarity matrix across all embeddings\n    2. For each query, find matches sorted by similarity\n    3. Compute Average Precision (AP) for each query\n    4. Average all AP values → mAP\n    \n    Args:\n        features: Normalized embedding vectors [N, D].\n        labels: Identity label for each sample [N].\n        \n    Returns:\n        (mAP, Rank-1): Mean average precision and top-1 accuracy.\n    \"\"\"\n    # Cosine similarity matrix (dot product — features are normalized)\n    sim_mat = np.dot(features, features.T)\n    \n    # Prevent self-matching (diagonal → -1)\n    np.fill_diagonal(sim_mat, -1)\n\n    # Sort each row from most similar to least similar\n    indices = np.argsort(-sim_mat, axis=1)\n    \n    # Match matrix: does the retrieved sample share the same identity?\n    matches = (labels[indices] == labels[:, None])\n\n    # Rank-1: Is the top-1 retrieved result correct?\n    rank1 = np.mean(matches[:, 0])\n\n    # Compute Average Precision for each query\n    aps = []\n    for i in range(len(matches)):\n        match = matches[i]\n        if match.sum() == 0:\n            aps.append(0.0)\n        else:\n            cumsum = np.cumsum(match)\n            precision = cumsum / np.arange(1, len(match) + 1)\n            aps.append((precision * match).sum() / match.sum())\n\n    return float(np.mean(aps)), float(rank1)\n\n\ndef validate(\n    model: nn.Module, \n    loader: DataLoader, \n    device: torch.device\n) -> Tuple[float, float]:\n    \"\"\"\n    Evaluates the model on the validation set.\n    \n    Extracts embeddings from all validation samples and computes mAP/Rank-1.\n    \n    Args:\n        model: Model to evaluate.\n        loader: Validation data loader.\n        device: Compute device.\n        \n    Returns:\n        (mAP, Rank-1): Validation metrics.\n    \"\"\"\n    model.eval()\n    feats_all, lbls_all = [], []\n    \n    with torch.no_grad():\n        for imgs, lbls in tqdm(loader, desc=\"Validating\", leave=False):\n            imgs = imgs.to(device)\n            f = model(imgs, labels=None)  # Embedding only\n            feats_all.append(f.cpu().numpy())\n            lbls_all.append(lbls.numpy())\n\n    if not feats_all:\n        return 0.0, 0.0\n\n    feats_all = np.concatenate(feats_all)\n    lbls_all = np.concatenate(lbls_all)\n    return compute_mAP(feats_all, lbls_all)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.189103Z","iopub.execute_input":"2026-03-27T08:30:19.190040Z","iopub.status.idle":"2026-03-27T08:30:19.204082Z","shell.execute_reply.started":"2026-03-27T08:30:19.190007Z","shell.execute_reply":"2026-03-27T08:30:19.203449Z"}},"outputs":[],"execution_count":null},{"id":"7bd91595","cell_type":"markdown","source":"## 4. Loss Functions & Custom Layers\n\nThis section contains the core building blocks of the Re-ID task:\n\n### 🔹 Batch Hard Triplet Loss\nSelects the **hardest positive** (same identity, farthest) and **hardest negative** \n(different identity, closest) pairs within each mini-batch to strengthen the model on edge cases.\n\n### 🔹 GeM (Generalized Mean Pooling)\nProvides more discriminative feature pooling than standard average pooling through a learnable \n`p` parameter. When `p > 1`, it assigns higher weight to strong activations.\n\n### 🔹 ArcFace Head\nAdds an angular margin to the classification layer, enforcing a minimum angular distance between \nclasses in the embedding space. This is the standard approach in Re-ID and face recognition tasks.\n","metadata":{}},{"id":"c4193ed1","cell_type":"code","source":"class BatchHardTripletLoss(nn.Module):\n    \"\"\"\n    Batch Hard Triplet Loss — metric learning loss function.\n    \n    For each sample:\n    - Hardest positive: Farthest sample with the same identity\n    - Hardest negative: Closest sample with a different identity\n    \n    Loss = max(0, margin + d(anchor, hard_pos) - d(anchor, hard_neg))\n    \n    Args:\n        margin: Minimum required gap between positive and negative distances.\n    \"\"\"\n    def __init__(self, margin: float = 0.3):\n        super().__init__()\n        self.margin = margin\n        self.ranking_loss = nn.MarginRankingLoss(margin=margin)\n\n    def forward(self, embeddings: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:\n        # Pairwise Euclidean distance matrix [B, B]\n        dist = torch.cdist(embeddings, embeddings, p=2)\n        \n        # Mask for pairs sharing the same identity\n        mask = targets.unsqueeze(1).eq(targets.unsqueeze(0))\n        \n        # Hardest positive: FARTHEST sample with the same identity\n        dist_ap = dist * mask.float()\n        hardest_pos, _ = torch.max(dist_ap, dim=1)\n        \n        # Hardest negative: CLOSEST sample with a different identity\n        # (Mask same-identity pairs with large penalty)\n        dist_an = dist + mask.float() * 1e6\n        hardest_neg, _ = torch.min(dist_an, dim=1)\n        \n        # Margin-based ranking loss\n        y = torch.ones_like(hardest_neg)\n        return self.ranking_loss(hardest_neg, hardest_pos, y)\n\n\nclass GeM(nn.Module):\n    \"\"\"\n    Generalized Mean Pooling (GeM) — learnable pooling layer.\n    \n    Formula: GeM(x) = (1/N * Σ x_i^p)^(1/p)\n    \n    p = 1 → Average Pooling\n    p → ∞ → Max Pooling\n    p is initialized as a learnable parameter.\n    \n    Args:\n        p: Initial p value (default: 3.0).\n        eps: Small value for numerical stability.\n    \"\"\"\n    def __init__(self, p: float = 3.0, eps: float = 1e-6):\n        super().__init__()\n        self.p = nn.Parameter(torch.ones(1) * p)  # Learnable parameter\n        self.eps = eps\n        \n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        # Clamp negatives → raise to power → average pool → take root\n        return F.avg_pool2d(\n            x.clamp(min=self.eps).pow(self.p), \n            (x.size(-2), x.size(-1))\n        ).pow(1. / self.p)\n\n\nclass ArcFaceHead(nn.Module):\n    \"\"\"\n    ArcFace (Additive Angular Margin) classification head.\n    \n    Instead of standard softmax, adds an angular margin (m) to the cosine\n    similarity of the correct class, increasing inter-class separation.\n    \n    Formula: L = -log(exp(s * cos(θ_y + m)) / (exp(s * cos(θ_y + m)) + Σ exp(s * cos(θ_j))))\n    \n    Args:\n        in_features: Input embedding dimension.\n        out_features: Number of classes (unique jaguar identities).\n        s: Scaling factor — controls logit magnitude.\n        m: Angular margin — minimum inter-class angle (radians).\n    \"\"\"\n    def __init__(self, in_features: int, out_features: int, s: float = 30.0, m: float = 0.5):\n        super().__init__()\n        self.s = s\n        self.m = m\n        # Weight matrix: each column represents a class center\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)  # Xavier initialization\n        \n    def forward(\n        self, \n        embeddings: torch.Tensor, \n        label: Optional[torch.Tensor] = None\n    ) -> torch.Tensor:\n        # Normalize both embeddings and weights → cosine similarity\n        cosine = F.linear(F.normalize(embeddings), F.normalize(self.weight))\n        \n        # Inference mode (no label): return scaled cosine\n        if label is None:\n            return cosine * self.s\n        \n        # Training mode: apply angular margin to the correct class\n        phi = cosine - self.m  # cos(θ) - m ≈ cos(θ + m) (for small angles)\n        \n        # One-hot vector: mark the correct class position\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, label.view(-1, 1), 1.0)\n        \n        # Correct class: cos(θ + m), other classes: cos(θ)\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        return output * self.s  # Scale\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.205275Z","iopub.execute_input":"2026-03-27T08:30:19.205524Z","iopub.status.idle":"2026-03-27T08:30:19.220265Z","shell.execute_reply.started":"2026-03-27T08:30:19.205495Z","shell.execute_reply":"2026-03-27T08:30:19.219602Z"}},"outputs":[],"execution_count":null},{"id":"677b7bc7","cell_type":"markdown","source":"## 5. EVA-02 Re-ID Model Architecture\n\nThe model follows a pipeline composed of the following components:\n\n```\nImage [B, 3, 448, 448]\n    ↓\nEVA-02 Backbone (ViT-Large)\n    ↓\nSpatial Features [B, 1024, 32, 32]   ← CLS token removed, reshaped to 2D grid\n    ↓\nGeM Pooling [B, 1024]                ← Learnable global pooling\n    ↓\nBatchNorm [B, 1024]                  ← Feature normalization\n    ↓\nArcFace Head [B, num_classes]        ← Classification logits (training)\n    or\nEmbedding [B, 1024]                  ← Normalized feature vector (inference)\n```\n","metadata":{}},{"id":"5c1c6261","cell_type":"code","source":"class EVAReIDModel(nn.Module):\n    \"\"\"\n    EVA-02 based Jaguar Re-Identification model.\n    \n    Architecture:\n        1. EVA-02 Large backbone: Feature extraction with pretrained ViT\n        2. GeM Pooling: Compress spatial features into a global vector\n        3. BatchNorm: Embedding normalization\n        4. ArcFace Head: Angular margin-based classification\n    \n    Training mode: returns (ArcFace logits, raw embedding) tuple\n    Inference mode: returns normalized embedding vector\n    \n    Args:\n        config: Model and training configuration.\n    \"\"\"\n    def __init__(self, config: Config):\n        super().__init__()\n        self.config = config\n        logger.info(f\"Initializing model: EVA-02 Large | Resolution: {config.img_size}x{config.img_size}\")\n\n        # ----- Backbone: EVA-02 Large (timm) -----\n        # num_classes=0 → remove classification layer\n        # global_pool=\"\" → we'll use our own pooling layer\n        self.backbone = timm.create_model(\n            config.model_name, \n            pretrained=True, \n            num_classes=0, \n            global_pool=\"\", \n            img_size=config.img_size\n        )\n        \n        # Gradient checkpointing: reduces VRAM usage (trades off speed)\n        if hasattr(self.backbone, \"set_grad_checkpointing\"):\n            self.backbone.set_grad_checkpointing(True)\n            logger.info(\"Gradient checkpointing enabled (memory optimization)\")\n\n        # Backbone output dimension (EVA-02 Large → 1024)\n        self.in_features = self.backbone.num_features\n        \n        # ----- Pooling and Normalization -----\n        self.pool = GeM(p=3)\n        \n        # BatchNorm: stabilizes embedding distribution\n        self.bn = nn.BatchNorm1d(self.in_features)\n        nn.init.constant_(self.bn.weight, 1.0)\n        nn.init.constant_(self.bn.bias, 0.0)\n        self.bn.bias.requires_grad_(False)  # Bias kept fixed (ArcFace best practice)\n        \n        # ----- Classification Head: ArcFace -----\n        self.arcface = ArcFaceHead(\n            self.in_features, \n            config.num_classes, \n            s=config.arcface_s, \n            m=config.arcface_m\n        )\n        \n        logger.info(f\"Model ready | Backbone: {self.in_features}D | Classes: {config.num_classes}\")\n\n    def forward_features(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Extracts spatial features from the backbone.\n        \n        ViT output shape: [B, seq_len, C].\n        The CLS token (first element) is removed and remaining patch tokens\n        are reshaped into 2D grid format: [B, C, H, W]\n        \n        Args:\n            x: Input image tensor [B, 3, 448, 448].\n            \n        Returns:\n            Spatial feature map [B, C, grid_size, grid_size].\n        \"\"\"\n        features = self.backbone.forward_features(x)\n        B, N, C = features.shape\n        \n        # Expected sequence length: grid² + 1 (including CLS token)\n        expected_seq_len = (self.config.grid_size * self.config.grid_size) + 1\n        \n        # Remove CLS token (first element)\n        if N == expected_seq_len:\n            features = features[:, 1:, :]\n            \n        # [B, N, C] → [B, C, H, W]\n        features = features.permute(0, 2, 1).reshape(\n            B, C, self.config.grid_size, self.config.grid_size\n        )\n        return features\n\n    def forward(\n        self, \n        x: torch.Tensor, \n        labels: Optional[torch.Tensor] = None\n    ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:\n        \"\"\"\n        Forward pass.\n        \n        Args:\n            x: Input images [B, 3, H, W].\n            labels: Training labels (None for inference mode).\n            \n        Returns:\n            Training: (ArcFace logits, raw embedding) tuple\n            Inference: Normalized embedding vector\n        \"\"\"\n        # Step 1: Backbone → spatial features\n        spatial_features = self.forward_features(x)\n        \n        # Step 2: GeM pooling → global feature vector\n        pooled_features = self.pool(spatial_features).flatten(start_dim=1)\n        \n        # Step 3: BatchNorm → embedding normalization\n        normalized_features = self.bn(pooled_features)\n        \n        # Training mode: return logits + embedding (for both CE and Triplet loss)\n        if labels is not None:\n            return self.arcface(normalized_features, labels), pooled_features\n        \n        # Inference mode: return normalized embedding only\n        return normalized_features\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.221170Z","iopub.execute_input":"2026-03-27T08:30:19.221451Z","iopub.status.idle":"2026-03-27T08:30:19.236311Z","shell.execute_reply.started":"2026-03-27T08:30:19.221428Z","shell.execute_reply":"2026-03-27T08:30:19.235698Z"}},"outputs":[],"execution_count":null},{"id":"0b12cba5","cell_type":"markdown","source":"## 6. Dataset, Augmentation & PK Sampler\n\n### Augmentation Strategy\n- **Training:** HorizontalFlip, RandomBrightnessContrast, Affine transforms, CoarseDropout\n- **Validation:** Resize + Normalize only (clean evaluation)\n\n### PK Sampler\nFor Triplet Loss to work effectively, each batch must contain at least 2 different identities \nwith at least 2 samples per identity. `PKSampler` is a custom batch sampler that guarantees \nthis constraint.\n","metadata":{}},{"id":"99f848af","cell_type":"code","source":"class JaguarDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset for loading Jaguar images.\n    \n    Responsibilities:\n    - Read image files from disk\n    - Apply albumentations transforms\n    - Handle corrupted files gracefully (fallback to black image)\n    \n    Args:\n        df: DataFrame containing file paths and labels.\n        root_dir: Root path for image directory.\n        img_size: Target image size (square).\n        transform: Albumentations pipeline to apply.\n    \"\"\"\n    def __init__(\n        self, \n        df: pd.DataFrame, \n        root_dir: Path, \n        img_size: int, \n        transform: Optional[A.Compose] = None\n    ):\n        self.df = df\n        self.root_dir = root_dir\n        self.img_size = img_size\n        self.transform = transform\n        \n        # Convert to lists for performance (DataFrame indexing is slow)\n        self.paths = df[\"filepath\"].tolist()\n        self.labels = df[\"label\"].tolist()\n        \n    def __len__(self) -> int:\n        return len(self.df)\n    \n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:\n        path = self.paths[idx]\n        \n        # Read image (produce black image on failure)\n        try:\n            image = np.array(Image.open(path).convert(\"RGB\"))\n        except Exception as e:\n            logger.warning(f\"Failed to load image: {path} — Error: {e}\")\n            image = np.zeros((self.img_size, self.img_size, 3), dtype=np.uint8)\n        \n        # Apply augmentation (zero tensor on failure)\n        if self.transform:\n            try:\n                image = self.transform(image=image)[\"image\"]\n            except Exception as e:\n                logger.error(f\"Transform error ({path}): {e}\")\n                image = torch.zeros((3, self.img_size, self.img_size))\n                \n        return image, torch.tensor(self.labels[idx], dtype=torch.long)\n\n\ndef get_transforms(img_size: int, mode: str = \"train\") -> A.Compose:\n    \"\"\"\n    Builds augmentation pipelines for training and validation.\n    \n    EVA-02 normalization values (ImageNet statistics):\n        mean = [0.481, 0.457, 0.408]\n        std  = [0.268, 0.261, 0.275]\n    \n    Args:\n        img_size: Target resolution.\n        mode: \"train\" → augmentation active, other → resize + normalize only.\n        \n    Returns:\n        Albumentations Compose pipeline.\n    \"\"\"\n    # ImageNet normalization values\n    mean = [0.481, 0.457, 0.408]\n    std = [0.268, 0.261, 0.275]\n\n    if mode == \"train\":\n        return A.Compose([\n            A.Resize(img_size, img_size),                     # Resize\n            A.HorizontalFlip(p=0.5),                          # Horizontal flip (50%)\n            A.RandomBrightnessContrast(p=0.2),                # Brightness/contrast\n            A.Affine(                                         # Geometric transforms\n                scale=(0.9, 1.1),                             #   Scale: 90-110%\n                translate_percent=(0.1, 0.1),                 #   Translation: 10%\n                rotate=(-15, 15),                             #   Rotation: ±15°\n                p=0.3\n            ),\n            A.CoarseDropout(                                  # Random region erasing\n                max_holes=8,                                  #   Max hole count\n                max_height=img_size // 10,                    #   Hole size\n                max_width=img_size // 10, \n                p=0.3\n            ),\n            A.Normalize(mean, std),                           # Normalize\n            ToTensorV2(),                                     # NumPy → Tensor\n        ])\n    \n    # Validation: resize and normalize only\n    return A.Compose([\n        A.Resize(img_size, img_size),\n        A.Normalize(mean, std),\n        ToTensorV2()\n    ])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.237072Z","iopub.execute_input":"2026-03-27T08:30:19.237345Z","iopub.status.idle":"2026-03-27T08:30:19.252418Z","shell.execute_reply.started":"2026-03-27T08:30:19.237310Z","shell.execute_reply":"2026-03-27T08:30:19.251701Z"}},"outputs":[],"execution_count":null},{"id":"5b21c9ad","cell_type":"code","source":"class PKSampler(Sampler):\n    \"\"\"\n    PK Batch Sampler — custom batching strategy for Triplet Loss.\n    \n    Each batch contains:\n    - P different identities (classes)\n    - K samples per identity\n    - Total batch size = P × K\n    \n    This structure guarantees that Batch Hard Triplet Loss can find\n    meaningful positive and negative pairs within each batch.\n    \n    Args:\n        dataset: JaguarDataset to sample from.\n        p_classes: Number of unique identities per batch.\n        k_instances: Number of samples per identity.\n    \"\"\"\n    def __init__(self, dataset: JaguarDataset, p_classes: int = 2, k_instances: int = 2):\n        super().__init__()\n        self.p_classes = p_classes\n        self.k_instances = k_instances\n        self.batch_size = self.p_classes * self.k_instances\n        \n        # Build label → index list mapping\n        self.label_to_indices = defaultdict(list)\n        for idx, label in enumerate(dataset.labels):\n            self.label_to_indices[label].append(idx)\n            \n        self.unique_labels = list(self.label_to_indices.keys())\n        self.num_batches = len(dataset) // self.batch_size\n\n    def __iter__(self):\n        final_indices = []\n        \n        for _ in range(self.num_batches):\n            # Randomly select P identities\n            sampled_classes = random.sample(\n                self.unique_labels, \n                min(self.p_classes, len(self.unique_labels))\n            )\n            batch_indices = []\n            \n            for cls in sampled_classes:\n                cls_indices = self.label_to_indices[cls]\n                \n                # Sample if enough, otherwise sample with replacement\n                if len(cls_indices) >= self.k_instances:\n                    sampled_indices = random.sample(cls_indices, self.k_instances)\n                else:\n                    sampled_indices = random.choices(cls_indices, k=self.k_instances)\n                    \n                batch_indices.extend(sampled_indices)\n                \n            final_indices.extend(batch_indices)\n            \n        return iter(final_indices)\n\n    def __len__(self) -> int:\n        return self.num_batches * self.batch_size\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.253277Z","iopub.execute_input":"2026-03-27T08:30:19.253521Z","iopub.status.idle":"2026-03-27T08:30:19.268310Z","shell.execute_reply.started":"2026-03-27T08:30:19.253502Z","shell.execute_reply":"2026-03-27T08:30:19.267793Z"}},"outputs":[],"execution_count":null},{"id":"a9a3d2cd","cell_type":"markdown","source":"## 7. Training Engine (Trainer)\n\nThe class responsible for the entire training loop. Key features:\n\n- **Mixed Precision (FP16):** Halves memory usage with `torch.amp.autocast`\n- **Gradient Accumulation:** Simulates large effective batches by accumulating small ones\n- **Gradient Clipping:** Prevents exploding gradients\n- **Dual Loss:** CrossEntropy (ArcFace) + Triplet Loss optimized jointly\n\n### Training Flow\n```\nFor each batch:\n  1. Forward pass (FP16) → logits + embedding\n  2. Compute CE Loss + Triplet Loss\n  3. Loss / accum_steps (gradient accumulation normalization)\n  4. Backward pass (gradient computation)\n  5. Every accum_steps: gradient clip → optimizer step → zero grad\n```\n","metadata":{}},{"id":"b5684b46","cell_type":"code","source":"class Trainer:\n    \"\"\"\n    Training loop manager — supports mixed precision and gradient accumulation.\n    \n    This class encapsulates all training process details, keeping\n    the main pipeline clean.\n    \n    Args:\n        model: EVAReIDModel to train.\n        train_loader: Training data loader.\n        optimizer: AdamW optimizer.\n        scheduler: Cosine annealing LR scheduler.\n        config: Training configuration.\n    \"\"\"\n    def __init__(\n        self,\n        model: nn.Module,\n        train_loader: DataLoader,\n        optimizer: torch.optim.Optimizer,\n        scheduler: torch.optim.lr_scheduler.LRScheduler,\n        config: Config\n    ):\n        self.model = model\n        self.train_loader = train_loader\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.config = config\n        self.device = config.device\n        \n        # Loss functions\n        self.criterion_ce = nn.CrossEntropyLoss()       # For ArcFace logits\n        self.criterion_triplet = BatchHardTripletLoss(margin=0.3)  # For embeddings\n        \n        # Mixed precision scaler (CUDA only)\n        self.scaler = torch.amp.GradScaler('cuda') if torch.cuda.is_available() else None\n        \n        if self.scaler:\n            logger.info(\"Mixed Precision (FP16) training enabled\")\n\n    def train_epoch(self, epoch: int) -> float:\n        \"\"\"\n        Runs a single training epoch.\n        \n        Args:\n            epoch: Current epoch number.\n            \n        Returns:\n            Average training loss for the epoch.\n        \"\"\"\n        self.model.train()\n        total_loss = 0.0\n        pbar = tqdm(self.train_loader, desc=f\"Epoch {epoch}/{self.config.epochs}\", leave=False)\n\n        for step, (images, labels) in enumerate(pbar, 1):\n            # Move data to GPU (non_blocking → async transfer)\n            images = images.to(self.device, non_blocking=True)\n            labels = labels.to(self.device, non_blocking=True)\n\n            if self.scaler is not None:\n                # ===== Mixed Precision (FP16) Training Path =====\n                with torch.amp.autocast('cuda'):\n                    logits, unnormalized_features = self.model(images, labels)\n                    \n                    # Dual loss: ArcFace CE + Triplet\n                    loss_ce = self.criterion_ce(logits, labels)\n                    loss_triplet = self.criterion_triplet(unnormalized_features, labels)\n                    \n                    # Gradient accumulation: divide loss by number of steps\n                    loss = (loss_ce + loss_triplet) / self.config.accum_steps\n\n                # Scaled backpropagation\n                self.scaler.scale(loss).backward()\n\n                # Update optimizer every accum_steps\n                if step % self.config.accum_steps == 0 or step == len(self.train_loader):\n                    self.scaler.unscale_(self.optimizer)\n                    torch.nn.utils.clip_grad_norm_(self.model.parameters(), 2.0)  # Grad clip\n                    self.scaler.step(self.optimizer)\n                    self.scaler.update()\n                    self.optimizer.zero_grad()\n            else:\n                # ===== Standard (FP32) Training Path =====\n                logits, unnormalized_features = self.model(images, labels)\n                loss_ce = self.criterion_ce(logits, labels)\n                loss_triplet = self.criterion_triplet(unnormalized_features, labels)\n                loss = (loss_ce + loss_triplet) / self.config.accum_steps\n                loss.backward()\n\n                if step % self.config.accum_steps == 0 or step == len(self.train_loader):\n                    torch.nn.utils.clip_grad_norm_(self.model.parameters(), 2.0)\n                    self.optimizer.step()\n                    self.optimizer.zero_grad()\n\n            # Update progress bar\n            total_loss += loss.item() * self.config.accum_steps\n            pbar.set_postfix({\n                \"Loss\": f\"{total_loss / step:.4f}\", \n                \"LR\": f\"{self.optimizer.param_groups[0]['lr']:.2e}\"\n            })\n\n        return total_loss / len(self.train_loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.270432Z","iopub.execute_input":"2026-03-27T08:30:19.270661Z","iopub.status.idle":"2026-03-27T08:30:19.284307Z","shell.execute_reply.started":"2026-03-27T08:30:19.270640Z","shell.execute_reply":"2026-03-27T08:30:19.283641Z"}},"outputs":[],"execution_count":null},{"id":"e1fbfe45","cell_type":"markdown","source":"## 8. Main Training Pipeline\n\nEnd-to-end training flow that brings all components together:\n\n1. **Data Preparation:** Read CSV → encode labels → train/val split\n2. **Model Creation:** EVA-02 backbone + ArcFace head\n3. **Gradual Unfreezing:**\n   - Epochs 1–3: Backbone frozen, only head trains (warm-up)\n   - Epoch 4+: Backbone unfrozen with lower LR (differential learning rate)\n4. **Training Loop:** Each epoch → train → validate → save\n5. **Finalization:** Generate visualizations → clean up memory\n","metadata":{}},{"id":"58297444","cell_type":"code","source":"def main() -> None:\n    \"\"\"\n    Main training pipeline — brings all components together for end-to-end execution.\n    \"\"\"\n    config = Config()\n    seed_everything(config.seed)\n    \n    logger.info(\"=\" * 60)\n    logger.info(\"   EVA-02 Jaguar Re-ID Training Pipeline Starting\")\n    logger.info(\"=\" * 60)\n\n    # =========================================================================\n    # Step 1: Data Reading & Label Encoding\n    # =========================================================================\n    if not config.csv_path.exists():\n        logger.error(f\"CSV not found: {config.csv_path}\")\n        return\n\n    df = pd.read_csv(config.csv_path)\n    \n    # Create filepath column if missing\n    if \"filepath\" not in df.columns:\n        df[\"filepath\"] = df[\"filename\"].apply(lambda x: str(config.train_dir / x))\n    \n    # Determine target column (support two naming conventions)\n    target_col = \"ground_truth\" if \"ground_truth\" in df.columns else \"label\"\n    if target_col not in df.columns:\n        logger.error(f\"Label column not found: '{target_col}'\")\n        return\n    \n    # Convert labels to sorted numerical indices\n    unique_labels = sorted(df[target_col].unique())\n    label_map = {label: idx for idx, label in enumerate(unique_labels)}\n    df[\"label\"] = df[target_col].map(label_map)\n    config.num_classes = len(unique_labels)\n    \n    # Stratified train/val split (90% train, 10% validation)\n    train_df, val_df = train_test_split(\n        df, \n        test_size=0.10, \n        random_state=config.seed, \n        stratify=df['label']\n    )\n    logger.info(f\"Train: {len(train_df)} | Val: {len(val_df)} | Total classes: {config.num_classes}\")\n\n    # =========================================================================\n    # Step 2: Dataset & DataLoader Creation\n    # =========================================================================\n    train_ds = JaguarDataset(\n        train_df, config.train_dir, config.img_size, \n        transform=get_transforms(config.img_size, \"train\")\n    )\n    val_ds = JaguarDataset(\n        val_df, config.train_dir, config.img_size, \n        transform=get_transforms(config.img_size, \"valid\")\n    )\n    \n    # PK Sampler: P=2 identities × K=2 samples per batch\n    train_sampler = PKSampler(train_ds, p_classes=2, k_instances=2)\n    \n    train_loader = DataLoader(\n        train_ds,\n        batch_size=config.batch_size,\n        sampler=train_sampler,      # PK sampler for batch construction\n        num_workers=config.num_workers,\n        drop_last=True              # Drop incomplete final batch\n    )\n    val_loader = DataLoader(\n        val_ds, \n        batch_size=config.batch_size, \n        shuffle=False,              # No shuffling for validation\n        num_workers=config.num_workers\n    )\n\n    # =========================================================================\n    # Step 3: Model Creation & Gradual Unfreezing Setup\n    # =========================================================================\n    model = EVAReIDModel(config).to(config.device)\n    \n    # Freeze backbone initially — only head trains\n    logger.info(\"Backbone frozen → Only classification head will train.\")\n    for p in model.backbone.parameters():\n        p.requires_grad = False\n    \n    # Only pass trainable parameters to optimizer\n    optimizer = AdamW(\n        filter(lambda p: p.requires_grad, model.parameters()), \n        lr=config.lr, \n        weight_decay=config.weight_decay\n    )\n    scheduler = CosineAnnealingLR(optimizer, T_max=config.epochs, eta_min=1e-6)\n    \n    # Training engine\n    trainer = Trainer(model, train_loader, optimizer, scheduler, config)\n\n    # =========================================================================\n    # Step 4: Training Loop\n    # =========================================================================\n    best_map = 0.0\n    history: Dict[str, List[float]] = {'loss': [], 'lr': [], 'map': []}\n    \n    for epoch in range(1, config.epochs + 1):\n        \n        # --- Unfreeze backbone at epoch 4 (with differential LR) ---\n        if epoch == 4:\n            logger.info(\"UNFREEZE: Backbone included in training (low LR: base_lr × 0.1)\")\n            for p in model.backbone.parameters():\n                p.requires_grad = True\n            # Add backbone as separate LR group (1/10th of base LR)\n            optimizer.add_param_group({\n                'params': model.backbone.parameters(), \n                'lr': config.lr * 0.1\n            })\n        \n        # Train + validate\n        loss = trainer.train_epoch(epoch)\n        scheduler.step()\n        val_map, val_rank1 = validate(model, val_loader, config.device)\n        \n        # Record history\n        history['loss'].append(loss)\n        history['lr'].append(optimizer.param_groups[0]['lr'])\n        history['map'].append(val_map)\n        \n        logger.info(\n            f\"Epoch {epoch}/{config.epochs} | \"\n            f\"Loss: {loss:.4f} | \"\n            f\"mAP: {val_map:.4f} | \"\n            f\"Rank-1: {val_rank1:.4f}\"\n        )\n        \n        # Save best model (only after warm-up phase)\n        if epoch > 3 and val_map > best_map:\n            best_map = val_map\n            best_model_path = config.output_dir / \"model_eva02_best.pth\"\n            torch.save(model.state_dict(), best_model_path)\n            logger.info(f\"NEW BEST MODEL SAVED → mAP: {best_map:.4f}\")\n        \n        # Save checkpoints for last 3 epochs (for ensemble)\n        if epoch >= config.epochs - 3:\n            epoch_model_path = config.output_dir / f\"model_eva02_epoch_{epoch}.pth\"\n            torch.save(model.state_dict(), epoch_model_path)\n\n    # =========================================================================\n    # Step 5: Post-Training — Visualization & Cleanup\n    # =========================================================================\n    logger.info(\"Training complete. Generating visualizations...\")\n    plot_training_curves(history, save_path=config.output_dir / \"training_log_eva02.png\")\n    \n    # Load best model and generate embedding visualization\n    best_model_path = config.output_dir / \"model_eva02_best.pth\"\n    if best_model_path.exists():\n        model.load_state_dict(\n            torch.load(best_model_path, map_location=config.device, weights_only=True)\n        )\n    visualize_3d_embeddings(model, val_loader, config.device, config.output_dir, max_points=1000)\n    \n    # Memory cleanup\n    logger.info(\"Cleaning up memory...\")\n    del model, optimizer, scheduler, trainer, train_loader, val_loader, train_ds, val_ds\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    logger.info(\"PROCESS COMPLETED SUCCESSFULLY!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.285485Z","iopub.execute_input":"2026-03-27T08:30:19.286019Z","iopub.status.idle":"2026-03-27T08:30:19.303378Z","shell.execute_reply.started":"2026-03-27T08:30:19.285996Z","shell.execute_reply":"2026-03-27T08:30:19.302696Z"}},"outputs":[],"execution_count":null},{"id":"a9437838","cell_type":"markdown","source":"## 🚀 Start Training\n\nRun the cell below to start the training pipeline.\nTraining loss, mAP, and learning rate can be tracked via logs.\n","metadata":{}},{"id":"64476fc9","cell_type":"code","source":"# =============================================================================\n# Start training\n# =============================================================================\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-27T08:30:19.304517Z","iopub.execute_input":"2026-03-27T08:30:19.304811Z","iopub.status.idle":"2026-03-27T11:54:37.959554Z","shell.execute_reply.started":"2026-03-27T08:30:19.304782Z","shell.execute_reply":"2026-03-27T11:54:37.958670Z"}},"outputs":[],"execution_count":null}]}