{"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":126777,"databundleVersionId":15314950}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":19421.609166,"end_time":"2026-02-27T18:36:19.874238","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-02-27T13:12:38.265072","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"dae0f33c","cell_type":"code","source":"!pip install pytorch-metric-learning --quiet","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:40.880676Z","iopub.status.busy":"2026-02-27T13:12:40.880391Z","iopub.status.idle":"2026-02-27T13:12:45.726407Z","shell.execute_reply":"2026-02-27T13:12:45.725693Z"},"papermill":{"duration":4.853011,"end_time":"2026-02-27T13:12:45.728077","exception":false,"start_time":"2026-02-27T13:12:40.875066","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8c571761","cell_type":"code","source":"import os\nimport pandas as pd \nfrom PIL import Image\nfrom torch.utils.data import Dataset\nimport torch\nimport numpy as np ","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:45.736784Z","iopub.status.busy":"2026-02-27T13:12:45.736530Z","iopub.status.idle":"2026-02-27T13:12:50.565539Z","shell.execute_reply":"2026-02-27T13:12:50.564972Z"},"papermill":{"duration":4.835375,"end_time":"2026-02-27T13:12:50.567338","exception":false,"start_time":"2026-02-27T13:12:45.731963","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"4c22a184","cell_type":"code","source":"class Config:\n    train_images_path = \"/kaggle/input/jaguar-re-id/train/train\"\n    test_images_path = \"/kaggle/input/jaguar-re-id/test/test\"\n    train_data_csv_path = \"/kaggle/input/jaguar-re-id/train.csv\"\n    test_data_csv_path = \"/kaggle/input/jaguar-re-id/test.csv\"\n    model_save_path = \"/kaggle/working/best_model.pth\"\n    random_seed = 42\n    # PK Sampler parameters\n    P = 16  # number of identities per batch\n    K = 4  # number of images per identity\n    batch_size = P * K  # 32\n    lr=0.001\n    num_epochs=60\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    embedding_dim = 1024\n    Input_size = (256, 256)\n    Arcface_s=64 #64 is the default\n    Arcface_m=0.5 #0.5 is the default\n    \n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:50.576160Z","iopub.status.busy":"2026-02-27T13:12:50.575573Z","iopub.status.idle":"2026-02-27T13:12:50.826997Z","shell.execute_reply":"2026-02-27T13:12:50.826414Z"},"papermill":{"duration":0.257358,"end_time":"2026-02-27T13:12:50.828442","exception":false,"start_time":"2026-02-27T13:12:50.571084","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7a4922cd","cell_type":"code","source":"\nprint(f\"Using device: {Config.device}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:50.837530Z","iopub.status.busy":"2026-02-27T13:12:50.836986Z","iopub.status.idle":"2026-02-27T13:12:50.840965Z","shell.execute_reply":"2026-02-27T13:12:50.840159Z"},"papermill":{"duration":0.009959,"end_time":"2026-02-27T13:12:50.842268","exception":false,"start_time":"2026-02-27T13:12:50.832309","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"16a65a28","cell_type":"code","source":"#train Data set definition\n\nclass JaguarReIDDataset(Dataset):\n    def __init__(self, csv_file, images_path, transform=None):\n        self.data = pd.read_csv(csv_file)\n        self.images_path = images_path\n        self.transform = transform\n        \n        # Create label mapping\n        unique_labels = sorted(self.data['ground_truth'].unique())\n        self.label2idx = {label: idx for idx, label in enumerate(unique_labels)}\n        self.idx2label = {idx: label for label, idx in self.label2idx.items()}\n    \n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        img_name = os.path.join(self.images_path, self.data.iloc[idx]['filename'])\n        image = Image.open(img_name)  # Load as RGBA\n        \n        # Apply alpha mask directly (zero out background, keep jaguar only)\n        if image.mode == 'RGBA':\n            img = np.array(image)\n            \n            rgb = img[:, :, :3]\n            alpha = img[:, :, 3]\n            \n            # Normalize alpha to [0, 1] and apply as mask\n            alpha_normalized = alpha / 255.0\n            \n            # Multiply RGB by alpha to keep only jaguar pixels (background becomes black)\n            rgb_masked = (rgb * alpha_normalized[:, :, np.newaxis]).astype(np.uint8)\n            \n            image = Image.fromarray(rgb_masked)\n        else:\n            image = image.convert('RGB')\n\n\n        label_str = self.data.iloc[idx]['ground_truth']\n        label = self.label2idx[label_str]  # Convert to integer\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label\n    ","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:50.850267Z","iopub.status.busy":"2026-02-27T13:12:50.850046Z","iopub.status.idle":"2026-02-27T13:12:50.857347Z","shell.execute_reply":"2026-02-27T13:12:50.856840Z"},"papermill":{"duration":0.012729,"end_time":"2026-02-27T13:12:50.858564","exception":false,"start_time":"2026-02-27T13:12:50.845835","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c29c4738","cell_type":"code","source":"#test dataset definition\nclass JaguarReIDTestDataset(Dataset):\n    \"\"\"Dataset for unique test images - load each image only once\"\"\"\n    def __init__(self, csv_file, images_path, transform=None):\n        self.data = pd.read_csv(csv_file)\n        self.images_path = images_path\n        self.transform = transform\n        \n        # Extract unique images from pairs\n        query_images = set(self.data['query_image'])\n        gallery_images = set(self.data['gallery_image'])\n        self.unique_images = sorted(query_images | gallery_images)\n        \n    def __len__(self):\n        return len(self.unique_images)  # ~371, not 137,270\n    \n    def __getitem__(self, idx):\n        img_name = os.path.join(self.images_path, self.unique_images[idx])\n        image = Image.open(img_name)  # Load as RGBA\n        \n        # Apply alpha mask directly (zero out background, keep jaguar only)\n        if image.mode == 'RGBA':\n            img = np.array(image)\n            \n            rgb = img[:, :, :3]\n            alpha = img[:, :, 3]\n            \n            # Normalize alpha to [0, 1] and apply as mask\n            alpha_normalized = alpha / 255.0\n            \n            # Multiply RGB by alpha to keep only jaguar pixels (background becomes black)\n            rgb_masked = (rgb * alpha_normalized[:, :, np.newaxis]).astype(np.uint8)\n            \n            image = Image.fromarray(rgb_masked)\n        else:\n            image = image.convert('RGB')\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, self.unique_images[idx]  # Return filename for lookup","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:50.866733Z","iopub.status.busy":"2026-02-27T13:12:50.866512Z","iopub.status.idle":"2026-02-27T13:12:50.873203Z","shell.execute_reply":"2026-02-27T13:12:50.872618Z"},"papermill":{"duration":0.012211,"end_time":"2026-02-27T13:12:50.874498","exception":false,"start_time":"2026-02-27T13:12:50.862287","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a4d9d62f","cell_type":"code","source":"# Define transforms\nfrom torchvision import transforms\n\ntrain_transform = transforms.Compose([\n    transforms.Resize(Config.Input_size),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomRotation(15),\n    transforms.RandomGrayscale(p=0.1),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\ntest_transform = transforms.Compose([\n    transforms.Resize(Config.Input_size),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:50.882383Z","iopub.status.busy":"2026-02-27T13:12:50.882168Z","iopub.status.idle":"2026-02-27T13:12:56.400316Z","shell.execute_reply":"2026-02-27T13:12:56.399318Z"},"papermill":{"duration":5.524134,"end_time":"2026-02-27T13:12:56.402120","exception":false,"start_time":"2026-02-27T13:12:50.877986","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"4d71fbc4","cell_type":"code","source":"# loading data and explore \ntrain_dataset = JaguarReIDDataset(\n    Config.train_data_csv_path, \n    Config.train_images_path,\n    transform=train_transform\n)\nprint(f\"Number of training samples: {len(train_dataset)}\")\n\nprint(f\"Unique individuals in training set: {len(train_dataset.label2idx)}\")\n# Display a few sample images and their labels\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n# Create a temporary dataset without transforms for visualization\ntemp_dataset = JaguarReIDDataset(\n    Config.train_data_csv_path, \n    Config.train_images_path,\n    transform=None\n)\n\nfor i in range(5):\n    image, label = temp_dataset[i]\n    plt.imshow(image)\n    plt.title(f\"Individual ID: {label}\")\n    plt.axis('off')\n    plt.show()\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:56.412253Z","iopub.status.busy":"2026-02-27T13:12:56.411853Z","iopub.status.idle":"2026-02-27T13:12:58.468258Z","shell.execute_reply":"2026-02-27T13:12:58.467545Z"},"papermill":{"duration":2.066569,"end_time":"2026-02-27T13:12:58.472606","exception":false,"start_time":"2026-02-27T13:12:56.406037","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"802a6a56","cell_type":"code","source":"# min, max, median number of images per individual\nindividual_counts = train_dataset.data['ground_truth'].value_counts()\nprint(f\"Min images per individual: {individual_counts.min()}\")\nprint(f\"Max images per individual: {individual_counts.max()}\")\nprint(f\"Median images per individual: {individual_counts.median()}\")\n\n#plot distribution of images per individual\nplt.figure(figsize=(10, 6))\nplt.hist(individual_counts, bins=30, edgecolor='black')\nplt.title('Distribution of Images per Individual')\nplt.xlabel('Number of Images')\nplt.ylabel('Frequency')\nplt.grid(axis='y', alpha=0.75)\nplt.show()\n\n\n\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:58.503515Z","iopub.status.busy":"2026-02-27T13:12:58.503035Z","iopub.status.idle":"2026-02-27T13:12:58.655088Z","shell.execute_reply":"2026-02-27T13:12:58.654397Z"},"papermill":{"duration":0.168314,"end_time":"2026-02-27T13:12:58.656353","exception":false,"start_time":"2026-02-27T13:12:58.488039","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8597867f","cell_type":"code","source":"for i in range(1):\n    image, label = temp_dataset[i]\n    print(Image.open(os.path.join(Config.train_images_path, temp_dataset.data.iloc[i]['filename'])).mode)\n    # plt.imshow(image)\n    # plt.title(f\"Individual ID: {label}\")\n    # plt.axis('off')\n    # plt.show()","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:58.684320Z","iopub.status.busy":"2026-02-27T13:12:58.684087Z","iopub.status.idle":"2026-02-27T13:12:58.859556Z","shell.execute_reply":"2026-02-27T13:12:58.858777Z"},"papermill":{"duration":0.191027,"end_time":"2026-02-27T13:12:58.861048","exception":false,"start_time":"2026-02-27T13:12:58.670021","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a6c1b82e","cell_type":"code","source":"#Model definition\nimport torch.nn as nn\nimport torchvision.models as models\nimport torch.nn.functional as F\n\n\n\nclass JaguarReIDModel(torch.nn.Module):\n    def __init__(self, embedding_dim):\n        super(JaguarReIDModel, self).__init__()\n        # Backbone for feature extraction (ResNet-50)\n        backbone=models.convnext_tiny(pretrained=True)\n\n        # Remove the final classification layer\n        self.backbone = torch.nn.Sequential(*list(backbone.children())[:-2])\n        self.backbone_out_channels=768\n\n        # Global Average Pooling\n        self.pool = torch.nn.AdaptiveAvgPool2d((1, 1))\n\n        # Fully connected layer to get the embedding\n        self.embedding= nn.Sequential(\n            nn.Linear(self.backbone_out_channels, embedding_dim),\n            nn.BatchNorm1d(embedding_dim)\n        )\n\n    def forward(self, x):\n        feats = self.backbone(x)          # [B, 768, H, W]\n        pooled = self.pool(feats)         # [B, 768, 1, 1]\n        pooled = pooled.view(pooled.size(0), -1)  # [B, 768]\n\n        emb = self.embedding(pooled)      # [B, 256]\n        emb = F.normalize(emb, p=2, dim=1)  # L2 normalize\n\n        return emb\n\n\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:58.889576Z","iopub.status.busy":"2026-02-27T13:12:58.889329Z","iopub.status.idle":"2026-02-27T13:12:58.895392Z","shell.execute_reply":"2026-02-27T13:12:58.894834Z"},"papermill":{"duration":0.021554,"end_time":"2026-02-27T13:12:58.896664","exception":false,"start_time":"2026-02-27T13:12:58.875110","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"82d49a42","cell_type":"code","source":"#training loop definition\n\ndef train_one_epoch(model, dataloader, criterion, optimizer, device):\n    model.train()\n    epoch_loss = 0.0\n    for images, labels in dataloader:\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        embeddings = model(images)\n        loss = criterion(embeddings, labels)\n        loss.backward()\n        optimizer.step()\n        \n        epoch_loss += loss.item() * images.size(0)\n    \n    return epoch_loss / len(dataloader.dataset)\n\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:58.925202Z","iopub.status.busy":"2026-02-27T13:12:58.924717Z","iopub.status.idle":"2026-02-27T13:12:58.929150Z","shell.execute_reply":"2026-02-27T13:12:58.928484Z"},"papermill":{"duration":0.020159,"end_time":"2026-02-27T13:12:58.930521","exception":false,"start_time":"2026-02-27T13:12:58.910362","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"96785b12","cell_type":"code","source":"# define data loaders, model, loss function, and optimizer\nfrom torch.utils.data import DataLoader\nfrom pytorch_metric_learning import losses\n\nnum_ids = len(train_dataset.label2idx)\n\nmodel = JaguarReIDModel(Config.embedding_dim).to(Config.device)\ncriterion = losses.ArcFaceLoss(\n    num_classes=num_ids,\n    embedding_size=Config.embedding_dim,\n    margin=Config.Arcface_m,\n    scale=Config.Arcface_s\n)\n#ArcFaceLoss has learnable weights, so we need to include them in the optimizer\noptimizer = torch.optim.Adam(\n    list(model.parameters()) + list(criterion.parameters()), \n    lr=Config.lr\n)\n\n# Learning rate scheduler for better convergence\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, \n    T_max=Config.num_epochs,\n    eta_min=1e-6\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:12:58.958564Z","iopub.status.busy":"2026-02-27T13:12:58.958100Z","iopub.status.idle":"2026-02-27T13:13:01.621217Z","shell.execute_reply":"2026-02-27T13:13:01.620600Z"},"papermill":{"duration":2.678991,"end_time":"2026-02-27T13:13:01.622877","exception":false,"start_time":"2026-02-27T13:12:58.943886","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ccc4ade7","cell_type":"code","source":"#validation loop definition\n# mAP based validation, because we want to save the best model based on the competition metric, not just loss\n\nimport torch\nimport numpy as np\nfrom sklearn.metrics import average_precision_score\n\n@torch.no_grad()\ndef extract_embeddings(model, dataloader, device):\n    \"\"\"Extract embeddings for all images in dataloader\"\"\"\n    model.eval()\n    embeddings_list = []\n    labels_list = []\n    \n    for images, labels in dataloader:\n        images = images.to(device)\n        embeddings = model(images)  # Already L2-normalized in your model\n        embeddings_list.append(embeddings.cpu())\n        labels_list.append(labels)\n    \n    embeddings = torch.cat(embeddings_list, dim=0)  # [N, embedding_dim]\n    labels = torch.cat(labels_list, dim=0)  # [N]\n    \n    return embeddings, labels\n\ndef compute_map(embeddings, labels):\n    \"\"\"\n    Compute identity-balanced mAP\n    - Compute AP per identity (average over all queries of that identity)\n    - Average across identities (not queries)\n    This prevents popular identities from dominating the metric\n    \"\"\"\n    similarities = embeddings @ embeddings.T  # [N, N] cosine similarity\n    \n    unique_labels = labels.unique()\n    ap_per_identity = []\n    \n    for identity in unique_labels:\n        # Get all queries for this identity\n        identity_mask = (labels == identity)\n        identity_indices = torch.where(identity_mask)[0]\n        \n        identity_aps = []\n        \n        for query_idx in identity_indices:\n            # Similarities for this query\n            sim_scores = similarities[query_idx].clone()\n            \n            # Ground truth: same identity or not\n            gt = (labels == identity).float()\n            \n            # Create mask to exclude the query itself\n            mask = torch.ones(len(labels), dtype=torch.bool)\n            mask[query_idx] = False\n            \n            # Filter out the query from both scores and ground truth\n            sim_scores_filtered = sim_scores[mask]\n            gt_filtered = gt[mask]\n            \n            # Skip if no other samples of this identity\n            if gt_filtered.sum() == 0:\n                continue\n            \n            # Sort by similarity (descending)\n            sorted_indices = torch.argsort(sim_scores_filtered, descending=True)\n            y_true = gt_filtered[sorted_indices].numpy()\n            y_scores = sim_scores_filtered[sorted_indices].numpy()\n            \n            # Compute AP for this query\n            ap = average_precision_score(y_true, y_scores)\n            identity_aps.append(ap)\n        \n        # Average AP over all queries of this identity\n        if len(identity_aps) > 0:\n            identity_ap = np.mean(identity_aps)\n            ap_per_identity.append(identity_ap)\n    \n    # Average across identities (not queries!) - this is the key difference\n    balanced_map = float(np.mean(ap_per_identity))\n    return balanced_map\n\n\n@torch.no_grad()\ndef validate_with_map(model, dataloader, device):\n    \"\"\"Validation using mAP instead of loss\"\"\"\n    embeddings, labels = extract_embeddings(model, dataloader, device)\n    mAP = compute_map(embeddings, labels)\n    return mAP\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:13:01.653879Z","iopub.status.busy":"2026-02-27T13:13:01.653474Z","iopub.status.idle":"2026-02-27T13:13:01.894323Z","shell.execute_reply":"2026-02-27T13:13:01.893745Z"},"papermill":{"duration":0.257646,"end_time":"2026-02-27T13:13:01.895989","exception":false,"start_time":"2026-02-27T13:13:01.638343","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b0026b05","cell_type":"code","source":"\n","metadata":{"papermill":{"duration":0.013786,"end_time":"2026-02-27T13:13:01.924131","exception":false,"start_time":"2026-02-27T13:13:01.910345","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e3da490f","cell_type":"code","source":"# Main training loop\nfrom sklearn.model_selection import train_test_split\nfrom pytorch_metric_learning.samplers import MPerClassSampler\n\n# Create two dataset instances: one for training (augmented), one for validation (clean)\ntrain_dataset_augmented = JaguarReIDDataset(\n    Config.train_data_csv_path, \n    Config.train_images_path,\n    transform=train_transform  # With RandomFlip, ColorJitter\n)\n\nval_dataset_clean = JaguarReIDDataset(\n    Config.train_data_csv_path, \n    Config.train_images_path,\n    transform=test_transform  # Clean, deterministic transforms only\n)\n\n# Split indices ONCE - apply to both datasets\ntrain_indices, val_indices = train_test_split(\n    list(range(len(train_dataset_augmented))), \n    test_size=0.2, \n    random_state=Config.random_seed, \n    stratify=train_dataset_augmented.data['ground_truth']\n)\n\n# Create subsets with appropriate transforms\ntrain_subset = torch.utils.data.Subset(train_dataset_augmented, train_indices)\nval_subset = torch.utils.data.Subset(val_dataset_clean, val_indices)  # ✅ Clean transforms\n\n# Get labels for train subset\ntrain_labels = train_dataset_augmented.data.iloc[train_indices]['ground_truth'].map(train_dataset_augmented.label2idx).values\n\n# PK Sampler: ensures P identities with K images each per batch\ntrain_loader = DataLoader(\n    train_subset,\n    batch_size=Config.batch_size,\n    sampler=MPerClassSampler(\n        labels=train_labels,\n        m=Config.K,  # K images per identity\n        length_before_new_iter=len(train_indices)\n    ),\n    num_workers=4,\n    pin_memory=True,\n    drop_last=True  # Drop incomplete batches to maintain P×K structure\n)\n\nval_loader = DataLoader(val_subset, batch_size=Config.batch_size, shuffle=False, pin_memory=True)\n\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, device, num_epochs):\n    best_val_map = 0.0  # Track best mAP, not loss\n    \n    for epoch in range(num_epochs):\n        # Training\n        train_loss = train_one_epoch(model, train_loader, criterion, optimizer, device)\n        \n        # Validation with mAP\n        val_map = validate_with_map(model, val_loader, device)\n        \n        # Get current learning rate\n        current_lr = optimizer.param_groups[0]['lr']\n        \n        print(f\"Epoch {epoch+1}/{num_epochs} - Train Loss: {train_loss:.4f} - Val mAP: {val_map:.4f} - LR: {current_lr:.6f}\")\n        \n        # Save best model based on mAP\n        if val_map > best_val_map:\n            best_val_map = val_map\n            torch.save(model.state_dict(), Config.model_save_path)\n            print(f\"✓ Best model saved with mAP: {best_val_map:.4f}\")\n        \n        # Step the scheduler\n        scheduler.step()\n\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:13:01.953266Z","iopub.status.busy":"2026-02-27T13:13:01.952869Z","iopub.status.idle":"2026-02-27T13:13:01.995549Z","shell.execute_reply":"2026-02-27T13:13:01.994852Z"},"papermill":{"duration":0.059256,"end_time":"2026-02-27T13:13:01.997036","exception":false,"start_time":"2026-02-27T13:13:01.937780","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"badf5486","cell_type":"code","source":"train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, Config.device, Config.num_epochs)\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T13:13:02.026129Z","iopub.status.busy":"2026-02-27T13:13:02.025709Z","iopub.status.idle":"2026-02-27T18:35:01.989559Z","shell.execute_reply":"2026-02-27T18:35:01.988694Z"},"papermill":{"duration":19319.995832,"end_time":"2026-02-27T18:35:02.007068","exception":false,"start_time":"2026-02-27T13:13:02.011236","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"9bb90cde","cell_type":"code","source":"# ========================================\n# TEST INFERENCE PIPELINE\n# ========================================\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:35:02.040239Z","iopub.status.busy":"2026-02-27T18:35:02.039961Z","iopub.status.idle":"2026-02-27T18:35:02.043297Z","shell.execute_reply":"2026-02-27T18:35:02.042738Z"},"papermill":{"duration":0.022081,"end_time":"2026-02-27T18:35:02.044584","exception":false,"start_time":"2026-02-27T18:35:02.022503","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"bbb48738","cell_type":"code","source":"# Load best trained model\nprint(\"Loading best model for inference...\")\nmodel = JaguarReIDModel(Config.embedding_dim).to(Config.device)\nmodel.load_state_dict(torch.load(Config.model_save_path, map_location=Config.device))\nmodel.eval()\nprint(\"✓ Model loaded successfully\")\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:35:02.077228Z","iopub.status.busy":"2026-02-27T18:35:02.077006Z","iopub.status.idle":"2026-02-27T18:35:02.625694Z","shell.execute_reply":"2026-02-27T18:35:02.624765Z"},"papermill":{"duration":0.566641,"end_time":"2026-02-27T18:35:02.627300","exception":false,"start_time":"2026-02-27T18:35:02.060659","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"aeb4e167","cell_type":"code","source":"# Create test dataset and dataloader\ntest_dataset = JaguarReIDTestDataset(\n    Config.test_data_csv_path,\n    Config.test_images_path,\n    transform=test_transform\n)\nprint(f\"Number of unique test images: {len(test_dataset)}\")\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=Config.batch_size,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:35:02.661797Z","iopub.status.busy":"2026-02-27T18:35:02.661269Z","iopub.status.idle":"2026-02-27T18:35:02.762273Z","shell.execute_reply":"2026-02-27T18:35:02.761497Z"},"papermill":{"duration":0.119779,"end_time":"2026-02-27T18:35:02.763789","exception":false,"start_time":"2026-02-27T18:35:02.644010","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"2d630893","cell_type":"code","source":"# Extract embeddings for all test images\n@torch.no_grad()\ndef extract_test_embeddings(model, dataloader, device):\n    \"\"\"Extract and store embeddings for all test images\"\"\"\n    model.eval()\n    embedding_dict = {}\n    \n    for images, filenames in dataloader:\n        images = images.to(device)\n        embeddings = model(images)  # [B, D], already L2-normalized\n        \n        # Store embeddings with filename as key\n        for emb, fname in zip(embeddings.cpu(), filenames):\n            embedding_dict[fname] = emb\n    \n    return embedding_dict\n\nprint(\"Extracting embeddings for test images...\")\ntest_embeddings = extract_test_embeddings(model, test_loader, Config.device)\nprint(f\"✓ Extracted embeddings for {len(test_embeddings)} unique test images\")\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:35:02.797612Z","iopub.status.busy":"2026-02-27T18:35:02.797018Z","iopub.status.idle":"2026-02-27T18:36:09.559296Z","shell.execute_reply":"2026-02-27T18:36:09.558426Z"},"papermill":{"duration":66.797565,"end_time":"2026-02-27T18:36:09.577763","exception":false,"start_time":"2026-02-27T18:35:02.780198","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"5ca41245","cell_type":"code","source":"# Compute similarity for each test pair\nprint(\"Computing pairwise similarities...\")\ntest_pairs = pd.read_csv(Config.test_data_csv_path)\nprint(f\"Number of test pairs: {len(test_pairs)}\")\n\nsimilarities = []\n\nfor idx, row in test_pairs.iterrows():\n    q_img = row[\"query_image\"]\n    g_img = row[\"gallery_image\"]\n    \n    # Get embeddings\n    q_emb = test_embeddings[q_img]\n    g_emb = test_embeddings[g_img]\n    \n    # Compute cosine similarity (dot product since embeddings are normalized)\n    sim = torch.dot(q_emb, g_emb).item()\n    similarities.append(sim)\n    \n    if (idx + 1) % 20000 == 0:\n        print(f\"  Processed {idx + 1}/{len(test_pairs)} pairs...\")\n\nprint(\"✓ All similarities computed\")\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:36:09.611711Z","iopub.status.busy":"2026-02-27T18:36:09.611209Z","iopub.status.idle":"2026-02-27T18:36:16.324580Z","shell.execute_reply":"2026-02-27T18:36:16.323878Z"},"papermill":{"duration":6.732444,"end_time":"2026-02-27T18:36:16.326117","exception":false,"start_time":"2026-02-27T18:36:09.593673","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d9aaed4d","cell_type":"code","source":"# Create submission file\n# Note: Cosine similarity ranges from [-1, 1], need to scale to [0, 1]\nscaled_similarities = [(sim + 1) / 2 for sim in similarities]\n\nsubmission = pd.DataFrame({\n    'row_id': range(len(test_pairs)),  # Competition expects 'row_id'\n    'similarity': scaled_similarities\n})\n\n# Verify submission format\nprint(f\"\\nSubmission shape: {submission.shape}\")\nprint(f\"Similarity range: [{min(scaled_similarities):.4f}, {max(scaled_similarities):.4f}]\")\nprint(f\"Missing values: {submission['similarity'].isna().sum()}\")\nassert submission['similarity'].min() >= 0.0 and submission['similarity'].max() <= 1.0, \"Similarity must be in [0, 1]\"\n\n# Display first few rows\nprint(\"\\nFirst few predictions:\")\nprint(submission.head(10))\n","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:36:16.361007Z","iopub.status.busy":"2026-02-27T18:36:16.360453Z","iopub.status.idle":"2026-02-27T18:36:16.401358Z","shell.execute_reply":"2026-02-27T18:36:16.400525Z"},"papermill":{"duration":0.059602,"end_time":"2026-02-27T18:36:16.402728","exception":false,"start_time":"2026-02-27T18:36:16.343126","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c31fa8b0","cell_type":"code","source":"# Save submission file\nsubmission_path = '/kaggle/working/submission.csv'\nsubmission.to_csv(\n    submission_path,\n    index=False,\n    encoding='utf-8',\n    lineterminator='\\n'\n)\nprint(f\"✓ Submission saved to: {submission_path}\")\n\n# Verify the saved file\nverify_submission = pd.read_csv(submission_path)\nprint(f\"✓ Verified submission file: {len(verify_submission)} rows\")","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:36:16.439004Z","iopub.status.busy":"2026-02-27T18:36:16.438735Z","iopub.status.idle":"2026-02-27T18:36:16.723405Z","shell.execute_reply":"2026-02-27T18:36:16.722656Z"},"papermill":{"duration":0.303605,"end_time":"2026-02-27T18:36:16.724911","exception":false,"start_time":"2026-02-27T18:36:16.421306","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"65fc209f","cell_type":"code","source":"cpu_count = os.cpu_count()","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:36:16.759539Z","iopub.status.busy":"2026-02-27T18:36:16.759123Z","iopub.status.idle":"2026-02-27T18:36:16.762519Z","shell.execute_reply":"2026-02-27T18:36:16.761866Z"},"papermill":{"duration":0.0219,"end_time":"2026-02-27T18:36:16.763885","exception":false,"start_time":"2026-02-27T18:36:16.741985","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"09db095a","cell_type":"code","source":"print(f\"Number of CPU cores available: {cpu_count}\")","metadata":{"execution":{"iopub.execute_input":"2026-02-27T18:36:16.799102Z","iopub.status.busy":"2026-02-27T18:36:16.798499Z","iopub.status.idle":"2026-02-27T18:36:16.802465Z","shell.execute_reply":"2026-02-27T18:36:16.801781Z"},"papermill":{"duration":0.02321,"end_time":"2026-02-27T18:36:16.803864","exception":false,"start_time":"2026-02-27T18:36:16.780654","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"00d7b610","cell_type":"code","source":"","metadata":{"papermill":{"duration":0.016036,"end_time":"2026-02-27T18:36:16.835996","exception":false,"start_time":"2026-02-27T18:36:16.819960","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}