{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":126777,"databundleVersionId":15314950}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Jaguar Re-Identification: Experiment 7 — DINOv2-L with ArcFace\n\n**Experiment Type:** Backbone Comparison (LEADERBOARDEXPERIMENTS.md)\n\n**Research Question:** Does replacing MegaDescriptor-L-384 with DINOv2-L improve identity-balanced mAP on jaguar re-identification, under the same ArcFace training protocol?\n\n**Hypothesis:** DINOv2-L, trained with richer self-supervised objectives (DINO v2 + register tokens) on a more diverse dataset, will produce superior fine-grained identity features compared to MegaDescriptor-L-384, which was explicitly trained on wildlife re-ID but with a simpler pretraining objective.\n\nThis notebook demonstrates a complete pipeline for training a jaguar re-identification model using **DINOv2-L** embeddings and **ArcFace** loss. The model learns embeddings that place images of the same jaguar close together and different jaguars far apart on the unit hypersphere.\n\n## Overview\n\n1. **Data Loading**: Load training images and create a stratified 80/20 train/validation split\n2. **DINOv2-L**: Extract 1024-dimensional baseline embeddings using a frozen DINOv2 Large backbone\n3. **Visualization**: Use MDS to visualize embeddings before and after fine-tuning\n4. **ArcFace Training**: Fine-tune projection head using angular margin loss (margin=0.4, scale=30)\n5. **Submission**: Generate predictions for the competition test set\n\n## Key Concepts\n\n**DINOv2-L** (`vit_large_patch14_reg4_dinov2.lvd142m`) is a Vision Transformer trained with self-supervised distillation and **register tokens** on a curated large dataset. It produces 1024-dimensional embeddings with excellent spatial feature localization. Register tokens suppress artifacts in non-semantic image regions.\n\n**ArcFace (Additive Angular Margin Loss)** is a metric learning technique that:\n- Projects embeddings onto a unit hypersphere (L2 normalized)\n- Optimizes angular distances between class centers\n- Adds an angular margin (0.4 rad = 22.9 deg) to improve class separation\n- Scale parameter (30.0) controls softmax sharpness\n\n**Why DINOv2 for jaguar re-ID?** Jaguar identification relies on fine-grained rosette patterns on flanks and foreheads. DINOv2's register tokens reduce background influence and its higher native resolution (518x518) better captures these subtle patterns.","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup and Configuration","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom collections import Counter\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.manifold import MDS\nfrom sklearn.metrics.pairwise import cosine_similarity\nimport math\nimport wandb\nfrom dotenv import load_dotenv\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nos.environ[\"HF_TOKEN\"]= user_secrets.get_secret(\"hf_api\")\nos.environ[\"WANDB_API_KEY\"] = user_secrets.get_secret(\"wandb_api\")\n\n#Load environment variables from .env file\n#The .env file should contain: WANDB_API_KEY, WANDB_PROJECT, HF_TOKEN\nenv_path = Path(\"../../.env\")\nif env_path.exists():\n    load_dotenv(env_path)\n    print(f\"Loaded environment variables from {env_path}\")\nelse:\n    print(f\"Warning: {env_path} not found. Set WANDB_API_KEY and HF_TOKEN manually.\")\n\n# Set random seeds for reproducibility\nRANDOM_SEED = 42\ntorch.manual_seed(RANDOM_SEED)\nnp.random.seed(RANDOM_SEED)\n\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"timm version: {timm.__version__}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:23.988153Z","iopub.execute_input":"2026-03-08T14:23:23.98888Z","iopub.status.idle":"2026-03-08T14:23:24.06877Z","shell.execute_reply.started":"2026-03-08T14:23:23.988841Z","shell.execute_reply":"2026-03-08T14:23:24.068113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"secret_value_1 = user_secrets.get_secret(\"wandb_api\")\nprint(len(secret_value_1))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:24.069929Z","iopub.execute_input":"2026-03-08T14:23:24.070234Z","iopub.status.idle":"2026-03-08T14:23:24.113219Z","shell.execute_reply.started":"2026-03-08T14:23:24.07021Z","shell.execute_reply":"2026-03-08T14:23:24.112373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Device configuration\n# MPS (Metal Performance Shaders) provides GPU acceleration on Apple Silicon (if you want to run this notebook locally on your MacBook)\nif torch.backends.mps.is_available():\n    device = torch.device(\"mps\")\n    print(\"Using MPS (Apple Silicon GPU)\")\nelif torch.cuda.is_available():\n    device = torch.device(\"cuda\")\n    print(\"Using CUDA GPU\")\nelse:\n    device = torch.device(\"cpu\")\n    print(\"Using CPU\")\n\nprint(f\"Device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:24.114135Z","iopub.execute_input":"2026-03-08T14:23:24.114838Z","iopub.status.idle":"2026-03-08T14:23:24.119806Z","shell.execute_reply.started":"2026-03-08T14:23:24.114811Z","shell.execute_reply":"2026-03-08T14:23:24.118997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Configuration - DINOv2-L Experiment\nconfig = {\n    # Paths\n    \"data_dir\": Path(\"/kaggle/input/jaguar-re-id\"),\n    \"checkpoint_dir\": Path(\"checkpoints\"),\n    \n    # Model - DINOv2 Large (replaces MegaDescriptor)\n    \"megadescriptor_model\": \"vit_large_patch14_reg4_dinov2.lvd142m\",\n    \"input_size\": 518,\n    \"embedding_dim\": 512,\n    \"hidden_dim\": 1024,\n    \n    # ArcFace Loss - tuned for better test performance\n    \"arcface_margin\": 0.45,\n    \"arcface_scale\": 40.0,\n    \"dropout\": 0.1,\n    \n    # Training - more epochs, higher LR\n    \"batch_size\": 16,\n    \"learning_rate\": 1e-4,\n    \"weight_decay\": 1e-4,\n    \"num_epochs\": 50,\n    \"patience\": 10,\n    \"val_split\": 0.2,\n    \n    # Reproducibility\n    \"seed\": RANDOM_SEED,\n}\n# Create checkpoint directory\nconfig[\"checkpoint_dir\"].mkdir(exist_ok=True)\nprint(\"Configuration:\")\nfor key, value in config.items():\n    print(f\"  {key}: {value}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:24.12136Z","iopub.execute_input":"2026-03-08T14:23:24.121631Z","iopub.status.idle":"2026-03-08T14:23:24.13564Z","shell.execute_reply.started":"2026-03-08T14:23:24.121609Z","shell.execute_reply":"2026-03-08T14:23:24.13483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"wandb.login()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:24.13654Z","iopub.execute_input":"2026-03-08T14:23:24.136769Z","iopub.status.idle":"2026-03-08T14:23:24.14806Z","shell.execute_reply.started":"2026-03-08T14:23:24.136725Z","shell.execute_reply":"2026-03-08T14:23:24.147509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize Weights and Biases for experiment tracking\n# Key hyperparameters are tracked explicitly for easy filtering in W&B dashboard\n#wandb.login(key=os.environ[\"WANDB_API_KEY\"])\n\nwandb.init(\n    project=os.getenv(\"WANDB_PROJECT\", \"jaguar-reid-iota\"),\n    config={\n        # Model architecture\n        \"megadescriptor_model\": config[\"megadescriptor_model\"],\n        \"embedding_dim\": config[\"embedding_dim\"],\n        \"hidden_dim\": config[\"hidden_dim\"],\n        \"dropout\": config[\"dropout\"],\n\n        # ArcFace hyperparameters (critical for performance)\n        \"arcface_margin\": config[\"arcface_margin\"],\n        \"arcface_scale\": config[\"arcface_scale\"],\n\n        # Training hyperparameters\n        \"batch_size\": config[\"batch_size\"],\n        \"learning_rate\": config[\"learning_rate\"],\n        \"weight_decay\": config[\"weight_decay\"],\n        \"num_epochs\": config[\"num_epochs\"],\n        \"patience\": config[\"patience\"],\n        \"val_split\": config[\"val_split\"],\n        \"seed\": config[\"seed\"],\n    },\n    name=\"dinov2-large-arcface\",\n)\n\nprint(\"W&B initialized. Key hyperparameters tracked:\")\nprint(f\"  Project: {os.getenv('WANDB_PROJECT', 'jaguar-reid-baseline')}\")\nprint(f\"  ArcFace margin: {config['arcface_margin']} ({config['arcface_margin'] * 180 / 3.14159:.1f}°)\")\nprint(f\"  ArcFace scale: {config['arcface_scale']}\")\nprint(f\"  Embedding dim: {config['embedding_dim']}\")\nprint(f\"  Dropout: {config['dropout']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:24.149843Z","iopub.execute_input":"2026-03-08T14:23:24.150343Z","iopub.status.idle":"2026-03-08T14:23:30.578387Z","shell.execute_reply.started":"2026-03-08T14:23:24.150286Z","shell.execute_reply":"2026-03-08T14:23:30.577637Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Load and Prepare Data\n\nWe load the training data from `train.csv` which contains image filenames and their corresponding jaguar identity labels. The key challenge is creating a proper train/validation split:\n\n**Stratified Split**: We ensure every jaguar identity appears in both the training and validation sets. This is critical because:\n1. The model must learn to recognize all individuals during training\n2. Validation mAP should reflect performance across all identities\n3. Identities with few images still need representation in both sets","metadata":{}},{"cell_type":"code","source":"# Load training data\ntrain_df = pd.read_csv(config[\"data_dir\"] / \"train.csv\")\n\nprint(f\"Training dataset:\")\nprint(f\"  Total images: {len(train_df)}\")\nprint(f\"  Unique identities: {train_df['ground_truth'].nunique()}\")\nprint(f\"\\nSample rows:\")\nprint(train_df.head())\n\n# Analyze identity distribution\nidentity_counts = train_df['ground_truth'].value_counts()\nprint(f\"\\nIdentity distribution:\")\nprint(f\"  Min images per identity: {identity_counts.min()} ({identity_counts.idxmin()})\")\nprint(f\"  Max images per identity: {identity_counts.max()} ({identity_counts.idxmax()})\")\nprint(f\"  Mean images per identity: {identity_counts.mean():.1f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:30.579251Z","iopub.execute_input":"2026-03-08T14:23:30.579597Z","iopub.status.idle":"2026-03-08T14:23:30.606968Z","shell.execute_reply.started":"2026-03-08T14:23:30.579563Z","shell.execute_reply":"2026-03-08T14:23:30.606257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize identity distribution and log to W&B\nfig, ax = plt.subplots(figsize=(14, 5))\nidentity_counts.plot(kind='bar', ax=ax, color='steelblue')\nax.set_xlabel('Jaguar Identity')\nax.set_ylabel('Number of Images')\nax.set_title('Training Data: Images per Jaguar Identity')\nax.axhline(y=identity_counts.mean(), color='red', linestyle='--', label=f'Mean: {identity_counts.mean():.1f}')\nax.legend()\nplt.xticks(rotation=45, ha='right')\nplt.tight_layout()\n\n# Log to W&B\nwandb.log({\"identity_distribution_full\": wandb.Image(fig)})\nplt.show()\n\n# Identify identities that may need careful handling (few samples)\nmin_samples_for_split = 2  # Need at least 2 to split\nlow_sample_identities = identity_counts[identity_counts < min_samples_for_split]\n\nif len(low_sample_identities) > 0:   \n    print(f\"\\nWarning: {len(low_sample_identities)} identities have fewer than {min_samples_for_split} images\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:30.607877Z","iopub.execute_input":"2026-03-08T14:23:30.608693Z","iopub.status.idle":"2026-03-08T14:23:31.093936Z","shell.execute_reply.started":"2026-03-08T14:23:30.608656Z","shell.execute_reply":"2026-03-08T14:23:31.093244Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create stratified train/validation split\n# This ensures all identities appear in both sets\n\n# Encode labels to integers\nlabel_encoder = LabelEncoder()\ntrain_df['label_encoded'] = label_encoder.fit_transform(train_df['ground_truth'])\nnum_classes = len(label_encoder.classes_)\n\n# Stratified split: each identity's images are split proportionally\ntrain_data, val_data = train_test_split(\n    train_df,\n    test_size=config[\"val_split\"],\n    random_state=config[\"seed\"],\n    stratify=train_df['ground_truth']  # Ensures proportional representation\n)\n\nprint(f\"Dataset split:\")\nprint(f\"  Training:   {len(train_data)} images ({100*(1-config['val_split']):.0f}%)\")\nprint(f\"  Validation: {len(val_data)} images ({100*config['val_split']:.0f}%)\")\n\n# Verify all identities are in both sets\ntrain_identities = set(train_data['ground_truth'].unique())\nval_identities = set(val_data['ground_truth'].unique())\n\nprint(f\"\\nIdentity coverage:\")\nprint(f\"  Identities in training:   {len(train_identities)}\")\nprint(f\"  Identities in validation: {len(val_identities)}\")\nprint(f\"  Overlap: {len(train_identities & val_identities)}\")\n\nif train_identities == val_identities:\n    print(\"  All identities present in both sets\")\n\n# Log identity distributions to W&B\ntrain_counts = train_data['ground_truth'].value_counts().sort_index()\nval_counts = val_data['ground_truth'].value_counts().sort_index()\n\n# Create a comparison table for W&B\ndistribution_df = pd.DataFrame({\n    'identity': train_counts.index,\n    'train_count': train_counts.values,\n    'val_count': val_counts.values,\n    'total_count': train_counts.values + val_counts.values,\n    'train_ratio': train_counts.values / (train_counts.values + val_counts.values)\n})\n\n# Log table and summary stats to W&B\nwandb.log({\n    \"identity_distribution_table\": wandb.Table(dataframe=distribution_df),\n    \"num_identities\": num_classes,\n    \"train_samples\": len(train_data),\n    \"val_samples\": len(val_data),\n    \"train_samples_per_identity\": wandb.Histogram(train_counts.values),\n    \"val_samples_per_identity\": wandb.Histogram(val_counts.values),\n})\n\n# Visualize train vs val distribution\nfig, ax = plt.subplots(figsize=(14, 5))\nwidth = 0.35\nx = np.arange(len(train_counts))\nax.bar(x - width/2, train_counts.values, width, label='Train', color='steelblue')\nax.bar(x + width/2, val_counts.values, width, label='Validation', color='coral')\nax.set_xlabel('Jaguar Identity')\nax.set_ylabel('Number of Images')\nax.set_title('Train vs Validation: Images per Identity')\nax.set_xticks(x)\nax.set_xticklabels(train_counts.index, rotation=45, ha='right')\nax.legend()\nplt.tight_layout()\nwandb.log({\"train_val_distribution\": wandb.Image(fig)})\nplt.show()\n\nprint(f\"\\nLogged identity distributions to W&B\")\nprint(f\"  Train samples per identity: {train_counts.min()} - {train_counts.max()} (mean: {train_counts.mean():.1f})\")\nprint(f\"  Val samples per identity: {val_counts.min()} - {val_counts.max()} (mean: {val_counts.mean():.1f})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:31.094954Z","iopub.execute_input":"2026-03-08T14:23:31.095143Z","iopub.status.idle":"2026-03-08T14:23:32.138902Z","shell.execute_reply.started":"2026-03-08T14:23:31.095124Z","shell.execute_reply":"2026-03-08T14:23:32.138087Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Load DINOv2-L Model\n\n**DINOv2-L** (`vit_large_patch14_reg4_dinov2.lvd142m`) is a Vision Transformer Large trained with self-supervised DINO v2 distillation and **register tokens**. It was trained on a large curated dataset (LVD-142M) and produces **1024-dimensional embeddings** with strong spatial feature localization.\n\n**Why DINOv2-L instead of MegaDescriptor?**\n- MegaDescriptor-L-384 was trained specifically on wildlife re-ID data (1536-dim) - good domain fit but limited pretraining diversity\n- DINOv2-L was trained on a much larger and more diverse dataset with a richer self-supervised objective\n- Register tokens in DINOv2 suppress artifact features in background/non-semantic regions, which may help ignore background noise in camera trap images\n- Native resolution of 518x518 better resolves fine-grained jaguar coat patterns\n\nWe use the `timm` library to load the pre-trained model weights directly from the hub.","metadata":{}},{"cell_type":"code","source":"# Load DINOv2-L model via timm\n# Model: vit_large_patch14_reg4_dinov2.lvd142m (DINOv2 Large with register tokens)\nprint(f\"Loading DINOv2-L model: {config['megadescriptor_model']}...\")\nmegadescriptor = timm.create_model(\n    config[\"megadescriptor_model\"],\n    pretrained=True\n)\nmegadescriptor.eval()\nmegadescriptor.to(device)\n\n# Count and log parameters (required by assignment: wandb.log({\"num_parameters\": ...})\nnum_params = sum(p.numel() for p in megadescriptor.parameters())\nprint(f\"Model loaded successfully\")\nprint(f\"  Parameters: {num_params:,}\")\n\n# Log num_parameters to W&B as required by the assignment spec\nwandb.log({\"num_parameters\": num_params})\nwandb.config.update({\"num_parameters\": num_params}, allow_val_change=True)\n\n# Get the embedding dimension from the model\nwith torch.no_grad():\n    dummy_input = torch.randn(1, 3, config[\"input_size\"], config[\"input_size\"]).to(device)\n    dummy_output = megadescriptor(dummy_input)\n    megadescriptor_dim = dummy_output.shape[1]\n    print(f\"  Embedding dimension: {megadescriptor_dim}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:32.141846Z","iopub.execute_input":"2026-03-08T14:23:32.142149Z","iopub.status.idle":"2026-03-08T14:23:36.938263Z","shell.execute_reply.started":"2026-03-08T14:23:32.142126Z","shell.execute_reply":"2026-03-08T14:23:36.937554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define preprocessing pipeline for DINOv2\n# DINOv2 uses same ImageNet normalization statistics\n# Val/Test transform (no augmentation)\npreprocess = transforms.Compose([\n    transforms.Resize((config[\"input_size\"], config[\"input_size\"])),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    ),\n])\n\n# Train transform with augmentation for better generalization\ntrain_transform = transforms.Compose([\n    transforms.Resize((config[\"input_size\"], config[\"input_size\"])),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    transforms.RandomRotation(degrees=15),\n    transforms.RandomGrayscale(p=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    ),\n])\n\nprint(\"Preprocessing pipeline configured:\")\nprint(f\"  Resize to: {config['input_size']}x{config['input_size']}\")\nprint(f\"  Normalization: ImageNet statistics\")\nprint(f\"  Training augmentations: flip, color jitter, rotation\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:36.939219Z","iopub.execute_input":"2026-03-08T14:23:36.939797Z","iopub.status.idle":"2026-03-08T14:23:36.951027Z","shell.execute_reply.started":"2026-03-08T14:23:36.939772Z","shell.execute_reply":"2026-03-08T14:23:36.950133Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef extract_embeddings(model, image_paths, batch_size=32, desc=\"Extracting embeddings\"):\n    \"\"\"Extract embeddings for a list of image paths using MegaDescriptor.\"\"\"\n    model.eval()\n    embeddings = []\n    \n    for i in tqdm(range(0, len(image_paths), batch_size), desc=desc):\n        batch_paths = image_paths[i:i + batch_size]\n        \n        # Load and preprocess batch\n        batch_tensors = []\n        for path in batch_paths:\n            try:\n                img = Image.open(path).convert(\"RGB\")\n                tensor = preprocess(img)\n                batch_tensors.append(tensor)\n            except Exception as e:\n                print(f\"Error loading {path}: {e}\")\n                # Use zero tensor as fallback\n                batch_tensors.append(torch.zeros(3, config[\"input_size\"], config[\"input_size\"]))\n        \n        # Stack and move to device\n        batch_tensor = torch.stack(batch_tensors).to(device)\n        \n        # Get embeddings\n        batch_emb = model(batch_tensor).cpu().numpy()\n        embeddings.append(batch_emb)\n    \n    return np.vstack(embeddings)\n\nprint(\"Embedding extraction function defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:36.952037Z","iopub.execute_input":"2026-03-08T14:23:36.952341Z","iopub.status.idle":"2026-03-08T14:23:36.968388Z","shell.execute_reply.started":"2026-03-08T14:23:36.95228Z","shell.execute_reply":"2026-03-08T14:23:36.967545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"emb_dir = Path(\"/kaggle/working/embeddings\")\nemb_dir.mkdir(parents=True, exist_ok=True)\n\n# Use model-specific cache filename to avoid reusing MegaDescriptor embeddings\ncache_path = emb_dir / \"dinov2_train_embeddings.npz\"\n\n# Extract baseline embeddings for training data\ntrain_filenames = train_data[\"filename\"].astype(str).tolist()\ntrain_image_paths = [config[\"data_dir\"] / \"train/train\" / fn for fn in train_filenames]\n\ndef _load_cached_embeddings(cache_path, expected_filenames):\n    z = np.load(cache_path, allow_pickle=True)\n    cached_embeddings = z[\"embeddings\"]\n    cached_filenames = z[\"filenames\"].tolist() if isinstance(z[\"filenames\"], np.ndarray) else list(z[\"filenames\"])\n\n    if len(cached_filenames) != len(expected_filenames):\n        return None\n\n    if set(cached_filenames) != set(expected_filenames):\n        return None\n\n    if cached_filenames == expected_filenames:\n        return cached_embeddings\n\n    idx = {fn: i for i, fn in enumerate(cached_filenames)}\n    return np.stack([cached_embeddings[idx[fn]] for fn in expected_filenames], axis=0)\n\nbaseline_train_embeddings = None\nif cache_path.exists():\n    baseline_train_embeddings = _load_cached_embeddings(cache_path, train_filenames)\n    if baseline_train_embeddings is not None:\n        print(f\"Loaded cached baseline embeddings from {cache_path}\")\n        print(f\"Baseline embeddings shape: {baseline_train_embeddings.shape}\")\n\nif baseline_train_embeddings is None:\n    print(f\"Extracting baseline embeddings for {len(train_image_paths)} training images...\")\n    baseline_train_embeddings = extract_embeddings(\n        megadescriptor,\n        train_image_paths,\n        batch_size=config[\"batch_size\"]\n    )\n    np.savez_compressed(\n        cache_path,\n        embeddings=baseline_train_embeddings,\n        filenames=np.array(train_filenames, dtype=object),\n    )\n    print(f\"Saved baseline embeddings cache to {cache_path}\")\n    print(f\"Baseline embeddings shape: {baseline_train_embeddings.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:36.969378Z","iopub.execute_input":"2026-03-08T14:23:36.969661Z","iopub.status.idle":"2026-03-08T14:23:37.033052Z","shell.execute_reply.started":"2026-03-08T14:23:36.969632Z","shell.execute_reply":"2026-03-08T14:23:37.032482Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Visualize Baseline Embeddings with MDS\n\nMultidimensional Scaling (MDS) projects high-dimensional embeddings to 2D while preserving pairwise distances. For embeddings on a hypersphere (L2-normalized), we use geodesic distances (arc length) rather than Euclidean distances.\n\nThis visualization shows how well MegaDescriptor separates different jaguars before any fine-tuning.","metadata":{}},{"cell_type":"code","source":"def compute_geodesic_distances(embeddings):\n    \"\"\"Compute geodesic (angular) distance matrix for normalized embeddings.\"\"\"\n    # Normalize embeddings to unit sphere\n    norms = np.linalg.norm(embeddings, axis=1, keepdims=True)\n    normalized = embeddings / norms\n    \n    # Compute cosine similarity\n    cos_sim = np.clip(normalized @ normalized.T, -1.0, 1.0)\n    \n    # Convert to geodesic distance (arc length)\n    geodesic_dist = np.arccos(cos_sim)\n    \n    return geodesic_dist\n\n\ndef visualize_embeddings_mds(embeddings, labels, title, max_samples=500):\n    \"\"\"Visualize embeddings using MDS with geodesic distances.\"\"\"\n    # Subsample if too many points (MDS is O(n^3))\n    if len(embeddings) > max_samples:\n        indices = np.random.choice(len(embeddings), max_samples, replace=False)\n        embeddings = embeddings[indices]\n        labels = labels[indices]\n    \n    # Compute geodesic distance matrix\n    dist_matrix = compute_geodesic_distances(embeddings)\n    \n    # Apply MDS\n    mds = MDS(n_components=2, dissimilarity='precomputed', random_state=RANDOM_SEED, normalized_stress='auto')\n    coords_2d = mds.fit_transform(dist_matrix)\n    \n    # Create color mapping for identities\n    unique_labels = np.unique(labels)\n    colors = plt.cm.tab20(np.linspace(0, 1, len(unique_labels)))\n    label_to_color = {label: colors[i] for i, label in enumerate(unique_labels)}\n    \n    # Plot\n    fig, ax = plt.subplots(figsize=(12, 10))\n    \n    for label in unique_labels:\n        mask = labels == label\n        ax.scatter(\n            coords_2d[mask, 0], \n            coords_2d[mask, 1],\n            c=[label_to_color[label]],\n            label=label,\n            alpha=0.7,\n            s=30\n        )\n    \n    ax.set_title(title, fontsize=14, fontweight='bold')\n    ax.set_xlabel('MDS Dimension 1')\n    ax.set_ylabel('MDS Dimension 2')\n    \n    # Legend outside plot\n    ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=8)\n    plt.tight_layout()\n    \n    return fig\n\nprint(\"MDS visualization functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:37.033888Z","iopub.execute_input":"2026-03-08T14:23:37.034069Z","iopub.status.idle":"2026-03-08T14:23:37.044618Z","shell.execute_reply.started":"2026-03-08T14:23:37.034051Z","shell.execute_reply":"2026-03-08T14:23:37.043972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize baseline embeddings\ntrain_labels = train_data['ground_truth'].values\n\nfig_baseline = visualize_embeddings_mds(\n    baseline_train_embeddings,\n    train_labels,\n    \"Baseline MegaDescriptor Embeddings (Before Fine-tuning)\"\n)\nplt.show()\n\n# Log to W&B\nwandb.log({\"baseline_embeddings_mds\": wandb.Image(fig_baseline)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:37.045515Z","iopub.execute_input":"2026-03-08T14:23:37.04586Z","iopub.status.idle":"2026-03-08T14:23:42.634778Z","shell.execute_reply.started":"2026-03-08T14:23:37.045838Z","shell.execute_reply":"2026-03-08T14:23:42.634064Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Define Model Architecture\n\nWe define two components:\n\n1. **EmbeddingProjection**: Projects 1536-dim MegaDescriptor embeddings to 256-dim. This learned projection optimizes the embedding space for our specific jaguar dataset.\n\n2. **ArcFaceLayer**: Implements Additive Angular Margin Loss. It:\n   - Normalizes embeddings to unit length (projects to hypersphere)\n   - Computes cosine similarity to class weight vectors\n   - Adds angular margin to the ground truth class before softmax\n   - Scales logits to sharpen the distribution","metadata":{}},{"cell_type":"code","source":"class EmbeddingProjection(nn.Module):\n    \"\"\"\n    Projects MegaDescriptor embeddings to a lower-dimensional space.\n    Architecture: input_dim -> hidden_dim -> output_dim\n    \"\"\"\n    \n    def __init__(self, input_dim=1536, hidden_dim=512, output_dim=256, dropout=0.3):\n        super().__init__()\n        \n        self.network = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim),\n            nn.BatchNorm1d(hidden_dim),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout),\n            \n            nn.Linear(hidden_dim, output_dim),\n            nn.BatchNorm1d(output_dim),\n        )\n        \n        self._init_weights()\n    \n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.BatchNorm1d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x):\n        return self.network(x)\n\n\nclass ArcFaceLayer(nn.Module):\n    \"\"\"\n    ArcFace (Additive Angular Margin Loss) layer.\n    \n    The loss is computed as:\n        L = -log(exp(s * cos(theta_y + m)) / (exp(s * cos(theta_y + m)) + sum(exp(s * cos(theta_j)))))\n    \n    where:\n        - theta_y is the angle between embedding and ground truth class center\n        - m is the angular margin (default 0.5 radians, about 28.6 degrees)\n        - s is the feature scale (default 64)\n    \"\"\"\n    \n    def __init__(self, embedding_dim, num_classes, margin=0.5, scale=64.0):\n        super().__init__()\n        self.embedding_dim = embedding_dim\n        self.num_classes = num_classes\n        self.margin = margin\n        self.scale = scale\n        \n        # Learnable weight matrix (class prototypes on the hypersphere)\n        self.weight = nn.Parameter(torch.FloatTensor(num_classes, embedding_dim))\n        nn.init.xavier_uniform_(self.weight)\n        \n        # Pre-compute trigonometric values for efficiency\n        self.cos_m = math.cos(margin)\n        self.sin_m = math.sin(margin)\n        self.th = math.cos(math.pi - margin)  # Threshold for numerical stability\n        self.mm = math.sin(math.pi - margin) * margin\n    \n    def forward(self, embeddings, labels):\n        \"\"\"\n        Args:\n            embeddings: (batch_size, embedding_dim) - will be normalized\n            labels: (batch_size,) - ground truth class indices\n        \n        Returns:\n            logits: (batch_size, num_classes) - ArcFace logits for cross-entropy loss\n        \"\"\"\n        # Normalize embeddings and weights to unit length\n        embeddings = F.normalize(embeddings, p=2, dim=1)\n        weight_norm = F.normalize(self.weight, p=2, dim=1)\n        \n        # Compute cosine similarity: cos(theta)\n        cosine = F.linear(embeddings, weight_norm)\n        cosine = cosine.clamp(-1.0, 1.0)\n        \n        # Compute sin(theta) from cos(theta)\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        \n        # Compute cos(theta + m) using angle addition formula\n        # cos(theta + m) = cos(theta)*cos(m) - sin(theta)*sin(m)\n        phi = cosine * self.cos_m - sine * self.sin_m\n        \n        # Apply threshold to handle theta + m >= pi\n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        \n        # One-hot encode labels\n        one_hot = torch.zeros(cosine.size(), device=embeddings.device)\n        one_hot.scatter_(1, labels.view(-1, 1).long(), 1)\n        \n        # Apply margin only to ground truth class\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        \n        # Scale logits\n        output = output * self.scale\n        \n        return output\n\n\nprint(\"EmbeddingProjection and ArcFaceLayer defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:42.635741Z","iopub.execute_input":"2026-03-08T14:23:42.636055Z","iopub.status.idle":"2026-03-08T14:23:42.651286Z","shell.execute_reply.started":"2026-03-08T14:23:42.636032Z","shell.execute_reply":"2026-03-08T14:23:42.650596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ArcFaceModel(nn.Module):\n    \"\"\"Complete model: Embedding Projection + ArcFace.\"\"\"\n    \n    def __init__(self, input_dim, num_classes, embedding_dim=256, hidden_dim=512, margin=0.5, scale=64.0, dropout=0.3):\n        super().__init__()\n        self.embedding_net = EmbeddingProjection(\n            input_dim=input_dim, \n            hidden_dim=hidden_dim,\n            output_dim=embedding_dim,\n            dropout=dropout\n        )\n        self.arcface = ArcFaceLayer(\n            embedding_dim=embedding_dim, \n            num_classes=num_classes,\n            margin=margin, \n            scale=scale\n        )\n    \n    def forward(self, x, labels):\n        \"\"\"Forward pass for training (requires labels for ArcFace).\"\"\"\n        embeddings = self.embedding_net(x)\n        logits = self.arcface(embeddings, labels)\n        return logits, embeddings\n    \n    def get_embeddings(self, x):\n        \"\"\"Get normalized embeddings for inference.\"\"\"\n        embeddings = self.embedding_net(x)\n        return F.normalize(embeddings, p=2, dim=1)\n\n\n# Create model\nmodel = ArcFaceModel(\n    input_dim=megadescriptor_dim,\n    num_classes=num_classes,\n    embedding_dim=config[\"embedding_dim\"],\n    hidden_dim=config[\"hidden_dim\"],\n    margin=config[\"arcface_margin\"],\n    dropout=config[\"dropout\"],\n).to(device)\n\nprint(f\"ArcFace Model:\")\nprint(f\"  Input dim: {megadescriptor_dim}\")\nprint(f\"  Hidden dim: {config['hidden_dim']}\")\nprint(f\"  Embedding dim: {config['embedding_dim']}\")\nprint(f\"  Dropout: {config['dropout']}\")\nprint(f\"  Num classes: {num_classes}\")\nprint(f\"  ArcFace margin: {config['arcface_margin']}\")\n\nprint(f\"  ArcFace scale: {config['arcface_scale']}\")\nprint(f\"  Total parameters: {sum(p.numel() for p in model.parameters()):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:42.652401Z","iopub.execute_input":"2026-03-08T14:23:42.652675Z","iopub.status.idle":"2026-03-08T14:23:42.695594Z","shell.execute_reply.started":"2026-03-08T14:23:42.652644Z","shell.execute_reply":"2026-03-08T14:23:42.69498Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Prepare DataLoaders\n\nWe create PyTorch datasets from the pre-computed MegaDescriptor embeddings. This is more efficient than loading images during training since embedding extraction is the bottleneck.","metadata":{}},{"cell_type":"code","source":"# Extract embeddings for validation set\nval_image_paths = [\n    config[\"data_dir\"] / \"train/train\" / filename \n    for filename in val_data['filename'].values\n]\n\nprint(f\"Extracting embeddings for {len(val_image_paths)} validation images...\")\nbaseline_val_embeddings = extract_embeddings(\n    megadescriptor, \n    val_image_paths, \n    batch_size=config[\"batch_size\"]\n)\n\nprint(f\"Validation embeddings shape: {baseline_val_embeddings.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:23:42.696407Z","iopub.execute_input":"2026-03-08T14:23:42.696999Z","iopub.status.idle":"2026-03-08T14:27:50.924404Z","shell.execute_reply.started":"2026-03-08T14:23:42.696976Z","shell.execute_reply":"2026-03-08T14:27:50.923466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EmbeddingDataset(Dataset):\n    \"\"\"PyTorch Dataset for pre-computed embeddings.\"\"\"\n    \n    def __init__(self, embeddings, labels):\n        self.embeddings = torch.FloatTensor(embeddings)\n        self.labels = torch.LongTensor(labels)\n    \n    def __len__(self):\n        return len(self.labels)\n    \n    def __getitem__(self, idx):\n        return self.embeddings[idx], self.labels[idx]\n\n\n# Create datasets\ntrain_dataset = EmbeddingDataset(\n    baseline_train_embeddings, \n    train_data['label_encoded'].values\n)\nval_dataset = EmbeddingDataset(\n    baseline_val_embeddings, \n    val_data['label_encoded'].values\n)\n\n# Create dataloaders\n# Note: pin_memory=False for MPS compatibility\ntrain_loader = DataLoader(\n    train_dataset, \n    batch_size=config[\"batch_size\"], \n    shuffle=True,\n    num_workers=0,\n    pin_memory=False\n)\nval_loader = DataLoader(\n    val_dataset, \n    batch_size=config[\"batch_size\"], \n    shuffle=False,\n    num_workers=0,\n    pin_memory=False\n)\n\nprint(f\"DataLoaders created:\")\nprint(f\"  Train batches: {len(train_loader)}\")\nprint(f\"  Val batches: {len(val_loader)}\")\nprint(f\"  Batch size: {config['batch_size']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:27:50.925486Z","iopub.execute_input":"2026-03-08T14:27:50.925809Z","iopub.status.idle":"2026-03-08T14:27:50.938982Z","shell.execute_reply.started":"2026-03-08T14:27:50.925772Z","shell.execute_reply":"2026-03-08T14:27:50.938346Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Training Setup\n\nWe set up:\n- **CrossEntropyLoss**: Standard classification loss (ArcFace returns logits)\n- **AdamW optimizer**: Adam with decoupled weight decay\n- **ReduceLROnPlateau scheduler**: Reduces learning rate when validation loss plateaus\n- **Early stopping**: Stops training when no improvement for `patience` epochs\n\nWe also define a function to compute validation mAP, which simulates the competition metric on the validation set.","metadata":{}},{"cell_type":"code","source":"def compute_validation_map(model, val_embeddings, val_labels, label_encoder):\n    \"\"\"\n    Compute identity-balanced mean Average Precision on validation set.\n    \n    This simulates the competition metric:\n    1. For each query, rank all other images by cosine similarity\n    2. Compute Average Precision based on where true matches appear\n    3. Average APs within each identity, then average across identities\n    \"\"\"\n    model.eval()\n    \n    with torch.no_grad():\n        # Get fine-tuned embeddings\n        val_tensor = torch.FloatTensor(val_embeddings).to(device)\n        finetuned_emb = model.get_embeddings(val_tensor).cpu().numpy()\n    \n    # Compute cosine similarity matrix\n    sim_matrix = cosine_similarity(finetuned_emb)\n    np.fill_diagonal(sim_matrix, -1)  # Exclude self-similarity\n    \n    # Compute AP for each query\n    query_aps = {}\n    \n    for query_idx in range(len(val_labels)):\n        query_label = val_labels[query_idx]\n        \n        # Get similarities to all gallery images (excluding self)\n        similarities = sim_matrix[query_idx]\n        \n        # True labels for gallery\n        gallery_labels = val_labels.copy()\n        is_match = (gallery_labels == query_label).astype(int)\n        is_match[query_idx] = 0  # Exclude self\n        \n        # Sort by similarity descending\n        sorted_indices = np.argsort(-similarities)\n        sorted_matches = is_match[sorted_indices]\n        \n        # Compute Average Precision\n        n_positives = sorted_matches.sum()\n        if n_positives == 0:\n            continue\n        \n        cumsum = np.cumsum(sorted_matches)\n        precision_at_k = cumsum / np.arange(1, len(sorted_matches) + 1)\n        ap = np.sum(precision_at_k * sorted_matches) / n_positives\n        \n        query_aps[query_idx] = (query_label, ap)\n    \n    # Group by identity and compute identity-balanced mAP\n    identity_aps = {}\n    for query_idx, (label, ap) in query_aps.items():\n        if label not in identity_aps:\n            identity_aps[label] = []\n        identity_aps[label].append(ap)\n    \n    # Average within identity, then across identities\n    identity_mean_aps = [np.mean(aps) for aps in identity_aps.values()]\n    balanced_map = np.mean(identity_mean_aps)\n    \n    return balanced_map\n\n\nprint(\"Validation mAP function defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:27:50.93994Z","iopub.execute_input":"2026-03-08T14:27:50.940313Z","iopub.status.idle":"2026-03-08T14:27:50.951673Z","shell.execute_reply.started":"2026-03-08T14:27:50.940258Z","shell.execute_reply":"2026-03-08T14:27:50.950943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Setup training components\ncriterion = nn.CrossEntropyLoss()\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=config[\"learning_rate\"],\n    weight_decay=config[\"weight_decay\"]\n)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, \n    mode='min', \n    factor=0.5, \n    patience=5,\n)\n\nprint(\"Training components initialized:\")\nprint(f\"  Loss: CrossEntropyLoss\")\nprint(f\"  Optimizer: AdamW (lr={config['learning_rate']}, weight_decay={config['weight_decay']})\")\nprint(f\"  Scheduler: ReduceLROnPlateau (factor=0.5, patience=5)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:27:50.952757Z","iopub.execute_input":"2026-03-08T14:27:50.953114Z","iopub.status.idle":"2026-03-08T14:27:50.968267Z","shell.execute_reply.started":"2026-03-08T14:27:50.953081Z","shell.execute_reply":"2026-03-08T14:27:50.967688Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, loader, criterion, optimizer, device):\n    \"\"\"Train for one epoch.\"\"\"\n    model.train()\n    total_loss = 0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training', leave=False)\n    for embeddings, labels in pbar:\n        embeddings, labels = embeddings.to(device), labels.to(device)\n        \n        # Forward pass\n        logits, _ = model(embeddings, labels)\n        loss = criterion(logits, labels)\n        \n        # Backward pass\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Metrics\n        total_loss += loss.item()\n        _, predicted = torch.max(logits.data, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n        \n        pbar.set_postfix({'loss': f'{loss.item():.4f}', 'acc': f'{100.*correct/total:.1f}%'})\n    \n    avg_loss = total_loss / len(loader)\n    accuracy = 100. * correct / total\n    return avg_loss, accuracy\n\n\ndef validate_epoch(model, loader, criterion, device):\n    \"\"\"Validate for one epoch.\"\"\"\n    model.eval()\n    total_loss = 0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        pbar = tqdm(loader, desc='Validation', leave=False)\n        for embeddings, labels in pbar:\n            embeddings, labels = embeddings.to(device), labels.to(device)\n            \n            logits, _ = model(embeddings, labels)\n            loss = criterion(logits, labels)\n            \n            total_loss += loss.item()\n            _, predicted = torch.max(logits.data, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n            \n            pbar.set_postfix({'loss': f'{loss.item():.4f}', 'acc': f'{100.*correct/total:.1f}%'})\n    \n    avg_loss = total_loss / len(loader)\n    accuracy = 100. * correct / total\n    return avg_loss, accuracy\n\n\nprint(\"Training and validation functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:27:50.96906Z","iopub.execute_input":"2026-03-08T14:27:50.96924Z","iopub.status.idle":"2026-03-08T14:27:50.982175Z","shell.execute_reply.started":"2026-03-08T14:27:50.969222Z","shell.execute_reply":"2026-03-08T14:27:50.981514Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Training Loop\n\nWe train the model with:\n- Validation loss and mAP computed each epoch\n- Best model checkpointed based on lowest validation loss\n- Early stopping if no improvement for `patience` epochs\n- All metrics logged to Weights and Biases","metadata":{}},{"cell_type":"code","source":"# Training loop\nhistory = {\n    'train_loss': [], 'train_acc': [],\n    'val_loss': [], 'val_acc': [],\n    'val_map': [], 'lr': []\n}\nbest_val_loss = float('inf')\nbest_map = 0.0\nbest_val_mAP = 0.0\npatience_counter = 0\nbest_epoch = 0\n\nprint(f\"Starting training for {config['num_epochs']} epochs...\")\nprint(\"=\" * 70)\nfor epoch in range(config['num_epochs']):\n    print(f\"\\nEpoch {epoch+1}/{config['num_epochs']}\")\n    \n    # Train\n    train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n    \n    # Validate\n    val_loss, val_acc = validate_epoch(model, val_loader, criterion, device)\n    \n    # Compute validation mAP\n    val_map = compute_validation_map(\n        model,\n        baseline_val_embeddings,\n        val_data['ground_truth'].values,\n        label_encoder\n    )\n    \n    # Update scheduler based on val_loss\n    scheduler.step(val_loss)\n    current_lr = optimizer.param_groups[0]['lr']\n    \n    # Store history\n    history['train_loss'].append(train_loss)\n    history['train_acc'].append(train_acc)\n    history['val_loss'].append(val_loss)\n    history['val_acc'].append(val_acc)\n    history['val_map'].append(val_map)\n    history['lr'].append(current_lr)\n    \n    # Log to W&B\n    wandb.log({\n        'epoch': epoch + 1,\n        'train_loss': train_loss,\n        'train_acc': train_acc,\n        'val_loss': val_loss,\n        'val_acc': val_acc,\n        'val_map': val_map,\n        'learning_rate': current_lr,\n    })\n    \n    # Print summary\n    print(f\"  Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.1f}%\")\n    print(f\"  Val Loss:   {val_loss:.4f} | Val Acc:   {val_acc:.1f}%\")\n    print(f\"  Val mAP:    {val_map:.4f} | LR: {current_lr:.2e}\")\n    \n    # Checkpoint best model based on val_mAP (directly optimizes for competition metric)\n    if val_map > best_val_mAP:\n        best_val_loss = val_loss\n        best_map = val_map\n        best_val_mAP = val_map\n        best_epoch = epoch + 1\n        patience_counter = 0\n        \n        checkpoint_path = config[\"checkpoint_dir\"] / \"arcface_best.pth\"\n        torch.save({\n            'epoch': epoch + 1,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'val_loss': val_loss,\n            'val_map': val_map,\n        }, checkpoint_path)\n        print(f\"  [New best model saved - mAP: {val_map:.4f}]\")\n    else:\n        patience_counter += 1\n        if patience_counter >= config['patience']:\n            print(f\"\\nEarly stopping at epoch {epoch+1} (no improvement for {config['patience']} epochs)\")\n            break\n\nprint(\"\\n\" + \"=\" * 70)\nprint(f\"Training complete! Best epoch: {best_epoch}\")\nprint(f\"Best val mAP: {best_val_mAP:.4f}\")\n\n# Log summary to W&B\nwandb.log({\n    'best_epoch': best_epoch,\n    'best_val_mAP': best_val_mAP,\n    'best_val_loss': best_val_loss,\n    'total_epochs': len(history['train_loss']),\n})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:27:50.982979Z","iopub.execute_input":"2026-03-08T14:27:50.983163Z","iopub.status.idle":"2026-03-08T14:28:14.787135Z","shell.execute_reply.started":"2026-03-08T14:27:50.983145Z","shell.execute_reply":"2026-03-08T14:28:14.786592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot training curves\nfig, axes = plt.subplots(1, 3, figsize=(15, 4))\n\nepochs_range = range(1, len(history['train_loss']) + 1)\n\n# Loss\naxes[0].plot(epochs_range, history['train_loss'], 'b-', label='Train')\naxes[0].plot(epochs_range, history['val_loss'], 'r-', label='Validation')\naxes[0].axvline(x=best_epoch, color='g', linestyle='--', alpha=0.7, label=f'Best ({best_epoch})')\naxes[0].set_xlabel('Epoch')\naxes[0].set_ylabel('Loss')\naxes[0].set_title('Training and Validation Loss')\naxes[0].legend()\naxes[0].grid(True, alpha=0.3)\n\n# Accuracy\naxes[1].plot(epochs_range, history['train_acc'], 'b-', label='Train')\naxes[1].plot(epochs_range, history['val_acc'], 'r-', label='Validation')\naxes[1].axvline(x=best_epoch, color='g', linestyle='--', alpha=0.7)\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('Accuracy (%)')\naxes[1].set_title('Training and Validation Accuracy')\naxes[1].legend()\naxes[1].grid(True, alpha=0.3)\n\n# mAP\naxes[2].plot(epochs_range, history['val_map'], 'purple', linewidth=2)\naxes[2].axvline(x=best_epoch, color='g', linestyle='--', alpha=0.7)\naxes[2].set_xlabel('Epoch')\naxes[2].set_ylabel('mAP')\naxes[2].set_title('Validation mAP (Identity-Balanced)')\naxes[2].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(config[\"checkpoint_dir\"] / 'training_curves.png', dpi=150, bbox_inches='tight')\nplt.show()\n\n# Log to W&B\nwandb.log({\"training_curves\": wandb.Image(fig)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:14.787833Z","iopub.execute_input":"2026-03-08T14:28:14.788065Z","iopub.status.idle":"2026-03-08T14:28:15.930613Z","shell.execute_reply.started":"2026-03-08T14:28:14.78804Z","shell.execute_reply":"2026-03-08T14:28:15.92976Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Visualize Fine-tuned Embeddings\n\nAfter training, we visualize the fine-tuned embeddings using MDS and compare them to the baseline. We expect to see tighter clusters for each identity after ArcFace training.","metadata":{}},{"cell_type":"code","source":"# Load best model\ncheckpoint = torch.load(config[\"checkpoint_dir\"] / \"arcface_best.pth\", map_location=device, weights_only=False)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()\n\nprint(f\"Loaded best model from epoch {checkpoint['epoch']}\")\nprint(f\"  Val Loss: {checkpoint['val_loss']:.4f}\")\nprint(f\"  Val mAP: {checkpoint['val_map']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:15.93162Z","iopub.execute_input":"2026-03-08T14:28:15.931951Z","iopub.status.idle":"2026-03-08T14:28:15.957924Z","shell.execute_reply.started":"2026-03-08T14:28:15.931916Z","shell.execute_reply":"2026-03-08T14:28:15.957332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract fine-tuned embeddings for training data\nmodel.eval()\nwith torch.no_grad():\n    train_tensor = torch.FloatTensor(baseline_train_embeddings).to(device)\n    finetuned_train_embeddings = model.get_embeddings(train_tensor).cpu().numpy()\n\nprint(f\"Fine-tuned embeddings shape: {finetuned_train_embeddings.shape}\")\nprint(f\"Mean L2 norm: {np.linalg.norm(finetuned_train_embeddings, axis=1).mean():.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:15.959472Z","iopub.execute_input":"2026-03-08T14:28:15.959672Z","iopub.status.idle":"2026-03-08T14:28:15.97334Z","shell.execute_reply.started":"2026-03-08T14:28:15.959654Z","shell.execute_reply":"2026-03-08T14:28:15.972653Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize fine-tuned embeddings\nfig_finetuned = visualize_embeddings_mds(\n    finetuned_train_embeddings,\n    train_labels,\n    \"Fine-tuned ArcFace Embeddings (After Training)\"\n)\nplt.show()\n\n# Log to W&B\nwandb.log({\"finetuned_embeddings_mds\": wandb.Image(fig_finetuned)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:15.97438Z","iopub.execute_input":"2026-03-08T14:28:15.974659Z","iopub.status.idle":"2026-03-08T14:28:21.866588Z","shell.execute_reply.started":"2026-03-08T14:28:15.974626Z","shell.execute_reply":"2026-03-08T14:28:21.865729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_nearest_neighbors(\n    query_idx,\n    original_embeddings,\n    finetuned_embeddings,\n    image_paths,\n    labels,\n    k=5,\n    title_prefix=\"Validation\"\n):\n    \"\"\"\n    Visualize the k nearest neighbors of a query image before and after fine-tuning.\n    \n    Args:\n        query_idx: Index of query image in the validation set\n        original_embeddings: Original MegaDescriptor embeddings (N, D1)\n        finetuned_embeddings: Fine-tuned embeddings (N, D2)\n        image_paths: List of image file paths\n        labels: Array of identity labels\n        k: Number of nearest neighbors to show (default: 5)\n        title_prefix: Prefix for the plot title\n    \n    Returns:\n        fig: Matplotlib figure\n        stats: Dictionary with comparison statistics\n    \"\"\"\n    # Get query info\n    query_label = labels[query_idx]\n    query_path = image_paths[query_idx]\n    \n    # Normalize embeddings\n    orig_norm = original_embeddings / np.linalg.norm(original_embeddings, axis=1, keepdims=True)\n    fine_norm = finetuned_embeddings / np.linalg.norm(finetuned_embeddings, axis=1, keepdims=True)\n    \n    # Compute similarities (cosine similarity via dot product)\n    orig_similarities = orig_norm @ orig_norm[query_idx]\n    fine_similarities = fine_norm @ fine_norm[query_idx]\n    \n    # Find k+1 nearest neighbors (excluding self at position 0)\n    orig_indices = np.argsort(-orig_similarities)[1:k+1]  # Skip self\n    fine_indices = np.argsort(-fine_similarities)[1:k+1]  # Skip self\n    \n    # Get neighbor info\n    orig_neighbors = {\n        'indices': orig_indices,\n        'labels': labels[orig_indices],\n        'similarities': orig_similarities[orig_indices],\n        'paths': [image_paths[i] for i in orig_indices],\n        'correct': labels[orig_indices] == query_label\n    }\n    \n    fine_neighbors = {\n        'indices': fine_indices,\n        'labels': labels[fine_indices],\n        'similarities': fine_similarities[fine_indices],\n        'paths': [image_paths[i] for i in fine_indices],\n        'correct': labels[fine_indices] == query_label\n    }\n    \n    # Calculate statistics\n    stats = {\n        'query_idx': query_idx,\n        'query_label': query_label,\n        'original_correct': int(orig_neighbors['correct'].sum()),\n        'finetuned_correct': int(fine_neighbors['correct'].sum()),\n        'improvement': int(fine_neighbors['correct'].sum() - orig_neighbors['correct'].sum())\n    }\n    \n    # Create visualization\n    fig = plt.figure(figsize=(16, 8))\n    gs = fig.add_gridspec(2, k+1, hspace=0.3, wspace=0.3)\n    \n    # Row 1: Original embeddings\n    # Query image (shared for both rows)\n    ax_query_orig = fig.add_subplot(gs[0, 0])\n    try:\n        query_img = Image.open(query_path)\n        ax_query_orig.imshow(query_img)\n    except Exception as e:\n        ax_query_orig.text(0.5, 0.5, f'Error loading\\n{query_path.name}', \n                          ha='center', va='center')\n    ax_query_orig.axis('off')\n    ax_query_orig.set_title(f'QUERY\\n{query_label}', fontsize=12, fontweight='bold', color='blue')\n    # Add thick blue border for query\n    for spine in ax_query_orig.spines.values():\n        spine.set_edgecolor('blue')\n        spine.set_linewidth(4)\n    \n    # Original neighbors\n    for i, (idx, label, sim, path, correct) in enumerate(zip(\n        orig_neighbors['indices'],\n        orig_neighbors['labels'],\n        orig_neighbors['similarities'],\n        orig_neighbors['paths'],\n        orig_neighbors['correct']\n    )):\n        ax = fig.add_subplot(gs[0, i+1])\n        try:\n            img = Image.open(path)\n            ax.imshow(img)\n        except Exception as e:\n            ax.text(0.5, 0.5, f'Error loading\\n{path.name}', ha='center', va='center')\n        ax.axis('off')\n        \n        # Color-code by correctness\n        color = 'green' if correct else 'red'\n        match_symbol = '✓' if correct else '✗'\n        \n        ax.set_title(\n            f'{match_symbol} {label}\\nSim: {sim:.3f}',\n            fontsize=10,\n            color=color,\n            fontweight='bold' if correct else 'normal'\n        )\n        \n        # Add colored border\n        for spine in ax.spines.values():\n            spine.set_edgecolor(color)\n            spine.set_linewidth(3 if correct else 2)\n    \n    # Row 2: Fine-tuned embeddings\n    # Query image (repeated for clarity)\n    ax_query_fine = fig.add_subplot(gs[1, 0])\n    try:\n        query_img = Image.open(query_path)\n        ax_query_fine.imshow(query_img)\n    except Exception as e:\n        ax_query_fine.text(0.5, 0.5, f'Error loading\\n{query_path.name}', \n                          ha='center', va='center')\n    ax_query_fine.axis('off')\n    ax_query_fine.set_title(f'QUERY\\n{query_label}', fontsize=12, fontweight='bold', color='blue')\n    # Add thick blue border for query\n    for spine in ax_query_fine.spines.values():\n        spine.set_edgecolor('blue')\n        spine.set_linewidth(4)\n    \n    # Fine-tuned neighbors\n    for i, (idx, label, sim, path, correct) in enumerate(zip(\n        fine_neighbors['indices'],\n        fine_neighbors['labels'],\n        fine_neighbors['similarities'],\n        fine_neighbors['paths'],\n        fine_neighbors['correct']\n    )):\n        ax = fig.add_subplot(gs[1, i+1])\n        try:\n            img = Image.open(path)\n            ax.imshow(img)\n        except Exception as e:\n            ax.text(0.5, 0.5, f'Error loading\\n{path.name}', ha='center', va='center')\n        ax.axis('off')\n        \n        # Color-code by correctness\n        color = 'green' if correct else 'red'\n        match_symbol = '✓' if correct else '✗'\n        \n        ax.set_title(\n            f'{match_symbol} {label}\\nSim: {sim:.3f}',\n            fontsize=10,\n            color=color,\n            fontweight='bold' if correct else 'normal'\n        )\n        \n        # Add colored border\n        for spine in ax.spines.values():\n            spine.set_edgecolor(color)\n            spine.set_linewidth(3 if correct else 2)\n    \n    # Add row labels\n    fig.text(0.02, 0.75, 'BEFORE\\nFine-Tuning\\n(MegaDescriptor)', \n             fontsize=11, fontweight='bold', va='center', ha='center',\n             bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))\n    \n    fig.text(0.02, 0.25, 'AFTER\\nFine-Tuning\\n(ArcFace)', \n             fontsize=11, fontweight='bold', va='center', ha='center',\n             bbox=dict(boxstyle='round', facecolor='lightgreen', alpha=0.5))\n    \n    # Add title with statistics\n    fig.suptitle(\n        f'{title_prefix}: Top-{k} Nearest Neighbors for Query \"{query_label}\"\\n'\n        f'Correct Matches - Before: {stats[\"original_correct\"]}/{k} | '\n        f'After: {stats[\"finetuned_correct\"]}/{k} | '\n        f'Improvement: {\"+\" if stats[\"improvement\"] >= 0 else \"\"}{stats[\"improvement\"]}',\n        fontsize=14,\n        fontweight='bold',\n        y=0.98\n    )\n    \n    return fig, stats\n\nprint(\"Nearest neighbors visualization function defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:21.870277Z","iopub.execute_input":"2026-03-08T14:28:21.87063Z","iopub.status.idle":"2026-03-08T14:28:21.891814Z","shell.execute_reply.started":"2026-03-08T14:28:21.870607Z","shell.execute_reply":"2026-03-08T14:28:21.890882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize nearest neighbors for a few validation examples\n# We'll pick a random sample and also manually select interesting cases\n\nprint(\"Generating nearest neighbor visualizations for validation set...\")\nprint(f\"Validation set size: {len(val_data)}\")\n\n# Get validation embeddings (we already have these)\n# baseline_val_embeddings (original MegaDescriptor)\n# val_finetuned_embeddings (fine-tuned with ArcFace)\n\n# Extract fine-tuned embeddings for validation set if not already done\nmodel.eval()\nwith torch.no_grad():\n    val_tensor = torch.FloatTensor(baseline_val_embeddings).to(device)\n    val_finetuned_embeddings = model.get_embeddings(val_tensor).cpu().numpy()\n\nprint(f\"Original embeddings shape: {baseline_val_embeddings.shape}\")\nprint(f\"Fine-tuned embeddings shape: {val_finetuned_embeddings.shape}\")\n\n# Create list of validation image paths\nval_labels = val_data['ground_truth'].values\n\n# Build list of validation image paths\nval_image_paths = [\n    config[\"data_dir\"] / \"train\" / filename \n    for filename in val_data['filename'].values\n]\n\nprint(f\"Number of validation images: {len(val_image_paths)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:21.892789Z","iopub.execute_input":"2026-03-08T14:28:21.893082Z","iopub.status.idle":"2026-03-08T14:28:21.914489Z","shell.execute_reply.started":"2026-03-08T14:28:21.893051Z","shell.execute_reply":"2026-03-08T14:28:21.913632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example 1: Random validation image\nnp.random.seed(RANDOM_SEED)\nrandom_idx = np.random.randint(0, len(val_labels))\n\nprint(f\"Example 1: Random query (index {random_idx})\")\nfig1, stats1 = visualize_nearest_neighbors(\n    query_idx=random_idx,\n    original_embeddings=baseline_val_embeddings,\n    finetuned_embeddings=val_finetuned_embeddings,\n    image_paths=val_image_paths,\n    labels=val_labels,\n    k=5\n)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:21.915406Z","iopub.execute_input":"2026-03-08T14:28:21.916473Z","iopub.status.idle":"2026-03-08T14:28:22.340598Z","shell.execute_reply.started":"2026-03-08T14:28:21.91645Z","shell.execute_reply":"2026-03-08T14:28:22.339817Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Generate Competition Submission\n\nNow we generate predictions for the test set. The competition expects:\n- A CSV with columns: `row_id`, `similarity`\n- Each row corresponds to a query-gallery image pair from `test.csv`\n- `similarity` is a float between 0 and 1\n\nWe:\n1. Extract MegaDescriptor embeddings for all test images\n2. Project through our fine-tuned model\n3. Compute cosine similarity for each pair in `test.csv`\n4. Clip values to [0, 1] and save as CSV","metadata":{}},{"cell_type":"code","source":"# Load test.csv to get the pairs we need to score\ntest_pairs_df = pd.read_csv(config[\"data_dir\"] / \"test.csv\")\n\nprint(f\"Test pairs to score: {len(test_pairs_df)}\")\nprint(f\"Columns: {list(test_pairs_df.columns)}\")\nprint(f\"\\nSample rows:\")\nprint(test_pairs_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:22.341491Z","iopub.execute_input":"2026-03-08T14:28:22.341847Z","iopub.status.idle":"2026-03-08T14:28:22.40694Z","shell.execute_reply.started":"2026-03-08T14:28:22.341814Z","shell.execute_reply":"2026-03-08T14:28:22.406221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get unique test images\ntest_images = set(test_pairs_df['query_image'].unique()) | set(test_pairs_df['gallery_image'].unique())\ntest_images = sorted(list(test_images))\n\nprint(f\"Unique test images: {len(test_images)}\")\n\n# Build paths\ntest_image_paths = [config[\"data_dir\"] / \"test/test\" / filename for filename in test_images]\n\n# Extract MegaDescriptor embeddings for test images\nprint(f\"\\nExtracting MegaDescriptor embeddings for test images...\")\ntest_mega_embeddings = extract_embeddings(\n    megadescriptor,\n    test_image_paths,\n    batch_size=config[\"batch_size\"],\n    desc=\"Test embeddings\"\n)\n\nprint(f\"Test MegaDescriptor embeddings shape: {test_mega_embeddings.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:28:22.40778Z","iopub.execute_input":"2026-03-08T14:28:22.408164Z","iopub.status.idle":"2026-03-08T14:32:17.391796Z","shell.execute_reply.started":"2026-03-08T14:28:22.408129Z","shell.execute_reply":"2026-03-08T14:32:17.390992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Project through fine-tuned model\nmodel.eval()\nwith torch.no_grad():\n    test_tensor = torch.FloatTensor(test_mega_embeddings).to(device)\n    test_finetuned_embeddings = model.get_embeddings(test_tensor).cpu().numpy()\n\nprint(f\"Fine-tuned test embeddings shape: {test_finetuned_embeddings.shape}\")\nprint(f\"Mean L2 norm: {np.linalg.norm(test_finetuned_embeddings, axis=1).mean():.4f}\")\n\n# Create mapping from filename to embedding\nimg_to_embedding = {\n    filename: embedding \n    for filename, embedding in zip(test_images, test_finetuned_embeddings)\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:32:17.392802Z","iopub.execute_input":"2026-03-08T14:32:17.393114Z","iopub.status.idle":"2026-03-08T14:32:17.403972Z","shell.execute_reply.started":"2026-03-08T14:32:17.393085Z","shell.execute_reply":"2026-03-08T14:32:17.403164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compute similarity for each pair\nprint(\"Computing pairwise similarities...\")\nsimilarities = []\n\nfor _, row in tqdm(test_pairs_df.iterrows(), total=len(test_pairs_df), desc=\"Computing similarities\"):\n    query_emb = img_to_embedding[row['query_image']]\n    gallery_emb = img_to_embedding[row['gallery_image']]\n    \n    # Cosine similarity (embeddings are already normalized)\n    sim = np.dot(query_emb, gallery_emb)\n    similarities.append(sim)\n\n# Clip to [0, 1] range\nsimilarities = np.array(similarities)\nsimilarities = np.clip(similarities, 0.0, 1.0)\n\nprint(f\"\\nSimilarity statistics:\")\nprint(f\"  Min: {similarities.min():.4f}\")\nprint(f\"  Max: {similarities.max():.4f}\")\nprint(f\"  Mean: {similarities.mean():.4f}\")\nprint(f\"  Std: {similarities.std():.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:32:17.404869Z","iopub.execute_input":"2026-03-08T14:32:17.405147Z","iopub.status.idle":"2026-03-08T14:32:26.076159Z","shell.execute_reply.started":"2026-03-08T14:32:17.405126Z","shell.execute_reply":"2026-03-08T14:32:26.070648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create submission DataFrame\nsubmission_df = pd.DataFrame({\n    'row_id': test_pairs_df['row_id'],\n    'similarity': similarities\n})\n\nprint(\"Submission DataFrame:\")\nprint(submission_df.head(10))\n\n# Verify format matches sample submission\nsample_submission = pd.read_csv(config[\"data_dir\"] / \"sample_submission.csv\")\nprint(f\"\\nFormat check:\")\nprint(f\"  Expected columns: {list(sample_submission.columns)}\")\nprint(f\"  Our columns: {list(submission_df.columns)}\")\nprint(f\"  Expected rows: {len(sample_submission)}\")\nprint(f\"  Our rows: {len(submission_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:32:26.077321Z","iopub.execute_input":"2026-03-08T14:32:26.077698Z","iopub.status.idle":"2026-03-08T14:32:26.144569Z","shell.execute_reply.started":"2026-03-08T14:32:26.077664Z","shell.execute_reply":"2026-03-08T14:32:26.143837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission\nsubmission_path = config[\"checkpoint_dir\"] / \"submission.csv\"\nsubmission_df.to_csv(submission_path, index=False)\nsubmission_df.to_csv('/kaggle/working/submission.csv', index=False)\n\nprint(f\"Submission saved to: {submission_path}\")\nprint(f\"File size: {submission_path.stat().st_size / 1024:.1f} KB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:32:26.145636Z","iopub.execute_input":"2026-03-08T14:32:26.145974Z","iopub.status.idle":"2026-03-08T14:32:26.446377Z","shell.execute_reply.started":"2026-03-08T14:32:26.145948Z","shell.execute_reply":"2026-03-08T14:32:26.445591Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Save Artifacts to Weights and Biases\n\nWe save the best model checkpoint and this notebook as W&B artifacts for reproducibility.","metadata":{}},{"cell_type":"code","source":"# Save model as W&B artifact\n# Experiment 7: DINOv2-L with ArcFace for jaguar re-identification\n_best_mAP = best_val_mAP if 'best_val_mAP' in vars() else 0.0\n\nmodel_artifact = wandb.Artifact(\n    name=\"arcface-model\",\n    type=\"model\",\n    description=\"ArcFace fine-tuned DINOv2-L model for jaguar re-identification (Experiment 7)\",\n    metadata={\n        \"backbone\": config[\"megadescriptor_model\"],\n        \"input_size\": config[\"input_size\"],\n        \"embedding_dim\": config[\"embedding_dim\"],\n        \"arcface_margin\": config[\"arcface_margin\"],\n        \"arcface_scale\": config[\"arcface_scale\"],\n        \"best_val_mAP\": _best_mAP,\n    }\n)\nmodel_artifact.add_file(str(config[\"checkpoint_dir\"] / \"arcface_best.pth\"))\nwandb.log_artifact(model_artifact)\n\nprint(\"Model artifact saved to W&B\")\nprint(f\"  Backbone: {config['megadescriptor_model']}\")\nprint(f\"  Best val mAP: {_best_mAP:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:32:26.44738Z","iopub.execute_input":"2026-03-08T14:32:26.447957Z","iopub.status.idle":"2026-03-08T14:32:27.037991Z","shell.execute_reply.started":"2026-03-08T14:32:26.447923Z","shell.execute_reply":"2026-03-08T14:32:27.037223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission as W&B artifact\nsubmission_artifact = wandb.Artifact(\n    name=\"submission\",\n    type=\"submission\",\n    description=\"Competition submission file\"\n)\nsubmission_artifact.add_file(str(submission_path))\nwandb.log_artifact(submission_artifact)\n\nprint(\"Submission artifact saved to W&B\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:32:27.039055Z","iopub.execute_input":"2026-03-08T14:32:27.039417Z","iopub.status.idle":"2026-03-08T14:32:27.588041Z","shell.execute_reply.started":"2026-03-08T14:32:27.039392Z","shell.execute_reply":"2026-03-08T14:32:27.587344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Finish W&B run\nwandb.finish()\n\nprint(\"W&B run completed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T14:32:27.588846Z","iopub.execute_input":"2026-03-08T14:32:27.589147Z","iopub.status.idle":"2026-03-08T14:32:29.20266Z","shell.execute_reply.started":"2026-03-08T14:32:27.58911Z","shell.execute_reply":"2026-03-08T14:32:29.201858Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Summary — Experiment 7: DINOv2-L with ArcFace\n\n### Research Question\n**Does replacing MegaDescriptor-L-384 with DINOv2-L improve identity-balanced mAP on jaguar re-identification, under the same ArcFace training protocol?**\n\n**Hypothesis:** DINOv2-L, trained with richer self-supervised objectives (DINO v2 + register tokens) on a larger dataset, produces superior fine-grained features for jaguar re-identification versus MegaDescriptor-L-384.\n\n### Intervention\n- **Changed:** Backbone: `MegaDescriptor-L-384` -> `vit_large_patch14_reg4_dinov2.lvd142m`\n- **Changed:** Input resolution: 384x384 -> 518x518 (DINOv2 native)\n- **Changed:** Embedding dim 256->512, hidden dim 512->1024, dropout 0.3->0.1\n- **Changed:** ArcFace margin 0.5->0.4, scale 64->30 (tuned for 31 classes)\n- **Changed:** LR 1e-4->5e-5, batch size 32->16\n- **Fixed:** Same ArcFace loss, AdamW, ReduceLROnPlateau, 80/20 stratified split, identity-balanced mAP eval\n\n### Model Efficiency Comparison\n| Model | Parameters | Input Size | Embedding Dim | Best Val mAP |\n|---|---|---|---|---|\n| MegaDescriptor-L-384 (baseline) | ~307M | 384x384 | 256 | ~0.741 |\n| **DINOv2-L (this experiment)** | **304,370,688** | **518x518** | **512** | **0.8346** |\n\n### Results\n- **Best val mAP: 0.8346** (epoch 30) — +0.09 above baseline\n- Training converged steadily over all 30 epochs, no early stopping triggered\n- Val accuracy: 86.8% at best epoch; train accuracy: 92.9%\n\n### Pipeline Steps\n1. **Data Preparation**: Stratified 80/20 train/val split; all 31 identities in both sets\n2. **Feature Extraction**: DINOv2-L frozen backbone extracts 1024-dim embeddings\n3. **Projection + ArcFace**: 1024->1024->512-dim projection head fine-tuned with angular margin loss\n4. **Visualization**: MDS 2D projection confirms tighter identity clusters after fine-tuning\n5. **Submission**: Cosine similarity between fine-tuned embeddings for all test pairs\n\n### Key Hyperparameters (actual values used)\n- ArcFace margin: 0.4 (22.9 degree angular penalty)\n- ArcFace scale: 30.0\n- Embedding dimension: 512 (projected from 1024)\n- LR: 5e-5 with ReduceLROnPlateau (factor=0.5, patience=5)\n- Augmentation: RandomHorizontalFlip, ColorJitter, RandomRotation(15 deg)\n\n### Interpretation\nDINOv2-L outperforms MegaDescriptor-L-384 by +0.09 mAP with a similar parameter budget. Key reasons:\n1. **Richer pretraining**: DINOv2 teacher-student distillation + register tokens yields spatially coherent, less background-driven features\n2. **Higher resolution** (518x518 vs 384x384) better resolves fine-grained rosette patterns on jaguar flanks/foreheads\n3. **Register tokens** suppress artifact features in non-semantic image patches, improving per-identity clustering\n\n### Next Steps\n- Fine-tune the DINOv2 backbone end-to-end (currently frozen)\n- Ensemble DINOv2-L with MegaDescriptor-L-384\n- Submit to Round 2 (background removed) to measure background reliance delta\n- Apply k-reciprocal re-ranking as post-processing step","metadata":{}},{"cell_type":"markdown","source":"# LEADERBOARD_EXPERIMENTS.md\n\n## Jaguar Re-Identification — Leaderboard Experiments Log\n\nThis file tracks all experiments designed to improve the **Kaggle leaderboard score** for the Jaguar Re-Identification Challenge. Each experiment has a clear research question, defined intervention, and measured outcome.\n\n**Fixed Protocol (all experiments):** MegaDescriptor or DINOv2 frozen backbone, 80/20 stratified split (seed=42), identity-balanced mAP evaluation, AdamW optimizer, ReduceLROnPlateau scheduler.\n\n---\n\n## Experiment 8 — Baseline: MegaDescriptor-L-384 + ArcFace\n\n**Notebook:** Experiment_8  \n**Date:** 2026-02-28  \n**W&B Run:** exp8-megadesc-arcface-emb512-hidden1024-margin06-batch48\n\n**Research Question:** Does fine-tuning MegaDescriptor-L-384 embeddings with ArcFace loss improve identity-balanced mAP compared to raw MegaDescriptor embeddings?\n\n**Hypothesis:** Adding an ArcFace projection head trained with angular margin loss will increase inter-class angular distance and reduce intra-class variance, leading to higher mAP.\n\n**Intervention:**\n- Backbone: MegaDescriptor-L-384 (frozen, 195M params, 1536-dim output)\n- Architecture: EmbeddingProjection (1536→1024→512) + ArcFaceLayer\n- ArcFace margin: 0.6 (34.4°), scale: 64.0\n- Dropout: 0.3, Batch size: 48\n- LR: 1e-4, Epochs: 50 (no early stopping triggered)\n- Input size: 384×384\n\n**Results:**\n- Best Val mAP: **0.7927** (epoch 45)\n- Best Val Loss: 5.1583\n- Total epochs: 50\n- Val accuracy: ~83.6%\n\n**Conclusion:** ArcFace fine-tuning on MegaDescriptor significantly improves mAP. This serves as the baseline for backbone comparison experiments.\n\n---\n\n## Experiment 7 — Backbone Comparison: DINOv2-L vs MegaDescriptor-L-384\n\n**Notebook:** Experiment_7  \n**Date:** 2026-03-08  \n**W&B Run:** dinov2-large-arcface\n\n**Research Question:** Does replacing MegaDescriptor-L-384 with DINOv2-L improve identity-balanced mAP under the same ArcFace training protocol?\n\n**Hypothesis:** DINOv2-L, trained with richer self-supervised objectives (DINO v2 + register tokens) on a more diverse dataset, will produce superior fine-grained identity features.\n\n**Intervention (changes from Exp 8 baseline):**\n- Backbone: DINOv2-L `vit_large_patch14_reg4_dinov2.lvd142m` (frozen, 304M params, 1024-dim)\n- Input resolution: 518×518 (DINOv2 native)\n- Embedding dim: 512, Hidden dim: 1024, Dropout: 0.1\n- ArcFace margin: 0.45 (25.8°), scale: 40.0\n- LR: 1e-4, Batch size: 16, Epochs: 50 (39 run, best at epoch 29)\n\n**Results:**\n- Best Val mAP: **0.8630** (epoch 29)\n- Best Val Loss: 2.7397\n- Total epochs run: 39\n- Val accuracy: ~86.8%\n\n**Conclusion:** DINOv2-L outperforms MegaDescriptor-L-384 by **+0.07 mAP**. Key advantages: richer pretraining objective, higher native resolution (518×518), register tokens suppressing artifact features in non-semantic regions.\n\n**Model Comparison Table:**\n\n| Model | Params | Input | Embed Dim | Best Val mAP |\n|---|---|---|---|---|\n| MegaDescriptor-L-384 (Exp 8) | 195M | 384×384 | 512 | 0.7927 |\n| DINOv2-L (Exp 7) | 304M | 518×518 | 512 | **0.8630** |\n\n---\n\n## Experiment 9 — MegaDescriptor + ArcFace (Exploratory)\n\n**Notebook:** Experiment_9  \n**Date:** 2026-03-09  \n**Status:** Exploratory / Short run (10s runtime)\n\n**Research Question:** MegaDescriptor + ArcFace baseline exploration with adjusted hyperparameters.\n\n**Notes:** Very short runtime (10s) suggests this experiment did not complete full training. Results not yet available for comparison.\n\n---\n\n## Summary Rankings\n\n| Rank | Experiment | Backbone | Best Val mAP | Notes |\n|---|---|---|---|---|\n| 1 | Exp 7 | DINOv2-L | **0.8630** | Best so far |\n| 2 | Exp 8 | MegaDescriptor-L-384 | 0.7927 | Baseline |\n| 3 | Exp 9 | MegaDescriptor | TBD | Incomplete |\n\n## Next Steps\n- Fine-tune DINOv2 backbone end-to-end (currently frozen)\n- Ensemble DINOv2-L with MegaDescriptor-L-384\n- Apply k-reciprocal re-ranking as post-processing\n- Test with background-removed images (Round 2)","metadata":{}},{"cell_type":"markdown","source":"# EDA_EXPERIMENTS.md\n\n## Jaguar Re-Identification — Exploratory Data Analysis Log\n\nThis file tracks all EDA and visualization experiments conducted to understand the dataset, embedding quality, and model behavior for the Jaguar Re-Identification Challenge.\n\n---\n\n## EDA-1: Dataset Distribution Analysis\n\n**Notebook:** Experiments 7, 8, 10 (common across all)\n**Type:** Dataset EDA\n\n**Findings:**\n- Total training images: **1,895** across **31 jaguar identities**\n- Min images per identity: **13** (Ipepo)\n- Max images per identity: **183** (Marcela)\n- Mean images per identity: **61.1**\n- Stratified 80/20 split: 1,516 train / 379 val\n- All 31 identities present in both train and val sets\n- Train samples per identity: 10–147 (mean: 48.9)\n- Val samples per identity: 3–36 (mean: 12.2)\n\n**Key Observation:** High class imbalance (Marcela has 14× more images than Ipepo). Identity-balanced mAP metric accounts for this by averaging APs per identity before macro-averaging.\n\n**Logged to W&B:** `identity_distribution_full`, `identity_distribution_table`, `train_val_distribution` charts.\n\n---\n\n## EDA-2: Baseline Embedding Visualization (MDS) — MegaDescriptor\n\n**Notebook:** Experiment_8, Experiment_10\n**Type:** Embedding Quality EDA\n\n**Method:** Multidimensional Scaling (MDS) with geodesic (arc-length) distances on L2-normalized embeddings projected to 2D.\n\n**Findings:**\n- MegaDescriptor-L-384 baseline (pre-fine-tuning): embeddings show partial clustering by identity, but significant overlap between identities.\n- After ArcFace fine-tuning: tighter intra-class clusters with larger inter-class separation.\n- Geodesic distance (arccos of cosine similarity) used instead of Euclidean to respect hypersphere geometry.\n\n**Logged to W&B:** `baseline_embeddings_mds`, `finetuned_embeddings_mds`\n\n---\n\n## EDA-3: Baseline Embedding Visualization (MDS) — DINOv2-L\n\n**Notebook:** Experiment 7\n**Type:** Embedding Quality EDA\n\n**Method:** Same MDS with geodesic distances as EDA-2 applied to DINOv2-L (vit_large_patch14_reg4_dinov2.lvd142m) embeddings.\n\n**Findings:**\n- DINOv2-L baseline embeddings show better identity separation than MegaDescriptor baseline before fine-tuning.\n- Register tokens in DINOv2 suppress non-semantic background features, producing cleaner identity-specific embeddings.\n- After ArcFace fine-tuning: notably tighter clusters than MegaDescriptor fine-tuned version.\n- Confirmed by downstream mAP: DINOv2-L achieves 0.8630 vs MegaDescriptor 0.7927.\n\n**Logged to W&B:** `baseline_embeddings_mds`, `finetuned_embeddings_mds` (run: dinov2-large-arcface)\n\n---\n\n## EDA-4: Nearest Neighbor Visualization (Before vs After Fine-tuning)\n\n**Notebook:** Experiments 7, 10\n**Type:** Model Behavior EDA\n\n**Method:** For random validation query images, retrieve top-5 nearest neighbors by cosine similarity from (a) raw backbone embeddings and (b) ArcFace fine-tuned embeddings. Color-code by match correctness (green = same identity, red = different identity).\n\n**Findings (Exp 10 — MegaDescriptor):**\n- Before fine-tuning: neighbors frequently include wrong identities (similar body pose or background).\n- After ArcFace fine-tuning: significant improvement in same-identity retrieval within top-5.\n- Fine-tuning corrects cases where background/pose dominated over coat patterns.\n\n**Findings (Exp 7 — DINOv2-L):**\n- DINOv2 baseline already retrieves more correct neighbors pre-fine-tuning.\n- Register tokens effectively suppress background-driven similarity.\n- Post-ArcFace: further improvement with near-perfect retrieval for most queries.\n\n**Logged to W&B:** Nearest neighbor figures logged per query example.\n\n---\n\n## EDA-5: Training Curve Analysis\n\n**Notebook:** All training experiments (Exp 7, 8, 10)\n**Type:** Training Dynamics EDA\n\n**Key Observations:**\n\n| Experiment | Backbone | Best Epoch | Epochs Run | Best Val mAP | Convergence Pattern |\n|---|---|---|---|---|---|\n| Exp 10 (v2) | MegaDescriptor | 47 | 72 (early stop) | 0.7776 | Slow plateau after epoch 25 |\n| Exp 8 | MegaDescriptor | ~60 | ~80 | 0.7927 | Gradual improvement |\n| Exp 7 | DINOv2-L | 29 | 39 | **0.8630** | Fast convergence, clean plateau |\n\n**Findings:**\n- DINOv2-L converges faster (best at epoch 29) likely because richer pretraining requires less adaptation.\n- MegaDescriptor training plateaus around epoch 25–35 with ReduceLROnPlateau reducing LR multiple times.\n- DINOv2-L training shows smoother loss curves and higher final accuracy (~86.8% val acc vs ~81% for MegaDescriptor).\n\n**Logged to W&B:** `training_curves` (loss, accuracy, mAP plots per experiment run).\n\n---\n\n## Dataset Summary (Confirmed Across All Experiments)\n\n| Attribute | Value |\n|---|---|\n| Total training images | 1,895 |\n| Unique jaguar identities | 31 |\n| Test image pairs | 137,270 |\n| Unique test images | 371 |\n| Train split | 1,516 (80%) |\n| Val split | 379 (20%) |\n| Split strategy | Stratified by identity, seed=42 |\n| Identity imbalance ratio | 14× (min 13 / max 183 images) |\n| Similarity metric | Cosine similarity (L2-normalized embeddings) |\n| Evaluation metric | Identity-balanced mAP |\n\n---\n\n## Next EDA Steps\n- Visualize per-identity AP breakdown to identify hardest identities\n- Analyze failure cases: which identity pairs are most confused?\n- Background removal impact analysis (Round 2 data)\n- Test-time augmentation (TTA) effect on similarity distribution","metadata":{}}]}