{"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":"gpu","dataSources":[{"sourceId":126777,"databundleVersionId":15314950,"isSourceIdPinned":false,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Jaguar Re-Identification Baseline with MegaDescriptor and ArcFace\n\nThis notebook demonstrates a complete pipeline for training a jaguar re-identification model using MegaDescriptor embeddings and ArcFace loss. The goal is to learn embeddings that place images of the same jaguar close together and images of different jaguars far apart.\n\n## Overview\n\n1. **Data Loading**: Load training images and create a stratified train/validation split\n2. **MegaDescriptor**: Extract baseline embeddings using a pre-trained vision transformer\n3. **Visualization**: Use MDS to visualize embeddings before and after fine-tuning\n4. **ArcFace Training**: Fine-tune embeddings using angular margin loss\n5. **Submission**: Generate predictions for the competition test set\n\n## Key Concepts\n\n**MegaDescriptor** is a vision transformer trained on wildlife re-identification datasets. It produces 1536-dimensional embeddings that capture visual features useful for distinguishing individual animals.\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 to improve class separation\n\nThe combination allows us to fine-tune MegaDescriptor for our specific jaguar dataset.","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\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\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\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-01-14T12:31:35.524957Z","iopub.execute_input":"2026-01-14T12:31:35.525293Z","iopub.status.idle":"2026-01-14T12:31:50.032523Z","shell.execute_reply.started":"2026-01-14T12:31:35.525264Z","shell.execute_reply":"2026-01-14T12:31:50.031777Z"}},"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-01-14T12:31:50.034244Z","iopub.execute_input":"2026-01-14T12:31:50.034699Z","iopub.status.idle":"2026-01-14T12:31:50.099372Z","shell.execute_reply.started":"2026-01-14T12:31:50.034671Z","shell.execute_reply":"2026-01-14T12:31:50.098475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Configuration\nconfig = {\n    # Paths\n    \"data_dir\": Path(\"/kaggle/input/jaguar-re-id\"),\n    \"checkpoint_dir\": Path(\"checkpoints\"),\n    \n    # Model\n    \"megadescriptor_model\": \"hf-hub:BVRA/MegaDescriptor-L-384\",\n    \"input_size\": 384,\n    \"embedding_dim\": 256,\n    \"hidden_dim\": 512,\n    \n    # ArcFace\n    \"arcface_margin\": 0.5,\n    \"arcface_scale\": 64.0,\n    \"dropout\": 0.3,\n    \n    # Training\n    \"batch_size\": 32,\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\n# Create checkpoint directory\nconfig[\"checkpoint_dir\"].mkdir(exist_ok=True)\n\nprint(\"Configuration:\")\nfor key, value in config.items():\n    print(f\"  {key}: {value}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T12:31:50.100702Z","iopub.execute_input":"2026-01-14T12:31:50.101163Z","iopub.status.idle":"2026-01-14T12:31:50.119366Z","shell.execute_reply.started":"2026-01-14T12:31:50.101119Z","shell.execute_reply":"2026-01-14T12:31:50.118499Z"}},"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\nwandb.login(key=os.environ[\"WANDB_API_KEY\"])\n\nwandb.init(\n    project=os.getenv(\"WANDB_PROJECT\", \"jaguar-reid-baseline\"),\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=\"megadescriptor-arcface-local\",\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-01-14T12:31:50.120307Z","iopub.execute_input":"2026-01-14T12:31:50.120622Z","iopub.status.idle":"2026-01-14T12:32:04.821366Z","shell.execute_reply.started":"2026-01-14T12:31:50.120585Z","shell.execute_reply":"2026-01-14T12:32:04.820480Z"}},"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-01-14T12:32:04.823225Z","iopub.execute_input":"2026-01-14T12:32:04.823557Z","iopub.status.idle":"2026-01-14T12:32:04.873084Z","shell.execute_reply.started":"2026-01-14T12:32:04.823525Z","shell.execute_reply":"2026-01-14T12:32:04.872242Z"}},"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-01-14T12:32:04.874040Z","iopub.execute_input":"2026-01-14T12:32:04.874314Z","iopub.status.idle":"2026-01-14T12:32:05.967649Z","shell.execute_reply.started":"2026-01-14T12:32:04.874289Z","shell.execute_reply":"2026-01-14T12:32:05.966890Z"}},"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-01-14T12:32:05.968635Z","iopub.execute_input":"2026-01-14T12:32:05.968959Z","iopub.status.idle":"2026-01-14T12:32:07.528215Z","shell.execute_reply.started":"2026-01-14T12:32:05.968924Z","shell.execute_reply":"2026-01-14T12:32:07.527443Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Load MegaDescriptor Model\n\nMegaDescriptor is a Vision Transformer (ViT-L/14) trained specifically for wildlife re-identification. It was trained on multiple species datasets and produces 1536-dimensional embeddings.\n\nWe use the `timm` library to load the pre-trained model from Hugging Face Hub.","metadata":{}},{"cell_type":"code","source":"# Load MegaDescriptor model\nprint(\"Loading MegaDescriptor-L-384 model...\")\nmegadescriptor = timm.create_model(\n    config[\"megadescriptor_model\"],\n    pretrained=True\n)\nmegadescriptor.eval()\nmegadescriptor.to(device)\n\nprint(f\"Model loaded successfully\")\nprint(f\"  Parameters: {sum(p.numel() for p in megadescriptor.parameters()):,}\")\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-01-14T12:32:07.529238Z","iopub.execute_input":"2026-01-14T12:32:07.529616Z","iopub.status.idle":"2026-01-14T12:32:20.430249Z","shell.execute_reply.started":"2026-01-14T12:32:07.529571Z","shell.execute_reply":"2026-01-14T12:32:20.429107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define preprocessing pipeline\n# MegaDescriptor expects 384x384 images normalized with ImageNet statistics\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\nprint(\"Preprocessing pipeline configured:\")\nprint(f\"  Resize to: {config['input_size']}x{config['input_size']}\")\nprint(f\"  Normalization: ImageNet statistics\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T12:32:20.431425Z","iopub.execute_input":"2026-01-14T12:32:20.431772Z","iopub.status.idle":"2026-01-14T12:32:20.444821Z","shell.execute_reply.started":"2026-01-14T12:32:20.431738Z","shell.execute_reply":"2026-01-14T12:32:20.443822Z"}},"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-01-14T12:32:20.446592Z","iopub.execute_input":"2026-01-14T12:32:20.446943Z","iopub.status.idle":"2026-01-14T12:32:23.367004Z","shell.execute_reply.started":"2026-01-14T12:32:20.446907Z","shell.execute_reply":"2026-01-14T12:32:23.366022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"emb_dir = Path(\"/kaggle/working/embeddings\")\nemb_dir.mkdir(parents=True, exist_ok=True)\n\ncache_path = emb_dir / \"baseline_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}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T12:32:23.368384Z","iopub.execute_input":"2026-01-14T12:32:23.368636Z","iopub.status.idle":"2026-01-14T12:43:22.713726Z","shell.execute_reply.started":"2026-01-14T12:32:23.368613Z","shell.execute_reply":"2026-01-14T12:43:22.712877Z"}},"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-01-14T12:43:22.715914Z","iopub.execute_input":"2026-01-14T12:43:22.716469Z","iopub.status.idle":"2026-01-14T12:43:22.729113Z","shell.execute_reply.started":"2026-01-14T12:43:22.716439Z","shell.execute_reply":"2026-01-14T12:43:22.728377Z"}},"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-01-14T12:43:22.730295Z","iopub.execute_input":"2026-01-14T12:43:22.730599Z","iopub.status.idle":"2026-01-14T12:43:30.212377Z","shell.execute_reply.started":"2026-01-14T12:43:22.730565Z","shell.execute_reply":"2026-01-14T12:43:30.211310Z"}},"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-01-14T12:43:30.216202Z","iopub.execute_input":"2026-01-14T12:43:30.216579Z","iopub.status.idle":"2026-01-14T12:43:30.232688Z","shell.execute_reply.started":"2026-01-14T12:43:30.216550Z","shell.execute_reply":"2026-01-14T12:43:30.231895Z"}},"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-01-14T12:43:30.234095Z","iopub.execute_input":"2026-01-14T12:43:30.234476Z","iopub.status.idle":"2026-01-14T12:43:30.275861Z","shell.execute_reply.started":"2026-01-14T12:43:30.234433Z","shell.execute_reply":"2026-01-14T12:43:30.274858Z"}},"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-01-14T12:43:30.277035Z","iopub.execute_input":"2026-01-14T12:43:30.277395Z","iopub.status.idle":"2026-01-14T12:46:19.185406Z","shell.execute_reply.started":"2026-01-14T12:43:30.277360Z","shell.execute_reply":"2026-01-14T12:46:19.184506Z"}},"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-01-14T12:46:19.186499Z","iopub.execute_input":"2026-01-14T12:46:19.186825Z","iopub.status.idle":"2026-01-14T12:46:19.201597Z","shell.execute_reply.started":"2026-01-14T12:46:19.186799Z","shell.execute_reply":"2026-01-14T12:46:19.200761Z"}},"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-01-14T12:46:19.203753Z","iopub.execute_input":"2026-01-14T12:46:19.204336Z","iopub.status.idle":"2026-01-14T12:46:19.217049Z","shell.execute_reply.started":"2026-01-14T12:46:19.204307Z","shell.execute_reply":"2026-01-14T12:46:19.216146Z"}},"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-01-14T12:46:19.218120Z","iopub.execute_input":"2026-01-14T12:46:19.218414Z","iopub.status.idle":"2026-01-14T12:46:19.238867Z","shell.execute_reply.started":"2026-01-14T12:46:19.218391Z","shell.execute_reply":"2026-01-14T12:46:19.237978Z"}},"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-01-14T12:46:19.240130Z","iopub.execute_input":"2026-01-14T12:46:19.240430Z","iopub.status.idle":"2026-01-14T12:46:19.252933Z","shell.execute_reply.started":"2026-01-14T12:46:19.240407Z","shell.execute_reply":"2026-01-14T12:46:19.252129Z"}},"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}\n\nbest_val_loss = float('inf')\nbest_map = 0.0\npatience_counter = 0\nbest_epoch = 0\n\nprint(f\"Starting training for {config['num_epochs']} epochs...\")\nprint(\"=\" * 70)\n\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\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\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        best_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            'config': config,\n            'label_encoder_classes': label_encoder.classes_.tolist(),\n            'num_classes': num_classes,\n        }, checkpoint_path)\n        \n        print(f\"  [New best model saved]\")\n    else:\n        patience_counter += 1\n        print(f\"  No improvement. Patience: {patience_counter}/{config['patience']}\")\n    \n    # Early stopping\n    if patience_counter >= config['patience']:\n        print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n        break\n\nprint(\"\\n\" + \"=\" * 70)\nprint(f\"Training complete!\")\nprint(f\"Best epoch: {best_epoch} (Val Loss: {best_val_loss:.4f}, Val mAP: {best_map:.4f})\")\n\n# Log best metrics as W&B summary for easy comparison across runs\nwandb.run.summary[\"best_val_mAP\"] = best_map\nwandb.run.summary[\"best_val_loss\"] = best_val_loss\nwandb.run.summary[\"best_epoch\"] = best_epoch\nwandb.run.summary[\"total_epochs\"] = len(history['train_loss'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T12:46:19.253872Z","iopub.execute_input":"2026-01-14T12:46:19.254143Z","iopub.status.idle":"2026-01-14T12:46:39.275938Z","shell.execute_reply.started":"2026-01-14T12:46:19.254119Z","shell.execute_reply":"2026-01-14T12:46:39.275125Z"}},"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-01-14T12:46:39.277092Z","iopub.execute_input":"2026-01-14T12:46:39.277587Z","iopub.status.idle":"2026-01-14T12:46:40.568164Z","shell.execute_reply.started":"2026-01-14T12:46:39.277549Z","shell.execute_reply":"2026-01-14T12:46:40.567476Z"}},"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-01-14T12:46:40.569187Z","iopub.execute_input":"2026-01-14T12:46:40.569481Z","iopub.status.idle":"2026-01-14T12:46:40.595704Z","shell.execute_reply.started":"2026-01-14T12:46:40.569455Z","shell.execute_reply":"2026-01-14T12:46:40.594924Z"}},"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-01-14T12:46:40.596650Z","iopub.execute_input":"2026-01-14T12:46:40.597361Z","iopub.status.idle":"2026-01-14T12:46:40.611359Z","shell.execute_reply.started":"2026-01-14T12:46:40.597333Z","shell.execute_reply":"2026-01-14T12:46:40.610452Z"}},"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-01-14T12:46:40.612713Z","iopub.execute_input":"2026-01-14T12:46:40.613146Z","iopub.status.idle":"2026-01-14T12:46:48.035944Z","shell.execute_reply.started":"2026-01-14T12:46:40.613112Z","shell.execute_reply":"2026-01-14T12:46:48.034885Z"}},"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-01-14T12:46:48.037569Z","iopub.execute_input":"2026-01-14T12:46:48.037849Z","iopub.status.idle":"2026-01-14T12:46:48.061949Z","shell.execute_reply.started":"2026-01-14T12:46:48.037822Z","shell.execute_reply":"2026-01-14T12:46:48.060998Z"}},"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-01-14T12:46:48.063131Z","iopub.execute_input":"2026-01-14T12:46:48.063452Z","iopub.status.idle":"2026-01-14T12:46:48.093559Z","shell.execute_reply.started":"2026-01-14T12:46:48.063420Z","shell.execute_reply":"2026-01-14T12:46:48.092531Z"}},"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-01-14T12:46:48.094622Z","iopub.execute_input":"2026-01-14T12:46:48.094975Z","iopub.status.idle":"2026-01-14T12:46:48.661812Z","shell.execute_reply.started":"2026-01-14T12:46:48.094939Z","shell.execute_reply":"2026-01-14T12:46:48.661139Z"}},"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-01-14T12:46:48.662804Z","iopub.execute_input":"2026-01-14T12:46:48.663239Z","iopub.status.idle":"2026-01-14T12:46:48.762310Z","shell.execute_reply.started":"2026-01-14T12:46:48.663198Z","shell.execute_reply":"2026-01-14T12:46:48.761494Z"}},"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-01-14T12:46:48.763383Z","iopub.execute_input":"2026-01-14T12:46:48.763680Z","iopub.status.idle":"2026-01-14T12:49:28.956035Z","shell.execute_reply.started":"2026-01-14T12:46:48.763639Z","shell.execute_reply":"2026-01-14T12:49:28.955212Z"}},"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-01-14T12:49:28.957065Z","iopub.execute_input":"2026-01-14T12:49:28.957417Z","iopub.status.idle":"2026-01-14T12:49:28.969508Z","shell.execute_reply.started":"2026-01-14T12:49:28.957386Z","shell.execute_reply":"2026-01-14T12:49:28.968779Z"}},"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-01-14T12:49:28.970493Z","iopub.execute_input":"2026-01-14T12:49:28.970721Z","iopub.status.idle":"2026-01-14T12:49:39.339164Z","shell.execute_reply.started":"2026-01-14T12:49:28.970700Z","shell.execute_reply":"2026-01-14T12:49:39.336428Z"}},"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-01-14T12:49:39.345472Z","iopub.execute_input":"2026-01-14T12:49:39.346290Z","iopub.status.idle":"2026-01-14T12:49:39.437645Z","shell.execute_reply.started":"2026-01-14T12:49:39.346261Z","shell.execute_reply":"2026-01-14T12:49:39.436908Z"}},"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)\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-01-14T12:49:39.438546Z","iopub.execute_input":"2026-01-14T12:49:39.438763Z","iopub.status.idle":"2026-01-14T12:49:39.623077Z","shell.execute_reply.started":"2026-01-14T12:49:39.438741Z","shell.execute_reply":"2026-01-14T12:49:39.622159Z"}},"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\nmodel_artifact = wandb.Artifact(\n    name=\"arcface-model\",\n    type=\"model\",\n    description=\"ArcFace fine-tuned MegaDescriptor model for jaguar re-identification\"\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\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T12:49:39.624174Z","iopub.execute_input":"2026-01-14T12:49:39.624524Z","iopub.status.idle":"2026-01-14T12:49:40.499062Z","shell.execute_reply.started":"2026-01-14T12:49:39.624485Z","shell.execute_reply":"2026-01-14T12:49:40.498121Z"}},"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-01-14T12:49:40.500103Z","iopub.execute_input":"2026-01-14T12:49:40.500392Z","iopub.status.idle":"2026-01-14T12:49:41.153356Z","shell.execute_reply.started":"2026-01-14T12:49:40.500367Z","shell.execute_reply":"2026-01-14T12:49:41.152510Z"}},"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-01-14T12:49:41.154486Z","iopub.execute_input":"2026-01-14T12:49:41.154846Z","iopub.status.idle":"2026-01-14T12:49:41.696570Z","shell.execute_reply.started":"2026-01-14T12:49:41.154796Z","shell.execute_reply":"2026-01-14T12:49:41.695810Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Summary\n\nThis notebook demonstrated a complete pipeline for jaguar re-identification:\n\n1. **Data Preparation**: Loaded training data and created a stratified train/validation split ensuring all identities appear in both sets.\n\n2. **Baseline Embeddings**: Extracted 1536-dimensional embeddings using MegaDescriptor-L-384, a vision transformer pre-trained for wildlife re-identification.\n\n3. **ArcFace Training**: Fine-tuned embeddings using ArcFace loss, which optimizes angular distances on a hypersphere. This encourages embeddings of the same jaguar to cluster together.\n\n4. **Visualization**: Used MDS to project embeddings to 2D, comparing baseline vs fine-tuned representations.\n\n5. **Submission**: Generated predictions by computing cosine similarity between fine-tuned embeddings for all test pairs.\n\n**Key Hyperparameters**:\n- ArcFace margin: 0.5 (adds 28.6 degrees angular penalty)\n- ArcFace scale: 64 (controls softmax sharpness)\n- Embedding dimension: 256 (projected from 1536)\n- Learning rate: 1e-4 with ReduceLROnPlateau scheduler\n\n**Next Steps**:\n- Experiment with different margins and scales\n- Try data augmentation during training\n- Ensemble multiple models\n- Fine-tune the MegaDescriptor backbone (more compute required)","metadata":{}}]}