{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"jupytext":{"cell_metadata_filter":"-all","notebook_metadata_filter":"kernelspec,language_info"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":129543,"databundleVersionId":15525987}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Jaguar Re-Identification Baseline with MiewID v3 and ArcFace\n\nThis notebook demonstrates a complete pipeline for training a jaguar re-identification model using **MiewID v3** backbone features (`conservationxlabs/miewid-msv3`) and an ArcFace projection head. 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. **MiewID v3**: Extract baseline features using a pre-trained re-identification backbone\n3. **Visualization**: Use MDS to visualize embeddings before and after fine-tuning\n4. **ArcFace Training**: Train a projection head with additive angular margin loss\n5. **Submission**: Generate predictions for the competition test set\n\n## Key Concepts\n\n**MiewID v3 (msv3)** is a wildlife re-identification feature extractor trained across many species. We use it as a frozen backbone to produce per-image feature vectors.\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 adapt the MiewID feature space for our specific jaguar dataset with lightweight training.\n","metadata":{}},{"cell_type":"markdown","source":"## Images have been stripped of background information","metadata":{}},{"cell_type":"code","source":"from PIL import Image\nImage.open(\"/kaggle/input/competitions/round-2-jaguar-reidentification-challenge/test/test_0014.png\").convert(\"RGB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:16.950559Z","iopub.execute_input":"2026-03-12T16:45:16.951269Z","iopub.status.idle":"2026-03-12T16:45:17.296248Z","shell.execute_reply.started":"2026-03-12T16:45:16.951231Z","shell.execute_reply":"2026-03-12T16:45:17.295010Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"Image.open(\"/kaggle/input/competitions/round-2-jaguar-reidentification-challenge/train/train_0003.png\").convert(\"RGB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:17.297822Z","iopub.execute_input":"2026-03-12T16:45:17.298156Z","iopub.status.idle":"2026-03-12T16:45:17.775370Z","shell.execute_reply.started":"2026-03-12T16:45:17.298125Z","shell.execute_reply":"2026-03-12T16:45:17.774318Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## GPU Acceleration\n\nMake sure to `Settings -> Accelerator -> GPU P100` to enable a CUDA GPU for this notebook. Instructions on enabling GPU usage on Kaggle are available here: https://github.com/andandandand/practical-computer-vision/blob/main/docs/kaggle-gpu-tpu-guide.md","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup and Configuration","metadata":{}},{"cell_type":"code","source":"!uv pip install wandb==0.25.0 \"transformers==4.45.2\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:17.776613Z","iopub.execute_input":"2026-03-12T16:45:17.777041Z","iopub.status.idle":"2026-03-12T16:45:21.003299Z","shell.execute_reply.started":"2026-03-12T16:45:17.776997Z","shell.execute_reply":"2026-03-12T16:45:21.002207Z"}},"outputs":[],"execution_count":null},{"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 transformers import AutoModel\n\n#from 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\n#env_path = Path(\"../../.env\")\n#if env_path.exists():\n#    load_dotenv(env_path)\n#    print(f\"Loaded environment variables from {env_path}\")\n#else:\n#    print(f\"Warning: {env_path} not found. Set WANDB_API_KEY and HF_TOKEN manually.\")\n\n# Obtain WANDB_API_KEY and HF_TOKEN from Kaggle secrets (optional)\n# - HF token is only needed for private models / higher rate limits\n# - W&B is optional: if WANDB_API_KEY is missing, we will run in disabled mode\ntry:\n    from kaggle_secrets import UserSecretsClient\n\n    user_secrets = UserSecretsClient()\n\n    def _get_secret(name: str):\n        try:\n            return user_secrets.get_secret(name)\n        except Exception:\n            return None\n\n    hf_token = _get_secret(\"hf_api\")\n    wandb_api_key = _get_secret(\"wandb_api\")\n\n    if hf_token:\n        os.environ.setdefault(\"HF_TOKEN\", hf_token)\n    if wandb_api_key:\n        os.environ.setdefault(\"WANDB_API_KEY\", wandb_api_key)\n\n    print(\"Kaggle secrets: loaded HF_TOKEN/WANDB_API_KEY (if configured).\")\nexcept Exception:\n    print(\"Kaggle secrets not available (running outside Kaggle or secrets not configured).\")\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":{"lines_to_next_cell":2,"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:21.005677Z","iopub.execute_input":"2026-03-12T16:45:21.006383Z","iopub.status.idle":"2026-03-12T16:45:39.553410Z","shell.execute_reply.started":"2026-03-12T16:45:21.006347Z","shell.execute_reply":"2026-03-12T16:45:39.552077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"wandb.__version__","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:39.554652Z","iopub.execute_input":"2026-03-12T16:45:39.555293Z","iopub.status.idle":"2026-03-12T16:45:39.561808Z","shell.execute_reply.started":"2026-03-12T16:45:39.555241Z","shell.execute_reply":"2026-03-12T16:45:39.560379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Device configuration\n# - Kaggle: use CUDA GPU (P100) if enabled in Notebook settings\n# - Local Apple Silicon: uses MPS if available\nif torch.cuda.is_available():\n    device = torch.device(\"cuda\")\n    print(\"Using CUDA GPU\")\n    gpu_name = torch.cuda.get_device_name(0)\n    total_mem_gb = torch.cuda.get_device_properties(0).total_memory / (1024**3)\n    print(f\"GPU: {gpu_name}\")\n    print(f\"GPU memory: {total_mem_gb:.1f} GB\")\n    if \"P100\" not in gpu_name:\n        print(\"Warning: GPU is not a P100. On Kaggle set Settings -> Accelerator -> GPU P100 for reproducibility.\")\nelif torch.backends.mps.is_available():\n    device = torch.device(\"mps\")\n    print(\"Using MPS (Apple Silicon GPU)\")\nelse:\n    device = torch.device(\"cpu\")\n    print(\"Using CPU\")\n    print(\"No CUDA GPU detected. On Kaggle set Settings -> Accelerator -> GPU P100.\")\n\nprint(f\"Device: {device}\")","metadata":{"lines_to_next_cell":2,"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:39.563233Z","iopub.execute_input":"2026-03-12T16:45:39.563739Z","iopub.status.idle":"2026-03-12T16:45:39.712886Z","shell.execute_reply.started":"2026-03-12T16:45:39.563691Z","shell.execute_reply":"2026-03-12T16:45:39.712068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Configuration\nconfig = {\n    # Paths\n    #\"data_dir\": Path(\"/kaggle/input/jaguar-re-id\"),\n    \"data_dir\": Path(\"/kaggle/input/competitions/round-2-jaguar-reidentification-challenge\"),\n    \"checkpoint_dir\": Path(\"checkpoints\"),\n\n    # Backbone (MiewID v3)\n    \"miewid_model_tag\": \"conservationxlabs/miewid-msv3\",\n    \"input_size\": 440,\n    \"extract_batch_size\": 16,\n\n    # Embedding head\n    \"embedding_dim\": 256,\n    \"hidden_dim\": 512,\n\n    # ArcFace\n    \"arcface_margin\": 0.1,\n    \"arcface_scale\": 16.0,\n    \"dropout\": 0.0,\n\n    # Training\n    \"batch_size\": 64,\n    \"learning_rate\": 3e-5,\n    \"weight_decay\": 5e-5,\n    \"num_epochs\": 50,\n    \"patience\": 15,\n    \"val_split\": 0.2,\n\n    # Stage 3: full fine-tuning (unfreeze all MiewID layers)\n    \"full_ft_batch_size\": 8,\n    \"full_ft_learning_rate\": 2e-6,\n    \"full_ft_weight_decay\": 1e-5,\n    \"full_ft_num_epochs\": 8,\n    \"full_ft_patience\": 3,\n    \"full_ft_extract_batch_size\": 8,\n    \"full_ft_checkpoint_name\": \"arcface_full_finetune_best.pth\",\n    \"full_ft_submission_name\": \"submission_full_fine_tuning.csv\",\n\n    # Stability / debugging controls\n    \"force_recompute_embeddings\": True,\n    \"strict_image_loading\": True,\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":{"lines_to_next_cell":2,"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:39.714315Z","iopub.execute_input":"2026-03-12T16:45:39.714705Z","iopub.status.idle":"2026-03-12T16:45:39.724502Z","shell.execute_reply.started":"2026-03-12T16:45:39.714664Z","shell.execute_reply":"2026-03-12T16:45:39.723318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize Weights and Biases for experiment tracking (optional)\n# If WANDB_API_KEY is missing, we run in disabled mode and all wandb.log calls become no-ops.\nwandb_api_key = os.getenv(\"WANDB_API_KEY\")\nwandb_mode = \"online\" if wandb_api_key else \"disabled\"\n\nif wandb_api_key:\n    wandb.login(key=wandb_api_key)\n\nwandb.init(\n    project=os.getenv(\"WANDB_PROJECT\", \"jaguar-reid-baseline\"),\n    config={\n        # Backbone\n        \"miewid_model_tag\": config[\"miewid_model_tag\"],\n        \"input_size\": config[\"input_size\"],\n        \"extract_batch_size\": config[\"extract_batch_size\"],\n\n        # Model head\n        \"embedding_dim\": config[\"embedding_dim\"],\n        \"hidden_dim\": config[\"hidden_dim\"],\n        \"dropout\": config[\"dropout\"],\n\n        # ArcFace hyperparameters\n        \"arcface_margin\": config[\"arcface_margin\"],\n        \"arcface_scale\": config[\"arcface_scale\"],\n\n        # Training\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        \"full_ft_batch_size\": config[\"full_ft_batch_size\"],\n        \"full_ft_learning_rate\": config[\"full_ft_learning_rate\"],\n        \"full_ft_weight_decay\": config[\"full_ft_weight_decay\"],\n        \"full_ft_num_epochs\": config[\"full_ft_num_epochs\"],\n        \"full_ft_patience\": config[\"full_ft_patience\"],\n        \"full_ft_extract_batch_size\": config[\"full_ft_extract_batch_size\"],\n        \"force_recompute_embeddings\": config[\"force_recompute_embeddings\"],\n        \"strict_image_loading\": config[\"strict_image_loading\"],\n        \"seed\": config[\"seed\"],\n    },\n    name=os.getenv(\"WANDB_RUN_NAME\", \"miewid-msv3-arcface-head\"),\n    mode=wandb_mode,\n)\n\nprint(f\"W&B initialized (mode={wandb_mode}). Key hyperparameters tracked:\")\nprint(f\"  Project: {os.getenv('WANDB_PROJECT', 'jaguar-reid-baseline')}\")\nprint(f\"  Backbone: {config['miewid_model_tag']} (input={config['input_size']}x{config['input_size']})\")\nprint(f\"  Extract batch size: {config['extract_batch_size']}\")\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":{"lines_to_next_cell":2,"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:45:39.725647Z","iopub.execute_input":"2026-03-12T16:45:39.725918Z","iopub.status.idle":"2026-03-12T16:45:59.614898Z","shell.execute_reply.started":"2026-03-12T16:45:39.725893Z","shell.execute_reply":"2026-03-12T16:45:59.612946Z"}},"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-12T16:45:59.616467Z","iopub.execute_input":"2026-03-12T16:45:59.617067Z","iopub.status.idle":"2026-03-12T16:45:59.681140Z","shell.execute_reply.started":"2026-03-12T16:45:59.617022Z","shell.execute_reply":"2026-03-12T16:45:59.679103Z"}},"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-12T16:45:59.686767Z","iopub.execute_input":"2026-03-12T16:45:59.687405Z","iopub.status.idle":"2026-03-12T16:46:01.639537Z","shell.execute_reply.started":"2026-03-12T16:45:59.687353Z","shell.execute_reply":"2026-03-12T16:46:01.638577Z"}},"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-12T16:46:01.640736Z","iopub.execute_input":"2026-03-12T16:46:01.641098Z","iopub.status.idle":"2026-03-12T16:46:04.329322Z","shell.execute_reply.started":"2026-03-12T16:46:01.641053Z","shell.execute_reply":"2026-03-12T16:46:04.328534Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Load MiewID v3 Model\n\nMiewID-msv3 is a feature extractor trained for wildlife re-identification using contrastive learning across many species. It uses an EfficientNetV2 backbone and is typically run with 440x440 inputs.\n\nWe load the pre-trained checkpoint from Hugging Face using `transformers.AutoModel` with `trust_remote_code=True`.\n","metadata":{}},{"cell_type":"code","source":"# Load MiewID model\nprint(f\"Loading MiewID model: {config['miewid_model_tag']}...\")\nmiewid = AutoModel.from_pretrained(\n    config[\"miewid_model_tag\"],\n    trust_remote_code=True,\n)\nmiewid.eval()\nmiewid.to(device)\n\nprint(f\"Model loaded successfully\")\nprint(f\"  Parameters: {sum(p.numel() for p in miewid.parameters()):,}\")\n\n# Get the feature 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 = miewid(dummy_input)\n\n    # Be defensive in case a future checkpoint returns a ModelOutput/tuple\n    if hasattr(dummy_output, \"last_hidden_state\"):\n        dummy_output = dummy_output.last_hidden_state\n    if isinstance(dummy_output, (tuple, list)):\n        dummy_output = dummy_output[0]\n\n    backbone_dim = dummy_output.shape[1]\n    print(f\"  Feature dimension: {backbone_dim}\")","metadata":{"lines_to_next_cell":2,"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:46:04.330394Z","iopub.execute_input":"2026-03-12T16:46:04.330918Z","iopub.status.idle":"2026-03-12T16:46:12.790751Z","shell.execute_reply.started":"2026-03-12T16:46:04.330885Z","shell.execute_reply":"2026-03-12T16:46:12.789700Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define preprocessing pipeline\n# MiewID-msv3 is typically run with 440x440 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-03-12T16:46:12.792099Z","iopub.execute_input":"2026-03-12T16:46:12.792511Z","iopub.status.idle":"2026-03-12T16:46:12.804920Z","shell.execute_reply.started":"2026-03-12T16:46:12.792465Z","shell.execute_reply":"2026-03-12T16:46:12.803999Z"}},"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 the frozen backbone.\"\"\"\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                if config.get(\"strict_image_loading\", True):\n                    raise RuntimeError(f\"Failed to load image: {path}\") from e\n                print(f\"Error loading {path}: {e}\")\n                # Fallback only when strict loading is disabled\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-12T16:46:12.806329Z","iopub.execute_input":"2026-03-12T16:46:12.806733Z","iopub.status.idle":"2026-03-12T16:46:12.861634Z","shell.execute_reply.started":"2026-03-12T16:46:12.806689Z","shell.execute_reply":"2026-03-12T16:46:12.860573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"emb_dir = Path(\"/kaggle/working/embeddings\")\nemb_dir.mkdir(parents=True, exist_ok=True)\n\nmodel_slug = config[\"miewid_model_tag\"].replace(\"/\", \"-\")\ncache_path = emb_dir / f\"baseline_train_embeddings_{model_slug}_{config['input_size']}.npz\"\n\n# Extract baseline embeddings for training data\ntrain_filenames = train_data[\"filename\"].astype(str).tolist()\ntrain_image_paths = [config[\"data_dir\"] / \"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 config.get(\"force_recompute_embeddings\", False) and cache_path.exists():\n    cache_path.unlink()\n    print(f\"Force recompute enabled. Removed cache: {cache_path}\")\n\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        miewid,\n        train_image_paths,\n        batch_size=config[\"extract_batch_size\"],\n        desc=\"Train embeddings\",\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-12T16:46:12.862869Z","iopub.execute_input":"2026-03-12T16:46:12.863251Z","iopub.status.idle":"2026-03-12T16:55:32.133253Z","shell.execute_reply.started":"2026-03-12T16:46:12.863209Z","shell.execute_reply":"2026-03-12T16:55:32.132282Z"}},"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 the MiewID backbone separates different jaguars before any fine-tuning.\n","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-12T16:55:32.134894Z","iopub.execute_input":"2026-03-12T16:55:32.135737Z","iopub.status.idle":"2026-03-12T16:55:32.149273Z","shell.execute_reply.started":"2026-03-12T16:55:32.135700Z","shell.execute_reply":"2026-03-12T16:55:32.148213Z"}},"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 MiewID Embeddings (Before Fine-tuning)\"\n)\nplt.show()\n\n# Log to W&B\nwandb.log({\"baseline_embeddings_miewid_mds\": wandb.Image(fig_baseline)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:55:32.150671Z","iopub.execute_input":"2026-03-12T16:55:32.151162Z","iopub.status.idle":"2026-03-12T16:55:41.193878Z","shell.execute_reply.started":"2026-03-12T16:55:32.151110Z","shell.execute_reply":"2026-03-12T16:55:41.192992Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Define Model Architecture\n\nWe define two components:\n\n1. **EmbeddingProjection**: Projects MiewID backbone features (auto-detected dimensionality) to a 256-d embedding space. 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\n","metadata":{}},{"cell_type":"code","source":"class EmbeddingProjection(nn.Module):\n    \"\"\"\n    Projects backbone features to a lower-dimensional space.\n    Architecture: input_dim -> hidden_dim -> output_dim\n    \"\"\"\n    \n    def __init__(self, input_dim, 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.1 radians, about 5.7 degrees)\n        - s is the feature scale (default 16)\n    \"\"\"\n    \n    def __init__(self, embedding_dim, num_classes, margin=0.1, scale=16.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-12T16:55:41.195244Z","iopub.execute_input":"2026-03-12T16:55:41.195565Z","iopub.status.idle":"2026-03-12T16:55:41.215084Z","shell.execute_reply.started":"2026-03-12T16:55:41.195536Z","shell.execute_reply":"2026-03-12T16:55:41.214080Z"}},"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.1, scale=16.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=backbone_dim,\n    num_classes=num_classes,\n    embedding_dim=config[\"embedding_dim\"],\n    hidden_dim=config[\"hidden_dim\"],\n    margin=config[\"arcface_margin\"],\n    scale=config[\"arcface_scale\"],\n    dropout=config[\"dropout\"],\n).to(device)\n\nprint(f\"ArcFace Model:\")\nprint(f\"  Input dim: {backbone_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-12T16:55:41.216295Z","iopub.execute_input":"2026-03-12T16:55:41.216677Z","iopub.status.idle":"2026-03-12T16:55:41.268615Z","shell.execute_reply.started":"2026-03-12T16:55:41.216633Z","shell.execute_reply":"2026-03-12T16:55:41.267726Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Prepare DataLoaders\n\nWe create PyTorch datasets from the pre-computed MiewID backbone embeddings. This is more efficient than loading images during training since embedding extraction is the bottleneck.\n","metadata":{}},{"cell_type":"code","source":"# Extract embeddings for validation set (with caching)\nmodel_slug = config[\"miewid_model_tag\"].replace(\"/\", \"-\")\ncache_path_val = emb_dir / f\"baseline_val_embeddings_{model_slug}_{config['input_size']}.npz\"\n\nval_filenames = val_data[\"filename\"].astype(str).tolist()\nval_image_paths = [config[\"data_dir\"] / \"train\" / fn for fn in val_filenames]\n\nbaseline_val_embeddings = None\nif config.get(\"force_recompute_embeddings\", False) and cache_path_val.exists():\n    cache_path_val.unlink()\n    print(f\"Force recompute enabled. Removed cache: {cache_path_val}\")\n\nif cache_path_val.exists():\n    baseline_val_embeddings = _load_cached_embeddings(cache_path_val, val_filenames)\n    if baseline_val_embeddings is not None:\n        print(f\"Loaded cached validation embeddings from {cache_path_val}\")\n        print(f\"Validation embeddings shape: {baseline_val_embeddings.shape}\")\n\nif baseline_val_embeddings is None:\n    print(f\"Extracting embeddings for {len(val_image_paths)} validation images...\")\n    baseline_val_embeddings = extract_embeddings(\n        miewid,\n        val_image_paths,\n        batch_size=config[\"extract_batch_size\"],\n        desc=\"Val embeddings\",\n    )\n    np.savez_compressed(\n        cache_path_val,\n        embeddings=baseline_val_embeddings,\n        filenames=np.array(val_filenames, dtype=object),\n    )\n    print(f\"Saved validation embeddings cache to {cache_path_val}\")\n    print(f\"Validation embeddings shape: {baseline_val_embeddings.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:55:41.269738Z","iopub.execute_input":"2026-03-12T16:55:41.270050Z","iopub.status.idle":"2026-03-12T16:58:03.611309Z","shell.execute_reply.started":"2026-03-12T16:55:41.270012Z","shell.execute_reply":"2026-03-12T16:58:03.610388Z"}},"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-12T16:58:03.613543Z","iopub.execute_input":"2026-03-12T16:58:03.613874Z","iopub.status.idle":"2026-03-12T16:58:03.627807Z","shell.execute_reply.started":"2026-03-12T16:58:03.613842Z","shell.execute_reply":"2026-03-12T16:58:03.626723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images, labels = next(iter(train_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:58:03.629236Z","iopub.execute_input":"2026-03-12T16:58:03.629644Z","iopub.status.idle":"2026-03-12T16:58:03.652728Z","shell.execute_reply.started":"2026-03-12T16:58:03.629606Z","shell.execute_reply":"2026-03-12T16:58:03.651517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:58:03.654197Z","iopub.execute_input":"2026-03-12T16:58:03.654543Z","iopub.status.idle":"2026-03-12T16:58:03.663357Z","shell.execute_reply.started":"2026-03-12T16:58:03.654513Z","shell.execute_reply":"2026-03-12T16:58:03.662528Z"}},"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-12T16:58:03.664449Z","iopub.execute_input":"2026-03-12T16:58:03.664758Z","iopub.status.idle":"2026-03-12T16:58:03.684572Z","shell.execute_reply.started":"2026-03-12T16:58:03.664727Z","shell.execute_reply":"2026-03-12T16:58:03.683169Z"}},"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-12T16:58:03.686145Z","iopub.execute_input":"2026-03-12T16:58:03.686869Z","iopub.status.idle":"2026-03-12T16:58:03.709196Z","shell.execute_reply.started":"2026-03-12T16:58:03.686819Z","shell.execute_reply":"2026-03-12T16:58:03.708068Z"}},"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-12T16:58:03.710393Z","iopub.execute_input":"2026-03-12T16:58:03.710693Z","iopub.status.idle":"2026-03-12T16:58:03.731694Z","shell.execute_reply.started":"2026-03-12T16:58:03.710664Z","shell.execute_reply":"2026-03-12T16:58:03.730454Z"}},"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-03-12T16:58:03.733037Z","iopub.execute_input":"2026-03-12T16:58:03.733335Z","iopub.status.idle":"2026-03-12T16:58:20.637539Z","shell.execute_reply.started":"2026-03-12T16:58:03.733307Z","shell.execute_reply":"2026-03-12T16:58:20.636723Z"}},"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-12T16:58:20.645466Z","iopub.execute_input":"2026-03-12T16:58:20.646596Z","iopub.status.idle":"2026-03-12T16:58:22.078903Z","shell.execute_reply.started":"2026-03-12T16:58:20.646534Z","shell.execute_reply":"2026-03-12T16:58:22.078163Z"}},"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-12T16:58:22.080175Z","iopub.execute_input":"2026-03-12T16:58:22.080453Z","iopub.status.idle":"2026-03-12T16:58:22.110507Z","shell.execute_reply.started":"2026-03-12T16:58:22.080427Z","shell.execute_reply":"2026-03-12T16:58:22.109527Z"}},"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-12T16:58:22.111750Z","iopub.execute_input":"2026-03-12T16:58:22.112130Z","iopub.status.idle":"2026-03-12T16:58:22.130088Z","shell.execute_reply.started":"2026-03-12T16:58:22.112082Z","shell.execute_reply":"2026-03-12T16:58:22.128841Z"}},"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_miewid_mds\": wandb.Image(fig_finetuned)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:58:22.131448Z","iopub.execute_input":"2026-03-12T16:58:22.131844Z","iopub.status.idle":"2026-03-12T16:58:30.765802Z","shell.execute_reply.started":"2026-03-12T16:58:22.131813Z","shell.execute_reply":"2026-03-12T16:58:30.765072Z"}},"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 MiewID backbone 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(MiewID)', \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-12T16:58:30.767144Z","iopub.execute_input":"2026-03-12T16:58:30.767462Z","iopub.status.idle":"2026-03-12T16:58:30.796262Z","shell.execute_reply.started":"2026-03-12T16:58:30.767432Z","shell.execute_reply":"2026-03-12T16:58:30.795137Z"}},"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 MiewID backbone)\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-12T16:58:30.797507Z","iopub.execute_input":"2026-03-12T16:58:30.797920Z","iopub.status.idle":"2026-03-12T16:58:30.828945Z","shell.execute_reply.started":"2026-03-12T16:58:30.797874Z","shell.execute_reply":"2026-03-12T16:58:30.827690Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example 1: Random validation image\nnp.random.seed(43)\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-12T16:58:30.830288Z","iopub.execute_input":"2026-03-12T16:58:30.830566Z","iopub.status.idle":"2026-03-12T16:58:38.014885Z","shell.execute_reply.started":"2026-03-12T16:58:30.830541Z","shell.execute_reply":"2026-03-12T16:58:38.013558Z"}},"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 MiewID backbone 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\n","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-12T16:58:38.016459Z","iopub.execute_input":"2026-03-12T16:58:38.016843Z","iopub.status.idle":"2026-03-12T16:58:38.149479Z","shell.execute_reply.started":"2026-03-12T16:58:38.016791Z","shell.execute_reply":"2026-03-12T16:58:38.148445Z"}},"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\" / filename for filename in test_images]\n\n# Extract MiewID backbone embeddings for test images (with caching)\nmodel_slug = config[\"miewid_model_tag\"].replace(\"/\", \"-\")\ncache_path_test = emb_dir / f\"baseline_test_embeddings_{model_slug}_{config['input_size']}.npz\"\n\nbaseline_test_embeddings = None\nif config.get(\"force_recompute_embeddings\", False) and cache_path_test.exists():\n    cache_path_test.unlink()\n    print(f\"Force recompute enabled. Removed cache: {cache_path_test}\")\n\nif cache_path_test.exists():\n    baseline_test_embeddings = _load_cached_embeddings(cache_path_test, test_images)\n    if baseline_test_embeddings is not None:\n        print(f\"Loaded cached test embeddings from {cache_path_test}\")\n        print(f\"Test embeddings shape: {baseline_test_embeddings.shape}\")\n\nif baseline_test_embeddings is None:\n    print(f\"\\nExtracting MiewID embeddings for test images...\")\n    baseline_test_embeddings = extract_embeddings(\n        miewid,\n        test_image_paths,\n        batch_size=config[\"extract_batch_size\"],\n        desc=\"Test embeddings\",\n    )\n    np.savez_compressed(\n        cache_path_test,\n        embeddings=baseline_test_embeddings,\n        filenames=np.array(test_images, dtype=object),\n    )\n    print(f\"Saved test embeddings cache to {cache_path_test}\")\n    print(f\"Test embeddings shape: {baseline_test_embeddings.shape}\")","metadata":{"lines_to_next_cell":2,"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T16:58:38.150622Z","iopub.execute_input":"2026-03-12T16:58:38.150897Z","iopub.status.idle":"2026-03-12T17:00:49.228432Z","shell.execute_reply.started":"2026-03-12T16:58:38.150871Z","shell.execute_reply":"2026-03-12T17:00:49.227400Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Project through fine-tuned model\nmodel.eval()\nwith torch.no_grad():\n    test_tensor = torch.FloatTensor(baseline_test_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-12T17:00:49.229601Z","iopub.execute_input":"2026-03-12T17:00:49.229932Z","iopub.status.idle":"2026-03-12T17:00:49.243829Z","shell.execute_reply.started":"2026-03-12T17:00:49.229885Z","shell.execute_reply":"2026-03-12T17:00:49.242748Z"}},"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-12T17:00:49.245145Z","iopub.execute_input":"2026-03-12T17:00:49.245516Z","iopub.status.idle":"2026-03-12T17:01:00.704564Z","shell.execute_reply.started":"2026-03-12T17:00:49.245474Z","shell.execute_reply":"2026-03-12T17:01:00.696845Z"}},"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(\"/kaggle/input/competitions/round-2-jaguar-reidentification-challenge/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-12T17:05:55.599311Z","iopub.execute_input":"2026-03-12T17:05:55.600428Z","iopub.status.idle":"2026-03-12T17:05:55.664144Z","shell.execute_reply.started":"2026-03-12T17:05:55.600369Z","shell.execute_reply":"2026-03-12T17:05:55.663024Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission\n# submission_path = save_path / \"submission.csv\"\n# submission_df.to_csv(submission_path, index=False)\n\n#print(f\"Submission saved to: {submission_path}\")\n#print(f\"File size: {submission_path.stat().st_size / 1024:.1f} KB\")","metadata":{"lines_to_next_cell":1,"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:06:15.446511Z","iopub.execute_input":"2026-03-12T17:06:15.446860Z","iopub.status.idle":"2026-03-12T17:06:15.641509Z","shell.execute_reply.started":"2026-03-12T17:06:15.446831Z","shell.execute_reply":"2026-03-12T17:06:15.640542Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Stage 3 — Full Fine-Tuning (All Layers Unfrozen)\n\nIn this third stage, we warm-start from the best ArcFace head checkpoint and then fine-tune **all MiewID layers** end-to-end.\n\nWe then:\n1. Recompute projected embeddings with the full fine-tuned model\n2. Analyze effects across **baseline vs head-only vs full fine-tuning**\n3. Generate a new submission file: `submission_full_fine_tuning.csv`","metadata":{}},{"cell_type":"code","source":"class ImageLabelDataset(Dataset):\n    \"\"\"Image dataset for end-to-end full fine-tuning.\"\"\"\n\n    def __init__(self, image_paths, labels, preprocess_fn):\n        self.image_paths = list(image_paths)\n        self.labels = np.array(labels)\n        self.preprocess_fn = preprocess_fn\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n        path = self.image_paths[idx]\n        try:\n            image = Image.open(path).convert(\"RGB\")\n            tensor = self.preprocess_fn(image)\n        except Exception as e:\n            if config.get(\"strict_image_loading\", True):\n                raise RuntimeError(f\"Failed to load image: {path}\") from e\n            tensor = torch.zeros(3, config[\"input_size\"], config[\"input_size\"])\n\n        label = int(self.labels[idx])\n        return tensor, torch.tensor(label, dtype=torch.long)\n\n\nclass EndToEndArcFaceModel(nn.Module):\n    \"\"\"MiewID backbone + projection head + ArcFace layer for full fine-tuning.\"\"\"\n\n    def __init__(self, backbone, projection_head, arcface_layer):\n        super().__init__()\n        self.backbone = backbone\n        self.embedding_net = projection_head\n        self.arcface = arcface_layer\n\n    def _extract_backbone_features(self, images):\n        features = self.backbone(images)\n\n        if hasattr(features, \"last_hidden_state\"):\n            features = features.last_hidden_state\n        if isinstance(features, (tuple, list)):\n            features = features[0]\n\n        if features.ndim == 4:\n            features = F.adaptive_avg_pool2d(features, output_size=1).flatten(1)\n        elif features.ndim == 3:\n            features = features.mean(dim=1)\n\n        return features\n\n    def forward(self, images, labels):\n        features = self._extract_backbone_features(images)\n        embeddings = self.embedding_net(features)\n        logits = self.arcface(embeddings, labels)\n        return logits, embeddings\n\n    def get_embeddings(self, images):\n        features = self._extract_backbone_features(images)\n        embeddings = self.embedding_net(features)\n        return F.normalize(embeddings, p=2, dim=1)\n\n\n@torch.no_grad()\ndef extract_projected_embeddings_from_paths(model, image_paths, batch_size=8, desc=\"Projecting embeddings\"):\n    \"\"\"Extract projected normalized embeddings from image paths.\"\"\"\n    model.eval()\n    outputs = []\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        batch_tensors = []\n        for path in batch_paths:\n            try:\n                image = Image.open(path).convert(\"RGB\")\n                batch_tensors.append(preprocess(image))\n            except Exception as e:\n                if config.get(\"strict_image_loading\", True):\n                    raise RuntimeError(f\"Failed to load image: {path}\") from e\n                batch_tensors.append(torch.zeros(3, config[\"input_size\"], config[\"input_size\"]))\n\n        batch_tensor = torch.stack(batch_tensors).to(device)\n        batch_embeddings = model.get_embeddings(batch_tensor).cpu().numpy()\n        outputs.append(batch_embeddings)\n\n    return np.vstack(outputs)\n\n\ndef compute_balanced_map_from_embeddings(embeddings, labels):\n    \"\"\"Compute identity-balanced mAP from embeddings and labels.\"\"\"\n    sim_matrix = cosine_similarity(embeddings)\n    np.fill_diagonal(sim_matrix, -1)\n\n    query_aps = {}\n    for query_idx in range(len(labels)):\n        query_label = labels[query_idx]\n\n        similarities = sim_matrix[query_idx]\n        gallery_labels = labels.copy()\n        is_match = (gallery_labels == query_label).astype(int)\n        is_match[query_idx] = 0\n\n        sorted_indices = np.argsort(-similarities)\n        sorted_matches = is_match[sorted_indices]\n\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        query_aps[query_idx] = (query_label, ap)\n\n    identity_aps = {}\n    for _, (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    identity_mean_aps = [np.mean(aps) for aps in identity_aps.values()]\n    return float(np.mean(identity_mean_aps))\n\n\ndef compute_topk_match_rate(embeddings, labels, k=5):\n    \"\"\"Compute top-1 and top-k match rates for retrieval analysis.\"\"\"\n    normalized = embeddings / np.linalg.norm(embeddings, axis=1, keepdims=True)\n    sim_matrix = normalized @ normalized.T\n    np.fill_diagonal(sim_matrix, -np.inf)\n\n    sorted_indices = np.argsort(-sim_matrix, axis=1)\n    top1_indices = sorted_indices[:, :1]\n    topk_indices = sorted_indices[:, :k]\n\n    labels_array = np.array(labels)\n    top1_matches = labels_array[top1_indices[:, 0]] == labels_array\n    topk_matches = (labels_array[topk_indices] == labels_array[:, None]).any(axis=1)\n\n    return float(top1_matches.mean()), float(topk_matches.mean())\n\n\ndef train_epoch_full(model, loader, criterion, optimizer, device):\n    \"\"\"Train one epoch for full fine-tuning on images.\"\"\"\n    model.train()\n    total_loss = 0.0\n    total = 0\n    correct = 0\n\n    pbar = tqdm(loader, desc=\"Full-FT Training\", leave=False)\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n\n        logits, _ = model(images, labels)\n        loss = criterion(logits, labels)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\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.0 * correct / max(total, 1):.1f}%\"})\n\n    return total_loss / len(loader), 100.0 * correct / max(total, 1)\n\n\ndef validate_epoch_full(model, loader, criterion, device):\n    \"\"\"Validate one epoch for full fine-tuning on images.\"\"\"\n    model.eval()\n    total_loss = 0.0\n    total = 0\n    correct = 0\n\n    with torch.no_grad():\n        pbar = tqdm(loader, desc=\"Full-FT Validation\", leave=False)\n        for images, labels in pbar:\n            images, labels = images.to(device), labels.to(device)\n\n            logits, _ = model(images, 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.0 * correct / max(total, 1):.1f}%\"})\n\n    return total_loss / len(loader), 100.0 * correct / max(total, 1)\n\n\nprint(\"Stage 3 helper classes/functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:01.106652Z","iopub.execute_input":"2026-03-12T17:01:01.106923Z","iopub.status.idle":"2026-03-12T17:01:01.139966Z","shell.execute_reply.started":"2026-03-12T17:01:01.106897Z","shell.execute_reply":"2026-03-12T17:01:01.138766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Build image-level datasets and dataloaders\ntrain_image_paths_full = [config[\"data_dir\"] / \"train\" / fn for fn in train_data[\"filename\"].astype(str).tolist()]\nval_image_paths_full = [config[\"data_dir\"] / \"train\" / fn for fn in val_data[\"filename\"].astype(str).tolist()]\n\ntrain_labels_full = train_data[\"label_encoded\"].values\nval_labels_full = val_data[\"label_encoded\"].values\n\ntrain_dataset_full = ImageLabelDataset(train_image_paths_full, train_labels_full, preprocess)\nval_dataset_full = ImageLabelDataset(val_image_paths_full, val_labels_full, preprocess)\n\ntrain_loader_full = DataLoader(\n    train_dataset_full,\n    batch_size=config[\"full_ft_batch_size\"],\n    shuffle=True,\n    num_workers=0,\n    pin_memory=False,\n)\nval_loader_full = DataLoader(\n    val_dataset_full,\n    batch_size=config[\"full_ft_batch_size\"],\n    shuffle=False,\n    num_workers=0,\n    pin_memory=False,\n)\n\n# Initialize from the best stage-2 ArcFace checkpoint\nhead_init_model = ArcFaceModel(\n    input_dim=backbone_dim,\n    num_classes=num_classes,\n    embedding_dim=config[\"embedding_dim\"],\n    hidden_dim=config[\"hidden_dim\"],\n    margin=config[\"arcface_margin\"],\n    scale=config[\"arcface_scale\"],\n    dropout=config[\"dropout\"],\n).to(device)\nhead_init_model.load_state_dict(checkpoint[\"model_state_dict\"])\n\nfull_model = EndToEndArcFaceModel(\n    backbone=miewid,\n    projection_head=head_init_model.embedding_net,\n    arcface_layer=head_init_model.arcface,\n).to(device)\n\nfor param in full_model.backbone.parameters():\n    param.requires_grad = True\n\ncriterion_full = nn.CrossEntropyLoss()\noptimizer_full = torch.optim.AdamW(\n    full_model.parameters(),\n    lr=config[\"full_ft_learning_rate\"],\n    weight_decay=config[\"full_ft_weight_decay\"],\n)\nscheduler_full = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer_full,\n    mode=\"min\",\n    factor=0.5,\n    patience=2,\n)\n\nprint(\"Full fine-tuning setup ready:\")\nprint(f\"  Train batches: {len(train_loader_full)}\")\nprint(f\"  Val batches: {len(val_loader_full)}\")\nprint(f\"  Batch size: {config['full_ft_batch_size']}\")\nprint(f\"  LR: {config['full_ft_learning_rate']}\")\nprint(f\"  Weight decay: {config['full_ft_weight_decay']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:01.141337Z","iopub.execute_input":"2026-03-12T17:01:01.141706Z","iopub.status.idle":"2026-03-12T17:01:01.210042Z","shell.execute_reply.started":"2026-03-12T17:01:01.141673Z","shell.execute_reply":"2026-03-12T17:01:01.209070Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Stage 3 training loop (all layers unfrozen)\nhistory_full = {\n    \"train_loss\": [],\n    \"train_acc\": [],\n    \"val_loss\": [],\n    \"val_acc\": [],\n    \"val_map\": [],\n    \"lr\": [],\n}\n\nfull_best_val_loss = float(\"inf\")\nfull_best_map = 0.0\nfull_best_epoch = 0\nfull_patience_counter = 0\nfull_checkpoint_path = config[\"checkpoint_dir\"] / config[\"full_ft_checkpoint_name\"]\n\nprint(f\"Starting full fine-tuning for {config['full_ft_num_epochs']} epochs...\")\nprint(\"=\" * 70)\n\nfor epoch in range(config[\"full_ft_num_epochs\"]):\n    print(f\"\\n[Stage 3] Epoch {epoch + 1}/{config['full_ft_num_epochs']}\")\n\n    train_loss_full, train_acc_full = train_epoch_full(\n        full_model,\n        train_loader_full,\n        criterion_full,\n        optimizer_full,\n        device,\n    )\n    val_loss_full, val_acc_full = validate_epoch_full(\n        full_model,\n        val_loader_full,\n        criterion_full,\n        device,\n    )\n\n    val_full_epoch_embeddings = extract_projected_embeddings_from_paths(\n        full_model,\n        val_image_paths_full,\n        batch_size=config[\"full_ft_extract_batch_size\"],\n        desc=\"Stage 3 val embeddings\",\n    )\n    val_map_full = compute_balanced_map_from_embeddings(\n        val_full_epoch_embeddings,\n        val_data[\"ground_truth\"].values,\n    )\n\n    scheduler_full.step(val_loss_full)\n    current_lr_full = optimizer_full.param_groups[0][\"lr\"]\n\n    history_full[\"train_loss\"].append(train_loss_full)\n    history_full[\"train_acc\"].append(train_acc_full)\n    history_full[\"val_loss\"].append(val_loss_full)\n    history_full[\"val_acc\"].append(val_acc_full)\n    history_full[\"val_map\"].append(val_map_full)\n    history_full[\"lr\"].append(current_lr_full)\n\n    wandb.log({\n        \"full_ft_epoch\": epoch + 1,\n        \"full_ft_train_loss\": train_loss_full,\n        \"full_ft_train_acc\": train_acc_full,\n        \"full_ft_val_loss\": val_loss_full,\n        \"full_ft_val_acc\": val_acc_full,\n        \"full_ft_val_map\": val_map_full,\n        \"full_ft_learning_rate\": current_lr_full,\n    })\n\n    print(f\"  Train Loss: {train_loss_full:.4f} | Train Acc: {train_acc_full:.1f}%\")\n    print(f\"  Val Loss:   {val_loss_full:.4f} | Val Acc:   {val_acc_full:.1f}%\")\n    print(f\"  Val mAP:    {val_map_full:.4f} | LR: {current_lr_full:.2e}\")\n\n    if val_loss_full < full_best_val_loss:\n        full_best_val_loss = val_loss_full\n        full_best_map = val_map_full\n        full_best_epoch = epoch + 1\n        full_patience_counter = 0\n\n        torch.save(\n            {\n                \"epoch\": epoch + 1,\n                \"model_state_dict\": full_model.state_dict(),\n                \"optimizer_state_dict\": optimizer_full.state_dict(),\n                \"val_loss\": val_loss_full,\n                \"val_map\": val_map_full,\n                \"config\": config,\n                \"num_classes\": num_classes,\n            },\n            full_checkpoint_path,\n        )\n        print(f\"  [New best full-finetune checkpoint saved to {full_checkpoint_path.name}]\")\n    else:\n        full_patience_counter += 1\n        print(f\"  No improvement. Patience: {full_patience_counter}/{config['full_ft_patience']}\")\n\n    if full_patience_counter >= config[\"full_ft_patience\"]:\n        print(f\"\\nStage 3 early stopping triggered after {epoch + 1} epochs\")\n        break\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"Stage 3 training complete!\")\nprint(\n    f\"Best Stage 3 epoch: {full_best_epoch} \"\n    f\"(Val Loss: {full_best_val_loss:.4f}, Val mAP: {full_best_map:.4f})\"\n)\n\nfull_checkpoint = torch.load(full_checkpoint_path, map_location=device, weights_only=False)\nfull_model.load_state_dict(full_checkpoint[\"model_state_dict\"])\nfull_model.eval()\n\nwandb.run.summary[\"full_ft_best_val_mAP\"] = full_best_map\nwandb.run.summary[\"full_ft_best_val_loss\"] = full_best_val_loss\nwandb.run.summary[\"full_ft_best_epoch\"] = full_best_epoch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:01.211431Z","iopub.execute_input":"2026-03-12T17:01:01.211842Z","iopub.status.idle":"2026-03-12T17:01:39.036505Z","shell.execute_reply.started":"2026-03-12T17:01:01.211790Z","shell.execute_reply":"2026-03-12T17:01:39.035124Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Re-extract embeddings with full fine-tuned model\nfull_finetuned_train_embeddings = extract_projected_embeddings_from_paths(\n    full_model,\n    train_image_paths_full,\n    batch_size=config[\"full_ft_extract_batch_size\"],\n    desc=\"Stage 3 train embeddings\",\n)\nval_full_finetuned_embeddings = extract_projected_embeddings_from_paths(\n    full_model,\n    val_image_paths_full,\n    batch_size=config[\"full_ft_extract_batch_size\"],\n    desc=\"Stage 3 val embeddings\",\n)\n\nprint(f\"Full fine-tuned train embeddings shape: {full_finetuned_train_embeddings.shape}\")\nprint(f\"Full fine-tuned val embeddings shape:   {val_full_finetuned_embeddings.shape}\")\nprint(\n    f\"Full fine-tuned train mean L2 norm: \"\n    f\"{np.linalg.norm(full_finetuned_train_embeddings, axis=1).mean():.4f}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:39.037339Z","iopub.status.idle":"2026-03-12T17:01:39.037670Z","shell.execute_reply.started":"2026-03-12T17:01:39.037514Z","shell.execute_reply":"2026-03-12T17:01:39.037533Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize full fine-tuned embeddings\nfig_full_finetuned = visualize_embeddings_mds(\n    full_finetuned_train_embeddings,\n    train_labels,\n    \"Full Fine-Tuned MiewID + ArcFace Embeddings (All Layers Unfrozen)\",\n)\nplt.show()\n\nwandb.log({\"full_finetuned_embeddings_miewid_mds\": wandb.Image(fig_full_finetuned)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:39.039116Z","iopub.status.idle":"2026-03-12T17:01:39.039439Z","shell.execute_reply.started":"2026-03-12T17:01:39.039292Z","shell.execute_reply":"2026-03-12T17:01:39.039310Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Analyze effects: baseline vs head-only vs full fine-tuning\nif \"val_finetuned_embeddings\" not in globals():\n    model.eval()\n    with 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\nval_labels_text = val_data[\"ground_truth\"].values\n\nbaseline_map = compute_balanced_map_from_embeddings(baseline_val_embeddings, val_labels_text)\nhead_map = compute_balanced_map_from_embeddings(val_finetuned_embeddings, val_labels_text)\nfull_map = compute_balanced_map_from_embeddings(val_full_finetuned_embeddings, val_labels_text)\n\nbaseline_top1, baseline_top5 = compute_topk_match_rate(baseline_val_embeddings, val_labels_text, k=5)\nhead_top1, head_top5 = compute_topk_match_rate(val_finetuned_embeddings, val_labels_text, k=5)\nfull_top1, full_top5 = compute_topk_match_rate(val_full_finetuned_embeddings, val_labels_text, k=5)\n\neffects_df = pd.DataFrame(\n    {\n        \"stage\": [\"baseline\", \"head_only_finetune\", \"full_finetune\"],\n        \"balanced_mAP\": [baseline_map, head_map, full_map],\n        \"top1_match_rate\": [baseline_top1, head_top1, full_top1],\n        \"top5_match_rate\": [baseline_top5, head_top5, full_top5],\n    }\n)\n\nprint(\"Three-way validation effects analysis:\")\nprint(effects_df.to_string(index=False))\nprint(\"\\nDeltas:\")\nprint(f\"  head - baseline mAP: {head_map - baseline_map:+.4f}\")\nprint(f\"  full - head mAP:     {full_map - head_map:+.4f}\")\nprint(f\"  full - baseline mAP: {full_map - baseline_map:+.4f}\")\n\nwandb.log(\n    {\n        \"effects_three_stage\": wandb.Table(dataframe=effects_df),\n        \"map_baseline\": baseline_map,\n        \"map_head_only\": head_map,\n        \"map_full_finetune\": full_map,\n        \"map_delta_full_vs_head\": full_map - head_map,\n    }\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:39.041381Z","iopub.status.idle":"2026-03-12T17:01:39.041801Z","shell.execute_reply.started":"2026-03-12T17:01:39.041579Z","shell.execute_reply":"2026-03-12T17:01:39.041599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate full fine-tuning submission\ntest_full_finetuned_embeddings = extract_projected_embeddings_from_paths(\n    full_model,\n    test_image_paths,\n    batch_size=config[\"full_ft_extract_batch_size\"],\n    desc=\"Stage 3 test embeddings\",\n)\n\nimg_to_full_embedding = {\n    filename: embedding\n    for filename, embedding in zip(test_images, test_full_finetuned_embeddings)\n}\n\nsimilarities_full = []\nfor _, row in tqdm(test_pairs_df.iterrows(), total=len(test_pairs_df), desc=\"Computing full-FT similarities\"):\n    query_emb = img_to_full_embedding[row[\"query_image\"]]\n    gallery_emb = img_to_full_embedding[row[\"gallery_image\"]]\n    similarities_full.append(np.dot(query_emb, gallery_emb))\n\nsimilarities_full = np.array(similarities_full)\nsimilarities_full = np.clip(similarities_full, 0.0, 1.0)\n\nsubmission_full_df = pd.DataFrame(\n    {\n        \"row_id\": test_pairs_df[\"row_id\"],\n        \"similarity\": similarities_full,\n    }\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:03:49.641822Z","iopub.execute_input":"2026-03-12T17:03:49.642445Z","iopub.status.idle":"2026-03-12T17:05:55.590832Z","shell.execute_reply.started":"2026-03-12T17:03:49.642401Z","shell.execute_reply":"2026-03-12T17:05:55.585176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"save_path = Path(\"/kaggle/working\")\n\nsubmission_full_path = save_path / \"submission.csv\"#config[\"full_ft_submission_name\"]\nsubmission_full_df.to_csv(submission_full_path, index=False)\n\nprint(f\"Full fine-tuning submission saved to: {submission_full_path}\")\nprint(f\"File size: {submission_full_path.stat().st_size / 1024:.1f} KB\")\nprint(f\"Min similarity: {similarities_full.min():.4f}\")\nprint(f\"Max similarity: {similarities_full.max():.4f}\")\nprint(f\"Mean similarity: {similarities_full.mean():.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:08:44.675704Z","iopub.execute_input":"2026-03-12T17:08:44.676114Z","iopub.status.idle":"2026-03-12T17:08:44.878465Z","shell.execute_reply.started":"2026-03-12T17:08:44.676083Z","shell.execute_reply":"2026-03-12T17:08:44.877502Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. 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 artifacts as W&B artifacts\nmodel_artifact_head = wandb.Artifact(\n    name=\"arcface-model-miewid-msv3-head-only\",\n    type=\"model\",\n    description=\"ArcFace projection head trained on frozen MiewID-msv3 backbone features\"\n)\nmodel_artifact_head.add_file(str(config[\"checkpoint_dir\"] / \"arcface_best.pth\"))\nwandb.log_artifact(model_artifact_head)\n\nprint(\"Head-only model artifact saved to W&B\")\n\nif \"full_checkpoint_path\" in globals() and Path(full_checkpoint_path).exists():\n    model_artifact_full = wandb.Artifact(\n        name=\"arcface-model-miewid-msv3-full-finetune\",\n        type=\"model\",\n        description=\"End-to-end full fine-tuned MiewID-msv3 + ArcFace model\"\n    )\n    model_artifact_full.add_file(str(full_checkpoint_path))\n    wandb.log_artifact(model_artifact_full)\n    print(\"Full fine-tuned model artifact saved to W&B\")\nelse:\n    print(\"Full fine-tuned model artifact skipped (checkpoint not found)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:39.045394Z","iopub.status.idle":"2026-03-12T17:01:39.045698Z","shell.execute_reply.started":"2026-03-12T17:01:39.045557Z","shell.execute_reply":"2026-03-12T17:01:39.045575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission artifacts as W&B artifacts\nsubmission_artifact_head = wandb.Artifact(\n    name=\"submission-miewid-msv3-head-only\",\n    type=\"submission\",\n    description=\"Competition submission from ArcFace head-only fine-tuning\"\n)\nsubmission_artifact_head.add_file(str(submission_path))\nwandb.log_artifact(submission_artifact_head)\n\nprint(\"Head-only submission artifact saved to W&B\")\n\nif \"submission_full_path\" in globals() and Path(submission_full_path).exists():\n    submission_artifact_full = wandb.Artifact(\n        name=\"submission-miewid-msv3-full-finetune\",\n        type=\"submission\",\n        description=\"Competition submission from full end-to-end fine-tuning\"\n    )\n    submission_artifact_full.add_file(str(submission_full_path))\n    wandb.log_artifact(submission_artifact_full)\n    print(\"Full fine-tuned submission artifact saved to W&B\")\nelse:\n    print(\"Full fine-tuned submission artifact skipped (file not found)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T17:01:39.047219Z","iopub.status.idle":"2026-03-12T17:01:39.047642Z","shell.execute_reply.started":"2026-03-12T17:01:39.047440Z","shell.execute_reply":"2026-03-12T17:01:39.047474Z"}},"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-12T17:01:39.049111Z","iopub.status.idle":"2026-03-12T17:01:39.049542Z","shell.execute_reply.started":"2026-03-12T17:01:39.049349Z","shell.execute_reply":"2026-03-12T17:01:39.049379Z"}},"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 backbone features using MiewID v3 (`conservationxlabs/miewid-msv3`), a pre-trained wildlife re-identification model.\n\n3. **Head-Only ArcFace Training**: Trained a projection head using ArcFace loss while keeping MiewID frozen.\n\n4. **Full Fine-Tuning (Stage 3)**: Warm-started from the best head-only checkpoint, unfroze all MiewID layers, and optimized the full model end-to-end.\n\n5. **Three-Way Effect Analysis**: Compared baseline vs head-only vs full fine-tuning with balanced mAP and top-k retrieval metrics.\n\n6. **Submissions**: Generated both `submission.csv` (head-only) and `submission_full_fine_tuning.csv` (full fine-tuned model).\n\n**Key Hyperparameters**:\n- ArcFace margin: 0.1 (adds 5.7 degrees angular penalty)\n- ArcFace scale: 16 (controls softmax sharpness)\n- Embedding dimension: 256 (projected from MiewID backbone feature dim)\n- Head-only learning rate: 3e-5 with ReduceLROnPlateau scheduler\n- Full fine-tuning learning rate: 2e-6 (all layers unfrozen)\n\n**Next Steps**:\n- Tune full-finetuning epochs and LR schedule per GPU budget\n- Add stronger augmentations for robust identity invariance\n- Evaluate ensembling between head-only and full-finetuned submissions\n- Track per-identity gains/losses to target hard individuals\n","metadata":{}}]}