{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":[{"sourceId":126777,"databundleVersionId":15314950,"sourceType":"competition"},{"sourceId":733318,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":558848,"modelId":571424}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Jaguar Re-ID: Inference Notebook\n\n## Train Notebook is here! **https://www.kaggle.com/code/kawaharataishi/jaguar-re-id-starter-set-training-notebook**\n\n## Overview\nThis notebook performs inference on the test data using a pre-trained model loaded as a Kaggle Dataset.\nSeparating the training and inference processes offers several benefits:\n1. **Fast Submission**: Skip the long training time (hours) and generate a submission file in minutes.\n2. **Resource Efficiency**: Focus GPU hours on training experiments.\n3. **Ensemble Flexibility**: Easily combine multiple pre-trained models to adjust the inference logic.\n\n## Workflow\n1. **Configuration**: Define the same settings (image size, model architecture) used during training.\n2. **Model Loading**: Load weights from the `.pth` file added as a Kaggle Model.\n3. **Feature Extraction**: Extract features for test images (Query & Gallery).\n4. **Re-ranking**: Apply k-Reciprocal Re-ranking to maximize the score.\n5. **Submission**: Generate the `submission.csv` file.","metadata":{}},{"cell_type":"code","source":"import os\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport timm\n\n# ==========================================\n# Configuration\n# ==========================================\nclass Config:\n    # --- Path Settings ---\n    DATA_DIR = '/kaggle/input/jaguar-re-id' \n        \n    TEST_CSV = os.path.join(DATA_DIR, 'test.csv')\n    TEST_IMG_DIR = os.path.join(DATA_DIR, 'test', 'test')\n    if not os.path.exists(TEST_IMG_DIR):\n         TEST_IMG_DIR = os.path.join(DATA_DIR, 'test')\n\n    # --- Model Settings ---\n    MODEL_NAME = 'convnext_base'\n    IMG_SIZE = (384, 384) # High resolution for better accuracy\n    BATCH_SIZE = 32\n    EMBEDDING_DIM = 512\n    \n    # --- Reranking Settings ---\n    USE_RERANKING = True\n    K1 = 20\n    K2 = 6\n    LAMBDA_VALUE = 0.3\n    \n    # --- Weights Path ---\n    # [IMPORTANT] Using pre-trained model\n    MODEL_WEIGHT_PATH = '/kaggle/input/m/kawaharataishi/jaguar-re-id/pytorch/default/1/jaguar_resnet50_arcface_ep15.pth'\n    \n    # --- Device ---\n    if torch.cuda.is_available():\n        DEVICE = torch.device('cuda')\n    else:\n        DEVICE = torch.device('cpu')\n        \nprint(f\"Device: {Config.DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T19:56:26.055160Z","iopub.execute_input":"2026-01-27T19:56:26.055458Z","iopub.status.idle":"2026-01-27T19:56:40.681462Z","shell.execute_reply.started":"2026-01-27T19:56:26.055434Z","shell.execute_reply":"2026-01-27T19:56:40.680659Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Definition\nDefine the exact same architecture used during training.\nHere we use `convnext_base` from the `timm` library as the backbone, connected to an `ArcFace` head.","metadata":{}},{"cell_type":"code","source":"class ArcMarginProduct(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.50, easy_margin=False):\n        super(ArcMarginProduct, self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(m)\n        self.sin_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        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        \n        one_hot = torch.zeros(cosine.size(), device=input.device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n        return output\n\nclass JaguarReIDModel(nn.Module):\n    def __init__(self, num_classes, embedding_dim=512, pretrained=False):\n        super(JaguarReIDModel, self).__init__()\n        # Backbone: ConvNeXt Base\n        # For inference, pretrained=False is fine (we will load weights later)\n        # This also allows it to run without internet access\n        self.backbone = timm.create_model('convnext_base', pretrained=pretrained, num_classes=0)\n        in_features = self.backbone.num_features\n        \n        # Head\n        self.embedding = nn.Sequential(\n            nn.Linear(in_features, embedding_dim),\n            nn.BatchNorm1d(embedding_dim),\n            nn.PReLU()\n        )\n\n        self.arcface = ArcMarginProduct(embedding_dim, num_classes, s=30.0, m=0.50)\n\n    def forward(self, x, labels=None):\n        features = self.backbone(x)\n        embeddings = self.embedding(features)\n\n        if labels is not None:\n            return self.arcface(embeddings, labels)\n        else:\n            return F.normalize(embeddings)\n\n# Dataset & Image Processing\ndef crop_alphachannel(img: Image.Image) -> Image.Image:\n    if img.mode in ('RGBA', 'LA') or (img.mode == 'P' and 'transparency' in img.info):\n        bbox = img.getbbox()\n        if bbox:\n            return img.crop(bbox)\n    return img\n\nclass JaguarDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None, is_train=True):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n        self.is_train = is_train\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.img_dir, row['filename'])\n\n        try:\n            image = Image.open(img_path).convert(\"RGBA\")\n            image = crop_alphachannel(image)\n            image = image.convert(\"RGB\")\n        except:\n             # Fallback in case image load fails\n             image = Image.new(\"RGB\", Config.IMG_SIZE)\n\n        if self.transform:\n            image = self.transform(image)\n\n        if self.is_train:\n            return image, 0\n        else:\n            return image, row['filename']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T19:56:40.682852Z","iopub.execute_input":"2026-01-27T19:56:40.683095Z","iopub.status.idle":"2026-01-27T19:56:40.698306Z","shell.execute_reply.started":"2026-01-27T19:56:40.683072Z","shell.execute_reply":"2026-01-27T19:56:40.697534Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Loading","metadata":{}},{"cell_type":"code","source":"# In the baseline implementation, the Jaguar dataset had 31 classes.\nNUM_CLASSES_TRAINING = 31 \n\nmodel = JaguarReIDModel(NUM_CLASSES_TRAINING, Config.EMBEDDING_DIM, pretrained=False).to(Config.DEVICE)\n\nif os.path.exists(Config.MODEL_WEIGHT_PATH):\n    # Simply loading; add weights_only=True if required by newer torch versions for security\n    state_dict = torch.load(Config.MODEL_WEIGHT_PATH, map_location=Config.DEVICE)\n    model.load_state_dict(state_dict)\n    print(f\"Loaded weights from {Config.MODEL_WEIGHT_PATH}\")\nelse:\n    print(f\"⚠️ Warning: Weight file not found at {Config.MODEL_WEIGHT_PATH}. Please check dataset path.\")\n    print(\"Please add your Model Dataset via the 'Add Input' button on the right sidebar.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T19:56:40.699343Z","iopub.execute_input":"2026-01-27T19:56:40.699659Z","iopub.status.idle":"2026-01-27T19:56:47.319131Z","shell.execute_reply.started":"2026-01-27T19:56:40.699628Z","shell.execute_reply":"2026-01-27T19:56:47.318492Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Feature Extraction","metadata":{}},{"cell_type":"code","source":"test_df_raw = pd.read_csv(Config.TEST_CSV)\nunique_test_images = sorted(list(set(test_df_raw['query_image']) | set(test_df_raw['gallery_image'])))\ntest_images_df = pd.DataFrame({'filename': unique_test_images})\n\nval_transform = transforms.Compose([\n    transforms.Resize(Config.IMG_SIZE),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\ntest_dataset = JaguarDataset(test_images_df, Config.TEST_IMG_DIR, transform=val_transform, is_train=False)\ntest_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n\ndef extract_embeddings(model, loader, device):\n    model.eval()\n    embeddings = {}\n    feats_list = []\n    fnames_list = []\n    \n    with torch.no_grad():\n        for images, filenames in tqdm(loader, desc=\"Extracting Features\"):\n            images = images.to(device)\n            feats = model(images, labels=None)\n            feats = feats.cpu().numpy()\n            \n            for fname, feat in zip(filenames, feats):\n                embeddings[fname] = feat\n                feats_list.append(feat)\n                fnames_list.append(fname)\n    \n    return embeddings, np.array(feats_list), fnames_list\n\nprint(\"Extracting embeddings...\")\nembeddings_dict, feats_array, fnames_list = extract_embeddings(model, test_loader, Config.DEVICE)\nfname_to_idx = {fname: i for i, fname in enumerate(fnames_list)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T19:56:47.320804Z","iopub.execute_input":"2026-01-27T19:56:47.321243Z","iopub.status.idle":"2026-01-27T19:58:09.378405Z","shell.execute_reply.started":"2026-01-27T19:56:47.321221Z","shell.execute_reply":"2026-01-27T19:58:09.377350Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Re-ranking & Submission\nWe apply **k-Reciprocal Re-ranking** to improve the baseline score.\nThis increases the score for image pairs that are mutually similar (reciprocal relationship), thereby enhancing accuracy.","metadata":{}},{"cell_type":"code","source":"def re_ranking(probFea, gallery_num, k1, k2, lambda_value):\n    print(\"Computing Euclidean distance matrix...\")\n    query_num = probFea.shape[0]\n    all_num = query_num\n    feat = torch.from_numpy(probFea)\n    # You could use chunk processing for memory efficiency if needed,\n    # but we compute it all at once given the small dataset size.\n    distmat = torch.pow(feat, 2).sum(dim=1, keepdim=True).expand(all_num, all_num) + \\\n              torch.pow(feat, 2).sum(dim=1, keepdim=True).expand(all_num, all_num).t()\n    distmat.addmm_(feat, feat.t(), beta=1, alpha=-2)\n    original_dist = distmat.cpu().numpy()\n    \n    del feat\n    original_dist = np.maximum(original_dist, 0)\n    original_dist = np.sqrt(original_dist)\n    \n    print(\"Starting k-reciprocal reranking calculation...\")\n    initial_rank = np.argsort(original_dist, axis=1)\n    \n    V = np.zeros((all_num, all_num), dtype=np.float32)\n    initial_rank = initial_rank.astype(np.int32)\n    \n    print(\"Computing Jaccard distance...\")\n    for i in tqdm(range(all_num)):\n        forward_k_neigh_index = initial_rank[i,:k1+1]\n        backward_k_neigh_index = initial_rank[forward_k_neigh_index,:k1+1]\n        fi = np.where(backward_k_neigh_index==i)[0]\n        k_reciprocal_index = forward_k_neigh_index[fi]\n        k_reciprocal_expansion_index = k_reciprocal_index\n        for j in range(len(k_reciprocal_index)):\n            candidate = k_reciprocal_index[j]\n            candidate_forward_k_neigh_index = initial_rank[candidate,:int(np.round(k1/2))+1]\n            candidate_backward_k_neigh_index = initial_rank[candidate_forward_k_neigh_index,:int(np.round(k1/2))+1]\n            fi_candidate = np.where(candidate_backward_k_neigh_index == candidate)[0]\n            candidate_k_reciprocal_index = candidate_forward_k_neigh_index[fi_candidate]\n            if len(np.intersect1d(candidate_k_reciprocal_index,k_reciprocal_index))> 2/3*len(candidate_k_reciprocal_index):\n                k_reciprocal_expansion_index = np.append(k_reciprocal_expansion_index,candidate_k_reciprocal_index)\n        \n        k_reciprocal_expansion_index = np.unique(k_reciprocal_expansion_index)\n        weight = np.exp(-original_dist[i,k_reciprocal_expansion_index])\n        V[i,k_reciprocal_expansion_index] = weight / np.sum(weight)\n        \n    original_dist = original_dist[:query_num,]\n    if k2 != 1:\n        V_qe = np.zeros_like(V, dtype=np.float32)\n        for i in range(all_num):\n            V_qe[i,:] = np.mean(V[initial_rank[i,:k2],:], axis=0)\n        V = V_qe\n        del V_qe\n    del initial_rank\n    \n    invIndex = []\n    for i in range(all_num):\n        invIndex.append(np.where(V[:,i] != 0)[0])\n        \n    jaccard_dist = np.zeros_like(original_dist, dtype=np.float32)\n    \n    print(\"Finalizing Jaccard distance...\")\n    for i in range(query_num):\n        temp_min = np.zeros(shape=[1,all_num], dtype=np.float32)\n        indNonZero = np.where(V[i,:] != 0)[0]\n        indImages = []\n        indImages = [invIndex[ind] for ind in indNonZero]\n        for j in range(len(indNonZero)):\n            temp_min[0,indImages[j]] = temp_min[0,indImages[j]] + np.minimum(V[i,indNonZero[j]], V[indImages[j],indNonZero[j]])\n        jaccard_dist[i] = 1 - temp_min / (2 - temp_min)\n        \n    final_dist = jaccard_dist * (1-lambda_value) + original_dist * lambda_value\n    final_sim = np.exp(-final_dist)\n    return final_sim","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T19:58:09.380649Z","iopub.execute_input":"2026-01-27T19:58:09.381091Z","iopub.status.idle":"2026-01-27T19:58:09.394385Z","shell.execute_reply.started":"2026-01-27T19:58:09.381046Z","shell.execute_reply":"2026-01-27T19:58:09.393643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate Submission\nif Config.USE_RERANKING:\n    print(\"Applying k-Reciprocal Reranking...\")\n    rerank_sim_matrix = re_ranking(feats_array, len(feats_array), Config.K1, Config.K2, Config.LAMBDA_VALUE)\n    \n    print(\"Calculating final submission scores...\")\n    similarities = []\n    for _, row in tqdm(test_df_raw.iterrows(), total=len(test_df_raw)):\n        q_idx = fname_to_idx[row['query_image']]\n        g_idx = fname_to_idx[row['gallery_image']]\n        sim = rerank_sim_matrix[q_idx, g_idx]\n        similarities.append(sim)\nelse:\n    print(\"Applying Cosine Similarity...\")\n    similarities = []\n    for _, row in tqdm(test_df_raw.iterrows(), total=len(test_df_raw)):\n        q_emb = embeddings_dict[row['query_image']]\n        g_emb = embeddings_dict[row['gallery_image']]\n        sim = np.dot(q_emb, g_emb) / (np.linalg.norm(q_emb) * np.linalg.norm(g_emb) + 1e-6)\n        sim = (sim + 1) / 2\n        similarities.append(sim)\n\nsubmission = test_df_raw[['row_id']].copy()\nsubmission['similarity'] = similarities\nsubmission.to_csv('submission.csv', index=False)\nprint(\"Saved submission.csv\")\nprint(\"Done! Generated submission.csv for upload.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T19:58:09.395356Z","iopub.execute_input":"2026-01-27T19:58:09.395620Z","iopub.status.idle":"2026-01-27T19:58:25.218143Z","shell.execute_reply.started":"2026-01-27T19:58:09.395574Z","shell.execute_reply":"2026-01-27T19:58:25.217357Z"}},"outputs":[],"execution_count":null}]}