{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":126777,"databundleVersionId":15314950,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":293652533,"isSourceIdPinned":false}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import gc\nimport os\nimport time\nimport math\nimport warnings\nfrom pathlib import Path\nfrom typing import Callable, Dict, Literal, Tuple, Union\nimport numpy as np\nimport pandas as pd\nimport timm\nimport wandb\nfrom dotenv import load_dotenv\nfrom PIL import Image, ImageFilter, UnidentifiedImageError\nfrom sklearn.calibration import LabelEncoder\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nfrom kaggle_secrets import UserSecretsClient\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\nimport torchvision.transforms.v2 as v2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:04:51.821726Z","iopub.execute_input":"2026-02-24T10:04:51.822324Z","iopub.status.idle":"2026-02-24T10:04:51.827467Z","shell.execute_reply.started":"2026-02-24T10:04:51.822294Z","shell.execute_reply":"2026-02-24T10:04:51.826825Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VielleichtguarModel(nn.Module):\n    def __init__(self, backbone: nn.Module, layers: nn.Sequential, criterion: nn.Module):\n        super().__init__()\n        self.backbone = backbone\n        self.layers = layers\n        self.criterion = criterion\n\n    def forward(self, x: torch.Tensor, labels: torch.Tensor = None) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:\n        backbone_features = self.backbone(x)\n        embeddings = self.layers(backbone_features)\n        if labels is not None:\n            loss = self.criterion(embeddings, labels)\n            return embeddings, loss\n        return embeddings\n\n    def save_model(self, path: str, with_backbone: bool = True, with_criterion: bool = True):\n        state = {\n            \"layers\": self.layers.state_dict(),\n        }\n        if with_backbone:\n            state[\"backbone\"] = self.backbone.state_dict()\n        if with_criterion:\n            state[\"criterion\"] = self.criterion.state_dict()\n        torch.save(state, path)\n\n    def load_model(self, path: str):\n        state = torch.load(path, map_location=\"cpu\")\n        if \"backbone\" in state:\n            self.backbone.load_state_dict(state[\"backbone\"])\n        self.layers.load_state_dict(state[\"layers\"])\n        if \"criterion\" in state:\n            self.criterion.load_state_dict(state[\"criterion\"])\n            \nclass EmbeddingProjection(nn.Module):\n\n    def __init__(self, input_dim=1536, hidden_dim=512, output_dim=256, dropout=0.3, n_layers=2):\n        super().__init__()\n\n        modules = []\n        if n_layers >= 2:\n            for i in range(n_layers - 1):\n                modules.extend(\n                    [\n                        nn.Linear(input_dim if i == 0 else hidden_dim, hidden_dim),\n                        nn.BatchNorm1d(hidden_dim),\n                        nn.ReLU(inplace=True),\n                        nn.Dropout(dropout),\n                    ]\n                )\n        modules.extend(\n            [\n                nn.Linear(hidden_dim if n_layers >= 2 else input_dim, output_dim),\n                nn.BatchNorm1d(output_dim),\n            ]\n        )\n\n        self.embedding_projection = nn.Sequential(*modules)\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.kaiming_normal_(m.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.BatchNorm1d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        return self.embedding_projection(x)\n\n\nclass EmbeddingModel(nn.Module):\n\n    def __init__(self, model_name: str, freeze=True, cache_folder=\"./embeddings\", use_caching: bool = True, *args, **kwargs):\n        super().__init__()\n        self.model = timm.create_model(model_name, *args, **kwargs)\n        self.device = self._update_device()\n        self.cache_folder = Path(cache_folder) / f\"{model_name.replace('/', '_')}_{self.__class__.__name__.lower()}\"\n        self.cache_folder.mkdir(parents=True, exist_ok=True)\n        self.freeze = freeze\n        self.use_caching = use_caching\n        if freeze:\n            self.freeze_weights()\n        if not freeze and use_caching:\n            warnings.warn(\"Warning: Caching a non frozen model may lead to inconsistent results.\")\n\n    def get_transforms(self, is_training=False) -> transforms.Compose:\n        data_config = timm.data.resolve_model_data_config(self.model)\n        transforms = timm.data.create_transform(**data_config, is_training=is_training)\n        return transforms\n\n    def out_dim(self) -> int:\n        self._update_device()\n        img = Image.new(\"RGB\", (512, 512), (0, 0, 0))\n        transform = self.get_transforms()\n        img_tensor = transform(img).unsqueeze(0).to(self.device)\n        dummy_output = self.forward(img_tensor)\n        return dummy_output.shape[1]\n\n    def to(self, device):\n        rt = super().to(device)\n        self.device = self._update_device()\n        return rt\n\n    def _update_device(self):\n        return next(self.model.parameters()).device\n\n    def _tensor_hash(self, x: torch.Tensor) -> int:\n        return torch.hash_tensor(x).item()\n\n    def _get_cache_path(self, input_hash: int) -> Path:\n        cache_file = self.cache_folder / f\"{input_hash}.pt\"\n        return cache_file\n\n    def freeze_weights(self):\n        for param in self.model.parameters():\n            param.requires_grad = False\n\n    def unfreeze_weights(self):\n        for param in self.model.parameters():\n            param.requires_grad = True\n\n    def train(self, mode=True):\n        super().train(mode)\n        if self.freeze:\n            self.model.eval()\n        return self\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        if not self.freeze:\n            return self.model(x)\n        if not self.use_caching:\n            with torch.no_grad():\n                return self.model(x)\n\n        batch_size = x.shape[0]\n        outputs = [None] * batch_size\n        indices_to_compute = []\n\n        for i in range(batch_size):\n            sample = x[i]\n            h = torch.hash_tensor(sample).item()\n            cache_path = self.cache_folder / f\"{h}.pt\"\n\n            if cache_path.exists():\n                outputs[i] = torch.load(cache_path, map_location=x.device, weights_only=True)\n            else:\n                indices_to_compute.append(i)\n\n        if indices_to_compute:\n            to_compute_tensor = x[indices_to_compute]\n\n            with torch.no_grad():\n                computed_outputs = self.model(to_compute_tensor)\n\n            for j, original_idx in enumerate(indices_to_compute):\n                single_output = computed_outputs[j].unsqueeze(0)\n                outputs[original_idx] = single_output\n\n                h = torch.hash_tensor(x[original_idx]).item()\n                torch.save(single_output, self.cache_folder / f\"{h}.pt\")\n\n        return torch.cat(outputs, dim=0)\n\nclass DINOv3(EmbeddingModel):\n\n    def __init__(self, *args, **kwargs):\n        super().__init__(\"vit_large_patch16_dinov3.lvd1689m\", pretrained=True, *args, **kwargs)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        output = super().forward(x)  # shape (B, C)\n        return output\n\nclass ArcFaceLayer(nn.Module):\n\n    def __init__(self, embedding_dim, num_classes, margin=0.5, scale=64.0):\n        super().__init__()\n        self.embedding_dim = embedding_dim\n        self.num_classes = num_classes\n        self.margin = margin\n        self.scale = scale\n\n        self.weight = nn.Parameter(torch.FloatTensor(num_classes, embedding_dim))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.cos_m = math.cos(margin)\n        self.sin_m = math.sin(margin)\n        self.th = math.cos(math.pi - margin)\n        self.mm = math.sin(math.pi - margin) * margin\n\n    def forward(self, embeddings, labels):\n\n        embeddings = F.normalize(embeddings, p=2, dim=1)\n        weight_norm = F.normalize(self.weight, p=2, dim=1)\n\n        cosine = F.linear(embeddings, weight_norm)\n        cosine = cosine.clamp(-1.0, 1.0)\n\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n\n        phi = cosine * self.cos_m - sine * self.sin_m\n\n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n\n        one_hot = torch.zeros(cosine.size(), device=embeddings.device)\n        one_hot.scatter_(1, labels.view(-1, 1).long(), 1)\n\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output = output * self.scale\n        return output\n\n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2.0, reduction=\"mean\"):\n        super(FocalLoss, self).__init__()\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, logits, labels):\n        # Calculate log probabilities\n        log_p = F.log_softmax(logits, dim=1)\n        p = torch.exp(log_p)\n\n        # Gather the log probabilities of the target classes\n        log_p_t = log_p.gather(1, labels.view(-1, 1))\n        p_t = p.gather(1, labels.view(-1, 1))\n\n        # Focal Loss formula\n        loss = -((1 - p_t) ** self.gamma) * log_p_t\n\n        if self.reduction == \"mean\":\n            return loss.mean()\n        elif self.reduction == \"sum\":\n            return loss.sum()\n        else:\n            return loss\n\n\n\nclass FocalArcFaceCriterion(nn.Module):\n    def __init__(self, embedding_dim, num_classes, margin, scale, gamma=2.0):\n        super().__init__()\n        self.arcface = ArcFaceLayer(embedding_dim=embedding_dim, num_classes=num_classes, margin=margin, scale=scale)\n        self.focal_loss = FocalLoss(gamma=gamma)\n\n    def forward(self, embeddings, labels):\n        logits = self.arcface(embeddings, labels)\n        loss = self.focal_loss(logits, labels)\n        return loss\n\n\ndef get_data(data_path: str, validation_split_size=0.2, seed=42) -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame, pd.DataFrame, int, LabelEncoder]:\n    \"\"\"Get dataset located at `path`.\n\n    Args:\n        data_path (str): Path to the dataset.\n\n    Returns:\n        Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame, int, LabelEncoder]: Training data, validation data, test data,\n        number of classes, and label encoder.\n    \"\"\"\n    data_path = Path(data_path) if isinstance(data_path, str) else data_path\n    train_df = pd.read_csv(data_path / \"train.csv\")\n    train_df[\"filename\"] = train_df[\"filename\"].apply(lambda x: str(data_path / \"train/train\" / x))\n    num_classes = train_df[\"ground_truth\"].nunique()\n    label_encoder = LabelEncoder()\n    train_df[\"label_encoded\"] = label_encoder.fit_transform(train_df[\"ground_truth\"])\n    train_data, val_data = train_test_split(\n        train_df,\n        test_size=validation_split_size,\n        random_state=seed,\n        stratify=train_df[\"ground_truth\"],\n    )\n\n    test_data_all = pd.read_csv(data_path / \"test.csv\")\n    test_data = test_data_all.copy().drop(columns=[\"gallery_image\"])\n    test_data = test_data.drop_duplicates(subset=[\"query_image\"]).reset_index(drop=True)\n    test_data[\"filename\"] = test_data[\"query_image\"].apply(lambda x: str(data_path / \"test/test\" / x))\n    test_data[\"label_encoded\"] = -1  # Dummy labels for test set\n\n    return train_data, val_data, test_data, num_classes, label_encoder\n\n\ndef get_dataloaders(\n    data_dir: str,\n    validation_split_size: float,\n    seed: int,\n    batch_size: int,\n    image_size: int,\n    cache_dir: str,\n    train_process_fn: Callable = None,\n    val_process_fn: Callable = None,\n    mode: Literal[\"segmented\", \"background\"] = \"background\",\n) -> Tuple[DataLoader, DataLoader, DataLoader, int, LabelEncoder]:\n    \"\"\"\n    Get dataloaders for training, validation, test query, and test gallery datasets.\n\n    Args:\n        data_dir (str): Path to the dataset.\n        validation_split_size (float): Proportion of the dataset to include in the validation split.\n        seed (int): Random seed for reproducibility.\n        batch_size (int): Batch size for the dataloaders.\n        image_size (int): Size to which images will be resized.\n        cache_dir (str): Directory to cache processed images.\n        process_fn (Callable, optional): Function to process images. Defaults to None.\n        mode (Literal[\"segmented\", \"background\"], optional): Mode for image processing. Defaults to \"segmented\".\n    Returns:\n        Tuple[DataLoader, DataLoader, DataLoader, int, LabelEncoder]: Training dataloader, validation dataloader,\n        test dataloader, number of classes, and label encoder.\n    \"\"\"\n    train_data, val_data, test_gallery, num_classes, label_encoder = get_data(\n        data_dir,\n        validation_split_size=validation_split_size,\n        seed=seed,\n    )\n\n    train_dataset = ImageDataset(train_data, prewarm=True, image_size=image_size, cache_dir=cache_dir, key=\"train\", process_fn=train_process_fn, mode=mode)\n    train_dataloader = DataLoader(\n        train_dataset, batch_size=batch_size, shuffle=True, pin_memory=True, num_workers=4, prefetch_factor=2, persistent_workers=True\n    )\n\n    validation_dataset = ImageDataset(val_data, prewarm=True, image_size=image_size, cache_dir=cache_dir, key=\"val\", process_fn=val_process_fn, mode=mode)\n    validation_dataloader = DataLoader(\n        validation_dataset, batch_size=batch_size, shuffle=False, pin_memory=True, num_workers=4, prefetch_factor=2, persistent_workers=True\n    )\n\n    test_dataset = ImageDataset(test_gallery, prewarm=True, image_size=image_size, cache_dir=cache_dir, key=\"test\", process_fn=val_process_fn, mode=mode)\n    test_dataloader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, pin_memory=True, num_workers=4, prefetch_factor=2, persistent_workers=True)\n    return train_dataloader, validation_dataloader, test_dataloader, num_classes, label_encoder\n\n\ndef get_base_transform():\n    base_transform = transforms.Compose(\n        [\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ]\n    )\n    return base_transform\n\n\ndef get_resize_transform(image_size: int):\n    resize_transform = transforms.Compose(\n        [\n            transforms.Resize((image_size, image_size)),\n        ]\n    )\n    return resize_transform\n\n\ndef segment_background(img, background_color=(0, 0, 0)):\n    bg = Image.new(\"RGB\", img.size, background_color)\n    bg.paste(img, mask=img.split()[3])\n    return bg\n\n\ndef with_background(img):\n    return img.convert(\"RGB\")\n\n\ndef blurred_background(img, blur_radius=1000):\n    blurred_img = img.convert(\"RGB\").filter(ImageFilter.GaussianBlur(radius=blur_radius))\n    r, g, b, a = img.split()\n    final_img = Image.composite(img, blurred_img, a)\n    return final_img\n\n\ndef noisy_background(img):\n    r, g, b, a = img.split()\n    img_np = np.array(img.convert(\"RGB\"))\n    img_rgb_np = np.array(img.convert(\"RGB\"))\n    noise = np.random.randint(0, 255, (img.height, img.width, 3), dtype=np.int16)\n    noisy_img_np = np.clip(img_rgb_np + noise, 0, 255).astype(np.uint8)\n    noisy_img_np = img_np + np.clip(noisy_img_np, 0, 255).astype(np.uint8)\n    noisy_img = Image.fromarray(noisy_img_np, mode=\"RGB\")\n    final_img = Image.composite(img.convert(\"RGB\"), noisy_img, a)\n    return final_img\n\n\ndef random_background(img):\n    r, g, b, a = img.split()\n    noise = np.random.randint(0, 256, (img.height, img.width, 3), dtype=np.uint8)\n    noisy_img = Image.fromarray(noise, mode=\"RGB\")\n    final_img = Image.composite(img.convert(\"RGB\"), noisy_img, a)\n    return final_img\n\n\nclass ImageDataset(Dataset):\n    def __init__(\n        self,\n        data: pd.DataFrame,\n        image_column: str = \"filename\",\n        label_column: str = \"label_encoded\",\n        image_size: int = 384,\n        process_fn: Callable = None,\n        mode: Literal[\"segmented\", \"background\", \"blurred\", \"noisy\", \"random\"] = \"segmented\",\n        cache_dir: str = \"./cache\",\n        key: str = \"train\",\n        prewarm: bool = True,\n        verbose: bool = True,\n    ):\n        self.data = data\n        self.image_column = image_column\n        self.image_paths = data[image_column].tolist()\n        self.label_column = label_column\n        self.labels = data[label_column].tolist()\n        self.process_fn = process_fn if process_fn is not None else get_base_transform()\n        self.resize_transform = get_resize_transform(image_size)\n        self.mode = mode\n        self.verbose = verbose\n\n        # Initialize Cache\n        self.cache_dir = Path(cache_dir) / f\"{key}_{mode}_{image_size}\"\n        self.cache_dir.mkdir(parents=True, exist_ok=True)\n\n        self.to_tensor = transforms.ToTensor()\n\n        if prewarm:\n            self._prewarm()\n\n    def _prewarm(self):\n        if not all(self._get_cache_path(p).exists() for p in self.image_paths):\n            prewarm_loader = DataLoader(\n                self,\n                batch_size=64,\n                shuffle=False,\n                num_workers=4,\n            )\n            for _ in tqdm(prewarm_loader, desc=f\"Prewarming cache for {self.cache_dir.name} dataset\", total=len(prewarm_loader), disable=not self.verbose):\n                pass\n\n    def _get_cache_path(self, image_path: str) -> Path:\n        source_path = Path(image_path)\n        cache_file = self.cache_dir / f\"{source_path.stem}.pt\"\n        return cache_file\n\n    def load_image(self, image_path: str) -> torch.Tensor:\n        cache_file = self._get_cache_path(image_path)\n\n        if cache_file.exists():\n            try:\n                return torch.load(cache_file, weights_only=True)\n            except (UnidentifiedImageError, OSError):\n                # If the cache is corrupt, remove it and fall through to recreation\n                cache_file.unlink(missing_ok=True)\n\n        # Re-create the image from source\n\n        if self.mode == \"segmented\":\n            img = segment_background(Image.open(image_path).convert(\"RGBA\"))\n            processed_img = self.resize_transform(img)\n        elif self.mode == \"blurred\":\n            img = blurred_background(Image.open(image_path).convert(\"RGBA\"))\n            processed_img = self.resize_transform(img)\n        elif self.mode == \"noisy\":\n            img = noisy_background(Image.open(image_path).convert(\"RGBA\"))\n            processed_img = self.resize_transform(img)\n        elif self.mode == \"random\":\n            img = random_background(Image.open(image_path).convert(\"RGBA\"))\n            processed_img = self.resize_transform(img)\n        elif self.mode == \"background\":\n            img = with_background(Image.open(image_path).convert(\"RGBA\"))\n            img = self.resize_transform(img)\n            processed_img = img\n\n        tensor_img = self.to_tensor(processed_img)\n        torch.save(tensor_img, cache_file, _use_new_zipfile_serialization=False)\n        return tensor_img\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n        label = self.labels[idx]\n\n        image = self.load_image(image_path)\n\n        if self.process_fn is not None:\n            image = self.process_fn(image)\n\n        return image, label\n\ndef build_submission(submission_path, model, test_dataloader, device, query_expansion_enabled=False, k_reciprocal_reranking_enabled=False, tta_enabled=False):\n\n    model.eval()\n    all_embeddings = []\n    submission_path = Path(submission_path)\n    if not submission_path.parent.exists():\n        submission_path.parent.mkdir(parents=True, exist_ok=True)\n    if submission_path.exists():\n        submission_path.unlink()\n\n    with torch.no_grad():\n        pbar = tqdm(test_dataloader, desc=\"Generating Embeddings\", leave=False)\n        for batch in pbar:\n            images, _ = batch\n            images = images.to(device)\n\n            if tta_enabled:\n                # Generate embeddings for both original and horizontally flipped images\n                embeddings_original = model(images)\n                embeddings_original = F.normalize(embeddings_original, p=2, dim=1)\n\n                # Flip images horizontally\n                images_flipped = torch.flip(images, dims=[3])  # Flip along width dimension\n                embeddings_flipped = model(images_flipped)\n                embeddings_flipped = F.normalize(embeddings_flipped, p=2, dim=1)\n\n                # Average the embeddings\n                embeddings = (embeddings_original + embeddings_flipped) / 2.0\n                embeddings = F.normalize(embeddings, p=2, dim=1)\n            else:\n                embeddings = model(images)\n                embeddings = F.normalize(embeddings, p=2, dim=1)\n\n            all_embeddings.append(embeddings)\n\n    all_embeddings = torch.cat(all_embeddings, dim=0)\n    all_embeddings = F.normalize(all_embeddings, p=2, dim=1)\n\n    if query_expansion_enabled:\n        all_embeddings = query_expansion(all_embeddings.cpu().numpy(), k=3)\n\n    sim_matrix = torch.mm(all_embeddings, all_embeddings.t()).cpu().numpy()\n\n    if k_reciprocal_reranking_enabled:\n        sim_matrix = k_reciprocal_reranking(all_embeddings.cpu().numpy(), k1=20, k2=6, lambda_value=0.3)\n\n    sim_matrix = np.clip(sim_matrix, a_min=0.0, a_max=1.0)\n\n    submission_rows = []\n    num_samples = all_embeddings.shape[0]\n    c = 0\n    for i in range(num_samples):\n        similarities = sim_matrix[i]\n        for j in range(num_samples):\n            if i == j:\n                continue\n            row_id = c\n            c += 1\n            similarity = similarities[j]\n            submission_rows.append({\"row_id\": row_id, \"similarity\": similarity})\n    submission_df = pd.DataFrame(submission_rows)\n    submission_df.to_csv(submission_path, index=False)\n\n\ndef query_expansion(emb, k=5):\n    \"\"\"Perform query expansion on the given embeddings.\n\n    Args:\n        emb (np.ndarray): The input embeddings of shape (N, D).\n        k (int): The number of nearest neighbors to consider for expansion.\n\n    Returns:\n        torch.Tensor: The expanded embeddings of shape (N, D).\n    \"\"\"\n    emb = torch.from_numpy(emb)\n    sim_matrix = torch.mm(emb, emb.T)\n    sim_matrix.fill_diagonal_(-1)\n\n    expanded_emb = []\n    for i in range(emb.size(0)):\n        similarities = sim_matrix[i]\n        topk_indices = torch.topk(similarities, k=k, largest=True).indices\n        neighbor_embs = emb[topk_indices]\n        expanded_vector = torch.mean(torch.cat([emb[i].unsqueeze(0), neighbor_embs], dim=0), dim=0)\n        expanded_vector = F.normalize(expanded_vector, p=2, dim=0)\n        expanded_emb.append(expanded_vector)\n\n    expanded_emb = torch.stack(expanded_emb, dim=0)\n    return expanded_emb\n\n\ndef k_reciprocal_reranking(embeddings, k1=20, k2=6, lambda_value=0.3):\n    \"\"\"Perform k-reciprocal re-ranking on the given embeddings.\n\n    Args:\n        embeddings (np.ndarray): The input embeddings of shape (N, D).\n        k1 (int): The number of nearest neighbors for initial ranking.\n        k2 (int): The number of nearest neighbors for expansion.\n        lambda_value (float): The weight for combining original and Jaccard distances.\n\n    Returns:\n        np.ndarray: The re-ranked similarity matrix of shape (N, N).\n    \"\"\"\n\n    # Compute original distance\n    original_dist = np.matmul(embeddings, embeddings.T)\n    original_dist = 1 - original_dist\n    original_dist = np.transpose(original_dist / np.max(original_dist, axis=0))\n\n    # Initialize variables\n    num_samples = embeddings.shape[0]\n    V = np.zeros_like(original_dist).astype(np.float32)\n\n    # k-reciprocal neighbors\n    for i in range(num_samples):\n        forward_k_neigh_index = np.argsort(original_dist[i])[: k1 + 1]\n        backward_k_neigh_index = np.array([np.argsort(original_dist[j])[: k1 + 1] for j in forward_k_neigh_index])\n        fi = np.where(backward_k_neigh_index == i)[0]\n        k_reciprocal_index = forward_k_neigh_index[fi]\n\n        # Expansion\n        for j in k_reciprocal_index:\n            candidate_forward_k_neigh_index = np.argsort(original_dist[j])[: int(np.around(k1 / 2)) + 1]\n            candidate_backward_k_neigh_index = np.array([np.argsort(original_dist[m])[: int(np.around(k1 / 2)) + 1] for m in candidate_forward_k_neigh_index])\n            fi_candidate = np.where(candidate_backward_k_neigh_index == j)[0]\n            candidate_k_reciprocal_index = candidate_forward_k_neigh_index[fi_candidate]\n            if len(np.intersect1d(candidate_k_reciprocal_index, k_reciprocal_index)) > 2 / 3 * len(candidate_k_reciprocal_index):\n                k_reciprocal_index = np.append(k_reciprocal_index, candidate_k_reciprocal_index)\n\n        k_reciprocal_index = np.unique(k_reciprocal_index)\n        weight = np.exp(-original_dist[i][k_reciprocal_index])\n        V[i][k_reciprocal_index] = weight / np.sum(weight)\n    # Jaccard distance\n    original_dist = original_dist[:num_samples, :]\n    if k2 != 1:\n        V_qe = np.zeros_like(V, dtype=np.float32)\n        for i in range(num_samples):\n            V_qe[i, :] = np.mean(V[np.argsort(original_dist[i])[:k2], :], axis=0)\n        V = V_qe\n        del V_qe\n    invIndex = []\n    for i in range(num_samples):\n        invIndex.append(np.where(V[:, i] != 0)[0])\n    jaccard_dist = np.zeros_like(original_dist, dtype=np.float32)\n    for i in range(num_samples):\n        temp_min = np.zeros((1, num_samples), dtype=np.float32)\n        indNonZero = np.where(V[i, :] != 0)[0]\n        for j in indNonZero:\n            temp_min[0, invIndex[j]] += np.minimum(V[i, j], V[invIndex[j], j])\n        jaccard_dist[i] = 1 - temp_min / (2 - temp_min)\n    # Final distance\n    final_dist = (1 - lambda_value) * jaccard_dist + lambda_value * original_dist\n    final_dist = final_dist[:num_samples,]\n    final_sim = 1 - final_dist\n    return final_sim\n\n\ndef train_epoch(model, data_loader, optimizer, device, lr_scheduler=None) -> float:\n    \"\"\"\n    Train the model for one epoch.\n\n    Args:\n        model: The model to train.\n        data_loader: DataLoader providing the training data.\n        optimizer: The optimizer to use for training.\n        device: The device to run the training on.\n        lr_scheduler: Optional learning rate scheduler to update after each batch.\n    Returns:\n        avg_loss (float): The average loss over the epoch.\n    \"\"\"\n    model.train()\n    total_loss = 0\n    pbar = tqdm(data_loader, desc=\"Training\", leave=False)\n    for batch in pbar:\n        images, labels = batch\n        images, labels = images.to(device), labels.to(device)\n\n        _, loss = model(images, labels)\n        optimizer.zero_grad()\n        if loss.dim() > 0:\n            loss = loss.mean()\n        loss.backward()\n        optimizer.step()\n        if lr_scheduler is not None:\n            lr_scheduler.step()\n\n        total_loss += loss.item()\n\n    avg_loss = total_loss / len(data_loader)\n    return avg_loss\n\n\ndef validate_epoch(model, data_loader, device) -> Tuple[float, float]:\n    \"\"\"\n    Validate the model.\n\n    Args:\n        model: The model to validate.\n        data_loader: DataLoader providing the validation data.\n        device: The device to run the validation on.\n    Returns:\n        avg_loss (float): The average loss over the validation set.\n        val_map (float): The balanced mean Average Precision (mAP) over the validation set.\n    \"\"\"\n    model.eval()\n    total_loss = 0\n    all_embeddings = []\n    all_labels = []\n\n    with torch.no_grad():\n        pbar = tqdm(data_loader, desc=\"Validation\", leave=False)\n        for batch in pbar:\n            images, labels = batch\n            images, labels = images.to(device), labels.to(device)\n\n            embeddings, loss = model(images, labels)\n            if loss.dim() > 0:\n                loss = loss.mean()\n            all_embeddings.append(embeddings)\n            all_labels.append(labels)\n\n            total_loss += loss.item()\n\n            pbar.set_postfix({\"loss\": f\"{loss.item():.4f}\"})\n\n    avg_loss = total_loss / len(data_loader)\n    avl_embeddings = torch.cat(all_embeddings, dim=0)\n    avl_labels = torch.cat(all_labels, dim=0)\n    val_map = compute_validation_map(avl_embeddings, avl_labels.cpu().numpy())\n    return avg_loss, val_map\n\n\ndef compute_validation_map(embeddings: torch.Tensor, val_labels: np.ndarray) -> float:\n    \"\"\"\n    Optimized version of v1 logic.\n    Matches the 'balanced' identity averaging of your winning version.\n\n    Args:\n        embeddings (torch.Tensor): Embeddings of the validation samples.\n        val_labels (np.ndarray): Ground truth labels for the validation samples.\n    Returns:\n        balanced_map (float): The balanced mean Average Precision (mAP).\n    \"\"\"\n    embeddings_norm = F.normalize(embeddings, p=2, dim=1)\n    sim_matrix = torch.mm(embeddings_norm, embeddings_norm.t())\n\n    sim_matrix = sim_matrix.cpu().numpy()\n    np.fill_diagonal(sim_matrix, -1)\n\n    query_aps = {}\n\n    for query_idx in range(len(val_labels)):\n        query_label = val_labels[query_idx]\n        similarities = sim_matrix[query_idx]\n\n        is_match = (val_labels == query_label).astype(int)\n        is_match[query_idx] = 0\n\n        n_positives = is_match.sum()\n        if n_positives == 0:\n            continue\n\n        sorted_indices = np.argsort(-similarities)\n        sorted_matches = is_match[sorted_indices]\n\n        cumsum = np.cumsum(sorted_matches)\n        precision_at_k = cumsum / np.arange(1, len(sorted_matches) + 1)\n        ap = np.sum(precision_at_k * sorted_matches) / n_positives\n\n        query_aps[query_idx] = (query_label, ap)\n\n    identity_aps = {}\n    for query_idx, (label, ap) in query_aps.items():\n        if label not in identity_aps:\n            identity_aps[label] = []\n        identity_aps[label].append(ap)\n\n    identity_mean_aps = [np.mean(aps) for aps in identity_aps.values()]\n    balanced_map = np.mean(identity_mean_aps)\n\n    return balanced_map\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:04:51.950516Z","iopub.execute_input":"2026-02-24T10:04:51.950982Z","iopub.status.idle":"2026-02-24T10:04:52.021096Z","shell.execute_reply.started":"2026-02-24T10:04:51.950959Z","shell.execute_reply":"2026-02-24T10:04:52.020505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PROJECT = \"jaguar-reid-josefandvincent\"\nGROUP = \"final\"\nRUN_NAME = f\"{GROUP}-final\"\n\nBASE_CONFIG = {\n    \"random_seed\": 456,\n    \"data_dir\": Path(\"/kaggle/input/competitions/jaguar-re-id\"),\n    \"checkpoint_dir\": Path(\"checkpoints\"),\n    \"cache_dir\": Path(\"./cache\"),\n    \"embeddings_dir\": Path(\"./embeddings\"),\n    \"validation_split_size\": 0.2,\n}\n\nEXPERIMENT_CONFIG = {\n    \"epochs\": 100,\n    \"batch_size\": 32,\n    \"image_size\": 256,\n    \"hidden_dim\": 1024,\n    \"n_layers\": 3,\n    \"output_dim\": 256,\n    \"dropout\": 0.2,\n    \"weight_decay\": 0.00014063618588945403,\n    \"learning_rate\": 0.00015993740640441972,\n    \"arcface_margin\": 0.3600732565980775,\n    \"arcface_scale\": 17.570478141393707,\n    \"patience\": 10,\n    \"backbone_lr_multiplier\": 0.054143381201692535,\n    \"background_intervention\": \"segmented\",\n    \"focal_arcface_gamma\": 2.111056877247217,\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:04:52.022481Z","iopub.execute_input":"2026-02-24T10:04:52.023064Z","iopub.status.idle":"2026-02-24T10:04:52.034479Z","shell.execute_reply.started":"2026-02-24T10:04:52.023041Z","shell.execute_reply":"2026-02-24T10:04:52.033698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ntorch.manual_seed(BASE_CONFIG[\"random_seed\"])\nnp.random.seed(BASE_CONFIG[\"random_seed\"])\n\nuser_secrets = UserSecretsClient()\nos.environ[\"HF_TOKEN\"]= user_secrets.get_secret(\"hf_token\")\nos.environ[\"WANDB_API_KEY\"] = user_secrets.get_secret(\"wandb_token\")\n\nwandb.login(key=os.getenv(\"WANDB_API_KEY\"))\nwandb.init(project=PROJECT, config=EXPERIMENT_CONFIG, group=GROUP, name=RUN_NAME)\n\nBASE_CONFIG[\"checkpoint_dir\"].mkdir(exist_ok=True)\ncheckpoint_path = BASE_CONFIG[\"checkpoint_dir\"] / f\"{RUN_NAME}_best.pth\"\nsubmission_path = BASE_CONFIG[\"checkpoint_dir\"] / f\"{RUN_NAME}_submission.csv\"\n\nnum_gpus = torch.cuda.device_count()\ndevice_ids = list(range(num_gpus))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:04:52.035551Z","iopub.execute_input":"2026-02-24T10:04:52.036227Z","iopub.status.idle":"2026-02-24T10:04:58.721931Z","shell.execute_reply.started":"2026-02-24T10:04:52.036202Z","shell.execute_reply":"2026-02-24T10:04:58.721341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"backbone = DINOv3(freeze=False, cache_folder=BASE_CONFIG[\"embeddings_dir\"], use_caching=False)\nbase_transforms = backbone.get_transforms()\naugmentation_transforms = v2.Compose(\n    [\n        v2.RandomHorizontalFlip(p=0.5),\n        v2.RandomAffine(degrees=10, translate=(0.05, 0.05), scale=(0.95, 1.05)),\n        v2.RandomErasing(p=0.5, scale=(0.02, 0.25), value=0),\n        *base_transforms.transforms,\n    ]\n)\n\ntrain_dataloader, validation_dataloader, test_dataloader, num_classes, label_encoder = get_dataloaders(\n    data_dir=BASE_CONFIG[\"data_dir\"],\n    validation_split_size=BASE_CONFIG[\"validation_split_size\"],\n    seed=BASE_CONFIG[\"random_seed\"],\n    batch_size=EXPERIMENT_CONFIG[\"batch_size\"],\n    image_size=EXPERIMENT_CONFIG[\"image_size\"],\n    cache_dir=BASE_CONFIG[\"cache_dir\"],\n    train_process_fn=augmentation_transforms,\n    val_process_fn=base_transforms,\n    mode=EXPERIMENT_CONFIG[\"background_intervention\"],\n)\n\nmodel = VielleichtguarModel(\n    backbone=backbone,\n    layers=nn.Sequential(\n        EmbeddingProjection(\n            input_dim=backbone.out_dim(),\n            hidden_dim=EXPERIMENT_CONFIG[\"hidden_dim\"],\n            output_dim=EXPERIMENT_CONFIG[\"output_dim\"],\n            dropout=EXPERIMENT_CONFIG[\"dropout\"],\n            n_layers=EXPERIMENT_CONFIG[\"n_layers\"],\n        ),\n    ),\n    criterion=FocalArcFaceCriterion(\n        embedding_dim=EXPERIMENT_CONFIG[\"output_dim\"],\n        num_classes=num_classes,\n        margin=EXPERIMENT_CONFIG[\"arcface_margin\"],\n        scale=EXPERIMENT_CONFIG[\"arcface_scale\"],\n        gamma=EXPERIMENT_CONFIG[\"focal_arcface_gamma\"],\n    ),\n)\n\nmodel = nn.DataParallel(model, device_ids).to(device)\n\nwandb.log(\n    {\"trainable_parameters\": sum(p.numel() for p in model.parameters() if p.requires_grad), \"total_parameters\": sum(p.numel() for p in model.parameters())}\n)\n# %%\n\nparams = [\n    {\"params\": model.module.backbone.parameters(), \"lr\": EXPERIMENT_CONFIG[\"learning_rate\"] * EXPERIMENT_CONFIG[\"backbone_lr_multiplier\"]},\n    {\"params\": model.module.layers.parameters(), \"lr\": EXPERIMENT_CONFIG[\"learning_rate\"]},\n    {\"params\": model.module.criterion.parameters(), \"lr\": EXPERIMENT_CONFIG[\"learning_rate\"]},\n]\noptimizer = AdamW(params, weight_decay=EXPERIMENT_CONFIG[\"weight_decay\"])\nlr_scheduler = ReduceLROnPlateau(optimizer, mode=\"min\", factor=0.5, patience=3)\n\nbest_epoch, best_map, patience_counter, total_duration, eta = 0, 0.0, 0, 0.0, 0.0\nfor epoch in range(EXPERIMENT_CONFIG[\"epochs\"]):\n\n    start = time.time()\n    train_loss = train_epoch(model, train_dataloader, optimizer, device)\n    validation_loss, validation_map = validate_epoch(model, validation_dataloader, device)\n    end = time.time()\n    epoch_time = end - start\n    total_duration += epoch_time\n    eta = (epoch_time if epoch == 0 else 0.9 * (eta / (EXPERIMENT_CONFIG[\"epochs\"] - epoch)) + 0.1 * epoch_time) * (EXPERIMENT_CONFIG[\"epochs\"] - epoch - 1)\n    lr_scheduler.step(validation_loss)\n\n    print(\n        f\"epoch: {epoch+1:>2}/{EXPERIMENT_CONFIG['epochs']} | \",\n        f\"train/loss: {train_loss:>8.4f} | \",\n        f\"val/loss: {validation_loss:>8.4f} | \",\n        f\"val/mAP: {validation_map:>7.4f} | \",\n        f\"lr: {optimizer.param_groups[0]['lr']:>7.1e} | \",\n        f\"eta: {max(0,eta)/60:.1f} min | \" if max(0, eta) > 60 else f\"eta: {max(0,eta):.1f} sec | \",\n        f\"patience: {patience_counter}/{EXPERIMENT_CONFIG['patience']} | \",\n        sep=\"\",\n    )\n\n    wandb.log(\n        {\n            \"epoch\": epoch + 1,\n            \"train/loss\": train_loss,\n            \"val/loss\": validation_loss,\n            \"val/mAP\": validation_map,\n            \"lr\": optimizer.param_groups[0][\"lr\"],\n        }\n    )\n\n    if validation_map > best_map:\n        best_map = validation_map\n        best_epoch = epoch + 1\n        patience_counter = 0\n        model.module.save_model(\n            checkpoint_path,\n            with_backbone=True,\n            with_criterion=False,\n        )\n    else:\n        patience_counter += 1\n\n    if patience_counter >= EXPERIMENT_CONFIG[\"patience\"]:\n        break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:04:58.723186Z","iopub.execute_input":"2026-02-24T10:04:58.723421Z","iopub.status.idle":"2026-02-24T10:10:30.711824Z","shell.execute_reply.started":"2026-02-24T10:04:58.723400Z","shell.execute_reply":"2026-02-24T10:10:30.710488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.module.load_model(checkpoint_path)\nbuild_submission(submission_path, model, test_dataloader, device, query_expansion_enabled=True, tta_enabled=True, k_reciprocal_reranking_enabled=True)\n\nwandb.run.summary[\"best_val_mAP\"] = best_map\nwandb.run.summary[\"best_epoch\"] = best_epoch\nwandb.run.summary[\"total_epochs\"] = epoch + 1\nwandb.run.summary[\"total_training_time\"] = total_duration\nmodel_artifact = wandb.Artifact(name=f\"{RUN_NAME}_model\", type=\"model\")\nmodel_artifact.add_file(str(checkpoint_path))\nwandb.log_artifact(model_artifact)\nsubmission_artifact = wandb.Artifact(name=f\"{RUN_NAME}_submission\", type=\"submission\")\nsubmission_artifact.add_file(str(submission_path))\nwandb.log_artifact(submission_artifact)\nwandb.finish()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:10:30.712740Z","iopub.status.idle":"2026-02-24T10:10:30.713099Z","shell.execute_reply.started":"2026-02-24T10:10:30.712922Z","shell.execute_reply":"2026-02-24T10:10:30.712943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"global model, optimizer\nif 'model' in globals(): del model\nif 'optimizer' in globals(): del optimizer\ngc.collect()\ntorch.cuda.empty_cache()\ntorch.cuda.reset_peak_memory_stats()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-24T10:10:30.713826Z","iopub.status.idle":"2026-02-24T10:10:30.716937Z","shell.execute_reply.started":"2026-02-24T10:10:30.716737Z","shell.execute_reply":"2026-02-24T10:10:30.716760Z"}},"outputs":[],"execution_count":null}]}