{"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":129543,"databundleVersionId":15525987},{"sourceType":"competition","sourceId":126777,"databundleVersionId":15314950}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q --upgrade wandb\n\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.metrics.pairwise import cosine_similarity\nfrom tqdm.notebook import tqdm\nimport wandb\nimport random\nfrom itertools import product\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nos.environ[\"HF_TOKEN\"] = user_secrets.get_secret(\"hf_api\")\nos.environ[\"WANDB_API_KEY\"] = user_secrets.get_secret(\"wandb_api\")\n\nSEED = 42\ntorch.manual_seed(SEED)\nnp.random.seed(SEED)\nrandom.seed(SEED)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {device}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = {\n    \"data_dir\": Path(\"/kaggle/input/competitions/jaguar-re-id\"),\n    \"checkpoint_dir\": Path(\"checkpoints\"),\n    \"megadescriptor_model\": \"hf-hub:BVRA/MegaDescriptor-L-384\",\n    \"input_size\": 384,\n    \"embedding_dim\": 256,\n    \"hidden_dim\": 512,\n    \"batch_size\": 32,\n    \"dropout\": 0.3,\n    \"val_split\": 0.2,\n    \"seed\": SEED,\n}\n\nconfig[\"checkpoint_dir\"].mkdir(exist_ok=True)\n\n# Load data — same split as all previous experiments\ntrain_df = pd.read_csv(config[\"data_dir\"] / \"train.csv\")\nlabel_encoder = LabelEncoder()\ntrain_df['label_encoded'] = label_encoder.fit_transform(train_df['ground_truth'])\nnum_classes = len(label_encoder.classes_)\n\ntrain_data, val_data = train_test_split(\n    train_df,\n    test_size=config[\"val_split\"],\n    random_state=config[\"seed\"],\n    stratify=train_df['ground_truth']\n)\n\nprint(f\"Train: {len(train_data)} | Val: {len(val_data)} | Classes: {num_classes}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define model\nclass EmbeddingProjection(nn.Module):\n    def __init__(self, input_dim, hidden_dim=512, output_dim=256, dropout=0.3):\n        super().__init__()\n        self.network = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim),\n            nn.BatchNorm1d(hidden_dim),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, output_dim),\n            nn.BatchNorm1d(output_dim),\n        )\n\n    def forward(self, x):\n        return self.network(x)\n\n    def get_embeddings(self, x):\n        return F.normalize(self.forward(x), p=2, dim=1)\n\n\n# preprocess\npreprocess = transforms.Compose([\n    transforms.Resize((config[\"input_size\"], config[\"input_size\"])),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225]),\n])\n\n@torch.no_grad()\ndef extract_embeddings(model, image_paths, batch_size=32, desc=\"Extracting\"):\n    model.eval()\n    embeddings = []\n    for i in tqdm(range(0, len(image_paths), batch_size), desc=desc):\n        batch_paths = image_paths[i:i+batch_size]\n        tensors = []\n        for p in batch_paths:\n            try:\n                img = Image.open(p).convert(\"RGB\")\n                tensors.append(preprocess(img))\n            except:\n                tensors.append(torch.zeros(3, config[\"input_size\"],\n                                           config[\"input_size\"]))\n        batch = torch.stack(tensors).to(device)\n        embeddings.append(model(batch).cpu().numpy())\n    return np.vstack(embeddings)\n\n\n# load MegaDescriptor\nprint(\"Loading MegaDescriptor...\")\nmegadescriptor = timm.create_model(\n    config[\"megadescriptor_model\"], pretrained=True)\nmegadescriptor.eval()\nmegadescriptor.to(device)\n\nwith torch.no_grad():\n    dummy = torch.randn(1, 3, config[\"input_size\"],\n                        config[\"input_size\"]).to(device)\n    megadescriptor_dim = megadescriptor(dummy).shape[1]\n\nprint(f\"MegaDescriptor loaded | Dim: {megadescriptor_dim}\")\n\n# extract and cache embeddings\ncache_dir = Path(\"embeddings\")\ncache_dir.mkdir(exist_ok=True)\n\nval_cache = cache_dir / \"val_embeddings.npz\"\ntest_cache = cache_dir / \"test_embeddings.npz\"\n\n# val embeddings\nif val_cache.exists():\n    val_embeddings = np.load(val_cache)[\"embeddings\"]\n    print(f\"Loaded val embeddings: {val_embeddings.shape}\")\nelse:\n    print(\"Extracting val embeddings...\")\n    val_paths = [config[\"data_dir\"] / \"train/train\" / f\n                 for f in val_data['filename'].values]\n    val_embeddings = extract_embeddings(megadescriptor, val_paths,\n                                        desc=\"Val\")\n    np.savez_compressed(val_cache, embeddings=val_embeddings)\n    print(f\"Saved val embeddings: {val_embeddings.shape}\")\n\n# test embeddings\ntest_pairs_df = pd.read_csv(config[\"data_dir\"] / \"test.csv\")\ntest_images = sorted(list(\n    set(test_pairs_df['query_image'].unique()) |\n    set(test_pairs_df['gallery_image'].unique())\n))\n\nif test_cache.exists():\n    test_mega_embeddings = np.load(test_cache)[\"embeddings\"]\n    print(f\"Loaded test embeddings: {test_mega_embeddings.shape}\")\nelse:\n    print(\"Extracting test embeddings...\")\n    test_paths = [config[\"data_dir\"] / \"test/test\" / f\n                  for f in test_images]\n    test_mega_embeddings = extract_embeddings(megadescriptor, test_paths,\n                                              desc=\"Test\")\n    np.savez_compressed(test_cache, embeddings=test_mega_embeddings)\n    print(f\"Saved test embeddings: {test_mega_embeddings.shape}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EmbeddingDataset(Dataset):\n    def __init__(self, embeddings, labels):\n        self.embeddings = torch.FloatTensor(embeddings)\n        self.labels = torch.LongTensor(labels)\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n        return self.embeddings[idx], self.labels[idx]\n\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, embedding_dim, num_classes, alpha=0.5):\n        super().__init__()\n        self.alpha = alpha\n        self.arcface_weight = nn.Parameter(\n            torch.FloatTensor(num_classes, embedding_dim))\n        nn.init.xavier_uniform_(self.arcface_weight)\n        self.margin = 0.5\n        self.scale = 64.0\n\n    def arcface_loss(self, embeddings, labels):\n        embeddings = F.normalize(embeddings, dim=1)\n        weight = F.normalize(self.arcface_weight, dim=1)\n        cosine = F.linear(embeddings, weight).clamp(-1+1e-6, 1-1e-6)\n        theta = torch.acos(cosine)\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, labels.view(-1, 1), 1.0)\n        output = self.scale * torch.cos(theta + one_hot * self.margin)\n        return F.cross_entropy(output, labels)\n\n    def triplet_loss(self, embeddings, labels):\n        embeddings = F.normalize(embeddings, dim=1)\n        dist = torch.cdist(embeddings, embeddings, p=2)\n        pos_mask = labels.unsqueeze(0) == labels.unsqueeze(1)\n        neg_mask = ~pos_mask\n        pos_mask.fill_diagonal_(False)\n        hardest_pos = (dist * pos_mask.float()).max(dim=1)[0]\n        hardest_neg = (dist + 1e6 * (~neg_mask).float()).min(dim=1)[0]\n        return F.relu(hardest_pos - hardest_neg + 0.3).mean()\n\n    def forward(self, embeddings, labels):\n        return (self.alpha * self.arcface_loss(embeddings, labels) +\n                (1 - self.alpha) * self.triplet_loss(embeddings, labels))\n\n\n# extract train embeddings\ntrain_cache = cache_dir / \"train_embeddings.npz\"\nif train_cache.exists():\n    train_embeddings = np.load(train_cache)[\"embeddings\"]\n    print(f\"Loaded train embeddings: {train_embeddings.shape}\")\nelse:\n    print(\"Extracting train embeddings...\")\n    train_paths = [config[\"data_dir\"] / \"train/train\" / f\n                   for f in train_data['filename'].values]\n    train_embeddings = extract_embeddings(megadescriptor, train_paths,\n                                          desc=\"Train\")\n    np.savez_compressed(train_cache, embeddings=train_embeddings)\n    print(f\"Saved train embeddings: {train_embeddings.shape}\")\n\n# train model\ntorch.manual_seed(SEED)\nmodel = EmbeddingProjection(\n    input_dim=megadescriptor_dim,\n    hidden_dim=config[\"hidden_dim\"],\n    output_dim=config[\"embedding_dim\"],\n    dropout=config[\"dropout\"]\n).to(device)\n\nloss_fn = CombinedLoss(config[\"embedding_dim\"], num_classes).to(device)\n\ntrain_dataset = EmbeddingDataset(\n    train_embeddings, train_data['label_encoded'].values)\ntrain_loader = DataLoader(train_dataset, batch_size=32,\n                          shuffle=True, num_workers=0)\n\noptimizer = torch.optim.AdamW(\n    list(model.parameters()) + list(loss_fn.parameters()),\n    lr=1e-4, weight_decay=1e-4)\n\nprint(\"Training best model (25 epochs)...\")\nfor epoch in range(25):\n    model.train()\n    loss_fn.train()\n    total_loss = 0\n    for embeddings, labels in tqdm(train_loader,\n                                   desc=f\"Epoch {epoch+1}\",\n                                   leave=False):\n        embeddings, labels = embeddings.to(device), labels.to(device)\n        optimizer.zero_grad()\n        projected = model(embeddings)\n        loss = loss_fn(projected, labels)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n    print(f\"Epoch {epoch+1:2d} | Loss: {total_loss/len(train_loader):.4f}\")\n\n# Get projected embeddings for val and test\nmodel.eval()\nwith torch.no_grad():\n    val_tensor = torch.FloatTensor(val_embeddings).to(device)\n    val_projected = model.get_embeddings(val_tensor).cpu().numpy()\n\n    test_tensor = torch.FloatTensor(test_mega_embeddings).to(device)\n    test_projected = model.get_embeddings(test_tensor).cpu().numpy()\n\nprint(f\"Val projected: {val_projected.shape}\")\nprint(f\"Test projected: {test_projected.shape}\")\nprint(\"Model ready\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def k_reciprocal_reranking(query_features, gallery_features, \n                            k1=20, k2=6, lambda_value=0.3):\n    \"\"\"\n    Re-ranks similarity scores using k-reciprocal neighbors.\n    Returns final distance matrix (lower = more similar).\n    \"\"\"\n    all_features = np.concatenate([query_features, gallery_features], axis=0)\n    n_query = query_features.shape[0]\n    n_all = all_features.shape[0]\n\n    # compute initial distance matrix\n    all_features = all_features / np.linalg.norm(\n        all_features, axis=1, keepdims=True)\n    dist = 1 - all_features @ all_features.T\n    dist = np.maximum(dist, 0)\n\n    # find k-reciprocal neighbors for each sample\n    k_reciprocal_neighbors = []\n    for i in range(n_all):\n        forward_k = np.argsort(dist[i])[:k1+1]\n        k_recip = []\n        for j in forward_k:\n            if i in np.argsort(dist[j])[:k1+1]:\n                k_recip.append(j)\n        k_reciprocal_neighbors.append(np.array(k_recip))\n\n    # build jaccard distance matrix\n    jaccard_dist = np.zeros((n_all, n_all))\n    for i in range(n_all):\n        for j in range(n_all):\n            set_i = set(k_reciprocal_neighbors[i])\n            set_j = set(k_reciprocal_neighbors[j])\n            intersection = len(set_i & set_j)\n            union = len(set_i | set_j)\n            jaccard_dist[i][j] = 1 - intersection / (union + 1e-6)\n\n    # final distance\n    final_dist = lambda_value * dist + (1 - lambda_value) * jaccard_dist\n    \n    return final_dist[:n_query, n_query:]\n\n\ndef compute_map_from_dist(dist_matrix, query_labels, gallery_labels):\n    \"\"\"Compute identity-balanced mAP from distance matrix.\"\"\"\n    identity_aps = {}\n    for query_idx in range(len(query_labels)):\n        query_label = query_labels[query_idx]\n        distances = dist_matrix[query_idx]\n        is_match = (gallery_labels == query_label).astype(int)\n\n        sorted_indices = np.argsort(distances)  # ascending distance\n        sorted_matches = is_match[sorted_indices]\n\n        n_positives = sorted_matches.sum()\n        if n_positives == 0:\n            continue\n\n        cumsum = np.cumsum(sorted_matches)\n        precision_at_k = cumsum / np.arange(1, len(sorted_matches) + 1)\n        ap = np.sum(precision_at_k * sorted_matches) / n_positives\n\n        if query_label not in identity_aps:\n            identity_aps[query_label] = []\n        identity_aps[query_label].append(ap)\n\n    return np.mean([np.mean(aps) for aps in identity_aps.values()])\n\n\n# baseline val mAP without re-ranking\nval_labels = val_data['ground_truth'].values\nbaseline_dist = 1 - cosine_similarity(val_projected)\nnp.fill_diagonal(baseline_dist, 1e6)  # exclude self\nbaseline_map = compute_map_from_dist(baseline_dist, val_labels, val_labels)\nprint(f\"Baseline val mAP (no re-ranking): {baseline_map:.4f}\")\n\nprint(\"\\nSearching best k1 and lambda parameters...\")\nprint(\"This may take 10-15 minutes...\\n\")\n\n# parameter search\nk1_values = [10, 15, 20, 25]\nlambda_values = [0.2, 0.3, 0.4, 0.5]\n\nsearch_results = []\n\nfor k1, lam in product(k1_values, lambda_values):\n    dist_matrix = k_reciprocal_reranking(\n        val_projected, val_projected, \n        k1=k1, k2=6, lambda_value=lam)\n    np.fill_diagonal(dist_matrix, 1e6)\n    val_map = compute_map_from_dist(dist_matrix, val_labels, val_labels)\n    search_results.append({\n        \"k1\": k1,\n        \"lambda\": lam,\n        \"val_map\": round(val_map, 4)\n    })\n    print(f\"k1={k1:2d} | lambda={lam} | val_mAP={val_map:.4f}\")\n\n# find best parameter\nsearch_df = pd.DataFrame(search_results).sort_values(\n    \"val_map\", ascending=False)\nbest = search_df.iloc[0]\nprint(f\"\\nBest params: k1={best['k1']} | lambda={best['lambda']} \"\n      f\"| val_mAP={best['val_map']}\")\nprint(f\"Improvement over baseline: \"\n      f\"{best['val_map'] - baseline_map:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# apply best re-ranking params to test set\nbest_k1 = int(best['k1'])\nbest_lambda = float(best['lambda'])\n\nprint(f\"Applying re-ranking with k1={best_k1}, lambda={best_lambda}...\")\n\n# Re-rank test embeddings\ntest_dist = k_reciprocal_reranking(\n    test_projected, test_projected,\n    k1=best_k1, k2=6, lambda_value=best_lambda\n)\n\n# convert distance to similarity\ntest_sim = 1 - test_dist\ntest_sim = np.clip(test_sim, 0.0, 1.0)\n\n# build image to index mapping\nimg_to_idx = {fn: idx for idx, fn in enumerate(test_images)}\n\nsimilarities = []\nfor _, row in tqdm(test_pairs_df.iterrows(), total=len(test_pairs_df)):\n    q_idx = img_to_idx[row['query_image']]\n    g_idx = img_to_idx[row['gallery_image']]\n    similarities.append(float(test_sim[q_idx, g_idx]))\n\nsimilarities = np.clip(similarities, 0.0, 1.0)\n\nsubmission_df = pd.DataFrame({\n    'row_id': test_pairs_df['row_id'],\n    'similarity': similarities\n})\n\nsubmission_df.to_csv(\"submission_reranked.csv\", index=False)\nprint(f\"Submission saved: {len(submission_df)} rows\")\n\n# log to W&B\nwandb.login(key=os.environ[\"WANDB_API_KEY\"])\nwandb.init(\n    project=\"jaguar-reid-mishank\",\n    name=\"reranking-k1-15-lambda-0.5\",\n    config={\n        \"k1\": best_k1,\n        \"k2\": 6,\n        \"lambda\": best_lambda,\n        \"baseline_val_map\": baseline_map,\n        \"reranked_val_map\": best['val_map'],\n        \"improvement\": best['val_map'] - baseline_map,\n    }\n)\nwandb.log({\n    \"baseline_val_map\": baseline_map,\n    \"best_reranked_val_map\": best['val_map'],\n    \"k1\": best_k1,\n    \"lambda\": best_lambda,\n})\n\n# log full search results\nwandb.log({\n    \"search_results\": wandb.Table(dataframe=search_df)\n})\nwandb.finish()\n\nfrom IPython.display import FileLink\nFileLink('submission_reranked.csv')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}