{"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"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":129543,"databundleVersionId":15525987,"isSourceIdPinned":false}],"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 DinoV2 Giant and Triplet Loss\n\nThis notebook demonstrates a complete pipeline for training a jaguar re-identification model using DinoV2 Giant embeddings and Triplet Loss. The goal is to learn embeddings that place images of the same jaguar close together and images of different jaguars far apart. However, we switched to the Adagrad optimizer, as it performed slightly better than AdamW in our experiments.\n\n## Overview\n\n1. **Data Loading**: Load training images and create a stratified train/validation split\n2. **DinoV2 Giant**: Extract baseline embeddings using a pre-trained vision transformer\n3. **Visualization**: Use MDS to visualize embeddings before and after fine-tuning\n4. **Triplet Loss Training**: Fine-tune embeddings using online hard triplet mining\n5. **Submission**: Generate predictions for the competition test set\n\n## Key Concepts\n\n**DinoV2 Giant** is a highly capable vision transformer. It produces 1536-dimensional embeddings that capture visual features.\n\n**Triplet Loss** is a metric learning technique that:\n- Selects triplets of (anchor, positive, negative) from each batch\n- Pulls anchor and positive (same identity) embeddings closer\n- Pushes anchor and negative (different identity) embeddings apart\n- Uses a margin to enforce a minimum gap between positive and negative distances\n\n**Online Hard Triplet Mining** selects the hardest triplets within each batch:\n- Hardest positive: the farthest same-identity sample from the anchor\n- Hardest negative: the closest different-identity sample to the anchor\n\nA **PK batch sampler** ensures each batch contains P identities with K samples each, guaranteeing enough same-identity pairs for triplet construction.\n\nThe combination allows us to fine-tune DinoV2 Giant for our specific jaguar dataset.","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup and Configuration","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom collections import Counter\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.manifold import MDS\nfrom sklearn.metrics.pairwise import cosine_similarity\nimport math\nimport wandb\nfrom dotenv import load_dotenv\n\n# Load environment variables from .env file\n# The .env file should contain: WANDB_API_KEY, WANDB_PROJECT, HF_TOKEN\nenv_path = Path(\"../../.env\")\nif env_path.exists():\n    load_dotenv(env_path)\n    print(f\"Loaded environment variables from {env_path}\")\nelse:\n    print(f\"Warning: {env_path} not found. Set WANDB_API_KEY and HF_TOKEN manually.\")\n\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nos.environ[\"HF_TOKEN\"]= user_secrets.get_secret(\"hf_api\")\nos.environ[\"WANDB_API_KEY\"] = user_secrets.get_secret(\"wandb_api\")\n\n\n# Set random seeds for reproducibility\nRANDOM_SEED = 42\ntorch.manual_seed(RANDOM_SEED)\nnp.random.seed(RANDOM_SEED)\n\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"timm version: {timm.__version__}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:41:24.825715Z","iopub.execute_input":"2026-03-14T20:41:24.826135Z","iopub.status.idle":"2026-03-14T20:41:24.906696Z","shell.execute_reply.started":"2026-03-14T20:41:24.826103Z","shell.execute_reply":"2026-03-14T20:41:24.905919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Device configuration\n# MPS (Metal Performance Shaders) provides GPU acceleration on Apple Silicon (if you want to run this notebook locally on your MacBook)\nif torch.backends.mps.is_available():\n    device = torch.device(\"mps\")\n    print(\"Using MPS (Apple Silicon GPU)\")\nelif torch.cuda.is_available():\n    device = torch.device(\"cuda\")\n    print(\"Using CUDA GPU\")\nelse:\n    device = torch.device(\"cpu\")\n    print(\"Using CPU\")\n\nprint(f\"Device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:34:54.601037Z","iopub.execute_input":"2026-03-14T20:34:54.601481Z","iopub.status.idle":"2026-03-14T20:34:54.858156Z","shell.execute_reply.started":"2026-03-14T20:34:54.601456Z","shell.execute_reply":"2026-03-14T20:34:54.857340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Configuration\nconfig = {\n    # Environment\n    \"environment\": \"kaggle\",  # \"local\" or \"kaggle\"\n    \n    # Model\n    \"backbone_model\": \"vit_giant_patch14_dinov2.lvd142m\",\n    \"input_size\": 518,\n    \"embedding_dim\": 256,\n    \"hidden_dim\": 512,\n    \n    # Triplet Loss\n    \"triplet_margin\": 1.0,\n    \"mining_type\": \"batch_hard\",\n    \"dropout\": 0.3,\n    \n    # PK Sampling: P identities x K samples per identity per batch\n    \"P\": 8,   # number of identities per batch\n    \"K\": 4,   # number of samples per identity per batch\n    \n    # Training\n    \"learning_rate\": 5e-2,\n    \"weight_decay\": 1e-4,\n    \"num_epochs\": 300,\n    \"patience\": 20,\n    \"val_split\": 0.2,\n    \n    # Reproducibility\n    \"seed\": RANDOM_SEED,\n}\n\n# Effective batch size = P * K\nconfig[\"batch_size\"] = config[\"P\"] * config[\"K\"]\n\n# Set paths based on environment\nif config[\"environment\"] == \"kaggle\":\n    config[\"data_dir\"] = Path(\"/kaggle/input/competitions/round-2-jaguar-reidentification-challenge\")\n    config[\"checkpoint_dir\"] = Path(\"/kaggle/working\")\n    config[\"embed_dir\"] = Path(\"/kaggle/working/embeddings\")\nelse:  # local\n    config[\"data_dir\"] = Path(\"data\")\n    config[\"checkpoint_dir\"] = Path(\"checkpoints\")\n    config[\"embed_dir\"] = Path(\"embeddings\")\n\n# Create checkpoint directory\nconfig[\"checkpoint_dir\"].mkdir(exist_ok=True)\n\n# Create embedding directory\nconfig[\"embed_dir\"].mkdir(exist_ok=True)\n\nprint(\"Configuration:\")\nfor key, value in config.items():\n    print(f\"  {key}: {value}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:34:54.859195Z","iopub.execute_input":"2026-03-14T20:34:54.859556Z","iopub.status.idle":"2026-03-14T20:34:54.875420Z","shell.execute_reply.started":"2026-03-14T20:34:54.859523Z","shell.execute_reply":"2026-03-14T20:34:54.874701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize Weights and Biases for experiment tracking\n# Key hyperparameters are tracked explicitly for easy filtering in W&B dashboard\nwandb.login(key=os.environ[\"WANDB_API_KEY\"])\n\nwandb.init(\n    project=os.getenv(\"WANDB_PROJECT\", \"jaguar-reid-jojojaguar\"),\n    config={\n        # Model architecture\n        \"backbone_model\": config[\"backbone_model\"],\n        \"embedding_dim\": config[\"embedding_dim\"],\n        \"hidden_dim\": config[\"hidden_dim\"],\n        \"dropout\": config[\"dropout\"],\n        \n        # Triplet Loss hyperparameters\n        \"triplet_margin\": config[\"triplet_margin\"],\n        \"mining_type\": config[\"mining_type\"],\n        \"P\": config[\"P\"],\n        \"K\": config[\"K\"],\n        \n        # Training hyperparameters\n        \"batch_size\": config[\"batch_size\"],\n        \"learning_rate\": config[\"learning_rate\"],\n        \"weight_decay\": config[\"weight_decay\"],\n        \"num_epochs\": config[\"num_epochs\"],\n        \"patience\": config[\"patience\"],\n        \"val_split\": config[\"val_split\"],\n        \"seed\": config[\"seed\"],\n    },\n    name=\"DinoV2-G-Adagrad-tripletLoss\",\n)\n\nprint(\"W&B initialized. Key hyperparameters tracked:\")\nprint(f\"  Project: {os.getenv('WANDB_PROJECT', 'jaguar-reid-jojojaguar')}\")\nprint(f\"  Triplet margin: {config['triplet_margin']}\")\nprint(f\"  Mining type: {config['mining_type']}\")\nprint(f\"  PK sampling: P={config['P']} x K={config['K']} = {config['batch_size']} per batch\")\nprint(f\"  Embedding dim: {config['embedding_dim']}\")\nprint(f\"  Dropout: {config['dropout']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:34:54.876413Z","iopub.execute_input":"2026-03-14T20:34:54.876776Z","iopub.status.idle":"2026-03-14T20:35:09.810590Z","shell.execute_reply.started":"2026-03-14T20:34:54.876752Z","shell.execute_reply":"2026-03-14T20:35:09.809899Z"}},"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-14T20:35:09.812356Z","iopub.execute_input":"2026-03-14T20:35:09.812945Z","iopub.status.idle":"2026-03-14T20:35:09.852536Z","shell.execute_reply.started":"2026-03-14T20:35:09.812904Z","shell.execute_reply":"2026-03-14T20:35:09.851841Z"}},"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-14T20:35:09.853439Z","iopub.execute_input":"2026-03-14T20:35:09.853815Z","iopub.status.idle":"2026-03-14T20:35:10.569044Z","shell.execute_reply.started":"2026-03-14T20:35:09.853784Z","shell.execute_reply":"2026-03-14T20:35:10.568322Z"}},"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-14T20:35:10.570086Z","iopub.execute_input":"2026-03-14T20:35:10.570484Z","iopub.status.idle":"2026-03-14T20:35:12.072784Z","shell.execute_reply.started":"2026-03-14T20:35:10.570454Z","shell.execute_reply":"2026-03-14T20:35:12.072137Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Load Backbone Model\n\nWe use the `timm` library to load the pre-trained model from Hugging Face Hub.","metadata":{}},{"cell_type":"code","source":"# Load Backbone model\nprint(\"Loading Backbone model...\")\nbackbone = timm.create_model(\n    config[\"backbone_model\"],\n    num_classes=0,\n    pretrained=True\n)\nbackbone.eval()\nbackbone.to(device)\n\nprint(f\"Model loaded successfully\")\nprint(f\"  Parameters: {sum(p.numel() for p in backbone.parameters()):,}\")\n\n# Get the embedding dimension from the model\nwith torch.no_grad():\n    dummy_input = torch.randn(1, 3, config[\"input_size\"], config[\"input_size\"]).to(device)\n    dummy_output = backbone(dummy_input)\n    backbone_dim = dummy_output.shape[1]\n    print(f\"  Embedding dimension: {backbone_dim}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:35:12.073741Z","iopub.execute_input":"2026-03-14T20:35:12.073998Z","iopub.status.idle":"2026-03-14T20:35:43.120272Z","shell.execute_reply.started":"2026-03-14T20:35:12.073974Z","shell.execute_reply":"2026-03-14T20:35:43.119645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define preprocessing pipeline\n# DinoV2 Giant expects 518x518 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-14T20:35:43.121342Z","iopub.execute_input":"2026-03-14T20:35:43.121875Z","iopub.status.idle":"2026-03-14T20:35:43.130288Z","shell.execute_reply.started":"2026-03-14T20:35:43.121850Z","shell.execute_reply":"2026-03-14T20:35:43.129777Z"}},"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 backbone model.\"\"\"\n    model.eval()\n    embeddings = []\n    \n    for i in tqdm(range(0, len(image_paths), batch_size), desc=desc):\n        batch_paths = image_paths[i:i + batch_size]\n        \n        # Load and preprocess batch\n        batch_tensors = []\n        for path in batch_paths:\n            try:\n                img = Image.open(path).convert(\"RGB\")\n                tensor = preprocess(img)\n                batch_tensors.append(tensor)\n            except Exception as e:\n                print(f\"Error loading {path}: {e}\")\n                # Use zero tensor as fallback\n                batch_tensors.append(torch.zeros(3, config[\"input_size\"], config[\"input_size\"]))\n        \n        # Stack and move to device\n        batch_tensor = torch.stack(batch_tensors).to(device)\n        \n        # Get embeddings\n        batch_emb = model(batch_tensor).cpu().numpy()\n        embeddings.append(batch_emb)\n    \n    return np.vstack(embeddings)\n\nprint(\"Embedding extraction function defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:35:43.131196Z","iopub.execute_input":"2026-03-14T20:35:43.131500Z","iopub.status.idle":"2026-03-14T20:35:52.804992Z","shell.execute_reply.started":"2026-03-14T20:35:43.131454Z","shell.execute_reply":"2026-03-14T20:35:52.804060Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"emb_dir = config[\"embed_dir\"]\n\ncache_path = emb_dir / f\"{config[\"backbone_model\"].split(\"/\")[-1]}-baseline_train_embeddings.npz\"\n\n# Extract baseline embeddings for training data\ntrain_filenames = train_data[\"filename\"].astype(str).tolist()\ntrain_image_paths = [config[\"data_dir\"] / \"train\" / fn for fn in train_filenames]\n\ndef _load_cached_embeddings(cache_path, expected_filenames):\n    z = np.load(cache_path, allow_pickle=True)\n    cached_embeddings = z[\"embeddings\"]\n    cached_filenames = z[\"filenames\"].tolist() if isinstance(z[\"filenames\"], np.ndarray) else list(z[\"filenames\"])\n\n    if len(cached_filenames) != len(expected_filenames):\n        return None\n\n    if set(cached_filenames) != set(expected_filenames):\n        return None\n\n    if cached_filenames == expected_filenames:\n        return cached_embeddings\n\n    idx = {fn: i for i, fn in enumerate(cached_filenames)}\n    return np.stack([cached_embeddings[idx[fn]] for fn in expected_filenames], axis=0)\n\nbaseline_train_embeddings = None\nif cache_path.exists():\n    baseline_train_embeddings = _load_cached_embeddings(cache_path, train_filenames)\n    if baseline_train_embeddings is not None:\n        print(f\"Loaded cached baseline embeddings from {cache_path}\")\n        print(f\"Baseline embeddings shape: {baseline_train_embeddings.shape}\")\n\nif baseline_train_embeddings is None:\n    print(f\"Extracting baseline embeddings for {len(train_image_paths)} training images...\")\n    baseline_train_embeddings = extract_embeddings(\n        backbone,\n        train_image_paths,\n        batch_size=config[\"batch_size\"]\n    )\n    np.savez_compressed(\n        cache_path,\n        embeddings=baseline_train_embeddings,\n        filenames=np.array(train_filenames, dtype=object),\n    )\n    print(f\"Saved baseline embeddings cache to {cache_path}\")\n    print(f\"Baseline embeddings shape: {baseline_train_embeddings.shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:35:52.806210Z","iopub.execute_input":"2026-03-14T20:35:52.806637Z","iopub.status.idle":"2026-03-14T20:40:49.936060Z","shell.execute_reply.started":"2026-03-14T20:35:52.806610Z","shell.execute_reply":"2026-03-14T20:40:49.933825Z"}},"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 Backbone separates different jaguars before any fine-tuning.","metadata":{}},{"cell_type":"code","source":"def compute_geodesic_distances(embeddings):\n    \"\"\"Compute geodesic (angular) distance matrix for normalized embeddings.\"\"\"\n    # Normalize embeddings to unit sphere\n    norms = np.linalg.norm(embeddings, axis=1, keepdims=True)\n    normalized = embeddings / norms\n    \n    # Compute cosine similarity\n    cos_sim = np.clip(normalized @ normalized.T, -1.0, 1.0)\n    \n    # Convert to geodesic distance (arc length)\n    geodesic_dist = np.arccos(cos_sim)\n    \n    return geodesic_dist\n\n\ndef visualize_embeddings_mds(embeddings, labels, title, max_samples=500):\n    \"\"\"Visualize embeddings using MDS with geodesic distances.\"\"\"\n    # Subsample if too many points (MDS is O(n^3))\n    if len(embeddings) > max_samples:\n        indices = np.random.choice(len(embeddings), max_samples, replace=False)\n        embeddings = embeddings[indices]\n        labels = labels[indices]\n    \n    # Compute geodesic distance matrix\n    dist_matrix = compute_geodesic_distances(embeddings)\n    \n    # Apply MDS\n    mds = MDS(n_components=2, dissimilarity='precomputed', random_state=RANDOM_SEED, normalized_stress='auto')\n    coords_2d = mds.fit_transform(dist_matrix)\n    \n    # Create color mapping for identities\n    unique_labels = np.unique(labels)\n    colors = plt.cm.tab20(np.linspace(0, 1, len(unique_labels)))\n    label_to_color = {label: colors[i] for i, label in enumerate(unique_labels)}\n    \n    # Plot\n    fig, ax = plt.subplots(figsize=(12, 10))\n    \n    for label in unique_labels:\n        mask = labels == label\n        ax.scatter(\n            coords_2d[mask, 0], \n            coords_2d[mask, 1],\n            c=[label_to_color[label]],\n            label=label,\n            alpha=0.7,\n            s=30\n        )\n    \n    ax.set_title(title, fontsize=14, fontweight='bold')\n    ax.set_xlabel('MDS Dimension 1')\n    ax.set_ylabel('MDS Dimension 2')\n    \n    # Legend outside plot\n    ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=8)\n    plt.tight_layout()\n    \n    return fig\n\nprint(\"MDS visualization functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.936667Z","iopub.status.idle":"2026-03-14T20:40:49.936939Z","shell.execute_reply.started":"2026-03-14T20:40:49.936814Z","shell.execute_reply":"2026-03-14T20:40:49.936830Z"}},"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 backbone Embeddings (Before Fine-tuning)\"\n)\nplt.show()\n\n# Log to W&B\nwandb.log({\"baseline_embeddings_mds\": wandb.Image(fig_baseline)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.937957Z","iopub.status.idle":"2026-03-14T20:40:49.938196Z","shell.execute_reply.started":"2026-03-14T20:40:49.938086Z","shell.execute_reply":"2026-03-14T20:40:49.938100Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Define Model Architecture and Triplet Mining\n\nWe define three components:\n\n1. **EmbeddingProjection**: Projects 1536-dim DinoV2 Giant embeddings to 256-dim. This learned projection optimizes the embedding space for our specific jaguar dataset.\n\n2. **Online Hard Triplet Mining**: Mines the hardest triplets within each batch:\n   - For each anchor, finds the hardest positive (farthest same-identity sample)\n   - For each anchor, finds the hardest negative (closest different-identity sample)\n   \n3. **TripletModel**: Wraps the projection network and provides L2-normalized embeddings for both training and inference.","metadata":{}},{"cell_type":"code","source":"class EmbeddingProjection(nn.Module):\n    \"\"\"\n    Projects Backbone embeddings to a lower-dimensional space.\n    Architecture: input_dim -> hidden_dim -> output_dim\n    \"\"\"\n    \n    def __init__(self, input_dim=1536, hidden_dim=512, output_dim=256, dropout=0.3):\n        super().__init__()\n        \n        self.network = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim),\n            nn.BatchNorm1d(hidden_dim),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout),\n            \n            nn.Linear(hidden_dim, output_dim),\n            nn.BatchNorm1d(output_dim),\n        )\n        \n        self._init_weights()\n    \n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.BatchNorm1d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x):\n        return self.network(x)\n\n\ndef mine_hard_triplets(embeddings, labels):\n    \"\"\"\n    Online batch-hard triplet mining.\n    \n    For each anchor, selects:\n      - Hardest positive: the farthest sample with the same label\n      - Hardest negative: the closest sample with a different label\n    \n    Args:\n        embeddings: (B, D) L2-normalized embeddings\n        labels: (B,) integer labels\n    \n    Returns:\n        anchors, positives, negatives: index tensors of shape (num_valid_triplets,)\n    \"\"\"\n    # Pairwise Euclidean distance matrix\n    dist_matrix = torch.cdist(embeddings, embeddings, p=2)  # (B, B)\n    \n    labels = labels.unsqueeze(0)  # (1, B)\n    same_identity = (labels == labels.T).float()  # (B, B)\n    diff_identity = 1.0 - same_identity\n    \n    # Mask out self-comparisons\n    eye = torch.eye(embeddings.size(0), device=embeddings.device)\n    same_identity = same_identity - eye  # exclude self\n    \n    anchor_indices = []\n    positive_indices = []\n    negative_indices = []\n    \n    for i in range(embeddings.size(0)):\n        # Hardest positive: max distance among same-identity (excluding self)\n        pos_mask = same_identity[i]\n        if pos_mask.sum() == 0:\n            continue  # no positive available\n        \n        neg_mask = diff_identity[i]\n        if neg_mask.sum() == 0:\n            continue  # no negative available\n        \n        # Hardest positive (farthest same-identity)\n        pos_dists = dist_matrix[i] * pos_mask + (-1e9) * (1 - pos_mask)\n        hardest_pos = pos_dists.argmax()\n        \n        # Hardest negative (closest different-identity)\n        neg_dists = dist_matrix[i] * neg_mask + 1e9 * (1 - neg_mask)\n        hardest_neg = neg_dists.argmin()\n        \n        anchor_indices.append(i)\n        positive_indices.append(hardest_pos.item())\n        negative_indices.append(hardest_neg.item())\n    \n    return (\n        torch.tensor(anchor_indices, device=embeddings.device),\n        torch.tensor(positive_indices, device=embeddings.device),\n        torch.tensor(negative_indices, device=embeddings.device),\n    )\n\n\nprint(\"EmbeddingProjection and hard triplet mining defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.939542Z","iopub.status.idle":"2026-03-14T20:40:49.939882Z","shell.execute_reply.started":"2026-03-14T20:40:49.939699Z","shell.execute_reply":"2026-03-14T20:40:49.939722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TripletModel(nn.Module):\n    \"\"\"Complete model: Embedding Projection with L2 normalization.\"\"\"\n    \n    def __init__(self, input_dim, embedding_dim=256, hidden_dim=512, 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    \n    def forward(self, x):\n        \"\"\"Forward pass: returns L2-normalized embeddings.\"\"\"\n        embeddings = self.embedding_net(x)\n        return F.normalize(embeddings, p=2, dim=1)\n    \n    def get_embeddings(self, x):\n        \"\"\"Get normalized embeddings for inference (same as forward).\"\"\"\n        return self.forward(x)\n\n\n# Create model\nmodel = TripletModel(\n    input_dim=backbone_dim,\n    embedding_dim=config[\"embedding_dim\"],\n    hidden_dim=config[\"hidden_dim\"],\n    dropout=config[\"dropout\"],\n).to(device)\n\nprint(f\"Triplet 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\"  Total parameters: {sum(p.numel() for p in model.parameters()):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.940873Z","iopub.status.idle":"2026-03-14T20:40:49.941206Z","shell.execute_reply.started":"2026-03-14T20:40:49.941070Z","shell.execute_reply":"2026-03-14T20:40:49.941090Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Prepare DataLoaders\n\nWe create PyTorch datasets from the pre-computed Backbone embeddings. For triplet loss, we use a **PK batch sampler** that constructs each batch with P randomly-chosen identities and K samples per identity. This guarantees every batch contains enough same-identity pairs for meaningful triplet mining.","metadata":{}},{"cell_type":"code","source":"# Extract embeddings for validation set\nval_image_paths = [\n    config[\"data_dir\"] / \"train\" / filename \n    for filename in val_data['filename'].values\n]\n\nprint(f\"Extracting embeddings for {len(val_image_paths)} validation images...\")\nbaseline_val_embeddings = extract_embeddings(\n    backbone, \n    val_image_paths, \n    batch_size=config[\"batch_size\"]\n)\n\nprint(f\"Validation embeddings shape: {baseline_val_embeddings.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.942554Z","iopub.status.idle":"2026-03-14T20:40:49.942854Z","shell.execute_reply.started":"2026-03-14T20:40:49.942723Z","shell.execute_reply":"2026-03-14T20:40:49.942740Z"}},"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\nclass PKBatchSampler:\n    \"\"\"\n    PK Batch Sampler for metric learning.\n    \n    Each batch contains P randomly-selected identities, each with K randomly-selected\n    samples. This guarantees that every batch has enough same-identity pairs for\n    triplet mining.\n    \n    If an identity has fewer than K samples, samples are repeated (with replacement).\n    \"\"\"\n    \n    def __init__(self, labels, P, K, num_batches=None):\n        self.labels = np.array(labels)\n        self.P = P\n        self.K = K\n        \n        # Build label -> indices mapping\n        self.label_to_indices = {}\n        for idx, label in enumerate(self.labels):\n            if label not in self.label_to_indices:\n                self.label_to_indices[label] = []\n            self.label_to_indices[label].append(idx)\n        \n        self.unique_labels = list(self.label_to_indices.keys())\n        \n        # Filter out labels with < 2 samples (can't form positive pairs)\n        self.valid_labels = [l for l in self.unique_labels if len(self.label_to_indices[l]) >= 2]\n        \n        if len(self.valid_labels) < P:\n            print(f\"Warning: Only {len(self.valid_labels)} identities have >= 2 samples, but P={P}\")\n            self.P = len(self.valid_labels)\n        \n        # Number of batches per epoch\n        if num_batches is None:\n            self.num_batches = len(self.labels) // (self.P * self.K)\n        else:\n            self.num_batches = num_batches\n    \n    def __iter__(self):\n        for _ in range(self.num_batches):\n            batch_indices = []\n            \n            # Sample P identities\n            selected_labels = np.random.choice(self.valid_labels, self.P, replace=False)\n            \n            for label in selected_labels:\n                indices = self.label_to_indices[label]\n                if len(indices) >= self.K:\n                    chosen = np.random.choice(indices, self.K, replace=False)\n                else:\n                    # If fewer than K samples, sample with replacement\n                    chosen = np.random.choice(indices, self.K, replace=True)\n                batch_indices.extend(chosen.tolist())\n            \n            yield batch_indices\n    \n    def __len__(self):\n        return self.num_batches\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 PK batch sampler for training\ntrain_pk_sampler = PKBatchSampler(\n    labels=train_data['label_encoded'].values,\n    P=config[\"P\"],\n    K=config[\"K\"],\n)\n\n# Create dataloaders\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_sampler=train_pk_sampler,\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 per epoch: {len(train_pk_sampler)}\")\nprint(f\"  PK sampling: P={config['P']} identities x K={config['K']} samples = {config['batch_size']} per batch\")\nprint(f\"  Val batches: {len(val_loader)}\")\nprint(f\"  Valid identities for sampling: {len(train_pk_sampler.valid_labels)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.944454Z","iopub.status.idle":"2026-03-14T20:40:49.944742Z","shell.execute_reply.started":"2026-03-14T20:40:49.944634Z","shell.execute_reply":"2026-03-14T20:40:49.944648Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Training Setup\n\nWe set up:\n- **TripletMarginLoss**: Pushes positive pairs closer and negative pairs farther apart with a margin\n- **Online hard triplet mining**: Selects hardest triplets within each PK-sampled batch\n- **Adagrad optimizer**: Per-parameter lr for more finegrained learning control\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-14T20:40:49.946936Z","iopub.status.idle":"2026-03-14T20:40:49.947286Z","shell.execute_reply.started":"2026-03-14T20:40:49.947113Z","shell.execute_reply":"2026-03-14T20:40:49.947135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Setup training components\ncriterion = nn.TripletMarginLoss(margin=config[\"triplet_margin\"], p=2)\n\noptimizer = torch.optim.Adagrad(\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: TripletMarginLoss (margin={config['triplet_margin']})\")\nprint(f\"  Mining: Online batch-hard\")\nprint(f\"  Optimizer: Adagrad (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-14T20:40:49.948673Z","iopub.status.idle":"2026-03-14T20:40:49.949329Z","shell.execute_reply.started":"2026-03-14T20:40:49.949165Z","shell.execute_reply":"2026-03-14T20:40:49.949196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, loader, criterion, optimizer, device):\n    \"\"\"Train for one epoch with online hard triplet mining.\"\"\"\n    model.train()\n    total_loss = 0\n    total_triplets = 0\n    num_batches = 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: get normalized embeddings\n        emb = model(embeddings)\n        \n        # Mine hard triplets from the batch\n        anchor_idx, pos_idx, neg_idx = mine_hard_triplets(emb, labels)\n        \n        if len(anchor_idx) == 0:\n            continue  # skip if no valid triplets\n        \n        # Compute triplet loss\n        loss = criterion(emb[anchor_idx], emb[pos_idx], emb[neg_idx])\n        \n        # Backward pass\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Metrics\n        total_loss += loss.item()\n        total_triplets += len(anchor_idx)\n        num_batches += 1\n        \n        pbar.set_postfix({\n            'loss': f'{loss.item():.4f}',\n            'triplets': len(anchor_idx),\n        })\n    \n    avg_loss = total_loss / max(num_batches, 1)\n    return avg_loss, total_triplets\n\n\ndef validate_epoch(model, loader, criterion, device):\n    \"\"\"Validate for one epoch with triplet loss.\"\"\"\n    model.eval()\n    total_loss = 0\n    total_triplets = 0\n    num_batches = 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            emb = model(embeddings)\n            \n            # Mine hard triplets\n            anchor_idx, pos_idx, neg_idx = mine_hard_triplets(emb, labels)\n            \n            if len(anchor_idx) == 0:\n                continue\n            \n            loss = criterion(emb[anchor_idx], emb[pos_idx], emb[neg_idx])\n            \n            total_loss += loss.item()\n            total_triplets += len(anchor_idx)\n            num_batches += 1\n            \n            pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n    \n    avg_loss = total_loss / max(num_batches, 1)\n    return avg_loss, total_triplets\n\n\nprint(\"Training and validation functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.950339Z","iopub.status.idle":"2026-03-14T20:40:49.950698Z","shell.execute_reply.started":"2026-03-14T20:40:49.950519Z","shell.execute_reply":"2026-03-14T20:40:49.950542Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Training Loop\n\nWe train the model with:\n- Online hard triplet mining within each PK-sampled batch\n- Validation mAP computed each epoch (the primary quality metric)\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_triplets': [],\n    'val_loss': [], 'val_triplets': [],\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_triplets = train_epoch(model, train_loader, criterion, optimizer, device)\n    \n    # Validate (use val_loader for loss; note: val_loader uses sequential batching,\n    # so triplet mining here is approximate -- mAP is the real metric)\n    val_loss, val_triplets = validate_epoch(model, val_loader, criterion, device)\n    \n    # Compute validation mAP (the true quality metric)\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_triplets'].append(train_triplets)\n    history['val_loss'].append(val_loss)\n    history['val_triplets'].append(val_triplets)\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_triplets': train_triplets,\n        'val_loss': val_loss,\n        'val_triplets': val_triplets,\n        'val_map': val_map,\n        'learning_rate': current_lr,\n    })\n    \n    # Print summary\n    print(f\"  Train Loss: {train_loss:.4f} | Triplets: {train_triplets}\")\n    print(f\"  Val Loss:   {val_loss:.4f} | Triplets: {val_triplets}\")\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\"] / \"triplet_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-14T20:40:49.952014Z","iopub.status.idle":"2026-03-14T20:40:49.952325Z","shell.execute_reply.started":"2026-03-14T20:40:49.952165Z","shell.execute_reply":"2026-03-14T20:40:49.952185Z"}},"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('Triplet Loss')\naxes[0].set_title('Training and Validation Loss')\naxes[0].legend()\naxes[0].grid(True, alpha=0.3)\n\n# Triplets mined\naxes[1].plot(epochs_range, history['train_triplets'], 'b-', label='Train')\naxes[1].plot(epochs_range, history['val_triplets'], 'r-', label='Validation')\naxes[1].axvline(x=best_epoch, color='g', linestyle='--', alpha=0.7)\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('Number of Triplets')\naxes[1].set_title('Hard Triplets Mined per Epoch')\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-14T20:40:49.953845Z","iopub.status.idle":"2026-03-14T20:40:49.954148Z","shell.execute_reply.started":"2026-03-14T20:40:49.953998Z","shell.execute_reply":"2026-03-14T20:40:49.954012Z"}},"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 triplet loss training.","metadata":{}},{"cell_type":"code","source":"# Load best model\ncheckpoint = torch.load(config[\"checkpoint_dir\"] / \"triplet_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-14T20:40:49.956164Z","iopub.status.idle":"2026-03-14T20:40:49.956955Z","shell.execute_reply.started":"2026-03-14T20:40:49.956822Z","shell.execute_reply":"2026-03-14T20:40:49.956839Z"}},"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-14T20:40:49.957910Z","iopub.status.idle":"2026-03-14T20:40:49.958258Z","shell.execute_reply.started":"2026-03-14T20:40:49.958087Z","shell.execute_reply":"2026-03-14T20:40:49.958108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize fine-tuned embeddings\nfig_finetuned = visualize_embeddings_mds(\n    finetuned_train_embeddings,\n    train_labels,\n    \"Fine-tuned ArcFace Embeddings (After Training)\"\n)\nplt.show()\n\n# Log to W&B\nwandb.log({\"finetuned_embeddings_mds\": wandb.Image(fig_finetuned)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.959350Z","iopub.status.idle":"2026-03-14T20:40:49.959697Z","shell.execute_reply.started":"2026-03-14T20:40:49.959525Z","shell.execute_reply":"2026-03-14T20:40:49.959546Z"}},"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 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    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    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    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    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(DinoV2 Backbone4)', \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-14T20:40:49.961229Z","iopub.status.idle":"2026-03-14T20:40:49.961595Z","shell.execute_reply.started":"2026-03-14T20:40:49.961401Z","shell.execute_reply":"2026-03-14T20:40:49.961422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize nearest neighbors for a few validation examples\n\nprint(\"Generating nearest neighbor visualizations for validation set...\")\nprint(f\"Validation set size: {len(val_data)}\")\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-14T20:40:49.962356Z","iopub.status.idle":"2026-03-14T20:40:49.962625Z","shell.execute_reply.started":"2026-03-14T20:40:49.962465Z","shell.execute_reply":"2026-03-14T20:40:49.962477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example 1: Random validation image\nnp.random.seed(RANDOM_SEED)\nrandom_idx = np.random.randint(0, len(val_labels))\n\nprint(f\"Example 1: Random query (index {random_idx})\")\nfig1, stats1 = visualize_nearest_neighbors(\n    query_idx=random_idx,\n    original_embeddings=baseline_val_embeddings,\n    finetuned_embeddings=val_finetuned_embeddings,\n    image_paths=val_image_paths,\n    labels=val_labels,\n    k=5\n)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.963744Z","iopub.status.idle":"2026-03-14T20:40:49.964130Z","shell.execute_reply.started":"2026-03-14T20:40:49.963949Z","shell.execute_reply":"2026-03-14T20:40:49.963971Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Generate Competition Submission\n\nNow we generate predictions for the test set. The competition expects:\n- A CSV with columns: `row_id`, `similarity`\n- Each row corresponds to a query-gallery image pair from `test.csv`\n- `similarity` is a float between 0 and 1\n\nWe:\n1. Extract MegaDescriptor embeddings for all test images\n2. Project through our fine-tuned model\n3. Compute cosine similarity for each pair in `test.csv`\n4. Clip values to [0, 1] and save as CSV","metadata":{}},{"cell_type":"code","source":"# Load test.csv to get the pairs we need to score\ntest_pairs_df = pd.read_csv(config[\"data_dir\"] / \"test.csv\")\n\nprint(f\"Test pairs to score: {len(test_pairs_df)}\")\nprint(f\"Columns: {list(test_pairs_df.columns)}\")\nprint(f\"\\nSample rows:\")\nprint(test_pairs_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.965469Z","iopub.status.idle":"2026-03-14T20:40:49.965840Z","shell.execute_reply.started":"2026-03-14T20:40:49.965663Z","shell.execute_reply":"2026-03-14T20:40:49.965684Z"}},"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 Backbone embeddings for test images\nprint(f\"\\nExtracting Backbone embeddings for test images...\")\ntest_backbone_embeddings = extract_embeddings(\n    backbone,\n    test_image_paths,\n    batch_size=config[\"batch_size\"],\n    desc=\"Test embeddings\"\n)\n\nprint(f\"Test Backbone embeddings shape: {test_backbone_embeddings.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.968126Z","iopub.status.idle":"2026-03-14T20:40:49.968532Z","shell.execute_reply.started":"2026-03-14T20:40:49.968331Z","shell.execute_reply":"2026-03-14T20:40:49.968347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Project through fine-tuned model\nmodel.eval()\nwith torch.no_grad():\n    test_tensor = torch.FloatTensor(test_backbone_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-14T20:40:49.970278Z","iopub.status.idle":"2026-03-14T20:40:49.970627Z","shell.execute_reply.started":"2026-03-14T20:40:49.970450Z","shell.execute_reply":"2026-03-14T20:40:49.970466Z"}},"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-14T20:40:49.971404Z","iopub.status.idle":"2026-03-14T20:40:49.971732Z","shell.execute_reply.started":"2026-03-14T20:40:49.971572Z","shell.execute_reply":"2026-03-14T20:40:49.971610Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create submission DataFrame\nsubmission_df = pd.DataFrame({\n    'row_id': test_pairs_df['row_id'],\n    'similarity': similarities\n})\n\nprint(\"Submission DataFrame:\")\nprint(submission_df.head(10))\n\n# Verify format matches sample submission\nsample_submission = pd.read_csv(config[\"data_dir\"] / \"sample_submission.csv\")\nprint(f\"\\nFormat check:\")\nprint(f\"  Expected columns: {list(sample_submission.columns)}\")\nprint(f\"  Our columns: {list(submission_df.columns)}\")\nprint(f\"  Expected rows: {len(sample_submission)}\")\nprint(f\"  Our rows: {len(submission_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.973143Z","iopub.status.idle":"2026-03-14T20:40:49.973509Z","shell.execute_reply.started":"2026-03-14T20:40:49.973316Z","shell.execute_reply":"2026-03-14T20:40:49.973337Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission\nsubmission_path = config[\"checkpoint_dir\"] / \"submission.csv\"\nsubmission_df.to_csv(submission_path, index=False)\n\nprint(f\"Submission saved to: {submission_path}\")\nprint(f\"File size: {submission_path.stat().st_size / 1024:.1f} KB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.974811Z","iopub.status.idle":"2026-03-14T20:40:49.975370Z","shell.execute_reply.started":"2026-03-14T20:40:49.975238Z","shell.execute_reply":"2026-03-14T20:40:49.975256Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Save Artifacts to Weights and Biases\n\nWe save the best model checkpoint and this notebook as W&B artifacts for reproducibility.","metadata":{}},{"cell_type":"code","source":"# Save model as W&B artifact\nmodel_artifact = wandb.Artifact(\n    name=\"triplet-model\",\n    type=\"model\",\n    description=\"Triplet loss fine-tuned Dinov2 Giant model for jaguar re-identification\"\n)\nmodel_artifact.add_file(str(config[\"checkpoint_dir\"] / \"triplet_best.pth\"))\nwandb.log_artifact(model_artifact)\n\nprint(\"Model artifact saved to W&B\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.976306Z","iopub.status.idle":"2026-03-14T20:40:49.976570Z","shell.execute_reply.started":"2026-03-14T20:40:49.976418Z","shell.execute_reply":"2026-03-14T20:40:49.976431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission as W&B artifact\nsubmission_artifact = wandb.Artifact(\n    name=\"submission\",\n    type=\"submission\",\n    description=\"Competition submission file\"\n)\nsubmission_artifact.add_file(str(submission_path))\nwandb.log_artifact(submission_artifact)\n\nprint(\"Submission artifact saved to W&B\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T20:40:49.977892Z","iopub.status.idle":"2026-03-14T20:40:49.978232Z","shell.execute_reply.started":"2026-03-14T20:40:49.978112Z","shell.execute_reply":"2026-03-14T20:40:49.978127Z"}},"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-14T20:40:49.979244Z","iopub.status.idle":"2026-03-14T20:40:49.979607Z","shell.execute_reply.started":"2026-03-14T20:40:49.979431Z","shell.execute_reply":"2026-03-14T20:40:49.979453Z"}},"outputs":[],"execution_count":null}]}