{"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\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,"execution":{"iopub.status.busy":"2026-03-03T22:36:42.877978Z","iopub.execute_input":"2026-03-03T22:36:42.878600Z","iopub.status.idle":"2026-03-03T22:37:12.185224Z","shell.execute_reply.started":"2026-03-03T22:36:42.878570Z","shell.execute_reply":"2026-03-03T22:37:12.184512Z"}},"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    \"val_split\": 0.2,\n    \"seed\": SEED,\n}\n\nconfig[\"checkpoint_dir\"].mkdir(exist_ok=True)\n\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,"execution":{"iopub.status.busy":"2026-03-03T22:37:16.318347Z","iopub.execute_input":"2026-03-03T22:37:16.319259Z","iopub.status.idle":"2026-03-03T22:37:16.347577Z","shell.execute_reply.started":"2026-03-03T22:37:16.319227Z","shell.execute_reply":"2026-03-03T22:37:16.346990Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preprocess = 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\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\ncache_dir = Path(\"embeddings\")\ncache_dir.mkdir(exist_ok=True)\n\ntrain_cache = cache_dir / \"train_embeddings.npz\"\nval_cache = cache_dir / \"val_embeddings.npz\"\n\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.shape}\")\n\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.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T22:37:48.945847Z","iopub.execute_input":"2026-03-03T22:37:48.946777Z","iopub.status.idle":"2026-03-03T22:51:34.593114Z","shell.execute_reply.started":"2026-03-03T22:37:48.946746Z","shell.execute_reply":"2026-03-03T22:51:34.592329Z"}},"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 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\nclass CombinedLoss(nn.Module):\n    def __init__(self, embedding_dim, num_classes, \n                 arcface_margin=0.5, arcface_scale=64.0,\n                 triplet_margin=0.3, 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.arcface_margin = arcface_margin\n        self.arcface_scale = arcface_scale\n        self.triplet_margin = triplet_margin\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.arcface_scale * torch.cos(\n            theta + one_hot * self.arcface_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 + self.triplet_margin).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\ndef compute_validation_map(model, val_embeddings, val_labels):\n    model.eval()\n    with torch.no_grad():\n        val_tensor = torch.FloatTensor(val_embeddings).to(device)\n        finetuned_emb = model.get_embeddings(val_tensor).cpu().numpy()\n\n    sim_matrix = cosine_similarity(finetuned_emb)\n    np.fill_diagonal(sim_matrix, -1)\n\n    identity_aps = {}\n    for query_idx in range(len(val_labels)):\n        query_label = val_labels[query_idx]\n        similarities = sim_matrix[query_idx]\n        is_match = (val_labels == query_label).astype(int)\n        is_match[query_idx] = 0\n        sorted_indices = np.argsort(-similarities)\n        sorted_matches = is_match[sorted_indices]\n        n_positives = sorted_matches.sum()\n        if n_positives == 0:\n            continue\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        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\nprint(\"Model, loss and validation ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T22:55:28.909532Z","iopub.execute_input":"2026-03-03T22:55:28.909910Z","iopub.status.idle":"2026-03-03T22:55:28.925613Z","shell.execute_reply.started":"2026-03-03T22:55:28.909881Z","shell.execute_reply":"2026-03-03T22:55:28.924896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sweep_config = {\n    \"method\": \"bayes\",\n    \"metric\": {\"name\": \"val_map\", \"goal\": \"maximize\"},\n    \"parameters\": {\n        \"learning_rate\": {\n            \"distribution\": \"log_uniform_values\",\n            \"min\": 1e-5,\n            \"max\": 1e-3\n        },\n        \"arcface_margin\": {\n            \"values\": [0.3, 0.4, 0.5, 0.6]\n        },\n        \"arcface_scale\": {\n            \"values\": [32, 48, 64]\n        },\n        \"triplet_margin\": {\n            \"values\": [0.2, 0.3, 0.4]\n        },\n        \"alpha\": {\n            \"values\": [0.3, 0.5, 0.7]\n        },\n        \"embedding_dim\": {\n            \"values\": [128, 256, 512]\n        },\n        \"dropout\": {\n            \"values\": [0.1, 0.3, 0.5]\n        },\n    }\n}\n\n\ndef train_sweep():\n    wandb.login(key=os.environ[\"WANDB_API_KEY\"])\n    \n    run = wandb.init()\n    sweep_cfg = run.config\n\n    # reset the seed\n    torch.manual_seed(SEED)\n    np.random.seed(SEED)\n    random.seed(SEED)\n\n    # datasets\n    train_dataset = EmbeddingDataset(\n        train_embeddings, train_data['label_encoded'].values)\n    val_dataset = EmbeddingDataset(\n        val_embeddings, val_data['label_encoded'].values)\n\n    train_loader = DataLoader(train_dataset, batch_size=32,\n                              shuffle=True, num_workers=0)\n\n    val_labels = val_data['ground_truth'].values\n\n    # model with sweep parameters\n    model = EmbeddingProjection(\n        input_dim=megadescriptor_dim,\n        hidden_dim=512,\n        output_dim=sweep_cfg.embedding_dim,\n        dropout=sweep_cfg.dropout\n    ).to(device)\n\n    loss_fn = CombinedLoss(\n        embedding_dim=sweep_cfg.embedding_dim,\n        num_classes=num_classes,\n        arcface_margin=sweep_cfg.arcface_margin,\n        arcface_scale=sweep_cfg.arcface_scale,\n        triplet_margin=sweep_cfg.triplet_margin,\n        alpha=sweep_cfg.alpha\n    ).to(device)\n\n    optimizer = torch.optim.AdamW(\n        list(model.parameters()) + list(loss_fn.parameters()),\n        lr=sweep_cfg.learning_rate,\n        weight_decay=1e-4\n    )\n\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='min', factor=0.5, patience=3)\n\n    best_val_map = 0.0\n    best_val_loss = float('inf')\n    patience_counter = 0\n\n    # training for 20 epochs / trial to save time\n    for epoch in range(20):\n        model.train()\n        loss_fn.train()\n        train_loss = 0\n\n        for embeddings, labels in train_loader:\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            train_loss += loss.item()\n\n        train_loss /= len(train_loader)\n\n        # validate\n        model.eval()\n        loss_fn.eval()\n        val_loss = 0\n        val_dataset_loader = DataLoader(val_dataset, batch_size=32,\n                                        shuffle=False, num_workers=0)\n\n        with torch.no_grad():\n            for embeddings, labels in val_dataset_loader:\n                embeddings, labels = embeddings.to(device), labels.to(device)\n                projected = model(embeddings)\n                loss = loss_fn(projected, labels)\n                val_loss += loss.item()\n\n        val_loss /= len(val_dataset_loader)\n        val_map = compute_validation_map(model, val_embeddings, val_labels)\n\n        scheduler.step(val_loss)\n\n        wandb.log({\n            \"epoch\": epoch + 1,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"val_map\": val_map,\n        })\n\n        if val_map > best_val_map:\n            best_val_map = val_map\n            best_val_loss = val_loss\n            patience_counter = 0\n            torch.save({\n                \"model_state_dict\": model.state_dict(),\n                \"loss_state_dict\": loss_fn.state_dict(),\n                \"val_map\": best_val_map,\n                \"config\": dict(sweep_cfg),\n            }, config[\"checkpoint_dir\"] / \"sweep_best.pth\")\n        else:\n            patience_counter += 1\n            if patience_counter >= 5:\n                break\n\n    wandb.log({\n        \"best_val_map\": best_val_map,\n        \"best_val_loss\": best_val_loss,\n    })\n\n    wandb.finish()\n\n\nprint(\"Sweep config and train function ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T22:57:51.971174Z","iopub.execute_input":"2026-03-03T22:57:51.972016Z","iopub.status.idle":"2026-03-03T22:57:51.985197Z","shell.execute_reply.started":"2026-03-03T22:57:51.971981Z","shell.execute_reply":"2026-03-03T22:57:51.984461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"wandb.login(key=os.environ[\"WANDB_API_KEY\"])\n\nsweep_id = wandb.sweep(\n    sweep_config,\n    project=\"jaguar-reid-mishank\"\n)\n\nprint(f\"Sweep ID: {sweep_id}\")\nprint(\"Starting sweep agent — running 15 trials...\")\nprint(\"Check your W&B project to see trials appearing live.\")\n\nwandb.agent(sweep_id, function=train_sweep, count=15)\n\nprint(\"\\nSweep complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T22:58:31.305539Z","iopub.execute_input":"2026-03-03T22:58:31.305853Z","iopub.status.idle":"2026-03-03T23:04:56.796441Z","shell.execute_reply.started":"2026-03-03T22:58:31.305827Z","shell.execute_reply":"2026-03-03T23:04:56.795763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import wandb\n\napi = wandb.Api()\n\n# get all runs from the sweep\nsweep = api.sweep(f\"jaguar-reid-mishank/{sweep_id}\")\nruns = sweep.runs\n\n# finding the best run\nbest_run = max(runs, key=lambda r: r.summary.get(\"best_val_map\", 0))\n\nprint(f\"Best run: {best_run.name}\")\nprint(f\"Best val mAP: {best_run.summary.get('best_val_map', 0):.4f}\")\nprint(f\"\\nBest hyperparameters:\")\nfor key, value in best_run.config.items():\n    print(f\"  {key}: {value}\")\n\n# show all runs sorted\nprint(f\"\\nAll runs sorted by val mAP:\")\nrun_results = []\nfor run in runs:\n    run_results.append({\n        \"name\": run.name,\n        \"val_map\": round(run.summary.get(\"best_val_map\", 0), 4),\n        \"lr\": round(run.config.get(\"learning_rate\", 0), 6),\n        \"embedding_dim\": run.config.get(\"embedding_dim\"),\n        \"arcface_margin\": run.config.get(\"arcface_margin\"),\n        \"dropout\": run.config.get(\"dropout\"),\n        \"alpha\": run.config.get(\"alpha\"),\n    })\n\nresults_df = pd.DataFrame(run_results).sort_values(\n    \"val_map\", ascending=False)\nprint(results_df.to_string(index=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T23:05:58.238886Z","iopub.execute_input":"2026-03-03T23:05:58.239754Z","iopub.status.idle":"2026-03-03T23:06:03.036065Z","shell.execute_reply.started":"2026-03-03T23:05:58.239712Z","shell.execute_reply":"2026-03-03T23:06:03.035359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_config = {\n    \"learning_rate\": 0.0007670550202151722,\n    \"arcface_margin\": 0.6,\n    \"arcface_scale\": 48,\n    \"triplet_margin\": 0.4,\n    \"alpha\": 0.3,\n    \"embedding_dim\": 256,\n    \"dropout\": 0.5,\n}\n\nprint(\"Training best config for 25 epochs...\")\n\ntorch.manual_seed(SEED)\nnp.random.seed(SEED)\nrandom.seed(SEED)\n\ntrain_dataset = EmbeddingDataset(\n    train_embeddings, train_data['label_encoded'].values)\nval_dataset = EmbeddingDataset(\n    val_embeddings, val_data['label_encoded'].values)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32,\n                          shuffle=True, num_workers=0)\nval_loader = DataLoader(val_dataset, batch_size=32,\n                        shuffle=False, num_workers=0)\n\nval_labels = val_data['ground_truth'].values\n\nmodel = EmbeddingProjection(\n    input_dim=megadescriptor_dim,\n    hidden_dim=512,\n    output_dim=best_config[\"embedding_dim\"],\n    dropout=best_config[\"dropout\"]\n).to(device)\n\nloss_fn = CombinedLoss(\n    embedding_dim=best_config[\"embedding_dim\"],\n    num_classes=num_classes,\n    arcface_margin=best_config[\"arcface_margin\"],\n    arcface_scale=best_config[\"arcface_scale\"],\n    triplet_margin=best_config[\"triplet_margin\"],\n    alpha=best_config[\"alpha\"]\n).to(device)\n\noptimizer = torch.optim.AdamW(\n    list(model.parameters()) + list(loss_fn.parameters()),\n    lr=best_config[\"learning_rate\"],\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='min', factor=0.5, patience=3)\n\n# W&B run for final training\nwandb.login(key=os.environ[\"WANDB_API_KEY\"])\nwandb.init(\n    project=\"jaguar-reid-mishank\",\n    name=\"sweep-best-final-training\",\n    config={**best_config, \"num_epochs\": 25, \"num_classes\": num_classes}\n)\n\nbest_val_map = 0.0\nbest_val_loss = float('inf')\npatience_counter = 0\n\nfor epoch in range(25):\n    model.train()\n    loss_fn.train()\n    train_loss = 0\n\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        train_loss += loss.item()\n\n    train_loss /= len(train_loader)\n\n    model.eval()\n    loss_fn.eval()\n    val_loss = 0\n\n    with torch.no_grad():\n        for embeddings, labels in val_loader:\n            embeddings, labels = embeddings.to(device), labels.to(device)\n            projected = model(embeddings)\n            loss = loss_fn(projected, labels)\n            val_loss += loss.item()\n\n    val_loss /= len(val_loader)\n    val_map = compute_validation_map(model, val_embeddings, val_labels)\n    scheduler.step(val_loss)\n\n    wandb.log({\n        \"epoch\": epoch + 1,\n        \"train_loss\": train_loss,\n        \"val_loss\": val_loss,\n        \"val_map\": val_map,\n    })\n\n    print(f\"Epoch {epoch+1:2d} | \"\n          f\"Train Loss: {train_loss:.4f} | \"\n          f\"Val Loss: {val_loss:.4f} | \"\n          f\"Val mAP: {val_map:.4f}\")\n\n    if val_map > best_val_map:\n        best_val_map = val_map\n        best_val_loss = val_loss\n        patience_counter = 0\n        torch.save({\n            \"model_state_dict\": model.state_dict(),\n            \"val_map\": best_val_map,\n        }, config[\"checkpoint_dir\"] / \"sweep_best_final.pth\")\n    else:\n        patience_counter += 1\n        if patience_counter >= 7:\n            print(f\"Early stopping at epoch {epoch+1}\")\n            break\n\nwandb.log({\"best_val_map\": best_val_map})\nwandb.finish()\n\nprint(f\"\\nBest val mAP: {best_val_map:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T23:07:24.962064Z","iopub.execute_input":"2026-03-03T23:07:24.962623Z","iopub.status.idle":"2026-03-03T23:07:39.683530Z","shell.execute_reply.started":"2026-03-03T23:07:24.962592Z","shell.execute_reply":"2026-03-03T23:07:39.682776Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# extract test embeddings\ntest_pairs_df = pd.read_csv(Path(\"/kaggle/input/competitions/jaguar-re-id\") / \"test.csv\")\ntest_images = sorted(list(\n    set(test_pairs_df['query_image'].unique()) |\n    set(test_pairs_df['gallery_image'].unique())\n))\n\ntest_cache = Path(\"embeddings\") / \"test_embeddings.npz\"\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 = [Path(\"/kaggle/input/competitions/jaguar-re-id\") / \"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_mega_embeddings.shape}\")\n\n# project test embeddings through best model\nmodel.eval()\nwith torch.no_grad():\n    test_tensor = torch.FloatTensor(test_mega_embeddings).to(device)\n    test_projected = model.get_embeddings(test_tensor).cpu().numpy()\n\nprint(f\"Test projected: {test_projected.shape}\")\n\n# apply re-ranking with best params from Experiment 4\ndef k_reciprocal_reranking(query_features, gallery_features,\n                            k1=15, k2=6, lambda_value=0.5):\n    all_features = np.concatenate([query_features, gallery_features], axis=0)\n    n_all = all_features.shape[0]\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    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    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_dist = lambda_value * dist + (1 - lambda_value) * jaccard_dist\n    return final_dist\n\nprint(\"Applying re-ranking...\")\ntest_dist = k_reciprocal_reranking(test_projected, test_projected,\n                                    k1=15, k2=6, lambda_value=0.5)\ntest_sim = np.clip(1 - test_dist, 0.0, 1.0)\n\n# generate the submission\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)\nsubmission_df = pd.DataFrame({\n    'row_id': test_pairs_df['row_id'],\n    'similarity': similarities\n})\n\nsubmission_df.to_csv(\"submission_sweep_reranked.csv\", index=False)\nprint(f\"Submission saved: {len(submission_df)} rows\")\n\nfrom IPython.display import FileLink\nFileLink('submission_sweep_reranked.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T23:11:35.919098Z","iopub.execute_input":"2026-03-03T23:11:35.919987Z","iopub.status.idle":"2026-03-03T23:14:17.794908Z","shell.execute_reply.started":"2026-03-03T23:11:35.919941Z","shell.execute_reply":"2026-03-03T23:14:17.794116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}