{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":126777,"databundleVersionId":15314950,"isSourceIdPinned":false}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🐆 Jaguar Re-ID: Semi-Supervised Learning with Pseudo-Labels\n\nIn this notebook, we implement a pipeline to expand our training dataset using pseudo-labels derived from the test set. \n \n## 🗺️ Pipeline Strategy:\n1.  **Calibration:** Determine the optimal hash distance threshold using the Train set to ensure zero identity collisions.\n2.  **Pseudo-Labeling:** Find Test images that match Train images using this safe threshold.\n3.  **Data Merge:** Combine Original Train + Pseudo Test data.\n4.  **Leakage Prevention:** Split the merged data using `GroupShuffleSplit` based on duplicate groups.\n5.  **Training:** Prepare data for the model training pipeline.","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\nimport sys\nimport random\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nfrom PIL import Image\nimport pandas as pd\nfrom multiprocessing import Pool, cpu_count\nfrom tqdm import tqdm\nimport imagehash\nfrom scipy.spatial.distance import cdist\nfrom scipy.sparse import csr_matrix\nfrom scipy.sparse.csgraph import connected_components\nfrom sklearn.model_selection import GroupShuffleSplit\nimport matplotlib.patches as patches\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Config\nBASE_DIR = Path(\"/kaggle/input/jaguar-re-id\")\nTRAIN_DIR = BASE_DIR / \"train/train\"\nTEST_DIR = BASE_DIR / \"test/test\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:46:35.132794Z","iopub.execute_input":"2026-02-24T10:46:35.133080Z","iopub.status.idle":"2026-02-24T10:46:36.686424Z","shell.execute_reply.started":"2026-02-24T10:46:35.133055Z","shell.execute_reply":"2026-02-24T10:46:36.685610Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔍 2. Calibration: Finding Optimal Threshold\n \nWe need to find the maximum Hamming distance that still guarantees that two images belong to the same jaguar.\n\nWe scan thresholds from 0 to 25 on the Train set. The optimal threshold is the **highest value that produces 0 collisions** (links between different jaguars).\n","metadata":{}},{"cell_type":"code","source":"def calculate_hash(args):\n    \"\"\"Worker for pHash calculation.\"\"\"\n    fn, img_dir = args\n    try:\n        path = Path(img_dir) / fn\n        with Image.open(path) as img:\n            if img.mode == \"RGBA\": img = img.convert(\"RGB\")\n            h = imagehash.phash(img) # default hash_size=8\n            return fn, h.hash.flatten().astype(np.int8)\n    except Exception:\n        return fn, None\n\nprint(\"🔍 Step 1: Calibrating Threshold on Train Set...\")\ntrain_df = pd.read_csv(BASE_DIR / \"train.csv\")\ntrain_args = [(f, str(TRAIN_DIR)) for f in train_df['filename']]\n\nwith Pool(cpu_count()) as pool:\n    results = list(tqdm(pool.imap(calculate_hash, train_args), total=len(train_df)))\n\nvalid = [r for r in results if r[1] is not None]\nfn_train = [r[0] for r in valid]\nhash_train = np.array([r[1] for r in valid], dtype=np.int8)\nlabel_map = train_df.set_index('filename')['ground_truth'].to_dict()\nlabels_train = np.array([label_map[f] for f in fn_train])\n\nprint(\"Computing distance matrix...\")\ndist_matrix = cdist(hash_train, hash_train, 'hamming') * hash_train.shape[1]\n\n# Find optimal threshold\noptimal_thr = 0\nfor t in range(0, 25):\n    triu_idx = np.triu_indices(len(fn_train), k=1)\n    mask = (dist_matrix[triu_idx] <= t)\n    \n    match_labels_0 = labels_train[triu_idx[0]][mask]\n    match_labels_1 = labels_train[triu_idx[1]][mask]\n    \n    collisions = (match_labels_0 != match_labels_1).sum()\n    \n    if collisions == 0:\n        optimal_thr = t\n    else:\n        break\n\nprint(f\"✅ Optimal Safe Threshold determined: {optimal_thr}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:46:36.687865Z","iopub.execute_input":"2026-02-24T10:46:36.688722Z","iopub.status.idle":"2026-02-24T10:51:13.354721Z","shell.execute_reply.started":"2026-02-24T10:46:36.688689Z","shell.execute_reply":"2026-02-24T10:51:13.353855Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔗 3. Pseudo-Labeling: Train <-> Test\n\nNow we apply the calibrated threshold to find Test images that are near-duplicates of Train images.\n\nWe build a global graph connecting all images (Train + Test) and extract connected components.","metadata":{}},{"cell_type":"code","source":"print(\"🔨 Step 2: Hashing Test Set...\")\ntest_df = pd.read_csv(BASE_DIR / \"test.csv\")\ntest_files = sorted(set(test_df['query_image']) | set(test_df['gallery_image']))\ntest_args = [(f, str(TEST_DIR)) for f in test_files]\n\nwith Pool(cpu_count()) as pool:\n    results_test = list(tqdm(pool.imap(calculate_hash, test_args), total=len(test_files)))\n\nvalid_test = [r for r in results_test if r[1] is not None]\nfn_test = [r[0] for r in valid_test]\nhash_test = np.array([r[1] for r in valid_test], dtype=np.int8)\n\n# Merge all data\nall_fns = fn_train + fn_test\nall_hashes = np.vstack([hash_train, hash_test])\ntrain_set_global = set(fn_train)\n\nprint(\"🔗 Building Global Graph...\")\ndist_all = cdist(all_hashes, all_hashes, 'hamming') * all_hashes.shape[1]\nadjacency = (dist_all <= optimal_thr)\nnp.fill_diagonal(adjacency, False)\n\ngraph = csr_matrix(adjacency)\nn_components, comp_labels = connected_components(csgraph=graph, directed=False)\n\npseudo_labels = {}\ngroups_data = []\n\nfor i in range(n_components):\n    indices = np.where(comp_labels == i)[0]\n    group_files = [all_fns[idx] for idx in indices]\n    \n    train_in_group = [f for f in group_files if f in train_set_global]\n    test_in_group = [f for f in group_files if f not in train_set_global]\n    \n    if train_in_group and test_in_group:\n        # We found a match!\n        anchor = train_in_group[0]\n        label = label_map[anchor]\n        \n        # Save for visualization\n        groups_data.append({\n            'label': label,\n            'train_anchor': anchor,\n            'test_candidates': test_in_group\n        })\n        \n        # Assign pseudo-labels\n        for f in test_in_group:\n            pseudo_labels[f] = label\n\nprint(f\"💎 Total Pseudo-Labels found: {len(pseudo_labels)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:51:13.355973Z","iopub.execute_input":"2026-02-24T10:51:13.356274Z","iopub.status.idle":"2026-02-24T10:52:02.578295Z","shell.execute_reply.started":"2026-02-24T10:51:13.356247Z","shell.execute_reply":"2026-02-24T10:52:02.577551Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🖼️ 4. Visual Verification of Pseudo-Labels\n \nLet's ensure the quality of our matches. \n- **Top Row (Green):** Train Anchor (Source of Truth).\n- **Bottom Rows (Blue):** Test Candidates (Assigned the same label).","metadata":{}},{"cell_type":"code","source":"def visualize_pseudo_groups(groups, n_show=10):\n    # Sort by group size (largest first to check risky cases)\n    groups = sorted(groups, key=lambda x: len(x['test_candidates']), reverse=True)\n    \n    for grp in groups[:n_show]:\n        n_test = len(grp['test_candidates'])\n        plt.figure(figsize=(4 * max(n_test, 1), 6))\n        grid = plt.GridSpec(2, n_test if n_test > 0 else 1, hspace=0.3)\n        \n        # Plot Train Anchor\n        try:\n            img = Image.open(TRAIN_DIR / grp['train_anchor'])\n            ax = plt.subplot(grid[0, 0])\n            ax.imshow(img)\n            ax.set_title(f\"TRAIN ANCHOR\\n{grp['label']}\", color='green', fontsize=12, weight='bold')\n            ax.axis('off')\n        except: pass\n        \n        # Plot Test Candidates\n        for idx, tfn in enumerate(grp['test_candidates']):\n            try:\n                img = Image.open(TEST_DIR / tfn)\n                ax = plt.subplot(grid[1, idx])\n                ax.imshow(img)\n                ax.set_title(f\"TEST\\n{tfn[:10]}...\", color='blue')\n                ax.axis('off')\n            except: pass\n        \n        plt.suptitle(f\"Group Size: {n_test + 1}\", fontsize=14)\n        plt.show()\n\nvisualize_pseudo_groups(groups_data, n_show=10)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:52:02.580279Z","iopub.execute_input":"2026-02-24T10:52:02.580529Z","iopub.status.idle":"2026-02-24T10:52:29.657472Z","shell.execute_reply.started":"2026-02-24T10:52:02.580504Z","shell.execute_reply":"2026-02-24T10:52:29.656799Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🗂️ 5. Data Merging & Group ID Assignment\n \nWe merge the original train data with the newly found pseudo-labeled test data. \n\nWe also assign a `group_id` to every image. This ID is crucial for the next step (Splitting) to prevent data leakage.","metadata":{}},{"cell_type":"code","source":"pseudo_df = pd.DataFrame(list(pseudo_labels.items()), columns=['filename', 'ground_truth'])\npseudo_df['is_pseudo'] = True\ntrain_df['is_pseudo'] = False\n\nmerged_df = pd.concat([train_df, pseudo_df], ignore_index=True)\nprint(f\"📦 Merged Dataset size: {len(merged_df)} (Original: {len(train_df)}, Pseudo: {len(pseudo_df)})\")\n\n# 2. Re-calculate Groups for Merged Set\n# Since we already have the graph components (`comp_labels`), we can map them directly.\n# Map filename -> component index\nfn_to_comp = {}\nfor i, fn in enumerate(all_fns):\n    fn_to_comp[fn] = comp_labels[i]\n\n# Assign group_id. \n# Note: Files not in `all_fns` (rare errors) get unique ID\nmerged_df['group_id'] = merged_df['filename'].map(fn_to_comp).fillna(merged_df['filename'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:52:29.658467Z","iopub.execute_input":"2026-02-24T10:52:29.658769Z","iopub.status.idle":"2026-02-24T10:52:29.675223Z","shell.execute_reply.started":"2026-02-24T10:52:29.658739Z","shell.execute_reply":"2026-02-24T10:52:29.674606Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 👀 6. Deep Dive: Merged Group Structure\n\nVisualizing the largest groups in the merged dataset to ensure that \"Train\" and \"Test\" images in the same group are actually the same jaguar.\n- **Green Border:** Train Image.\n- **Blue Border:** Test Image (Pseudo).","metadata":{}},{"cell_type":"code","source":"def visualize_mixed_groups(df, train_dir, test_dir, train_fns_set, n_groups=5):\n    # Filter groups with size > 1\n    group_ids = df['group_id'].unique()\n    \n    # Prepare groups list\n    groups_list = []\n    for gid in group_ids:\n        grp = df[df['group_id'] == gid]\n        if len(grp) > 1:\n            groups_list.append(grp)\n            \n    # Sort by size descending (check biggest merges first)\n    groups_list.sort(key=lambda x: len(x), reverse=True)\n    \n    print(f\"Displaying top {min(n_groups, len(groups_list))} largest groups...\")\n    \n    for grp in groups_list[:n_groups]:\n        files = grp['filename'].tolist()\n        label = grp['ground_truth'].iloc[0]\n        \n        print(f\"Group Label: {label} | Size: {len(files)}\")\n        \n        cols = min(len(files), 5) # Show max 5 images\n        rows = (len(files) + cols - 1) // cols\n        \n        fig, axs = plt.subplots(rows, cols, figsize=(cols*3, rows*3))\n        if rows == 1 and cols == 1: axs = np.array([[axs]])\n        else: axs = np.array(axs).reshape(rows, cols)\n        \n        axs = axs.flatten()\n        \n        for i, fn in enumerate(files):\n            if i >= len(axs): break\n            \n            # Load image from correct directory\n            is_train = fn in train_fns_set\n            path = train_dir / fn if is_train else test_dir / fn\n            \n            try:\n                img = Image.open(path)\n                axs[i].imshow(img.convert('RGB'))\n            except:\n                axs[i].text(0.5, 0.5, \"Missing\", ha='center')\n            \n            # Styling\n            axs[i].axis('off')\n            title_color = 'green' if is_train else 'blue'\n            border_color = 'lime' if is_train else 'blue'\n            \n            rect = patches.Rectangle((0,0), 1, 1, transform=axs[i].transAxes, \n                                      fill=False, edgecolor=border_color, linewidth=4)\n            axs[i].add_patch(rect)\n            axs[i].set_title(f\"{'TRAIN' if is_train else 'TEST'}\\n{fn[:10]}...\", fontsize=9, color=title_color)\n            \n        # Hide unused\n        for i in range(len(files), len(axs)): axs[i].axis('off')\n            \n        plt.suptitle(f\"Merged Group Verification\", fontsize=14)\n        plt.tight_layout()\n        plt.show()\n        print(\"-\" * 60)\n\ntrain_fns_set = set(train_df['filename'])\nvisualize_mixed_groups(merged_df, TRAIN_DIR, TEST_DIR, train_fns_set, n_groups=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:52:29.675975Z","iopub.execute_input":"2026-02-24T10:52:29.676266Z","iopub.status.idle":"2026-02-24T10:53:51.284482Z","shell.execute_reply.started":"2026-02-24T10:52:29.676240Z","shell.execute_reply":"2026-02-24T10:53:51.283707Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 📊 7. Statistics: Merged Dataset Composition\n \nLet's see how the \"Unique vs Duplicates\" balance looks after adding Pseudo-Labels.\n","metadata":{}},{"cell_type":"code","source":"merged_df['group_size'] = merged_df.groupby('group_id')['filename'].transform('count')\nmerged_df['is_duplicate'] = merged_df['group_size'] > 1\n\nstats_merged = merged_df.groupby('ground_truth').agg(\n    total=('filename', 'count'),\n    duplicates=('is_duplicate', 'sum')\n).reset_index()\n\nstats_merged['unique'] = stats_merged['total'] - stats_merged['duplicates']\nstats_merged = stats_merged.sort_values('total', ascending=False)\n\nplt.figure(figsize=(16, 8))\nplt.bar(stats_merged['ground_truth'], stats_merged['unique'], label='Unique Images', color='skyblue')\nplt.bar(stats_merged['ground_truth'], stats_merged['duplicates'], bottom=stats_merged['unique'], label='Duplicates', color='orange')\nplt.title(\"Merged Dataset Composition (Train + Pseudo)\", fontsize=16)\nplt.xlabel(\"Jaguar ID\")\nplt.ylabel(\"Count\")\nplt.xticks(rotation=90)\nplt.legend()\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:53:51.285757Z","iopub.execute_input":"2026-02-24T10:53:51.285999Z","iopub.status.idle":"2026-02-24T10:53:51.894190Z","shell.execute_reply.started":"2026-02-24T10:53:51.285979Z","shell.execute_reply":"2026-02-24T10:53:51.893447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ✂️ 8. Clean Split Strategy: GroupShuffleSplit\n \nWe split the merged dataset into Train and Val.\n \n**Critical:** We split by `group_id` to ensure that if a Train image is in the Train set, its Test duplicates are NOT in the Val set (and vice versa).","metadata":{}},{"cell_type":"code","source":"gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=21)\ntrain_idx, val_idx = next(gss.split(merged_df, groups=merged_df['group_id']))\n\ntrain_split = merged_df.iloc[train_idx].reset_index(drop=True)\nval_split = merged_df.iloc[val_idx].reset_index(drop=True)\n\n# Leakage Check\nval_groups = set(val_split['group_id'])\ntrain_groups = set(train_split['group_id'])\nleakage = train_groups.intersection(val_groups)\n\nprint(f\"✅ Train set size: {len(train_split)}\")\nprint(f\"✅ Val set size:   {len(val_split)}\")\nprint(f\"🛡️ Leakage Check: {len(leakage)} shared groups (Must be 0)\")\n\nif len(leakage) > 0:\n    print(\"❌ WARNING: Leakage detected!\")\nelse:\n    print(\"✅ SUCCESS: Data is clean and ready for training.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:53:51.895301Z","iopub.execute_input":"2026-02-24T10:53:51.895535Z","iopub.status.idle":"2026-02-24T10:53:51.906568Z","shell.execute_reply.started":"2026-02-24T10:53:51.895516Z","shell.execute_reply":"2026-02-24T10:53:51.905932Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 📉 9. Final Split Statistics\n \nDistribution of classes in the final Train and Validation sets.","metadata":{}},{"cell_type":"code","source":"def get_split_stats(df_split):\n    stats = df_split.groupby('ground_truth').agg(\n        total=('filename', 'count'),\n        dups=('is_duplicate', 'sum')\n    ).reset_index()\n    stats['unique'] = stats['total'] - stats['dups']\n    stats = stats.sort_values('total', ascending=False)\n    return stats\n\ntrain_stats = get_split_stats(train_split)\nval_stats = get_split_stats(val_split)\n\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(24, 8), sharey=True)\n\n# Train Plot\nax1.bar(train_stats['ground_truth'], train_stats['unique'], label='Unique', color='skyblue')\nax1.bar(train_stats['ground_truth'], train_stats['dups'], bottom=train_stats['unique'], label='Duplicates', color='orange')\nax1.set_title('TRAIN SET Distribution', fontsize=16)\nax1.set_ylabel('Image Count', fontsize=12)\nax1.tick_params(axis='x', rotation=90)\n\n# Val Plot\nax2.bar(val_stats['ground_truth'], val_stats['unique'], label='Unique', color='skyblue')\nax2.bar(val_stats['ground_truth'], val_stats['dups'], bottom=val_stats['unique'], label='Duplicates', color='orange')\nax2.set_title('VALIDATION SET Distribution', fontsize=16)\nax2.tick_params(axis='x', rotation=90)\n\n# Common Legend\nhandles, labels = ax1.get_legend_handles_labels()\nfig.legend(handles, labels, loc='upper center', bbox_to_anchor=(0.5, 1.02), ncol=2, fontsize=14)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:53:51.907591Z","iopub.execute_input":"2026-02-24T10:53:51.907956Z","iopub.status.idle":"2026-02-24T10:53:52.442787Z","shell.execute_reply.started":"2026-02-24T10:53:51.907924Z","shell.execute_reply":"2026-02-24T10:53:52.442213Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 🚀 Part 2: Training & Inference (Baseline)","metadata":{}},{"cell_type":"code","source":"import math\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport torch.optim as optim\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint, TQDMProgressBar\nfrom pytorch_lightning.loggers import CSVLogger\nimport gc\n\n# --- CONFIG ---\nclass Config:\n    seed = 21\n    BASE_DIR = Path(\"/kaggle/input/jaguar-re-id\") \n    TRAIN_DIR = BASE_DIR / \"train/train\"\n    TEST_DIR = BASE_DIR / \"test/test\"\n    \n    model_name = \"eva02_base_patch14_448.mim_in22k_ft_in22k_in1k\"\n    img_size = 512\n    embedding_dim = 1024\n    \n    # Training\n    epochs = 11\n    batch_size = 4\n    grad_accum = 4\n    lr = 2e-5\n    weight_decay = 1e-3\n    arcface_s = 30.0\n    arcface_m = 0.5\n    \n    num_workers = 4\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# --- SEED & TRANSFORMS ---\npl.seed_everything(Config.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T11:05:43.406861Z","iopub.execute_input":"2026-02-24T11:05:43.407594Z","iopub.status.idle":"2026-02-24T11:05:43.416850Z","shell.execute_reply.started":"2026-02-24T11:05:43.407565Z","shell.execute_reply":"2026-02-24T11:05:43.416223Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧠 Model & Dataset Definitions","metadata":{}},{"cell_type":"code","source":"def visualize_batch(df, dirs, transform=None, n_samples=8):\n    \"\"\"Visualizes a batch of images after applying transformations.\"\"\"\n    dataset = JaguarDataset(df, dirs, transform=transform, is_test=False)\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    \n    fig, axes = plt.subplots(2, n_samples // 2, figsize=(16, 6))\n    axes = axes.flatten()\n    \n    for i in range(n_samples):\n        idx = np.random.randint(0, len(dataset))\n        img_tensor, label_idx = dataset[idx]\n        \n        # Denormalize: (img * std) + mean\n        img_np = img_tensor.permute(1, 2, 0).numpy()\n        img_np = (img_np * std) + mean\n        img_np = np.clip(img_np, 0, 1)\n        \n        # Get label name from dataframe\n        label_name = df['ground_truth'].iloc[idx]\n        \n        axes[i].imshow(img_np)\n        axes[i].set_title(f\"Label: {label_name}\\nIdx: {idx}\", fontsize=9)\n        axes[i].axis('off')\n        \n    plt.suptitle(\"Training Samples (After Augmentations)\", fontsize=14)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T11:05:45.994801Z","iopub.execute_input":"2026-02-24T11:05:45.995087Z","iopub.status.idle":"2026-02-24T11:05:46.001325Z","shell.execute_reply.started":"2026-02-24T11:05:45.995063Z","shell.execute_reply":"2026-02-24T11:05:46.000600Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mean, std = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]\n\ntrain_tf = A.Compose([\n    A.LongestMaxSize(max_size=Config.img_size, interpolation=cv2.INTER_CUBIC),\n    A.PadIfNeeded(min_height=Config.img_size, min_width=Config.img_size, border_mode=0, value=(0,0,0)),\n    A.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.1, p=0.7),\n    A.ToGray(p=0.1),\n    A.HorizontalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5, border_mode=0, value=(0,0,0)),\n    A.CoarseDropout(max_holes=1, max_height=int(Config.img_size * 0.1), max_width=int(Config.img_size * 0.1), fill_value=0, p=0.3),\n    A.Normalize(mean=mean, std=std), ToTensorV2(),\n])\n\ntest_tf = A.Compose([\n    A.LongestMaxSize(max_size=Config.img_size, interpolation=cv2.INTER_CUBIC),\n    A.PadIfNeeded(min_height=Config.img_size, min_width=Config.img_size, border_mode=0, value=(0,0,0)),\n    A.Normalize(mean=mean, std=std), ToTensorV2()\n])\n\n# --- DATASET & SAMPLER ---\nclass JaguarDataset(Dataset):\n    def __init__(self, df, dirs, transform=None, is_test=False):\n        self.df = df.reset_index(drop=True)\n        self.dirs = dirs if isinstance(dirs, list) else [dirs]\n        self.transform = transform\n        self.is_test = is_test\n\n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        fn = row['filename']\n        final_img = None\n        \n        for d in self.dirs:\n            p = d / fn\n            if p.exists():\n                try:\n                    with Image.open(p) as img:\n                        if img.mode == 'RGBA':\n                            bg = Image.new(\"RGB\", img.size, (0, 0, 0))\n                            bg.paste(img, mask=img.split()[3])\n                            img = bg\n                        else: img = img.convert(\"RGB\")\n                        final_img = np.array(img)\n                        break\n                except: pass\n        \n        if final_img is None: final_img = np.zeros((Config.img_size, Config.img_size, 3), dtype=np.uint8)\n        if self.transform: final_img = self.transform(image=final_img)['image']\n        if self.is_test: return final_img, fn\n        return final_img, torch.tensor(row['label_idx'], dtype=torch.long)\n\ndef create_sampler(df):\n    class_counts = df['ground_truth'].value_counts()\n    class_weights = {k: 1.0 / v for k, v in class_counts.items()}\n    sample_weights = df['ground_truth'].map(class_weights).values\n    return WeightedRandomSampler(weights=sample_weights, num_samples=len(df), replacement=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T11:05:48.484509Z","iopub.execute_input":"2026-02-24T11:05:48.484802Z","iopub.status.idle":"2026-02-24T11:05:48.502928Z","shell.execute_reply.started":"2026-02-24T11:05:48.484777Z","shell.execute_reply":"2026-02-24T11:05:48.502197Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super().__init__()\n        self.p = nn.Parameter(torch.ones(1) * p)\n        self.eps = eps\n    def forward(self, x):\n        return F.avg_pool2d(F.relu(x).clamp(min=self.eps).pow(self.p), (x.size(-2), x.size(-1))).pow(1.0 / self.p)\n\nclass ArcFaceLayer(nn.Module):\n    def __init__(self, in_f, out_f, s=30.0, m=0.5):\n        super().__init__()\n        self.s, self.m = s, m\n        self.weight = nn.Parameter(torch.FloatTensor(out_f, in_f))\n        nn.init.xavier_uniform_(self.weight)\n        self.cos_m, self.sin_m = math.cos(m), math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n\n    def forward(self, input, label):\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt((1.0 - torch.pow(cosine, 2)).clamp(0, 1))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        one_hot = torch.zeros_like(cosine).scatter_(1, label.view(-1, 1).long(), 1)\n        return ((one_hot * phi) + ((1.0 - one_hot) * cosine)) * self.s\n\n# --- LIGHTNING MODULE ---\nclass JaguarLightning(pl.LightningModule):\n    def __init__(self, num_classes, train_df, val_df, config):\n        super().__init__()\n        self.save_hyperparameters(ignore=['train_df', 'val_df'])\n        \n        self.config = config\n        self.num_classes = num_classes\n        \n        # Model\n        self.backbone = timm.create_model(\n            config.model_name, \n            pretrained=True, \n            num_classes=0, \n            global_pool='',\n            img_size=self.config.img_size\n        )\n        self.feat_dim = self.backbone.num_features\n        self.neck = nn.Sequential(nn.Linear(self.feat_dim, config.embedding_dim), nn.BatchNorm1d(config.embedding_dim))\n        self.gem = GeM()\n        self.head = ArcFaceLayer(config.embedding_dim, num_classes, config.arcface_s, config.arcface_m)\n        \n        # Loss\n        self.criterion = nn.CrossEntropyLoss()\n        \n        # Data refs for dataloaders\n        self.train_df = train_df\n        self.val_df = val_df\n        \n        # Storage for validation embeddings\n        self.val_embeddings = []\n        self.val_labels = []\n\n    def forward(self, x, label=None):\n        features = self.backbone(x)\n        \n        # Universal Logic: Swin gives 4D, ViT gives 3D\n        if features.dim() == 4:\n            if features.shape[1] != self.feat_dim: features = features.permute(0, 3, 1, 2)\n            emb = self.gem(features).flatten(1)\n        elif features.dim() == 3:\n            B, N, C = features.shape\n            h = w = int(math.sqrt(N))\n            if h*w == N: features = features.transpose(1, 2).reshape(B, C, h, w)\n            elif (h*h)+1 == N: features = features[:, 1:, :].transpose(1, 2).reshape(B, C, h, w)\n            else: features = features.mean(dim=1)\n            \n            if features.dim() == 4: emb = self.gem(features).flatten(1)\n            else: emb = features\n        else:\n            emb = features\n            \n        emb = self.neck(emb)\n        \n        if label is not None:\n            return self.head(emb, label)\n        return F.normalize(emb)\n\n    def training_step(self, batch, batch_idx):\n        imgs, lbls = batch\n        outputs = self(imgs, lbls)\n        loss = self.criterion(outputs, lbls)\n        \n        # Logging\n        self.log('train_loss', loss, prog_bar=True)\n            \n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        imgs, lbls = batch\n        # Get embeddings (normalized)\n        emb = self(imgs)\n        self.val_embeddings.append(emb)\n        self.val_labels.append(lbls)\n\n    def on_validation_epoch_end(self):\n        # Concatenate all embeddings\n        embeddings = torch.cat(self.val_embeddings)\n        labels = torch.cat(self.val_labels)\n\n        # IMPORTANT: Print status to reset Kaggle's IOPub timeout timer\n        print(f\"⏳ Validating epoch {self.current_epoch}...\", flush=True)\n        \n        mAP, top1 = self.calculate_metrics_fast(embeddings, labels)\n\n        msg = f\"✅ Epoch {self.current_epoch} COMPLETE | val_mAP: {mAP:.4f} | val_acc: {top1:.4f}\"\n        print(msg, flush=True)\n        sys.stdout.flush()\n        \n        self.log('val_mAP', mAP, prog_bar=True, logger=True)\n        self.log('val_acc', top1, prog_bar=True)\n        \n        # Clear memory\n        self.val_embeddings.clear()\n        self.val_labels.clear()\n\n    def configure_optimizers(self):\n        optimizer = optim.AdamW(self.parameters(), lr=self.config.lr, weight_decay=self.config.weight_decay)\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=self.config.epochs)\n        return [optimizer], [scheduler]\n\n    def train_dataloader(self):\n        ds = JaguarDataset(self.train_df, [self.config.TRAIN_DIR, self.config.TEST_DIR], transform=train_tf)\n        return DataLoader(ds, batch_size=self.config.batch_size, sampler=create_sampler(self.train_df),\n                          num_workers=self.config.num_workers, pin_memory=True, drop_last=True)\n\n    def val_dataloader(self):\n        ds = JaguarDataset(self.val_df, [self.config.TRAIN_DIR, self.config.TEST_DIR], transform=test_tf)\n        return DataLoader(ds, batch_size=self.config.batch_size, shuffle=False,\n                          num_workers=self.config.num_workers, pin_memory=True)\n\n    # Metrics Helper\n    def calculate_metrics_fast(self, embeds, labels):\n        \"\"\"\n        Calculates retrieval metrics (mAP and Top-1 Accuracy)\n        \"\"\"\n        # 1. Compute Cosine Similarity Matrix (N x N)\n        # Embeddings are assumed to be normalized\n        sims = embeds @ embeds.t()\n        \n        # 2. Sort indices by similarity\n        sorted_indices = sims.argsort(dim=1, descending=True)\n        \n        # --- Top-1 Accuracy Calculation ---\n        # Get the labels of the top-1 prediction (excluding the query itself at index 0)\n        top1_preds = labels[sorted_indices[:, 1]]\n        top1_acc = (top1_preds == labels).float().mean().item()\n\n        # --- mAP Calculation (Vectorized) ---\n        # Create a mask where ground truth labels match\n        matches = (labels[sorted_indices] == labels.unsqueeze(1)).float()\n        \n        # Exclude the diagonal (self-matches) for correct AP calculation.\n        # Column 0 is the query itself (sim=1.0), so we exclude it.\n        matches = matches[:, 1:]\n        \n        # Range of ranks (1 to N-1)\n        ranks = torch.arange(1, matches.size(1) + 1, device=embeds.device).float().unsqueeze(0)\n        \n        # Cumulative sum of relevant items found\n        relevant_counts = matches.cumsum(dim=1)\n        \n        # Precision at each rank: (Relevant Found) / (Rank)\n        precisions = relevant_counts / ranks\n        \n        # Average Precision = Sum(Precision@k * Match@k) / Total Relevant\n        # Total relevant items per query is the sum of matches\n        num_relevant = matches.sum(dim=1, keepdim=True).clamp(min=1e-6)\n        ap_per_image = (precisions * matches).sum(dim=1) / num_relevant.squeeze()\n        \n        # Mean Average Precision\n        mAP = ap_per_image.mean().item()\n        \n        return mAP, top1_acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T11:05:51.848841Z","iopub.execute_input":"2026-02-24T11:05:51.849234Z","iopub.status.idle":"2026-02-24T11:05:51.871969Z","shell.execute_reply.started":"2026-02-24T11:05:51.849192Z","shell.execute_reply":"2026-02-24T11:05:51.871220Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🏋️ Model training","metadata":{}},{"cell_type":"code","source":"# Prepare Labels\nall_labels = sorted(merged_df['ground_truth'].unique())\nlabel_to_idx = {l: i for i, l in enumerate(all_labels)}\nConfig.num_classes = len(all_labels)\nmerged_df['label_idx'] = merged_df['ground_truth'].map(label_to_idx)\ntrain_split['label_idx'] = train_split['ground_truth'].map(label_to_idx)\nval_split['label_idx'] = val_split['ground_truth'].map(label_to_idx)\n\nvisualize_batch(\n    train_split, \n    [Config.TRAIN_DIR, Config.TEST_DIR], \n    transform=train_tf,\n    n_samples=6\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T11:05:57.623185Z","iopub.execute_input":"2026-02-24T11:05:57.623788Z","iopub.status.idle":"2026-02-24T11:06:00.526646Z","shell.execute_reply.started":"2026-02-24T11:05:57.623758Z","shell.execute_reply":"2026-02-24T11:06:00.525937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize Lightning Model\nmodel = JaguarLightning(\n    num_classes=Config.num_classes, \n    train_df=train_split, \n    val_df=val_split, \n    config=Config\n)\n\n# Callbacks\ncheckpoint_callback = ModelCheckpoint(\n    monitor='val_mAP',\n    filename='best_model-{epoch:02d}-{val_mAP:.4f}',\n    save_top_k=1,\n    mode='max',\n    save_weights_only=True\n)\n\nclass KaggleStdoutProgressBar(TQDMProgressBar):\n    \"\"\"\n    Fixes: `CellTimeoutError: A cell timed out while it was being executed, after 4 seconds.`\n    Cause: Kaggle kills the kernel if no output is printed to IOPub stdout for ~4 seconds.\n    \n    Solution:\n    1. Redirect tqdm output to `sys.stdout` (default is stderr).\n    2. Force `sys.stdout.flush()` after every batch update to ensure immediate output.\n    \"\"\"\n    def __init__(self):\n        # refresh_rate=1 ensures the bar attempts to update every batch\n        super().__init__(refresh_rate=1)\n\n    def init_train_tqdm(self) -> tqdm:\n        # file=sys.stdout: Redirects output to standard output stream (like print)\n        bar = tqdm(\n            desc=\"Training\", \n            position=0, \n            leave=True, \n            file=sys.stdout, \n            dynamic_ncols=True\n        )\n        return bar\n\n    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):\n        super().on_train_batch_end(trainer, pl_module, outputs, batch, batch_idx)\n        # Force flush: Sends the output to Kaggle logs immediately, resetting the timeout timer\n        sys.stdout.flush()\n\n    def init_validation_tqdm(self) -> tqdm:\n        bar = tqdm(\n            desc=\"Validation\", \n            position=0, \n            leave=True, \n            file=sys.stdout, \n            dynamic_ncols=True\n        )\n        return bar\n\n    def on_validation_batch_end(self, trainer, pl_module, outputs, batch, batch_idx, dataloader_idx=0):\n        super().on_validation_batch_end(trainer, pl_module, outputs, batch, batch_idx, dataloader_idx)\n        sys.stdout.flush()\n        \nprogress_bar = KaggleStdoutProgressBar()\n\n# Trainer\ntrainer = pl.Trainer(\n    max_epochs=Config.epochs,\n    accelerator=\"auto\", # Auto-detect GPU\n    devices=1,\n    precision=\"16-mixed\",\n    accumulate_grad_batches=Config.grad_accum,\n    callbacks=[checkpoint_callback, progress_bar],\n    logger=CSVLogger(save_dir=\"logs/\"),\n    enable_progress_bar=True,\n    log_every_n_steps=2\n)\n\nprint(f\"🚀 Starting Training: {Config.model_name} @ {Config.img_size}\")\ntrainer.fit(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T11:06:00.719711Z","iopub.execute_input":"2026-02-24T11:06:00.720322Z","iopub.status.idle":"2026-02-24T11:14:43.490373Z","shell.execute_reply.started":"2026-02-24T11:06:00.720294Z","shell.execute_reply":"2026-02-24T11:14:43.489072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"✅ Training Complete. Best mAP: {checkpoint_callback.best_model_score:.4f}\")\nprint(f\"📂 Best model saved at: {checkpoint_callback.best_model_path}\")\n\n# Load best checkpoint\nbest_model = JaguarLightning.load_from_checkpoint(\n    checkpoint_callback.best_model_path,\n    num_classes=Config.num_classes,\n    train_df=train_split,\n    val_df=val_split,\n    config=Config,\n    weights_only=False\n)\nbest_model.eval().to(Config.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:54:51.890527Z","iopub.status.idle":"2026-02-24T10:54:51.890802Z","shell.execute_reply.started":"2026-02-24T10:54:51.890680Z","shell.execute_reply":"2026-02-24T10:54:51.890698Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🏁 Inference & Submission (Basic)","metadata":{}},{"cell_type":"code","source":"print(\"## FINAL INFERENCE ##\")\n\n# Prepare Test Data\ntest_df = pd.read_csv(Config.BASE_DIR / \"test.csv\")\nunique_files = sorted(set(test_df['query_image']) | set(test_df['gallery_image']))\ndataset = JaguarDataset(pd.DataFrame({'filename': unique_files}), [Config.TEST_DIR], transform=test_tf, is_test=True)\nloader = DataLoader(dataset, batch_size=Config.batch_size, shuffle=False, num_workers=Config.num_workers)\n\n# Get Embeddings\nembeddings = []\nfnames = []\nwith torch.no_grad():\n    for imgs, names in tqdm(loader, desc=\"Embedding Test Set\"):\n        imgs = imgs.to(Config.device)\n        emb = best_model(imgs)\n        embeddings.append(emb.cpu())\n        fnames.extend(names)\n\nembeddings = torch.cat(embeddings)\nembeddings = F.normalize(embeddings, p=2, dim=1).numpy()\nfname_to_idx = {n: i for i, n in enumerate(fnames)}\n\n# Calculate Similarity\nprint(\"📝 Generating submission...\")\npreds = []\nfor _, row in tqdm(test_df.iterrows(), total=len(test_df)):\n    q_idx = fname_to_idx.get(row[\"query_image\"])\n    g_idx = fname_to_idx.get(row[\"gallery_image\"])\n    \n    if q_idx is not None and g_idx is not None:\n        sim = embeddings[q_idx] @ embeddings[g_idx]\n        sim = (sim + 1.0) / 2.0\n        preds.append(float(sim))\n    else:\n        preds.append(0.0)\n\nsub = pd.DataFrame({\"row_id\": test_df[\"row_id\"], \"similarity\": preds})\nsub.to_csv(\"submission.csv\", index=False)\nprint(\"✅ Submission saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:54:51.892216Z","iopub.status.idle":"2026-02-24T10:54:51.892530Z","shell.execute_reply.started":"2026-02-24T10:54:51.892404Z","shell.execute_reply":"2026-02-24T10:54:51.892424Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 🏁 Conclusion & Next Steps \n\nIn this notebook, we built a robust Re-ID pipeline with a focus on **Semi-Supervised Learning**.\n\n**Key takeaways:**\n1. **Pseudo-Labeling:** We successfully expanded our training dataset by finding near-duplicates in the test set using perceptual hashing.\n2. **Leakage Prevention:** Using `GroupShuffleSplit` ensured that augmented data integrity was maintained during validation.\n3. **Metric Learning:** ArcFace loss proved effective for learning discriminative embeddings.\n \n**Potential Improvements:**\n- **Test Time Augmentation (TTA):** Average embeddings from multiple augmented views of the test images to boost accuracy.\n- **Hard Sample Mining:** Implementing a mining strategy could further refine the embedding space.\n- **Ensembling:** Combining predictions from different architectures (e.g., ConvNeXt + EVA) often yields better results.\n\n\n\nIf you found this notebook helpful, please consider giving it an **Upvote**! 👍\n\nGood luck with competition! 🐆","metadata":{}}]}