{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":126777,"databundleVersionId":15314950},{"sourceType":"datasetVersion","sourceId":14957993,"datasetId":9574063,"databundleVersionId":15828903},{"sourceType":"datasetVersion","sourceId":14960344,"datasetId":9575709,"databundleVersionId":15831475}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# DATA EXPLORATION & PREPARATION","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom sklearn.model_selection import StratifiedShuffleSplit\nimport torch\nfrom torch.utils.data import Dataset, DataLoader, Sampler\nimport torchvision.transforms as transforms\nimport random\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\nclass Config:\n    # Paths\n    BASE_PATH = '/kaggle/input/competitions/jaguar-re-id'  \n    TRAIN_CSV = os.path.join(BASE_PATH, 'train.csv')\n    TEST_CSV = os.path.join(BASE_PATH, 'test.csv')\n    TRAIN_IMG_DIR = os.path.join(BASE_PATH, 'train/train')\n    TEST_IMG_DIR = os.path.join(BASE_PATH, 'test/test')\n    \n    # Image settings\n    IMG_SIZE = 384  # Larger size preserves fine-grained spot patterns\n    USE_ALPHA = True  # Use alpha mask to remove background\n    \n    # Training settings\n    BATCH_SIZE = 64  # Will be overridden by PK sampler\n    P = 16  # Number of identities per batch\n    K = 4   # Number of images per identity (batch_size = P * K = 64)\n    NUM_WORKERS = 4\n    PIN_MEMORY = True\n    \n    # Validation split\n    VAL_SIZE = 0.2\n    RANDOM_SEED = 42\n\nconfig = Config()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:10.232484Z","iopub.execute_input":"2026-02-25T17:02:10.233217Z","iopub.status.idle":"2026-02-25T17:02:10.241516Z","shell.execute_reply.started":"2026-02-25T17:02:10.233179Z","shell.execute_reply":"2026-02-25T17:02:10.240866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. Data Exploration\nprint(\"=\"*60)\nprint(\"JAGUAR RE-ID - DATA EXPLORATION\")\nprint(\"=\"*60)\n\n# Load training metadata\ntrain_df = pd.read_csv(config.TRAIN_CSV)\nprint(f\"\\n✅ Loaded training data: {len(train_df)} images\")\nprint(f\"✅ Unique jaguars: {train_df['ground_truth'].nunique()}\")\n\n# Class distribution\nclass_counts = train_df['ground_truth'].value_counts()\nprint(\"\\n📊 Class distribution (images per jaguar):\")\nprint(class_counts)\n\n# Plot distribution\nplt.figure(figsize=(15, 6))\nclass_counts.plot(kind='bar')\nplt.title('Number of Images per Jaguar (Training Set)', fontsize=14)\nplt.xlabel('Jaguar ID', fontsize=12)\nplt.ylabel('Count', fontsize=12)\nplt.xticks(rotation=45, ha='right')\nplt.tight_layout()\nplt.savefig('class_distribution.png')\nplt.show()\nprint(\"✅ Saved class distribution plot to 'class_distribution.png'\")\n\n# Display sample images with alpha channel visualization\ndef show_samples_with_alpha(df, img_dir, num_samples=3):\n    \"\"\"Show RGB image and alpha mask side by side\"\"\"\n    unique_ids = df['ground_truth'].unique()[:num_samples]\n    fig, axes = plt.subplots(num_samples, 3, figsize=(12, 4*num_samples))\n    \n    for i, jaguar_id in enumerate(unique_ids):\n        # Get first image for this jaguar\n        img_file = df[df['ground_truth'] == jaguar_id].iloc[0]['filename']\n        img_path = os.path.join(img_dir, img_file)\n        \n        # Load RGBA to see alpha channel\n        img_rgba = Image.open(img_path).convert('RGBA')\n        rgb = img_rgba.convert('RGB')\n        alpha = img_rgba.split()[3]  # Get alpha channel\n        \n        # Composite with black background (our training approach)\n        bg = Image.new('RGB', img_rgba.size, (0, 0, 0))\n        composite = Image.composite(rgb, bg, alpha)\n        \n        # Display\n        axes[i, 0].imshow(rgb)\n        axes[i, 0].set_title(f\"{jaguar_id}\\nRGB (with background)\", fontsize=10)\n        axes[i, 0].axis('off')\n        \n        axes[i, 1].imshow(alpha, cmap='gray')\n        axes[i, 1].set_title(f\"Alpha mask\\n(white = jaguar)\", fontsize=10)\n        axes[i, 1].axis('off')\n        \n        axes[i, 2].imshow(composite)\n        axes[i, 2].set_title(f\"Composite over black\\n(what model sees)\", fontsize=10)\n        axes[i, 2].axis('off')\n    \n    plt.tight_layout()\n    plt.savefig('alpha_mask_demo.png')\n    plt.show()\n    print(\"✅ Saved alpha mask demo to 'alpha_mask_demo.png'\")\n\nprint(\"\\n🖼️  Visualizing alpha mask usage...\")\nshow_samples_with_alpha(train_df, config.TRAIN_IMG_DIR, num_samples=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:12.556526Z","iopub.execute_input":"2026-02-25T17:02:12.557121Z","iopub.status.idle":"2026-02-25T17:02:19.304034Z","shell.execute_reply.started":"2026-02-25T17:02:12.557095Z","shell.execute_reply":"2026-02-25T17:02:19.303192Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3. Train/Validation Split (Stratified)\nprint(\"\\n\" + \"=\"*60)\nprint(\"CREATING STRATIFIED TRAIN/VAL SPLIT\")\nprint(\"=\"*60)\n\nX = train_df['filename']\ny = train_df['ground_truth']\n\nsss = StratifiedShuffleSplit(n_splits=1, test_size=config.VAL_SIZE, random_state=config.RANDOM_SEED)\ntrain_idx, val_idx = next(sss.split(X, y))\n\ntrain_split = train_df.iloc[train_idx].reset_index(drop=True)\nval_split = train_df.iloc[val_idx].reset_index(drop=True)\n\nprint(f\"\\n✅ Training set: {len(train_split)} images\")\nprint(f\"✅ Validation set: {len(val_split)} images\")\n\n# Verify class distribution in validation\nval_class_counts = val_split['ground_truth'].value_counts()\nprint(\"\\n📊 Validation set distribution (should mirror overall):\")\nprint(val_class_counts.head(10))\nprint(\"...\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:19.305246Z","iopub.execute_input":"2026-02-25T17:02:19.305481Z","iopub.status.idle":"2026-02-25T17:02:19.318241Z","shell.execute_reply.started":"2026-02-25T17:02:19.305458Z","shell.execute_reply":"2026-02-25T17:02:19.317666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4. Label Encoding\n# Create mapping from jaguar name to integer label\nunique_ids = sorted(train_df['ground_truth'].unique())\nid_to_label = {name: idx for idx, name in enumerate(unique_ids)}\nlabel_to_id = {idx: name for name, idx in id_to_label.items()}\nnum_classes = len(id_to_label)\n\n# Add label column\ntrain_split['label'] = train_split['ground_truth'].map(id_to_label)\nval_split['label'] = val_split['ground_truth'].map(id_to_label)\n\nprint(f\"\\n✅ Encoded {num_classes} jaguar identities to labels 0-{num_classes-1}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:19.319293Z","iopub.execute_input":"2026-02-25T17:02:19.319875Z","iopub.status.idle":"2026-02-25T17:02:19.335826Z","shell.execute_reply.started":"2026-02-25T17:02:19.319853Z","shell.execute_reply":"2026-02-25T17:02:19.335173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5. Image Transforms\n# ImageNet stats for normalization (suitable for most pre-trained models)\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]\n\n# Training transforms with augmentation\ntrain_transforms = transforms.Compose([\n    transforms.Resize((config.IMG_SIZE, config.IMG_SIZE)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomRotation(degrees=10),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD)\n])\n\n# Validation transforms (no augmentation)\nval_transforms = transforms.Compose([\n    transforms.Resize((config.IMG_SIZE, config.IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD)\n])\n\n# 6. Custom Dataset with Alpha Mask Support\nclass JaguarDataset(Dataset):\n    \"\"\"\n    Dataset for jaguar re-identification.\n    \n    If use_alpha=True, the alpha channel is used to composite the jaguar over\n    a black background, effectively removing the original background.\n    \"\"\"\n    def __init__(self, dataframe, img_dir, transform=None, use_alpha=True):\n        self.dataframe = dataframe\n        self.img_dir = img_dir\n        self.transform = transform\n        self.use_alpha = use_alpha\n        \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, idx):\n        row = self.dataframe.iloc[idx]\n        filename = row['filename']\n        label = row['label']\n        img_path = os.path.join(self.img_dir, filename)\n        \n        # Load image - PNG with alpha channel\n        img = Image.open(img_path)\n        \n        if self.use_alpha:\n            # Convert to RGBA to ensure alpha channel is present\n            img = img.convert('RGBA')\n            \n            # Split into RGB and alpha\n            r, g, b, alpha = img.split()\n            \n            # Create black background\n            bg = Image.new('RGB', img.size, (0, 0, 0))\n            \n            # Composite RGB over black using alpha as mask\n            # This effectively removes the original background\n            img = Image.composite(img.convert('RGB'), bg, alpha)\n        else:\n            # Just use RGB (original background remains)\n            img = img.convert('RGB')\n        \n        if self.transform:\n            img = self.transform(img)\n            \n        return img, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:19.443982Z","iopub.execute_input":"2026-02-25T17:02:19.444731Z","iopub.status.idle":"2026-02-25T17:02:19.453241Z","shell.execute_reply.started":"2026-02-25T17:02:19.444707Z","shell.execute_reply":"2026-02-25T17:02:19.452540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 7. PK Sampler for Balanced Batches\nclass PKSampler(Sampler):\n    \"\"\"\n    PK sampler for re-identification.\n    \n    Each batch contains P identities, each with K images.\n    This ensures balanced representation of all classes.\n    \n    Args:\n        labels: list of integer labels for each sample\n        P: number of identities per batch\n        K: number of images per identity\n        num_iterations: number of batches per epoch (optional)\n    \"\"\"\n    def __init__(self, labels, P, K, num_iterations=None):\n        self.labels = labels\n        self.P = P\n        self.K = K\n        \n        # Create dictionary mapping label to list of indices\n        self.label_to_indices = {}\n        for idx, label in enumerate(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        # Get list of unique labels\n        self.unique_labels = list(self.label_to_indices.keys())\n        \n        # Ensure we have at least P identities\n        assert len(self.unique_labels) >= P, f\"Need at least {P} identities, but have {len(self.unique_labels)}\"\n        \n        # Calculate number of iterations per epoch\n        if num_iterations is None:\n            # Rough estimate: total_samples / batch_size\n            total_samples = len(labels)\n            batch_size = P * K\n            self.num_iterations = total_samples // batch_size\n        else:\n            self.num_iterations = num_iterations\n    \n    def __iter__(self):\n        for _ in range(self.num_iterations):\n            # Randomly select P identities without replacement\n            selected_labels = random.sample(self.unique_labels, self.P)\n            \n            batch_indices = []\n            for label in selected_labels:\n                indices = self.label_to_indices[label]\n                \n                # Sample K images from this identity\n                # If fewer than K images, sample with replacement\n                if len(indices) >= self.K:\n                    selected_indices = random.sample(indices, self.K)\n                else:\n                    selected_indices = random.choices(indices, k=self.K)\n                \n                batch_indices.extend(selected_indices)\n            \n            yield batch_indices\n    \n    def __len__(self):\n        return self.num_iterations\n\n# 8. Create Datasets and DataLoaders\nprint(\"\\n\" + \"=\"*60)\nprint(\"CREATING DATASETS AND DATALOADERS\")\nprint(\"=\"*60)\n\n# Create datasets\ntrain_dataset = JaguarDataset(\n    dataframe=train_split,\n    img_dir=config.TRAIN_IMG_DIR,\n    transform=train_transforms,\n    use_alpha=config.USE_ALPHA\n)\n\nval_dataset = JaguarDataset(\n    dataframe=val_split,\n    img_dir=config.TRAIN_IMG_DIR,\n    transform=val_transforms,\n    use_alpha=config.USE_ALPHA\n)\n\nprint(f\"\\n✅ Train dataset: {len(train_dataset)} samples\")\nprint(f\"✅ Val dataset: {len(val_dataset)} samples\")\nprint(f\"✅ Using alpha mask: {config.USE_ALPHA}\")\n\n# Create PK sampler for training\ntrain_labels = train_split['label'].tolist()\npk_sampler = PKSampler(\n    labels=train_labels,\n    P=config.P,\n    K=config.K,\n    num_iterations=len(train_dataset) // (config.P * config.K)\n)\n\n# Create dataloaders\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_sampler=pk_sampler,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=config.PIN_MEMORY\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=config.P * config.K,  # Same batch size for simplicity\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=config.PIN_MEMORY\n)\n\nprint(f\"\\n✅ Train loader: {len(train_loader)} batches per epoch\")\nprint(f\"✅ Val loader: {len(val_loader)} batches\")\nprint(f\"✅ Batch size: {config.P * config.K} ({config.P} identities × {config.K} images)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:22.212334Z","iopub.execute_input":"2026-02-25T17:02:22.212619Z","iopub.status.idle":"2026-02-25T17:02:22.224678Z","shell.execute_reply.started":"2026-02-25T17:02:22.212596Z","shell.execute_reply":"2026-02-25T17:02:22.224049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 9. Inspect a Batch\nprint(\"\\n\" + \"=\"*60)\nprint(\"INSPECTING A TRAINING BATCH\")\nprint(\"=\"*60)\n\n# Get one batch\ndata_iter = iter(train_loader)\nimages, labels = next(data_iter)\n\nprint(f\"\\n📦 Batch image tensor shape: {images.shape}\")  # (batch, channels, H, W)\nprint(f\"📦 Batch labels shape: {labels.shape}\")\nprint(f\"📦 Labels in this batch: {labels.tolist()}\")\n\n# Count unique identities in batch\nunique_in_batch = len(torch.unique(labels))\nprint(f\"📦 Unique identities in batch: {unique_in_batch} (should be {config.P})\")\n\n# Show sample images from batch\nfig, axes = plt.subplots(2, 4, figsize=(16, 8))\naxes = axes.flatten()\n\nfor i in range(8):  # Show first 8 images\n    img = images[i].permute(1, 2, 0).numpy()\n    # Denormalize for display\n    img = img * np.array(IMAGENET_STD) + np.array(IMAGENET_MEAN)\n    img = np.clip(img, 0, 1)\n    \n    axes[i].imshow(img)\n    axes[i].set_title(f\"Label: {labels[i].item()} ({label_to_id[labels[i].item()]})\", fontsize=10)\n    axes[i].axis('off')\n\nplt.suptitle(\"Sample Batch Images (with alpha mask applied)\", fontsize=14)\nplt.tight_layout()\nplt.savefig('sample_batch.png')\nplt.show()\nprint(\"✅ Saved sample batch to 'sample_batch.png'\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T11:32:33.541797Z","iopub.execute_input":"2026-02-25T11:32:33.542086Z","iopub.status.idle":"2026-02-25T11:33:13.059714Z","shell.execute_reply.started":"2026-02-25T11:32:33.542062Z","shell.execute_reply":"2026-02-25T11:33:13.055380Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 10. Validation Metric Function (Identity-Balanced mAP)\ndef compute_identity_balanced_map(embeddings, labels):\n    \"\"\"\n    Compute identity-balanced mean Average Precision.\n    \n    Args:\n        embeddings: numpy array of shape (n_samples, feat_dim) - L2 normalized\n        labels: numpy array of shape (n_samples,) - integer labels\n    \n    Returns:\n        balanced_mAP: float, identity-balanced mean Average Precision\n        per_class_ap: dict mapping label -> AP\n    \"\"\"\n    from sklearn.metrics import average_precision_score\n    \n    # Ensure embeddings are normalized\n    embeddings = embeddings / np.linalg.norm(embeddings, axis=1, keepdims=True)\n    \n    # Compute cosine similarity matrix\n    sim_matrix = np.dot(embeddings, embeddings.T)\n    \n    unique_labels = np.unique(labels)\n    n_samples = len(labels)\n    \n    # Store AP for each query\n    query_aps = []\n    \n    for query_idx in range(n_samples):\n        query_label = labels[query_idx]\n        \n        # Get similarities to all gallery images\n        similarities = sim_matrix[query_idx]\n        \n        # Create binary relevance: 1 if same identity, 0 otherwise\n        relevance = (labels == query_label).astype(int)\n        \n        # Exclude self from gallery\n        mask = np.ones(n_samples, dtype=bool)\n        mask[query_idx] = False\n        similarities = similarities[mask]\n        relevance = relevance[mask]\n        \n        # Compute AP using sklearn\n        ap = average_precision_score(relevance, similarities)\n        query_aps.append(ap)\n    \n    # Average per identity\n    per_class_ap = {}\n    for label in unique_labels:\n        # Get indices of all queries with this label\n        label_indices = np.where(labels == label)[0]\n        label_aps = [query_aps[idx] for idx in label_indices]\n        per_class_ap[label] = np.mean(label_aps)\n    \n    # Macro-average across identities\n    balanced_mAP = np.mean(list(per_class_ap.values()))\n    \n    return balanced_mAP, per_class_ap\n\nprint(\"\\n✅ Validation metric function defined (identity-balanced mAP)\")\nprint(\"   Will compute per-class AP and macro-average\")\n\n# 11. Save Processed Data Info\n# Save class mapping for later use\nmapping_df = pd.DataFrame({\n    'jaguar_id': list(id_to_label.keys()),\n    'label': list(id_to_label.values())\n})\nmapping_df.to_csv('class_mapping.csv', index=False)\nprint(\"\\n✅ Saved class mapping to 'class_mapping.csv'\")\n\n# Save train/val split indices\ntrain_split[['filename', 'ground_truth', 'label']].to_csv('train_split.csv', index=False)\nval_split[['filename', 'ground_truth', 'label']].to_csv('val_split.csv', index=False)\nprint(\"✅ Saved train/val splits to CSV\")\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"🎉 DATA PREPARATION COMPLETE\")\nprint(\"=\"*60)\nprint(\"\\nReady for Stage 2: Model Building with ArcFace!\")\nprint(f\"- Train batches: {len(train_loader)} per epoch\")\nprint(f\"- Validation batches: {len(val_loader)}\")\nprint(f\"- Number of classes: {num_classes}\")\nprint(f\"- Image size: {config.IMG_SIZE}×{config.IMG_SIZE}\")\nprint(f\"- Using alpha mask: {config.USE_ALPHA}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:30.340378Z","iopub.execute_input":"2026-02-25T17:02:30.341057Z","iopub.status.idle":"2026-02-25T17:02:30.363278Z","shell.execute_reply.started":"2026-02-25T17:02:30.341032Z","shell.execute_reply":"2026-02-25T17:02:30.362704Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# FEATURE EXTRACTION BACKBONE","metadata":{}},{"cell_type":"code","source":"pip install timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:33.673343Z","iopub.execute_input":"2026-02-25T17:02:33.674083Z","iopub.status.idle":"2026-02-25T17:02:38.070053Z","shell.execute_reply.started":"2026-02-25T17:02:33.674057Z","shell.execute_reply":"2026-02-25T17:02:38.069076Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nimport torchvision.models as models\nfrom timm import create_model \n\n# 2. Feature Extraction Backbone\n# This stage builds the model based on the Config from stage 1.\n# We assume the Config class from your stage 1 code is available.\n\nclass JaguarReIDModel(nn.Module):\n    \"\"\"\n    Feature extraction model for Jaguar Re-ID.\n    \n    Uses a pre-trained backbone (default: MegaDescriptor-Large) and adds a\n    custom metric learning head (e.g., FC layer) to produce embeddings.\n    \"\"\"\n    def __init__(self, config, num_classes, embedding_size=512):\n        \"\"\"\n        Args:\n            config: The Config object from stage 1 (provides IMG_SIZE, etc.)\n            num_classes: Number of unique jaguar identities (for the classification layer in ArcFace)\n            embedding_size: Dimensionality of the output embedding (e.g., 512 or 1024)\n        \"\"\"\n        super(JaguarReIDModel, self).__init__()\n        self.config = config\n        self.embedding_size = embedding_size\n        self.num_classes = num_classes\n\n        # --- 1. Choose and load the backbone ---\n        # Option A: MegaDescriptor-Large (recommended baseline for animal Re-ID)\n        # Requires 'timm' library. Install with: pip install timm\n        # Model info: https://huggingface.co/BVRA/MegaDescriptor-L-384\n        print(f\"🔍 Building backbone: MegaDescriptor-L-384\")\n        # Using timm to create the model. It expects 384x384 input.\n        self.backbone = create_model(\n            'hf_hub:BVRA/MegaDescriptor-L-384',  # Path from Hugging Face hub\n            pretrained=True,\n            num_classes=0,  # Remove the original classification head\n            global_pool='avg',  # Use average pooling\n        )\n        backbone_output_dim = self.backbone.num_features  # Should be 1024 for this model\n\n        # Option B: DINOv2 (ViT-based, great for fine-grained features)\n        # from transformers import AutoModel\n        # self.backbone = AutoModel.from_pretrained('facebook/dinov2-large')\n        # backbone_output_dim = 1024 # Depends on the variant\n\n        # Option C: ConvNeXt-Large (CNN alternative)\n        # self.backbone = models.convnext_large(weights=models.ConvNeXt_Large_Weights.IMAGENET1K_V1)\n        # self.backbone = nn.Sequential(*list(self.backbone.children())[:-2]) # Remove classifier and pool\n        # backbone_output_dim = 1536 # Feature dimension for ConvNeXt-Large\n\n        print(f\"✅ Backbone loaded. Output feature dimension: {backbone_output_dim}\")\n\n        # --- 2. Add Metric Learning Head ---\n        # This head projects backbone features to a lower-dimensional embedding space.\n        self.neck = nn.Sequential(\n            nn.BatchNorm1d(backbone_output_dim),\n            nn.Linear(backbone_output_dim, self.embedding_size),\n            nn.BatchNorm1d(self.embedding_size),\n        )\n        # Note: The final embedding is not activated here. It will be normalized during loss computation.\n\n        # --- 3. (Optional) Classification Head for ArcFace ---\n        # ArcFace loss typically uses a separate, special layer (not a standard nn.Linear).\n        # We will define the ArcFace head in Stage 3 when we set up the loss.\n        # For now, we have the neck to produce embeddings.\n\n    def forward(self, x):\n        \"\"\"\n        Forward pass to extract embeddings.\n        \n        Args:\n            x: Input tensor of shape (batch_size, 3, IMG_SIZE, IMG_SIZE)\n        \n        Returns:\n            embeddings: L2-normalized embeddings of shape (batch_size, embedding_size)\n        \"\"\"\n        # Extract features from backbone\n        features = self.backbone(x)  # Shape: (batch_size, backbone_output_dim)\n        \n        # Pass through neck to get embeddings\n        embeddings = self.neck(features)  # Shape: (batch_size, embedding_size)\n        \n        # L2-normalize embeddings for use with cosine similarity\n        embeddings = nn.functional.normalize(embeddings, p=2, dim=1)\n        \n        return embeddings","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:38.071795Z","iopub.execute_input":"2026-02-25T17:02:38.072060Z","iopub.status.idle":"2026-02-25T17:02:42.275776Z","shell.execute_reply.started":"2026-02-25T17:02:38.072034Z","shell.execute_reply":"2026-02-25T17:02:42.274960Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example Usage and Integration\ndef create_model_for_training(config, num_classes):\n    \"\"\"\n    Factory function to create the model, print its structure, and show trainable params.\n    \"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"STAGE 2: BUILDING FEATURE EXTRACTION BACKBONE\")\n    print(\"=\"*60)\n    \n    # Instantiate the model\n    model = JaguarReIDModel(\n        config=config,\n        num_classes=num_classes,\n        embedding_size=512  # You can make this configurable\n    )\n    \n    # Count total and trainable parameters\n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    \n    print(f\"\\n✅ Model created successfully!\")\n    print(f\"   Total parameters: {total_params:,}\")\n    print(f\"   Trainable parameters: {trainable_params:,}\")\n    print(f\"   Embedding size: {model.embedding_size}\")\n    print(f\"   Backbone output dim: {model.backbone.num_features}\")\n    \n    # Show model architecture (first few layers)\n    print(\"\\n📐 Model Architecture (first few layers):\")\n    print(model.backbone)\n    print(\"\\n   ... (truncated) ...\")\n    print(model.neck)\n    \n    return model\n\n# -------------------------------\n# Testing the Model with a Dummy Batch\n# -------------------------------\n# This block demonstrates how the model integrates with your stage 1 DataLoader.\n# It should run after your stage 1 code that defines train_loader.\n\n# Check if we're in a context where train_loader is defined (e.g., after running stage 1)\nif 'train_loader' in locals() or 'train_loader' in globals():\n    print(\"\\n\" + \"=\"*60)\n    print(\"TESTING MODEL WITH A BATCH FROM TRAIN LOADER\")\n    print(\"=\"*60)\n    \n    # Get the number of classes from your stage 1 variable\n    if 'num_classes' in locals() or 'num_classes' in globals():\n        # Create model\n        model = create_model_for_training(config, num_classes)\n        \n        # Move model to appropriate device\n        device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        model = model.to(device)\n        print(f\"✅ Model moved to device: {device}\")\n        \n        # Get a sample batch\n        data_iter = iter(train_loader)\n        images, labels = next(data_iter)\n        images = images.to(device)\n        \n        print(f\"\\n📦 Input batch shape: {images.shape}\")\n        \n        # Forward pass\n        with torch.no_grad():\n            embeddings = model(images)\n        \n        print(f\"📤 Output embedding shape: {embeddings.shape}\")\n        print(f\"   Embedding norm (first 5): {torch.norm(embeddings, dim=1)[:5].cpu().numpy()}\")\n        print(f\"   (Should be close to 1.0 due to L2 normalization)\")\n        \n        print(\"\\n✅ Model forward pass successful!\")\n    else:\n        print(\"\\n⚠️  'num_classes' not found. Run stage 1 code first to define it.\")\nelse:\n    print(\"\\n⚠️  'train_loader' not found. Run stage 1 code first to create data loaders.\")\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"🎉 STAGE 2 COMPLETE - READY FOR METRIC LEARNING\")\nprint(\"=\"*60)\nprint(\"\\nNext steps for Stage 3:\")\nprint(\"1. Implement ArcFace loss (or another metric learning loss)\")\nprint(\"2. Set up the optimizer (e.g., AdamW with different LRs for backbone and head)\")\nprint(\"3. Create the training loop with the PK sampler batches\")\nprint(\"4. Integrate the validation metric (identity-balanced mAP)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:02:43.717523Z","iopub.execute_input":"2026-02-25T17:02:43.718495Z","iopub.status.idle":"2026-02-25T17:03:34.264696Z","shell.execute_reply.started":"2026-02-25T17:02:43.718466Z","shell.execute_reply":"2026-02-25T17:03:34.263709Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# METRIC-LEARNING AND FINE-TUNING","metadata":{}},{"cell_type":"code","source":"import torch.optim as optim\nfrom torch.optim import lr_scheduler\nimport copy\n\n# ArcFace Loss Implementation\nclass ArcFaceLoss(nn.Module):\n    \"\"\"\n    ArcFace loss: https://arxiv.org/abs/1801.07698\n    \"\"\"\n    def __init__(self, num_classes, embedding_size, margin=0.3, scale=30):\n        super(ArcFaceLoss, self).__init__()\n        self.num_classes = num_classes\n        self.embedding_size = embedding_size\n        self.margin = margin\n        self.scale = scale\n\n        # Learnable class centers (weights)\n        self.W = nn.Parameter(torch.FloatTensor(num_classes, embedding_size))\n        nn.init.xavier_normal_(self.W)\n\n        self.cos_m = np.cos(margin)\n        self.sin_m = np.sin(margin)\n        self.th = np.cos(np.pi - margin)\n        self.mm = np.sin(np.pi - margin) * margin\n\n    def forward(self, embeddings, labels):\n        \"\"\"\n        Args:\n            embeddings: (batch_size, embedding_size) L2-normalized\n            labels: (batch_size,)\n        Returns:\n            loss: scalar tensor\n        \"\"\"\n        # Normalize weights as well (for cosine similarity)\n        W_norm = nn.functional.normalize(self.W, dim=1)  # (num_classes, embedding_size)\n        \n        # Cosine similarity between embeddings and weights\n        cos_theta = torch.mm(embeddings, W_norm.t())  # (batch_size, num_classes)\n        cos_theta = torch.clamp(cos_theta, -1.0 + 1e-7, 1.0 - 1e-7)\n\n        # Get the cosine of the angle for the ground truth class\n        batch_size = labels.size(0)\n        gt_cos_theta = cos_theta[torch.arange(batch_size), labels].view(-1, 1)\n\n        # Compute sin(θ) from cos(θ) using the trigonometric identity\n        sin_theta = torch.sqrt(1.0 - gt_cos_theta ** 2)\n        # Compute cos(θ + margin) = cosθ cos_m - sinθ sin_m\n        cos_theta_m = gt_cos_theta * self.cos_m - sin_theta * self.sin_m\n        # Handle the case where cos(θ) > cos(π - margin) to avoid numerical issues\n        # (For angles beyond π, the margin addition doesn't make sense; we clip)\n        cond = gt_cos_theta > self.th\n        cos_theta_m = torch.where(cond, cos_theta_m, gt_cos_theta - self.mm)\n\n        # Replace the ground truth logit with cos(θ + margin)\n        new_cos_theta = cos_theta.clone()\n        new_cos_theta[torch.arange(batch_size), labels] = cos_theta_m.squeeze()\n\n        # Apply scale and compute cross-entropy loss\n        logits = self.scale * new_cos_theta\n        loss = nn.functional.cross_entropy(logits, labels)\n        return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:03:34.267192Z","iopub.execute_input":"2026-02-25T17:03:34.272869Z","iopub.status.idle":"2026-02-25T17:03:34.299912Z","shell.execute_reply.started":"2026-02-25T17:03:34.272821Z","shell.execute_reply":"2026-02-25T17:03:34.294904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_phase1(model, train_loader, val_loader, num_classes, device,\n                 epochs_head=10, lr_head=1e-3, save_path='best_model_phase1.pth'):\n    \"\"\"\n    Phase 1: train only the neck and ArcFace weights, backbone frozen.\n    Saves the best model (state_dict of model and arcface) to save_path.\n    \"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"PHASE 1: METRIC LEARNING – HEAD ONLY\")\n    print(\"=\"*60)\n\n    # ArcFace loss\n    arcface = ArcFaceLoss(\n        num_classes=num_classes,\n        embedding_size=model.embedding_size,\n        margin=0.3,\n        scale=30\n    ).to(device)\n\n    # Freeze backbone\n    for param in model.backbone.parameters():\n        param.requires_grad = False\n\n    # Optimiser for neck + arcface\n    params_to_optimize = list(model.neck.parameters()) + list(arcface.parameters())\n    optimizer = optim.Adam(params_to_optimize, lr=lr_head, weight_decay=5e-4)\n    scheduler = lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.5)\n\n    # Validation function (same as before)\n    def validate(model, val_loader, device):\n        model.eval()\n        all_embeddings = []\n        all_labels = []\n        with torch.no_grad():\n            for images, labels in tqdm(val_loader, desc=\"Validating\"):\n                images = images.to(device)\n                embeddings = model(images)\n                all_embeddings.append(embeddings.cpu().numpy())\n                all_labels.append(labels.numpy())\n        embeddings = np.concatenate(all_embeddings, axis=0)\n        labels = np.concatenate(all_labels, axis=0)\n        balanced_map, per_class_ap = compute_identity_balanced_map(embeddings, labels)\n        return balanced_map, per_class_ap\n\n    best_map = 0.0\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_arcface_wts = copy.deepcopy(arcface.state_dict())\n\n    for epoch in range(1, epochs_head + 1):\n        model.train()\n        arcface.train()\n        running_loss = 0.0\n        num_batches = 0\n\n        for images, labels in tqdm(train_loader, desc=f\"Epoch {epoch}/{epochs_head}\"):\n            images, labels = images.to(device), labels.to(device)\n            optimizer.zero_grad()\n            embeddings = model(images)\n            loss = arcface(embeddings, labels)\n            loss.backward()\n            optimizer.step()\n            running_loss += loss.item()\n            num_batches += 1\n\n        epoch_loss = running_loss / num_batches\n        scheduler.step()\n\n        val_map, _ = validate(model, val_loader, device)\n        print(f\"Epoch {epoch}: Loss = {epoch_loss:.4f}, Val mAP = {val_map:.4f}\")\n\n        if val_map > best_map:\n            best_map = val_map\n            best_model_wts = copy.deepcopy(model.state_dict())\n            best_arcface_wts = copy.deepcopy(arcface.state_dict())\n            print(f\"  -> New best model! (mAP: {val_map:.4f})\")\n\n    # Save best model\n    torch.save({\n        'model_state_dict': best_model_wts,\n        'arcface_state_dict': best_arcface_wts,\n        'val_map': best_map,\n    }, save_path)\n    print(f\"\\n✅ Phase 1 complete. Best mAP: {best_map:.4f} – saved to {save_path}\")\n    return model, arcface\n\n\nif __name__ == \"__main__\":\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n\n    # Assume config, train_loader, val_loader, num_classes are already defined\n    # (from previous stages)\n    model = JaguarReIDModel(config, num_classes, embedding_size=512).to(device)\n\n    train_phase1(\n        model=model,\n        train_loader=train_loader,\n        val_loader=val_loader,\n        num_classes=num_classes,\n        device=device,\n        epochs_head=10,\n        lr_head=1e-3,\n        save_path='best_model_phase1.pth'\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T11:34:50.531501Z","iopub.execute_input":"2026-02-25T11:34:50.535608Z","iopub.status.idle":"2026-02-25T12:27:51.061050Z","shell.execute_reply.started":"2026-02-25T11:34:50.535575Z","shell.execute_reply":"2026-02-25T12:27:51.060220Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nimport copy\nimport numpy as np\nfrom tqdm import tqdm\n\n# Existing performance settings\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\ntorch.backends.cudnn.benchmark = True\n\nimport os\nos.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'\n\n\ndef train_phase2(model, arcface, train_dataset, val_loader, num_classes, device,\n                 epochs_full=30, lr_head=1e-3, lr_backbone=1e-5,\n                 phase2_batch_size=8, grad_accum=4, use_amp=True,\n                 label_to_id=None, save_path='best_model_phase2.pth',\n                 early_stop_patience=5):\n    \"\"\"\n    Phase 2: fine‑tune entire model with memory optimisations + early stopping.\n    Includes:\n      - fused AdamW optimizer\n      - inference_mode in validation\n      - TF32 precision for float32 ops\n      - optimised DataLoader with persistent_workers and prefetch_factor\n    (torch.compile removed to avoid CUDA graph instability on small GPUs)\n    \"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"PHASE 2: FULL FINE‑TUNING (with memory optimisations + early stopping)\")\n    print(\"=\"*60)\n\n    # ----- 5. Enable TF32 for float32 ops -----\n    torch.set_float32_matmul_precision('medium')\n\n    # Move models to device\n    model = model.to(device)\n    arcface = arcface.to(device)\n\n    # Unfreeze backbone\n    for param in model.backbone.parameters():\n        param.requires_grad = True\n\n    # ----- 6. Optimised DataLoader -----\n    from torch.utils.data import DataLoader\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=phase2_batch_size,\n        shuffle=True,\n        num_workers=12,               # adjust based on CPU cores\n        pin_memory=True,\n        persistent_workers=True,       # keep workers alive between epochs\n        prefetch_factor=2               # prefetch 2 batches per worker\n    )\n    print(f\"Created new train loader with batch size = {phase2_batch_size}, \"\n          f\"num_workers=12, persistent_workers=True, prefetch_factor=2\")\n\n    # ----- 2. Fused AdamW optimizer -----\n    optimizer = optim.AdamW([\n        {'params': model.backbone.parameters(), 'lr': lr_backbone},\n        {'params': model.neck.parameters(), 'lr': lr_head},\n        {'params': arcface.parameters(), 'lr': lr_head}\n    ], weight_decay=5e-4, fused=True)        # fused=True for faster kernel\n    scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs_full)\n\n    # Mixed precision scaler\n    scaler = torch.cuda.amp.GradScaler() if use_amp else None\n\n    # ----- 4. Validation with inference_mode -----\n    @torch.inference_mode()\n    def validate(model, val_loader, device):\n        model.eval()\n        all_embeddings = []\n        all_labels = []\n        for images, labels in tqdm(val_loader, desc=\"Validating\"):\n            images = images.to(device)\n            embeddings = model(images)\n            all_embeddings.append(embeddings.cpu().numpy())\n            all_labels.append(labels.numpy())\n        embeddings = np.concatenate(all_embeddings, axis=0)\n        labels = np.concatenate(all_labels, axis=0)\n        balanced_map, per_class_ap = compute_identity_balanced_map(embeddings, labels)\n        return balanced_map, per_class_ap\n\n    best_map = 0.0\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_arcface_wts = copy.deepcopy(arcface.state_dict())\n\n    # Early stopping variables\n    epochs_no_improve = 0\n    best_epoch = 0\n\n    # Clear cache before starting\n    torch.cuda.empty_cache()\n\n    # Optional: validate less frequently (suggestion: uncomment if you want to save time)\n    # validate_every = 2   # validate every 2 epochs\n    # Adjust early_stop_patience accordingly if you change this.\n\n    for epoch in range(1, epochs_full + 1):\n        model.train()\n        arcface.train()\n        running_loss = 0.0\n        num_batches = 0\n        optimizer.zero_grad()\n\n        for i, (images, labels) in enumerate(tqdm(train_loader, desc=f\"Epoch {epoch}/{epochs_full}\")):\n            images, labels = images.to(device), labels.to(device)\n\n            if use_amp:\n                with torch.cuda.amp.autocast():\n                    embeddings = model(images)          # half\n                embeddings = embeddings.float()          # convert to float32 for loss\n                loss = arcface(embeddings, labels)\n                loss = loss / grad_accum\n                scaler.scale(loss).backward()\n            else:\n                embeddings = model(images)\n                loss = arcface(embeddings, labels) / grad_accum\n                loss.backward()\n\n            # Gradient accumulation step\n            if (i + 1) % grad_accum == 0:\n                if use_amp:\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    optimizer.step()\n                optimizer.zero_grad()\n\n            running_loss += loss.item() * grad_accum\n            num_batches += 1\n\n        # Handle any remaining gradients at epoch end\n        if num_batches % grad_accum != 0:\n            if use_amp:\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                optimizer.step()\n            optimizer.zero_grad()\n\n        epoch_loss = running_loss / num_batches\n        scheduler.step()\n\n        # Validation (can be reduced in frequency as per comment above)\n        val_map, per_class_ap = validate(model, val_loader, device)\n        print(f\"Epoch {epoch}: Loss = {epoch_loss:.4f}, Val mAP = {val_map:.4f}\")\n\n        # Optional per‑class AP logging\n        if epoch % 5 == 0 and label_to_id is not None:\n            sorted_ap = sorted(per_class_ap.items(), key=lambda x: x[1])\n            worst_ids = sorted_ap[:5]\n            print(\"  Worst performing classes (AP):\")\n            for label, ap in worst_ids:\n                jaguar_name = label_to_id[label]\n                print(f\"    {jaguar_name}: {ap:.3f}\")\n\n        # Check for improvement\n        if val_map > best_map:\n            best_map = val_map\n            best_epoch = epoch\n            best_model_wts = copy.deepcopy(model.state_dict())\n            best_arcface_wts = copy.deepcopy(arcface.state_dict())\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'arcface_state_dict': arcface.state_dict(),\n                'optimizer': optimizer.state_dict(),\n                'val_map': val_map,\n            }, save_path)\n            print(f\"  -> New best full model! (mAP: {val_map:.4f})\")\n            epochs_no_improve = 0\n        else:\n            epochs_no_improve += 1\n            print(f\"  -> No improvement for {epochs_no_improve} epoch(s). Best was {best_map:.4f} at epoch {best_epoch}.\")\n\n        # Early stopping check\n        if early_stop_patience is not None and epochs_no_improve >= early_stop_patience:\n            print(f\"\\n🛑 Early stopping triggered after {epoch} epochs (no improvement for {early_stop_patience} epochs).\")\n            break\n\n    print(f\"\\n✅ Phase 2 complete. Best mAP: {best_map:.4f} at epoch {best_epoch}\")\n    model.load_state_dict(best_model_wts)\n    arcface.load_state_dict(best_arcface_wts)\n    return model, arcface\n\n\n# Example usage (assuming your model classes are defined)\nif __name__ == \"__main__\":\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Using device: {device}\")\n\n    # Load model and arcface from phase 1 checkpoint\n    checkpoint = torch.load('/kaggle/input/datasets/bubbletide/model-phase-1/best_model_phase1.pth', map_location='cpu', weights_only=False)\n    model = JaguarReIDModel(config, num_classes, embedding_size=512)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    arcface = ArcFaceLoss(num_classes, embedding_size=512, margin=0.3, scale=30)\n    arcface.load_state_dict(checkpoint['arcface_state_dict'])\n\n    # Assume train_dataset, val_loader, num_classes, label_to_id are defined\n    model, arcface = train_phase2(\n        model=model,\n        arcface=arcface,\n        train_dataset=train_dataset,\n        val_loader=val_loader,\n        num_classes=num_classes,\n        device=device,\n        epochs_full=30,\n        lr_head=1e-3,\n        lr_backbone=1e-5,\n        phase2_batch_size=8,\n        grad_accum=4,\n        use_amp=True,\n        label_to_id=label_to_id,\n        save_path='best_model_phase2.pth',\n        early_stop_patience=5\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T13:58:08.356928Z","iopub.execute_input":"2026-02-25T13:58:08.360284Z","iopub.status.idle":"2026-02-25T16:43:53.134858Z","shell.execute_reply.started":"2026-02-25T13:58:08.360253Z","shell.execute_reply":"2026-02-25T16:43:53.131757Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# SIMILARITY PREDICTION & SUBMISSION","metadata":{}},{"cell_type":"code","source":"# After loading the model and before the test loop\ntorch.cuda.empty_cache()\n# Optional: If using other libraries that might use GPU\nimport gc\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T16:50:30.011077Z","iopub.execute_input":"2026-02-25T16:50:30.011734Z","iopub.status.idle":"2026-02-25T16:50:30.324498Z","shell.execute_reply.started":"2026-02-25T16:50:30.011700Z","shell.execute_reply":"2026-02-25T16:50:30.323747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Stage 5: Similarity Prediction & Submission\nSUBMISSION_PATH = 'submission.csv'\nBEST_MODEL_PATH = '/kaggle/input/datasets/bubbletide/model-phase-2/best_model_phase2.pth' \n\n# Use validation transforms (no augmentation)\ntest_transforms = transforms.Compose([\n    transforms.Resize((config.IMG_SIZE, config.IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD)\n])\n\n# Dataset for Test Images\nclass TestDataset(Dataset):\n    \"\"\"Dataset for test images (no labels).\"\"\"\n    def __init__(self, img_dir, transform=None, use_alpha=True):\n        self.img_dir = img_dir\n        self.transform = transform\n        self.use_alpha = use_alpha\n        # Get list of all test images (sorted to ensure consistent ordering)\n        self.image_files = sorted([f for f in os.listdir(img_dir) if f.endswith('.png')])\n        # Optional: filter to only those appearing in test.csv, but all 371 should be present.\n\n    def __len__(self):\n        return len(self.image_files)\n\n    def __getitem__(self, idx):\n        filename = self.image_files[idx]\n        img_path = os.path.join(self.img_dir, filename)\n        img = Image.open(img_path)\n\n        if self.use_alpha:\n            img = img.convert('RGBA')\n            r, g, b, alpha = img.split()\n            bg = Image.new('RGB', img.size, (0, 0, 0))\n            img = Image.composite(img.convert('RGB'), bg, alpha)\n        else:\n            img = img.convert('RGB')\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, filename  # return filename for mapping\n\n# Load Trained Model\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Recreate model architecture (must match training)\n# We need num_classes even for inference? Not really, but model definition requires it.\n# If num_classes is not available, we can load from checkpoint directly.\n# Safer: load checkpoint and extract state_dict, then instantiate model.\n# We'll assume num_classes is known (31 from stage 1).\nnum_classes = 31  # or retrieve from saved mapping\n\nmodel = JaguarReIDModel(config, num_classes=num_classes, embedding_size=512).to(device)\n\n# Load checkpoint\n\ncheckpoint = torch.load(BEST_MODEL_PATH, map_location=device, weights_only=False)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()\nprint(f\"✅ Model loaded from {BEST_MODEL_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:04:21.503677Z","iopub.execute_input":"2026-02-25T17:04:21.504362Z","iopub.status.idle":"2026-02-25T17:04:40.564594Z","shell.execute_reply.started":"2026-02-25T17:04:21.504328Z","shell.execute_reply":"2026-02-25T17:04:40.563920Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract Embeddings for All Test Images\ntest_dataset = TestDataset(\n    img_dir=config.TEST_IMG_DIR,\n    transform=test_transforms,\n    use_alpha=config.USE_ALPHA\n)\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=64,  # adjust based on GPU memory\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=config.PIN_MEMORY\n)\n\nall_embeddings = []\nall_filenames = []\n\nwith torch.no_grad():\n    for images, filenames in tqdm(test_loader, desc=\"Extracting test embeddings\"):\n        images = images.to(device)\n        embeddings = model(images)  # already L2-normalized\n        all_embeddings.append(embeddings.cpu().numpy())\n        all_filenames.extend(filenames)\n\nembeddings = np.concatenate(all_embeddings, axis=0)  # shape (371, embedding_size)\nprint(f\"✅ Extracted embeddings for {len(embeddings)} test images.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:04:42.523326Z","iopub.execute_input":"2026-02-25T17:04:42.524060Z","iopub.status.idle":"2026-02-25T17:05:53.881805Z","shell.execute_reply.started":"2026-02-25T17:04:42.524033Z","shell.execute_reply":"2026-02-25T17:05:53.881040Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create mapping from filename to index\nfilename_to_idx = {fname: idx for idx, fname in enumerate(all_filenames)}\nassert len(filename_to_idx) == 371, f\"Expected 371 images, got {len(filename_to_idx)}\"\n\n# Compute Cosine Similarity Matrix\nprint(\"Computing cosine similarity matrix...\")\n# Since embeddings are L2-normalized, dot product equals cosine similarity\nsimilarity_matrix = np.dot(embeddings, embeddings.T)  # (371, 371)\n# Clip to [0,1] as per competition requirement\nsimilarity_matrix = np.clip(similarity_matrix, 0, 1)\nprint(\"✅ Similarity matrix computed.\")\n\n# Map to Test Pairs and Create Submission\n# Load test.csv to get row_id, query_image, gallery_image\ntest_df = pd.read_csv(config.TEST_CSV)\nprint(f\"Loaded test pairs: {len(test_df)} rows.\")\n\n# Prepare submission DataFrame (copy sample_submission format)\n# If sample_submission.csv is available, we can use it as template.\n# Otherwise, create from scratch.\nsubmission = pd.DataFrame()\nsubmission['row_id'] = test_df['row_id']\n\n# Retrieve similarity for each pair\nsimilarities = []\nfor _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Mapping pairs\"):\n    query_img = row['query_image']\n    gallery_img = row['gallery_image']\n    query_idx = filename_to_idx[query_img]\n    gallery_idx = filename_to_idx[gallery_img]\n    sim = similarity_matrix[query_idx, gallery_idx]\n    similarities.append(sim)\n\nsubmission['similarity'] = similarities\n\n# Final checks\nassert len(submission) == 137270, f\"Expected 137270 rows, got {len(submission)}\"\nassert submission['similarity'].min() >= 0, \"Min similarity < 0\"\nassert submission['similarity'].max() <= 1, \"Max similarity > 1\"\n\n# Save\nsubmission.to_csv(SUBMISSION_PATH, index=False)\nprint(f\"✅ Submission saved to {SUBMISSION_PATH}\")\nprint(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:05:59.804417Z","iopub.execute_input":"2026-02-25T17:05:59.804855Z","iopub.status.idle":"2026-02-25T17:06:04.941102Z","shell.execute_reply.started":"2026-02-25T17:05:59.804823Z","shell.execute_reply":"2026-02-25T17:06:04.940382Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# SAVING ENSEMBLE EQUIPMENTS","metadata":{}},{"cell_type":"code","source":"# CONTINUATION: Save ensemble materials (with val_map from checkpoint)\nimport pickle\n\n# Configuration for saving\n# You can set MODEL_ID manually, e.g., \"seed42\", \"fold1\", etc.\n# Or automatically extract from checkpoint filename (optional)\nMODEL_ID = os.path.splitext(os.path.basename(BEST_MODEL_PATH))[0]  # e.g., \"best_model_phase2\"\n\n# Directory to save ensemble materials\nsave_dir = \"ensemble_materials\"\nos.makedirs(save_dir, exist_ok=True)\n\n# Extract validation mAP from checkpoint\nval_map = checkpoint['val_map']\nprint(f\"✅ Validation mAP extracted from checkpoint: {val_map:.4f}\")\n\n# Save similarity matrix\nsim_path = os.path.join(save_dir, f\"similarity_matrix_{MODEL_ID}.npy\")\nnp.save(sim_path, similarity_matrix)\nprint(f\"✅ Similarity matrix saved to {sim_path}\")\n\n# Save filename-to-index mapping (only once)\nmapping_path = os.path.join(save_dir, \"filename_to_idx.pkl\")\nif not os.path.exists(mapping_path):\n    with open(mapping_path, 'wb') as f:\n        pickle.dump(filename_to_idx, f)\n    print(f\"✅ Filename mapping saved to {mapping_path}\")\nelse:\n    print(f\"⚠️ Mapping file already exists, skipping (delete if you want to overwrite).\")\n\n# Save validation mAP (if available)\nif val_map is not None:\n    val_path = os.path.join(save_dir, f\"val_map_{MODEL_ID}.txt\")\n    with open(val_path, 'w') as f:\n        f.write(f\"{val_map:.6f}\")\n    print(f\"✅ Validation mAP saved to {val_path}\")\n\nprint(\"\\n🎉 Ensemble materials saved. Run this script for each model, changing BEST_MODEL_PATH and (optionally) MODEL_ID.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T17:06:14.672836Z","iopub.execute_input":"2026-02-25T17:06:14.673130Z","iopub.status.idle":"2026-02-25T17:06:14.682229Z","shell.execute_reply.started":"2026-02-25T17:06:14.673106Z","shell.execute_reply":"2026-02-25T17:06:14.681621Z"}},"outputs":[],"execution_count":null}]}