{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":29762,"databundleVersionId":2541532,"sourceType":"competition"},{"sourceId":11766825,"sourceType":"datasetVersion","datasetId":7374262},{"sourceId":11786457,"sourceType":"datasetVersion","datasetId":7385034},{"sourceId":11915250,"sourceType":"datasetVersion","datasetId":7419283},{"sourceId":12011827,"sourceType":"datasetVersion","datasetId":7393981}],"dockerImageVersionId":30124,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Import libraries and some directories ##","metadata":{}},{"cell_type":"code","source":"import pathlib\n\nimport torch\nimport torch.utils.data\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\n\nimport PIL.Image\nimport albumentations.pytorch\nimport cv2\nimport matplotlib.pyplot as plt\nfrom torchvision import models\n\nfrom tqdm.notebook import tqdm\nfrom typing import List, Tuple\nfrom pytorch_lightning.loggers import WandbLogger","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:19:34.928169Z","iopub.execute_input":"2025-06-01T15:19:34.928483Z","iopub.status.idle":"2025-06-01T15:19:36.300375Z","shell.execute_reply.started":"2025-06-01T15:19:34.928446Z","shell.execute_reply":"2025-06-01T15:19:36.299556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install wandb","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:19:46.698120Z","iopub.execute_input":"2025-06-01T15:19:46.698854Z","iopub.status.idle":"2025-06-01T15:19:54.753745Z","shell.execute_reply.started":"2025-06-01T15:19:46.698797Z","shell.execute_reply":"2025-06-01T15:19:54.752969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import wandb\nwandb.login()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:19:54.755282Z","iopub.execute_input":"2025-06-01T15:19:54.755525Z","iopub.status.idle":"2025-06-01T15:23:31.003241Z","shell.execute_reply.started":"2025-06-01T15:19:54.755494Z","shell.execute_reply":"2025-06-01T15:23:31.002419Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Data Preprocessing ###","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport os","metadata":{"execution":{"iopub.status.busy":"2025-06-01T15:23:34.042172Z","iopub.execute_input":"2025-06-01T15:23:34.042699Z","iopub.status.idle":"2025-06-01T15:23:34.046521Z","shell.execute_reply.started":"2025-06-01T15:23:34.042661Z","shell.execute_reply":"2025-06-01T15:23:34.045629Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.utils.data import Dataset, DataLoader\n\nIMAGE_ROOT = \"/kaggle/input/landmark-recognition-2021/train\"\nCSV_PATH = \"/kaggle/input/landmark-labels/train_with_landmark_names_fixed.csv\"\nIMAGE_SIZE = 224\n\n# === Load CSV ===\ndf = pd.read_csv(CSV_PATH,encoding = \"Latin1\")\n\n# Drop missing or malformed rows\ndf = df.dropna(subset=['id', 'landmark_id'])\n\n# Ensure landmark_id is int\ndf['landmark_id'] = df['landmark_id'].astype(int)\n\n# === Map landmark_id to class indices ===\nlandmark_id_to_idx = {lid: idx for idx, lid in enumerate(sorted(df['landmark_id'].unique()))}\ndf['class_idx'] = df['landmark_id'].map(landmark_id_to_idx)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:23:36.522542Z","iopub.execute_input":"2025-06-01T15:23:36.522843Z","iopub.status.idle":"2025-06-01T15:23:36.589194Z","shell.execute_reply.started":"2025-06-01T15:23:36.522794Z","shell.execute_reply":"2025-06-01T15:23:36.588607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Build image paths ===\ndef get_image_path(image_id):\n    return os.path.join(\n        IMAGE_ROOT, image_id[0], image_id[1], image_id[2], f\"{image_id}.jpg\"\n    )\n\ndf['image_path'] = df['id'].apply(get_image_path)\n\n# === Filter out missing files ===\ndf = df[df['image_path'].apply(os.path.exists)].reset_index(drop=True)\n\n# === Train/Validation split ===\ntrain_df, val_df = train_test_split(df, stratify=df['class_idx'], test_size=0.1, random_state=42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:23:38.956640Z","iopub.execute_input":"2025-06-01T15:23:38.956928Z","iopub.status.idle":"2025-06-01T15:25:01.946463Z","shell.execute_reply.started":"2025-06-01T15:23:38.956900Z","shell.execute_reply":"2025-06-01T15:25:01.945855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_id = \"111033371fee2d05\"\nprint(get_image_path(image_id))\n# Output: /kaggle/input/landmark-recognition-2021/train/1/1/1/111033371fee2d05.jpg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:25:12.019599Z","iopub.execute_input":"2025-06-01T15:25:12.020366Z","iopub.status.idle":"2025-06-01T15:25:12.024839Z","shell.execute_reply.started":"2025-06-01T15:25:12.020329Z","shell.execute_reply":"2025-06-01T15:25:12.024116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Transforms ===\ntrain_transform = A.Compose([\n    A.RandomResizedCrop(IMAGE_SIZE, IMAGE_SIZE, scale=(0.8, 1.0)),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.2),\n    A.ImageCompression(quality_lower=99, quality_upper=100),\n    A.RandomBrightnessContrast(p=0.2),\n    A.HueSaturationValue(p=0.2),\n    A.CLAHE(p=0.1),\n    A.GaussianBlur(p=0.1),\n    A.Normalize(),\n    ToTensorV2()\n])\n\nval_transform = A.Compose([\n    A.Resize(IMAGE_SIZE, IMAGE_SIZE),\n    A.Normalize(),\n    ToTensorV2()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:25:14.768407Z","iopub.execute_input":"2025-06-01T15:25:14.768679Z","iopub.status.idle":"2025-06-01T15:25:14.775129Z","shell.execute_reply.started":"2025-06-01T15:25:14.768650Z","shell.execute_reply":"2025-06-01T15:25:14.774319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Custom Dataset ===\nfrom PIL import Image\nimport torch\n\nclass LandmarkDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = Image.open(row['image_path']).convert(\"RGB\")\n        img = np.array(img)\n\n        if self.transform:\n            img = self.transform(image=img)['image']\n\n        label = row['class_idx']\n        return img, label\n\n# === Datasets & Loaders ===\ntrain_dataset = LandmarkDataset(train_df, transform=train_transform)\nval_dataset = LandmarkDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:25:17.780686Z","iopub.execute_input":"2025-06-01T15:25:17.781333Z","iopub.status.idle":"2025-06-01T15:25:17.790532Z","shell.execute_reply.started":"2025-06-01T15:25:17.781301Z","shell.execute_reply":"2025-06-01T15:25:17.789941Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## DOLG\nSingle-Stage Image Retrieval with Deep Orthogonal Fusion of\nLocal and Global Features\n","metadata":{}},{"cell_type":"markdown","source":"## Torch","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\n\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import DataLoader\nfrom pytorch_lightning import LightningModule","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:13:42.718576Z","iopub.execute_input":"2025-05-22T14:13:42.718866Z","iopub.status.idle":"2025-05-22T14:13:42.724178Z","shell.execute_reply.started":"2025-05-22T14:13:42.718836Z","shell.execute_reply":"2025-05-22T14:13:42.723290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ArcFace(nn.Module):\n    def __init__(self, in_features, out_features, scale_factor=64.0, margin=0.50, criterion=None):\n        super().__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = scale_factor\n        self.m = margin\n\n        # Initialize weights for the ArcFace layer\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.criterion = criterion if criterion else nn.CrossEntropyLoss()\n\n        # Precompute margin constants\n        self.cos_m = math.cos(self.m)\n        self.sin_m = math.sin(self.m)\n        self.th = math.cos(math.pi - self.m)\n        self.mm = math.sin(math.pi - self.m) * self.m\n\n    def forward(self, input, label):\n        # Project the input to the correct size (e.g., 512)\n        input_norm = F.normalize(input, p=2, dim=1)  # [batch_size, in_features]\n        weight_norm = F.normalize(self.weight, p=2, dim=1)  # [out_features, in_features]\n\n        # Cosine similarity between input and weights\n        cosine = F.linear(input_norm, weight_norm)  # [batch_size, num_classes]\n        cosine = cosine.clamp(-1.0, 1.0)  # numerical stability\n\n        # ArcFace margin adjustment\n        sine = torch.sqrt(1.0 - cosine ** 2 + 1e-6)\n        phi = cosine * self.cos_m - sine * self.sin_m\n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n\n        # One-hot encode labels\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, label.view(-1, 1), 1.0)\n\n        # Final logits\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s  # Apply scale factor\n\n        # Compute loss\n        loss = self.criterion(output, label)\n\n        return loss, output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:13:44.296819Z","iopub.execute_input":"2025-05-22T14:13:44.297103Z","iopub.status.idle":"2025-05-22T14:13:44.308978Z","shell.execute_reply.started":"2025-05-22T14:13:44.297069Z","shell.execute_reply":"2025-05-22T14:13:44.308013Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\n# Define input size and number of classes\nbatch_size = 8  # Example batch size\nin_features = 512  # Example input feature size (e.g., from a backbone network)\nout_features = 10  # Example number of classes (e.g., 10 classes)\n\n# Create a test input (random tensor with shape [batch_size, in_features])\ninput_tensor = torch.randn(batch_size, in_features)\n\n# Create a test label tensor (random labels, with shape [batch_size])\nlabel_tensor = torch.randint(0, out_features, (batch_size,))\n\n# Initialize ArcFace model\narcface = ArcFace(in_features=in_features, out_features=out_features)\n\n# Pass the input through the ArcFace model\nloss, output = arcface(input_tensor, label_tensor)\n\n# Check the loss and output shapes\nprint(\"Loss:\", loss.item())  # Print loss value\nprint(\"Output shape:\", output.shape)  # Check output shape (should be [batch_size, num_classes])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T08:13:04.569870Z","iopub.execute_input":"2025-05-14T08:13:04.570598Z","iopub.status.idle":"2025-05-14T08:13:04.583002Z","shell.execute_reply.started":"2025-05-14T08:13:04.570557Z","shell.execute_reply":"2025-05-14T08:13:04.582044Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#import timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\n\nfrom torch.utils.data import DataLoader\nfrom pytorch_lightning import LightningModule\n\n\nclass MultiAtrous(nn.Module):\n    def __init__(self, in_channel, out_channel, size, dilation_rates=[3, 6, 9]):\n        super().__init__()\n        self.dilated_convs = [\n            nn.Conv2d(in_channel, int(out_channel/4),\n                      kernel_size=3, dilation=rate, padding=rate)\n            for rate in dilation_rates\n        ]\n        self.gap_branch = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(in_channel, int(out_channel/4), kernel_size=1),\n            nn.ReLU(),\n            nn.Upsample(size=(size, size), mode='bilinear')\n        )\n        self.dilated_convs.append(self.gap_branch)\n        self.dilated_convs = nn.ModuleList(self.dilated_convs)\n\n    def forward(self, x):\n        local_feat = []\n        for dilated_conv in self.dilated_convs:\n            local_feat.append(dilated_conv(x))\n        local_feat = torch.cat(local_feat, dim=1)\n        return local_feat\n\n\nclass DolgLocalBranch(nn.Module):\n    def __init__(self, in_channel, out_channel, hidden_channel=2048):\n        super().__init__()\n        self.multi_atrous = MultiAtrous(in_channel, hidden_channel, size=int(224/8))\n        self.conv1x1_1 = nn.Conv2d(hidden_channel, out_channel, kernel_size=1)\n        self.conv1x1_2 = nn.Conv2d(\n            out_channel, out_channel, kernel_size=1, bias=False)\n        self.conv1x1_3 = nn.Conv2d(out_channel, out_channel, kernel_size=1)\n\n        self.relu = nn.ReLU()\n        self.bn = nn.BatchNorm2d(out_channel)\n        self.softplus = nn.Softplus()\n\n    def forward(self, x):\n        local_feat = self.multi_atrous(x)\n\n        local_feat = self.conv1x1_1(local_feat)\n        local_feat = self.relu(local_feat)\n        local_feat = self.conv1x1_2(local_feat)\n        local_feat = self.bn(local_feat)\n\n        attention_map = self.relu(local_feat)\n        attention_map = self.conv1x1_3(attention_map)\n        attention_map = self.softplus(attention_map)\n\n        local_feat = F.normalize(local_feat, p=2, dim=1)\n        local_feat = local_feat * attention_map\n        return local_feat\n\n\nclass OrthogonalFusion(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n    def forward(self, local_feat, global_feat):\n        global_feat_norm = torch.norm(global_feat, p=2, dim=1)\n        projection = torch.bmm(global_feat.unsqueeze(1), torch.flatten(local_feat, start_dim=2))\n        projection = torch.bmm(global_feat.unsqueeze(2), projection).view(local_feat.size())\n        projection = projection / (global_feat_norm * global_feat_norm).view(-1, 1, 1, 1)\n        orthogonal_comp = local_feat - projection\n        global_feat = global_feat.unsqueeze(-1).unsqueeze(-1)\n        return torch.cat([global_feat.expand(orthogonal_comp.size()), orthogonal_comp], dim=1)\n\nfrom pytorch_lightning import LightningModule\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import optim\nfrom torch.utils.data import DataLoader\n\n\nclass DolgNet(LightningModule):\n    def __init__(self, input_dim, hidden_dim, output_dim, num_of_classes, train_dataset, val_dataset):\n        super().__init__()\n        self.save_hyperparameters()  # Optional: logs args automatically\n\n        self.cnn = ResNetBackbone()\n        self.local_branch = DolgLocalBranch(512, hidden_dim)\n        self.orthogonal_fusion = OrthogonalFusion()\n\n        self.gem_pool = GeM()\n        self.gap = nn.AdaptiveAvgPool2d(1)\n\n        self.fc_1 = nn.Linear(1024, hidden_dim)  # for global feature processing\n\n        # Removed fc_2 — you’re using ArcFace\n        self.criterion = ArcFace(\n            in_features=2 * hidden_dim,\n            out_features=num_of_classes,\n            scale_factor=30,\n            margin=0.15,\n            criterion=nn.CrossEntropyLoss()\n        )\n\n        self.lr = Config.lr\n        self.train_dataset = train_dataset\n        self.val_dataset = val_dataset\n\n        self.freeze_resnet_layers()\n\n    def forward(self, x):\n        local_feat, global_feat = self.cnn(x)\n\n        local_feat = self.local_branch(local_feat)\n\n        global_feat = self.gem_pool(global_feat).squeeze()\n        global_feat = self.fc_1(global_feat)\n        global_feat = F.normalize(global_feat, p=2, dim=1)\n\n        feat = self.orthogonal_fusion(local_feat, global_feat)\n        feat = self.gap(feat).squeeze()\n        return feat\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        embeddings = self(x)\n        loss = self.criterion(embeddings, y)\n        self.log(\"train_loss\", loss)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        embeddings = self(x)\n        _, logits = self.criterion(embeddings, y)\n        confs, preds = torch.max(torch.softmax(logits, dim=1), dim=1)\n\n        return {\n            \"preds\": preds.cpu().numpy(),\n            \"labels\": y.cpu().numpy(),\n            \"confs\": confs.cpu().numpy()\n        }\n\n    def training_epoch_end(self, outputs):\n        # outputs is list of scalars\n        avg_loss = torch.stack(outputs).mean()\n        self.log(\"avg_train_loss\", avg_loss, prog_bar=True)\n\n    def validation_epoch_end(self, outputs):\n        all_preds, all_labels, all_confs = [], [], []\n        for out in outputs:\n            all_preds.extend(out[\"preds\"])\n            all_labels.extend(out[\"labels\"])\n            all_confs.extend(out[\"confs\"])\n\n        gap = compute_gap(all_preds, all_confs, all_labels)\n        self.log(\"val_gap\", gap, prog_bar=True, logger=True)\n\n    def configure_optimizers(self):\n        optimizer = optim.SGD(self.parameters(), lr=self.lr, momentum=0.9, weight_decay=1e-5)\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.epochs)\n        return [optimizer], [scheduler]\n\n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, batch_size=32,\n                          shuffle=True, num_workers=2)\n\n    def val_dataloader(self):\n        return DataLoader(self.val_dataset, batch_size=32,\n                          shuffle=False, num_workers=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T08:12:50.751010Z","iopub.execute_input":"2025-05-14T08:12:50.751277Z","iopub.status.idle":"2025-05-14T08:12:50.780355Z","shell.execute_reply.started":"2025-05-14T08:12:50.751250Z","shell.execute_reply":"2025-05-14T08:12:50.779329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DolgNet(LightningModule):\n    def __init__(self, input_dim, hidden_dim, output_dim, num_of_classes, train_dataset, val_dataset):\n        super().__init__()\n        self.cnn = ResNetBackbone()  # Returns (local_feat, global_feat)\n        self.train_dataset = train_dataset\n        self.val_dataset = val_dataset\n\n        self.local_branch = DolgLocalBranch(512, hidden_dim)\n        self.gem_pool = GeM()\n        self.gap = nn.AdaptiveAvgPool2d(1)\n\n        # Fusion & projection\n        self.orthogonal_fusion = OrthogonalFusion()\n        self.fc_proj = nn.Linear(hidden_dim + 512, 512)  # local + global -> ArcFace\n\n        # ArcFace loss\n        self.criterion = ArcFace(\n            in_features=512,\n            out_features=num_of_classes,\n            scale_factor=30,\n            margin=0.15,\n            criterion=nn.CrossEntropyLoss()\n        )\n        self.lr = 2e-4\n        self.freeze_resnet_layers()\n\n    def forward(self, x):\n        # Backbone output\n        local_feat, global_feat = self.cnn(x)  # local: [B, 512, 28, 28], global: [B, 2048, H, W]\n\n        # Local branch\n        local_feat = self.local_branch(local_feat)  # [B, hidden_dim, H, W]\n        local_feat = self.gap(local_feat).squeeze(-1).squeeze(-1)  # [B, hidden_dim]\n\n        # Global branch\n        global_feat = self.gem_pool(global_feat).view(global_feat.size(0), -1)  # [B, 2048]\n        global_feat = F.normalize(global_feat, p=2, dim=1)\n        global_feat = nn.Linear(2048, 512)(global_feat)  # Inline or define self.global_proj if reused\n        global_feat = F.normalize(global_feat, p=2, dim=1)\n\n        # Fusion\n        fused_feat = self.orthogonal_fusion(local_feat, global_feat)  # [B, hidden_dim + 512]\n        fused_feat = self.fc_proj(fused_feat)  # [B, 512]\n        fused_feat = F.normalize(fused_feat, p=2, dim=1)\n\n        return fused_feat\n\n    def freeze_resnet_layers(self):\n        layers_to_freeze = ['stem', 'layer1', 'layer2']\n        \n        for layer_name in layers_to_freeze:\n            layer = getattr(self.cnn, layer_name)\n            for param in layer.parameters():\n                param.requires_grad = False\n            print(f\"Froze {layer_name}\")\n    \n        # Ensure layer3 is trainable\n        for param in self.cnn.layer3.parameters():\n            param.requires_grad = True\n        print(\"Layer3 is trainable\")\n    \n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        feats = self.forward(x)\n        loss, logits = self.criterion(feats, y)\n        self.log('train_loss', loss)\n        return loss\n\n\n    def validation_step(self, batch, batch_idx):\n        img, label = batch\n        embd = self(img)\n        _, logits = self.criterion(embd, label)\n    \n        confs, preds = torch.max(torch.softmax(logits, dim=1), dim=1)\n    \n        return {\n            \"preds\": preds.cpu().numpy(),\n            \"labels\": label.cpu().numpy(),\n            \"confs\": confs.cpu().numpy()\n        }\n    \n    def configure_optimizers(self):\n        optimizer = optim.SGD(self.parameters(), lr=self.lr,\n                              momentum=0.9, weight_decay=1e-5)\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=1000)\n        return [optimizer], [scheduler]\n\n    def training_epoch_end(self, outputs):\n        avg_loss = torch.stack(outputs).mean()\n        self.log('avg_train_loss', avg_loss, prog_bar=True)\n\n    def validation_epoch_end(self, outputs):\n        all_preds = []\n        all_labels = []\n        all_confs = []\n    \n        for out in outputs:\n            all_preds.extend(out['preds'])\n            all_labels.extend(out['labels'])\n            all_confs.extend(out['confs'])\n\n        gap = compute_gap(all_preds, all_confs, all_labels)\n        self.log(\"val_gap\", gap, prog_bar=True, logger=True)\n        \n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, batch_size=32, shuffle=True, num_workers=2)\n\n    def val_dataloader(self):\n        return DataLoader(self.val_dataset, batch_size=32, shuffle=False, num_workers=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T08:23:09.582673Z","iopub.execute_input":"2025-05-14T08:23:09.583040Z","iopub.status.idle":"2025-05-14T08:23:09.603284Z","shell.execute_reply.started":"2025-05-14T08:23:09.583005Z","shell.execute_reply":"2025-05-14T08:23:09.602151Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Rebuild The project starts from here","metadata":{}},{"cell_type":"markdown","source":"### Loading data and images with albumentations","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\n\n# === Config ===\nIMAGE_ROOT = \"/kaggle/input/landmark-recognition-2021/train\"\nCSV_PATH = \"/kaggle/input/landmark-labels/train_with_landmark_names_fixed.csv\nIMAGE_SIZE = 224\nBATCH_SIZE = 32\nNUM_WORKERS = 4\n\n# === Load and preprocess CSV ===\ndf = pd.read_csv(CSV_PATH, encoding=\"Latin1\")\ndf = df.dropna(subset=['id', 'landmark_id'])\ndf['landmark_id'] = df['landmark_id'].astype(int)\n\n# Map landmark IDs to class indices\nlandmark_id_to_idx = {lid: idx for idx, lid in enumerate(sorted(df['landmark_id'].unique()))}\ndf['class_idx'] = df['landmark_id'].map(landmark_id_to_idx)\n\n# Generate full image paths\ndef get_image_path(image_id):\n    return os.path.join(IMAGE_ROOT, image_id[0], image_id[1], image_id[2], f\"{image_id}.jpg\")\n\ndf['image_path'] = df['id'].apply(get_image_path)\ndf = df[df['image_path'].apply(os.path.exists)].reset_index(drop=True)\n\n# === Train/Val Split ===\ntrain_df, val_df = train_test_split(df, stratify=df['class_idx'], test_size=0.1, random_state=42)\n\n# === Custom Dataset ===\nclass LandmarkDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.df = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        try:\n            img = Image.open(row['image_path']).convert(\"RGB\")\n            img = np.array(img)\n        except Exception as e:\n            print(f\"Failed to load image {row['image_path']}: {e}\")\n            # Skip this sample if image loading fails\n            return self.__getitem__((idx + 1) % len(self.df))  # Skip this sample\n\n        if self.transform:\n            img = self.transform(image=img)['image']\n\n        label = row['class_idx']\n        return img, label\n\ntrain_transform = A.Compose([\n    A.RandomResizedCrop(IMAGE_SIZE, IMAGE_SIZE, scale=(0.8, 1.0)),\n    A.HorizontalFlip(p=0.3),\n    A.RandomBrightnessContrast(p=0.2),\n    A.HueSaturationValue(p=0.2),\n    A.CLAHE(p=0.1),\n    A.GaussianBlur(p=0.1),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),  # Using ResNet mean & std\n    ToTensorV2()\n])\n\nval_transform = A.Compose([\n    A.Resize(IMAGE_SIZE, IMAGE_SIZE),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),  # Same mean & std as train\n    ToTensorV2()\n])\n\n# === DataLoaders ===\ntrain_dataset = LandmarkDataset(train_df, transform=train_transform)\nval_dataset = LandmarkDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\n\n# === DataLoaders ===\ntrain_dataset = LandmarkDataset(train_df, transform=train_transform)\nval_dataset = LandmarkDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:03:01.328797Z","iopub.execute_input":"2025-06-01T15:03:01.329112Z","iopub.status.idle":"2025-06-01T15:03:32.399453Z","shell.execute_reply.started":"2025-06-01T15:03:01.329081Z","shell.execute_reply":"2025-06-01T15:03:32.398849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check if train and val datasets are non-empty\nprint(f\"Training dataset size: {len(train_loader.dataset)}\")\nprint(f\"Validation dataset size: {len(val_loader.dataset)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:28:22.443423Z","iopub.execute_input":"2025-06-01T15:28:22.444016Z","iopub.status.idle":"2025-06-01T15:28:22.448721Z","shell.execute_reply.started":"2025-06-01T15:28:22.443978Z","shell.execute_reply":"2025-06-01T15:28:22.447919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check if data is loading properly by fetching one batch\ntrain_iter = iter(train_loader)\nval_iter = iter(val_loader)\n\n# Try loading the first batch from the train set\ntrain_images, train_labels = next(train_iter)\nprint(f\"Train batch - Images shape: {train_images.shape}, Labels shape: {train_labels.shape}\")\n\n# Try loading the first batch from the validation set\nval_images, val_labels = next(val_iter)\nprint(f\"Val batch - Images shape: {val_images.shape}, Labels shape: {val_labels.shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:28:37.197646Z","iopub.execute_input":"2025-06-01T15:28:37.198282Z","iopub.status.idle":"2025-06-01T15:28:38.318971Z","shell.execute_reply.started":"2025-06-01T15:28:37.198249Z","shell.execute_reply":"2025-06-01T15:28:38.317909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Get a sample from the training set\ntrain_img, train_label = train_loader.dataset[0]  # First sample from train dataset\nplt.imshow(train_img.permute(1, 2, 0))  # Convert from (C, H, W) to (H, W, C)\nplt.title(f\"Train Sample - Label: {train_label}\")\nplt.show()\n\n# Get a sample from the validation set\nval_img, val_label = val_loader.dataset[0]  # First sample from val dataset\nplt.imshow(val_img.permute(1, 2, 0))  # Convert from (C, H, W) to (H, W, C)\nplt.title(f\"Val Sample - Label: {val_label}\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:30:00.715368Z","iopub.execute_input":"2025-06-01T15:30:00.715687Z","iopub.status.idle":"2025-06-01T15:30:01.124206Z","shell.execute_reply.started":"2025-06-01T15:30:00.715647Z","shell.execute_reply":"2025-06-01T15:30:01.123551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\n\n# === File paths ===\nTRAIN_CSV = '/kaggle/input/landmark-recognition-2021/train.csv'\nTRAIN_DIR = '/kaggle/input/landmark-recognition-2021/train'\nOUTPUT_CSV = 'train_top_100_classes.csv'  # This will save to the working directory\n\n# === Load the full training CSV ===\ndf = pd.read_csv(TRAIN_CSV)\n\n# === Drop rows with missing landmark_id ===\ndf = df.dropna(subset=['landmark_id'])\n\n# === Get top 100 most frequent landmark_ids ===\ntop_100_ids = df['landmark_id'].value_counts().nlargest(100).index\n\n# === Filter dataset to keep only top 100 landmark_ids ===\ndf = df[df['landmark_id'].isin(top_100_ids)].copy()\n\n# === Remap landmark_id to new_id (0 to 99) ===\nunique_ids = sorted(df['landmark_id'].unique())\nid_to_new_id = {old: new for new, old in enumerate(unique_ids)}\ndf['new_id'] = df['landmark_id'].map(id_to_new_id)\n\n# === Build image paths ===\ndef make_path(img_id):\n    return os.path.join(TRAIN_DIR, f'{img_id[0]}/{img_id[1]}/{img_id[2]}/{img_id}.jpg')\n\ndf['path'] = df['id'].apply(make_path)\n\n# === Remove rows with non-existing image paths (optional but recommended) ===\ndf = df[df['path'].apply(os.path.exists)]\n\n# === Save to CSV for reproducibility ===\ndf.to_csv(OUTPUT_CSV, index=False)\nprint(f\"Saved filtered dataset with top 100 classes to: {OUTPUT_CSV}\")\n\n# === Train/Val/Test split ===\nX_train, X_temp, y_train, y_temp = train_test_split(\n    df[['id', 'path']], df['new_id'],\n    train_size=0.8,\n    stratify=df['new_id'],\n    random_state=123,\n    shuffle=True\n)\n\nX_val, X_test, y_val, y_test = train_test_split(\n    X_temp, y_temp,\n    train_size=0.5,\n    stratify=y_temp,\n    random_state=123,\n    shuffle=True\n)\n\n# === Confirm splits ===\nassert X_train.shape[0] + X_val.shape[0] + X_test.shape[0] == df.shape[0]\nprint(f\"Train: {X_train.shape[0]} samples, {y_train.nunique()} classes\")\nprint(f\"Val:   {X_val.shape[0]} samples, {y_val.nunique()} classes\")\nprint(f\"Test:  {X_test.shape[0]} samples, {y_test.nunique()} classes\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Class distribution in train set:\")\nprint(y_train.value_counts(normalize=True).head())\n\nprint(\"\\nClass distribution in val set:\")\nprint(y_val.value_counts(normalize=True).head())\n\nprint(\"\\nClass distribution in test set:\")\nprint(y_test.value_counts(normalize=True).head())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Classes distribution on training, validation and test sets\nplt.figure(figsize = (10, 3))\nax = sns.histplot(y_train, bins=75, kde = True)\nax.set_title('Distribution of Landmarks on training set')\nplt.tight_layout()\n\nplt.figure(figsize = (10, 3))\nax = sns.histplot(y_val, bins=75, kde = True)\nax.set_title('Distribution of Landmarks on validation set')\nplt.tight_layout()\n\nplt.figure(figsize = (10, 3))\nax = sns.histplot(y_test, bins=75, kde = True)\nax.set_title('Distribution of Landmarks on test set')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:47:37.548947Z","iopub.execute_input":"2025-06-01T15:47:37.549259Z","iopub.status.idle":"2025-06-01T15:47:37.577024Z","shell.execute_reply.started":"2025-06-01T15:47:37.549230Z","shell.execute_reply":"2025-06-01T15:47:37.576180Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\n# Creating image directories for classes subset\nNEW_BASE_DIR = \"/kaggle/working\"\n\n# Training set directory\nfor file, path, landmark in tqdm(zip(X_train['id'], X_train['path'], y_train)):\n    dir = f\"{NEW_BASE_DIR}/train_sub/{str(landmark)}\"\n    os.makedirs(dir, exist_ok = True)\n    fname = f\"{file}.jpg\"\n    shutil.copyfile(src = path, dst = f\"{dir}/{fname}\")\n\n# Validation set directory    \nfor file, path, landmark in tqdm(zip(X_val['id'], X_val['path'], y_val)):\n    dir = f\"{NEW_BASE_DIR}/val_sub/{str(landmark)}\"\n    os.makedirs(dir, exist_ok = True)\n    fname = f\"{file}.jpg\"\n    shutil.copyfile(src = path, dst = f\"{dir}/{fname}\")\n\n# Testing set directory\nfor file, path, landmark in tqdm(zip(X_test['id'], X_test['path'], y_test)):\n    dir = f\"{NEW_BASE_DIR}/test_sub/{str(landmark)}\"\n    os.makedirs(dir, exist_ok = True)\n    fname = f\"{file}.jpg\"\n    shutil.copyfile(src = path, dst = f\"{dir}/{fname}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Creating tensorflow tf.data.Dataset\nfrom tensorflow.keras.utils import image_dataset_from_directory\n\nIMG_SIZE = 224\nBATCH_SIZE = 16\n\nprint(\"Building training dataset...\")\n# Training tf.data.Dataset\ntrain_ds = image_dataset_from_directory(f\"{NEW_BASE_DIR}/train_sub\",\n                                        label_mode = 'int',\n                                        shuffle = True,\n                                        image_size = (IMG_SIZE, IMG_SIZE),\n                                        batch_size = BATCH_SIZE)\n\nprint(\"Building validation dataset...\")\n# Validation tf.data.Dataset\nval_ds = image_dataset_from_directory(f\"{NEW_BASE_DIR}/val_sub\",\n                                        label_mode = 'int',\n                                        shuffle = True,\n                                        image_size = (IMG_SIZE, IMG_SIZE),\n                                        batch_size = BATCH_SIZE)\n\nprint(\"Building test dataset...\")\n# Test tf.data.Dataset\ntest_ds = image_dataset_from_directory(f\"{NEW_BASE_DIR}/test_sub\",\n                                        label_mode = 'int',\n                                        shuffle = True,\n                                        image_size = (IMG_SIZE, IMG_SIZE),\n                                        batch_size = BATCH_SIZE)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualizing a random batch from training dataset\nfor data_batch, labels_batch in train_ds.take(1):\n    ncols = 4\n    nrows = int(data_batch.shape[0]/ncols)\n    fig, ax = plt.subplots(nrows = nrows, ncols = ncols, figsize=(10, 11),\n                           sharex = True, sharey = True)\n    img_counter = 0\n    for image, label in zip(data_batch, labels_batch):\n        axi = ax.flat[img_counter]\n        axi.imshow(image/255.)\n        label = label.numpy()\n#         axi.set_title(np.where(label == 1)[0])\n        axi.set_title(label)\n        img_counter += 1\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:41:11.379068Z","iopub.status.idle":"2025-06-01T15:41:11.379491Z","shell.execute_reply.started":"2025-06-01T15:41:11.379280Z","shell.execute_reply":"2025-06-01T15:41:11.379300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms, datasets\nfrom torch.utils.data import DataLoader\n\nIMG_SIZE = 224\nBATCH_SIZE = 16\n\n# === Transforms ===\ntransform = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),  # Convert PIL image to PyTorch tensor\n    transforms.Normalize([0.485, 0.456, 0.406],  # ImageNet mean\n                         [0.229, 0.224, 0.225])  # ImageNet std\n])\n\n# === PyTorch Dataset Loaders ===\ntrain_dataset = datasets.ImageFolder(root=f\"{NEW_BASE_DIR}/train_sub\", transform=transform)\nval_dataset = datasets.ImageFolder(root=f\"{NEW_BASE_DIR}/val_sub\", transform=transform)\ntest_dataset = datasets.ImageFolder(root=f\"{NEW_BASE_DIR}/test_sub\", transform=transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\ntest_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The DOLG model","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\nimport torch.nn.functional as F\n\n\nclass AttentionFusion(nn.Module):\n    def __init__(self, channels):\n        super().__init__()\n        self.attn = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(channels, channels, 1),\n            nn.ReLU(),\n            nn.Conv2d(channels, channels, 1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, local_feat, global_feat):\n        b, c, h, w = local_feat.shape\n        global_feat_expanded = global_feat.view(b, c, 1, 1).expand(-1, -1, h, w)\n        attn = self.attn(local_feat + global_feat_expanded)\n        fused = local_feat * attn + global_feat_expanded * (1 - attn)\n        return fused","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class NonLocalBlock(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.inter_channels = in_channels // 2\n\n        self.g = nn.Conv2d(in_channels, self.inter_channels, kernel_size=1)\n        self.theta = nn.Conv2d(in_channels, self.inter_channels, kernel_size=1)\n        self.phi = nn.Conv2d(in_channels, self.inter_channels, kernel_size=1)\n        self.W = nn.Conv2d(self.inter_channels, in_channels, kernel_size=1)\n        nn.init.constant_(self.W.weight, 0)\n        nn.init.constant_(self.W.bias, 0)\n\n    def forward(self, x):\n        batch_size, C, H, W = x.size()\n\n        g_x = self.g(x).view(batch_size, self.inter_channels, -1)  # [B, C', N]\n        g_x = g_x.permute(0, 2, 1)  # [B, N, C']\n\n        theta_x = self.theta(x).view(batch_size, self.inter_channels, -1)  # [B, C', N]\n        theta_x = theta_x.permute(0, 2, 1)  # [B, N, C']\n\n        phi_x = self.phi(x).view(batch_size, self.inter_channels, -1)  # [B, C', N]\n\n        f = torch.matmul(theta_x, phi_x)  # [B, N, N]\n        f_div_C = F.softmax(f, dim=-1)\n\n        y = torch.matmul(f_div_C, g_x)  # [B, N, C']\n        y = y.permute(0, 2, 1).contiguous().view(batch_size, self.inter_channels, H, W)\n        W_y = self.W(y)\n\n        return W_y + x  # residual connection","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import math\n\nclass ArcFace(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.3):\n        super(ArcFace, self).__init__()\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n        self.s = s\n        self.m = m\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, labels):\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))  # [B, C]\n        sine = torch.sqrt((1.0 - cosine ** 2).clamp(min=0.0))\n        phi = cosine * self.cos_m - sine * self.sin_m  # cos(θ + m)\n\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, labels.view(-1, 1), 1.0)\n\n        logits = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        logits *= self.s\n        return logits","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DOLG_ArcFace(nn.Module):\n    def __init__(self, embedding_dim=512):\n        super().__init__()\n        resnet = models.resnet50(pretrained=True)\n\n        # Shared layers\n        self.backbone_common = nn.Sequential(\n            resnet.conv1, resnet.bn1, resnet.relu,\n            resnet.maxpool, resnet.layer1,\n            resnet.layer2, resnet.layer3\n        )\n\n        # Global branch\n        self.backbone_global = resnet.layer4\n        self.global_pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.global_fc = nn.Linear(2048, embedding_dim)\n\n        # Local branch with multi-atrous conv + 1x1 projection\n        self.local_branch = nn.Sequential(\n            nn.Conv2d(1024, 512, kernel_size=3, padding=1, dilation=1),\n            nn.ReLU(),\n            nn.Conv2d(1024, 512, kernel_size=3, padding=2, dilation=2),\n            nn.ReLU(),\n            nn.Conv2d(1024, 512, kernel_size=3, padding=3, dilation=3),\n            nn.ReLU(),\n        )\n        self.local_proj = nn.Conv2d(512 * 3, embedding_dim, kernel_size=1)\n\n        # Non-local block for self-attention\n        self.self_attn = NonLocalBlock(embedding_dim)\n\n        self.fusion = AttentionFusion(embedding_dim)\n\n        self.head = nn.Sequential(\n            nn.Conv2d(embedding_dim, embedding_dim, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten()\n        )\n\n    def forward(self, x):\n        shared_feat = self.backbone_common(x)  # Output: [B, 1024, H, W]\n\n        # Global branch\n        global_feat_map = self.backbone_global(shared_feat)  # Output: [B, 2048, H/2, W/2]\n        global_feat = self.global_pool(global_feat_map).view(x.size(0), -1)  # [B, 2048]\n        global_feat = self.global_fc(global_feat)  # [B, 512]\n\n        # Local branch with atrous conv\n        feat1 = F.relu(self.local_branch[0](shared_feat))\n        feat2 = F.relu(self.local_branch[2](shared_feat))\n        feat3 = F.relu(self.local_branch[4](shared_feat))\n        local_feat = torch.cat([feat1, feat2, feat3], dim=1)  # [B, 512*3, H, W]\n        local_feat = self.local_proj(local_feat)  # [B, 512, H, W]\n\n        # Apply self-attention (only once)\n        local_feat = self.self_attn(local_feat)\n\n        # Fusion\n        fused_feat = self.fusion(local_feat, global_feat)  # [B, 512, H, W]\n        emb = self.head(fused_feat)  # [B, 512]\n        return emb","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DOLG_ArcFace(nn.Module):\n    def __init__(self, embedding_dim=512):\n        super().__init__()\n        resnet = models.resnet50(pretrained=True)\n\n        # Shared layers\n        self.backbone_common = nn.Sequential(\n            resnet.conv1, resnet.bn1, resnet.relu,\n            resnet.maxpool, resnet.layer1,\n            resnet.layer2, resnet.layer3\n        )\n\n        # ResNet layer4 expects input with 1024 channels (not from local_conv!)\n        self.backbone_global = resnet.layer4\n\n        self.global_pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.global_fc = nn.Linear(2048, embedding_dim)\n\n        self.local_branch = nn.Sequential(\n        nn.Conv2d(1024, 512, kernel_size=3, padding=1, dilation=1),\n        nn.Conv2d(1024, 512, kernel_size=3, padding=2, dilation=2),\n        nn.Conv2d(1024, 512, kernel_size=3, padding=3, dilation=3),\n    # Concatenate all, then project back with 1x1 conv\n)\n\n        self.fusion = AttentionFusion(embedding_dim)\n\n        self.head = nn.Sequential(\n            nn.Conv2d(embedding_dim, embedding_dim, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten()\n        )\n\n    def forward(self, x):\n        shared_feat = self.backbone_common(x)  # Output: [B, 1024, H, W]\n\n        # Global branch\n        global_feat_map = self.backbone_global(shared_feat)  # Output: [B, 2048, H/2, W/2]\n        global_feat = self.global_pool(global_feat_map).view(x.size(0), -1)  # [B, 2048]\n        global_feat = self.global_fc(global_feat)  # [B, 512]\n\n        # Local branch\n        local_feat = self.local_conv(shared_feat)  # [B, 512, H, W]\n\n        # Fuse\n        fused_feat = self.fusion(local_feat, global_feat)  # [B, 512, H, W]\n        emb = self.head(fused_feat)  # [B, 512]\n        return emb","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T01:28:39.366990Z","iopub.execute_input":"2025-05-23T01:28:39.367256Z","iopub.status.idle":"2025-05-23T01:28:39.375023Z","shell.execute_reply.started":"2025-05-23T01:28:39.367227Z","shell.execute_reply":"2025-05-23T01:28:39.374193Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Tạo tập train/val  ","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\n# === Train/Val Split ===\ntrain_df, val_df = train_test_split(df, test_size=0.1, stratify=df['class_idx'], random_state=42)\ntrain_df = train_df.copy()\nval_df = val_df.copy()\n\ntrain_df['image_path'] = train_df['id'].apply(lambda x: os.path.join(IMAGE_ROOT, f\"{x[0]}/{x[1]}/{x[2]}/{x}.jpg\"))\nval_df['image_path'] = val_df['id'].apply(lambda x: os.path.join(IMAGE_ROOT, f\"{x[0]}/{x[1]}/{x[2]}/{x}.jpg\"))\n\n# === Reload Datasets with correct DataFrames ===\ntrain_dataset = LandmarkDataset(train_df, transform=train_transform)\nval_dataset = LandmarkDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:03:50.732202Z","iopub.execute_input":"2025-06-01T15:03:50.732453Z","iopub.status.idle":"2025-06-01T15:03:50.839568Z","shell.execute_reply.started":"2025-06-01T15:03:50.732426Z","shell.execute_reply":"2025-06-01T15:03:50.839027Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training process","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\n# Load the CSV file\ndf = pd.read_csv(\"/kaggle/input/landmark-labels/train_with_landmark_names_fixed.csv\")\n\n# Get the number of unique landmark_id\nunique_landmarks = df['landmark_id'].nunique()\n\nprint(f\"Number of unique landmark classes: {unique_landmarks}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:06:22.395621Z","iopub.execute_input":"2025-06-01T15:06:22.396345Z","iopub.status.idle":"2025-06-01T15:06:22.439646Z","shell.execute_reply.started":"2025-06-01T15:06:22.396315Z","shell.execute_reply":"2025-06-01T15:06:22.438906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = 102","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:06:35.833625Z","iopub.execute_input":"2025-06-01T15:06:35.834412Z","iopub.status.idle":"2025-06-01T15:06:35.838046Z","shell.execute_reply.started":"2025-06-01T15:06:35.834377Z","shell.execute_reply":"2025-06-01T15:06:35.837265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\nmodel = DOLG_ArcFace().to(device)\narcface_head = ArcFace(512, num_classes).to(device)\noptimizer = torch.optim.Adam(list(model.parameters()) + list(arcface_head.parameters()), lr=1e-4)\n\n# Training loop:\nfor imgs, labels in train_loader:\n    imgs, labels = imgs.to(device), labels.to(device)\n    optimizer.zero_grad()\n    embeddings = model(imgs)  # [B, 512]\n    logits = arcface_head(embeddings, labels)  # [B, num_classes]\n    #logits = head(embeddings)\n    loss = criterion(logits, labels)\n    loss.backward()\n    optimizer.step()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import wandb\n\nwandb.init(project=\"landmark-recognition\", name=\"resnet-arcface-final\", config={\n    \"epochs\": 30,\n    \"batch_size\": 32,\n    \"image_size\": IMAGE_SIZE,\n    \"embedding_dim\": 512,\n    \"optimizer\": \"Adam\",\n    \"lr\": 1e-4\n})","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Evaluation define","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score, f1_score, top_k_accuracy_score, confusion_matrix\nimport numpy as np","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_gap(preds, confs, targets):\n    df = pd.DataFrame({\n        \"pred\": preds,\n        \"conf\": confs,\n        \"target\": targets\n    })\n    df = df.sort_values(\"conf\", ascending=False).reset_index(drop=True)\n\n    correct = 0\n    total_precision = 0.0\n\n    for i, row in df.iterrows():\n        if row[\"pred\"] == row[\"target\"]:\n            correct += 1\n            total_precision += correct / (i + 1)\n\n    return total_precision / len(df)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### New train code => Final results","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\n\nle = LabelEncoder()\ndf['class_idx'] = le.fit_transform(df['class_idx'])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:35:23.716282Z","iopub.execute_input":"2025-06-01T15:35:23.716599Z","iopub.status.idle":"2025-06-01T15:35:23.724715Z","shell.execute_reply.started":"2025-06-01T15:35:23.716550Z","shell.execute_reply":"2025-06-01T15:35:23.723938Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = len(df['class_idx'].unique())\narcface_head = ArcFace(in_features=512, out_features=num_classes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T15:35:32.272770Z","iopub.execute_input":"2025-06-01T15:35:32.273584Z","iopub.status.idle":"2025-06-01T15:35:32.279980Z","shell.execute_reply.started":"2025-06-01T15:35:32.273537Z","shell.execute_reply":"2025-06-01T15:35:32.279147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import wandb\nimport torch.nn.functional as F\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\n# Initialize wandb\n#wandb.init(project=\"landmark-recognition\", name=\"resnet-dolg-arcface\")\n\nbest_gap = 0.0\nnum_epochs = 30\n\ntrain_losses, val_losses = [], []\ntrain_accuracies, val_accuracies = [], []\ngap_scores = []\n\n# Scheduler: Reduce LR when GAP plateaus\nscheduler = ReduceLROnPlateau(\n    optimizer,\n    mode='max',\n    factor=0.5,      # Reduce LR by half\n    patience=2,      # Wait 2 epochs without improvement\n    min_lr=1e-6,     # Prevent LR from becoming too small\n    verbose=True\n)\n\nfor epoch in range(num_epochs):\n    model.train()\n    arcface_head.train()\n\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    loop = tqdm(train_loader, desc=f\"Epoch [{epoch+1}/{num_epochs}]\")\n\n    for imgs, labels in loop:\n        imgs, labels = imgs.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        embeddings = model(imgs)\n        logits = arcface_head(embeddings, labels)\n        loss = criterion(logits, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        preds = torch.argmax(logits, dim=1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n        loop.set_postfix(loss=loss.item(), acc=100 * correct / total)\n\n    train_acc = 100 * correct / total\n    train_loss = running_loss / len(train_loader)\n    print(f\"Epoch {epoch+1}, Loss: {train_loss:.4f}, Accuracy: {train_acc:.2f}%\")\n\n    # === Validation + GAP ===\n    model.eval()\n    arcface_head.eval()\n\n    val_loss = 0.0\n    val_correct = 0\n    val_total = 0\n\n    all_preds = []\n    all_confs = []\n    all_labels = []\n\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            imgs, labels = imgs.to(device), labels.to(device)\n\n            embeddings = model(imgs)\n            logits = arcface_head(embeddings, labels)\n            loss = criterion(logits, labels)\n\n            val_loss += loss.item()\n\n            probs = F.softmax(logits, dim=1)\n            confs, preds = torch.max(probs, dim=1)\n\n            all_preds.extend(preds.cpu().tolist())\n            all_confs.extend(confs.cpu().tolist())\n            all_labels.extend(labels.cpu().tolist())\n\n            val_correct += (preds == labels).sum().item()\n            val_total += labels.size(0)\n\n    val_acc = 100 * val_correct / val_total\n    val_loss /= len(val_loader)\n    gap_score = compute_gap(all_preds, all_confs, all_labels)\n\n    print(f\"Validation Loss: {val_loss:.4f}, Accuracy: {val_acc:.2f}%, GAP: {gap_score:.4f}\")\n\n    # Step the LR scheduler\n    scheduler.step(gap_score)\n\n    # Log to wandb\n    wandb.log({\n        \"train_loss\": train_loss,\n        \"train_acc\": train_acc,\n        \"val_loss\": val_loss,\n        \"val_acc\": val_acc,\n        \"gap\": gap_score,\n        \"lr\": optimizer.param_groups[0]['lr']  # log current LR\n    })\n\n    # Save best model\n    if gap_score > best_gap:\n        best_gap = gap_score\n        torch.save({\n            \"model_state_dict\": model.state_dict(),\n            \"arcface_state_dict\": arcface_head.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"gap\": gap_score,\n            \"epoch\": epoch + 1\n        }, \"best_model.pth\")\n        print(\"✅ Saved new best model with GAP:\", best_gap)\n\n    train_losses.append(train_loss)\n    train_accuracies.append(train_acc)\n    val_losses.append(val_loss)\n    val_accuracies.append(val_acc)\n    gap_scores.append(gap_score)\n\n    # === Plotting ===\n    plt.figure(figsize=(12, 4))\n\n    # Loss\n    plt.subplot(1, 3, 1)\n    plt.plot(train_losses, label='Train Loss')\n    plt.plot(val_losses, label='Val Loss')\n    plt.title('Loss')\n    plt.xlabel('Epoch')\n    plt.legend()\n\n    # Accuracy\n    plt.subplot(1, 3, 2)\n    plt.plot(train_accuracies, label='Train Acc')\n    plt.plot(val_accuracies, label='Val Acc')\n    plt.title('Accuracy')\n    plt.xlabel('Epoch')\n    plt.legend()\n\n    # GAP Score\n    plt.subplot(1, 3, 3)\n    plt.plot(gap_scores, label='GAP')\n    plt.title('GAP Score')\n    plt.xlabel('Epoch')\n    plt.legend()\n\n    plt.suptitle(f\"Training Metrics - Epoch {epoch+1}\", fontsize=16)\n    plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Save checkpoint","metadata":{}},{"cell_type":"code","source":"checkpoint = torch.load(\"best_model.pth\", map_location=device)\n\nmodel = DOLG_ArcFace().to(device)\narcface_head = ArcFace(512, num_classes).to(device)\n\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\narcface_head.load_state_dict(checkpoint[\"arcface_state_dict\"])\noptimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T14:43:33.447410Z","iopub.status.idle":"2025-05-14T14:43:33.447683Z","shell.execute_reply.started":"2025-05-14T14:43:33.447524Z","shell.execute_reply":"2025-05-14T14:43:33.447536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cd /kaggle/working\nfrom IPython.display import FileLink\nFileLink('/kaggle/working/best_model.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T00:21:44.443045Z","iopub.execute_input":"2025-05-23T00:21:44.443943Z","iopub.status.idle":"2025-05-23T00:21:44.454518Z","shell.execute_reply.started":"2025-05-23T00:21:44.443891Z","shell.execute_reply":"2025-05-23T00:21:44.453744Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\ndst_dir = 'kaggle/working/'\n!zip -r folder.zip /kaggle/working/best_model.pth","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T00:21:47.170069Z","iopub.execute_input":"2025-05-23T00:21:47.170500Z","iopub.status.idle":"2025-05-23T00:22:07.349640Z","shell.execute_reply.started":"2025-05-23T00:21:47.170463Z","shell.execute_reply":"2025-05-23T00:22:07.348892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink(r'folder.zip')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T00:22:14.403403Z","iopub.execute_input":"2025-05-23T00:22:14.404231Z","iopub.status.idle":"2025-05-23T00:22:14.411765Z","shell.execute_reply.started":"2025-05-23T00:22:14.404193Z","shell.execute_reply":"2025-05-23T00:22:14.410767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DOLG_model_dict = \"/kaggle/input/resnet-dolg/kaggle/working/best_model.pth\"\ntest_image_link = \"/kaggle/input/test-landmark-images/eiffel.jpg\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T02:53:43.006800Z","iopub.execute_input":"2025-05-22T02:53:43.007570Z","iopub.status.idle":"2025-05-22T02:53:43.010873Z","shell.execute_reply.started":"2025-05-22T02:53:43.007534Z","shell.execute_reply":"2025-05-22T02:53:43.009992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms\nfrom PIL import Image\nimport torch\nfrom torch.utils.data import Dataset\nimport pandas as pd\nimport torch.nn as nn\nfrom torchvision import models\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Custom dataset wrapper for a single image\nclass SingleImageDataset(Dataset):\n    def __init__(self, image_path, label=0, transform=None):\n        self.image_path = image_path\n        self.label = label\n        self.transform = transform\n\n    def __len__(self):\n        return 1\n\n    def __getitem__(self, idx):\n        image = Image.open(self.image_path).convert(\"RGB\")\n        if self.transform:\n            image = self.transform(image)\n        return image, self.label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T02:59:27.688318Z","iopub.execute_input":"2025-05-22T02:59:27.689036Z","iopub.status.idle":"2025-05-22T02:59:27.695222Z","shell.execute_reply.started":"2025-05-22T02:59:27.689001Z","shell.execute_reply":"2025-05-22T02:59:27.694438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint = torch.load(DOLG_model_dict, map_location=device)\n\n# Load model weights\nmodel = DOLG_ArcFace(embedding_dim=512).to(device)\nmodel.load_state_dict(checkpoint['model_state_dict'])\n\n# Load ArcFace head (you must define this first with correct output classes)\narcface_head = ArcFace(in_features=512, out_features=102).to(device)\narcface_head.load_state_dict(checkpoint['arcface_state_dict'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T02:57:19.268312Z","iopub.execute_input":"2025-05-22T02:57:19.268827Z","iopub.status.idle":"2025-05-22T02:57:20.375644Z","shell.execute_reply.started":"2025-05-22T02:57:19.268795Z","shell.execute_reply":"2025-05-22T02:57:20.374891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize([0.5]*3, [0.5]*3)  # Match training normalization\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T02:59:18.484335Z","iopub.execute_input":"2025-05-22T02:59:18.485002Z","iopub.status.idle":"2025-05-22T02:59:18.491289Z","shell.execute_reply.started":"2025-05-22T02:59:18.484969Z","shell.execute_reply":"2025-05-22T02:59:18.490548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Path to test image\ntest_image_link = \"/kaggle/input/test-landmark-images/eiffel.jpg\"\n\n# Dummy label (for ArcFace input; only needed to compute logits)\ntest_dataset = SingleImageDataset(test_image_link, label=0, transform=transform)\n\n# Run visualization\nshow_predictions(model, arcface_head, test_dataset, num_images=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T02:56:49.380435Z","iopub.execute_input":"2025-05-22T02:56:49.381221Z","iopub.status.idle":"2025-05-22T02:56:49.770280Z","shell.execute_reply.started":"2025-05-22T02:56:49.381187Z","shell.execute_reply":"2025-05-22T02:56:49.769586Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Demo\nĐây là các phần cần có để có thể chạy demo cho ResNet+DOLG","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\nimport torch.nn.functional as F\n\n\nclass AttentionFusion(nn.Module):\n    def __init__(self, channels):\n        super().__init__()\n        self.attn = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(channels, channels, 1),\n            nn.ReLU(),\n            nn.Conv2d(channels, channels, 1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, local_feat, global_feat):\n        b, c, h, w = local_feat.shape\n        global_feat_expanded = global_feat.view(b, c, 1, 1).expand(-1, -1, h, w)\n        attn = self.attn(local_feat + global_feat_expanded)\n        fused = local_feat * attn + global_feat_expanded * (1 - attn)\n        return fused","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import math\n\nclass ArcFace(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.3):\n        super(ArcFace, self).__init__()\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n        self.s = s\n        self.m = m\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, labels):\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))  # [B, C]\n        sine = torch.sqrt((1.0 - cosine ** 2).clamp(min=0.0))\n        phi = cosine * self.cos_m - sine * self.sin_m  # cos(θ + m)\n\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, labels.view(-1, 1), 1.0)\n\n        logits = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        logits *= self.s\n        return logits","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DOLG_ArcFace(nn.Module):\n    def __init__(self, embedding_dim=512):\n        super().__init__()\n        resnet = models.resnet50(pretrained=True)\n\n        # Shared layers\n        self.backbone_common = nn.Sequential(\n            resnet.conv1, resnet.bn1, resnet.relu,\n            resnet.maxpool, resnet.layer1,\n            resnet.layer2, resnet.layer3\n        )\n\n        # ResNet layer4 expects input with 1024 channels (not from local_conv!)\n        self.backbone_global = resnet.layer4\n\n        self.global_pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.global_fc = nn.Linear(2048, embedding_dim)\n\n        self.local_conv = nn.Conv2d(1024, embedding_dim, kernel_size=1)\n\n        self.fusion = AttentionFusion(embedding_dim)\n\n        self.head = nn.Sequential(\n            nn.Conv2d(embedding_dim, embedding_dim, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten()\n        )\n\n    def forward(self, x):\n        shared_feat = self.backbone_common(x)  # Output: [B, 1024, H, W]\n\n        # Global branch\n        global_feat_map = self.backbone_global(shared_feat)  # Output: [B, 2048, H/2, W/2]\n        global_feat = self.global_pool(global_feat_map).view(x.size(0), -1)  # [B, 2048]\n        global_feat = self.global_fc(global_feat)  # [B, 512]\n\n        # Local branch\n        local_feat = self.local_conv(shared_feat)  # [B, 512, H, W]\n\n        # Fuse\n        fused_feat = self.fusion(local_feat, global_feat)  # [B, 512, H, W]\n        emb = self.head(fused_feat)  # [B, 512]\n        return emb","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torchvision import transforms\nfrom PIL import Image\nimport torch.nn.functional as F\n\ndef predict_image(image_path, model, arcface_head, transform, device, idx_to_landmark_id=None, id_to_name=None):\n    model.eval()\n    arcface_head.eval()\n\n    image = Image.open(image_path).convert(\"RGB\")\n    input_tensor = transform(image).unsqueeze(0).to(device)\n\n    label_tensor = torch.tensor([0], device=device)\n\n    # Forward pass\n    with torch.no_grad():\n        embedding = model(input_tensor)\n        logits = arcface_head(embedding, label_tensor)\n        probs = F.softmax(logits, dim=1)\n        conf, pred = torch.max(probs, dim=1)\n\n    pred_idx = pred.item()\n    confidence = conf.item()\n\n    if idx_to_landmark_id:\n        pred_landmark_id = idx_to_landmark_id.get(pred_idx, \"Unknown ID\")\n    else:\n        pred_landmark_id = pred_idx\n\n    if id_to_name:\n        pred_name = id_to_name.get(pred_landmark_id, \"Unknown Landmark\")\n    else:\n        pred_name = f\"Class {pred_idx}\"\n\n    print(f\"Predicted: {pred_name} (ID: {pred_landmark_id}) | Confidence: {confidence:.2f}\")\n    return pred_idx, confidence","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T01:27:30.143417Z","iopub.execute_input":"2025-05-23T01:27:30.144236Z","iopub.status.idle":"2025-05-23T01:27:34.750495Z","shell.execute_reply.started":"2025-05-23T01:27:30.144203Z","shell.execute_reply":"2025-05-23T01:27:34.749785Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---- Setup ----\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Load model\nNUM_CLASSES = 102\nmodel = DOLG_ArcFace(embedding_dim=512).to(device)\narcface_head = ArcFace(in_features=512, out_features=NUM_CLASSES).to(device)  # Replace NUM_CLASSES\n\n# Load checkpoint\ncheckpoint = torch.load(\"/kaggle/input/resnet-dolg/ResNetDOLG.pth\", map_location=device)\nmodel.load_state_dict(checkpoint['model_state_dict'])\narcface_head.load_state_dict(checkpoint['arcface_state_dict'])\n\n# Image transform (must match training)\ntransform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize([0.5]*3, [0.5]*3)\n])\n\n# Optional mappings\nidx_to_landmark_id = {v: k for k, v in landmark_id_to_idx.items()}\nlandmark_id_to_name = {\n    27: \"Isa Khan Niyazi's tomb\",\n    # ... fill from your metadata\n}\n\n# ---- Predict ----\nimage_path = \"/kaggle/input/test-landmark-images/golden_gate_foggy.jpg\"\npredict_image(\n    image_path,\n    model,\n    arcface_head,\n    transform,\n    device,\n    idx_to_landmark_id=idx_to_landmark_id,\n    id_to_name=landmark_id_to_name\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T01:36:08.021611Z","iopub.execute_input":"2025-05-23T01:36:08.022134Z","iopub.status.idle":"2025-05-23T01:36:09.037974Z","shell.execute_reply.started":"2025-05-23T01:36:08.022099Z","shell.execute_reply":"2025-05-23T01:36:09.037207Z"}},"outputs":[],"execution_count":null}]}