{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":124685,"databundleVersionId":14664296,"sourceType":"competition"},{"sourceId":228781,"sourceType":"modelInstanceVersion","modelInstanceId":195042,"modelId":216938}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Import libraries","metadata":{"_uuid":"783db5f1-a506-4ebf-8188-b3c62483aa96","_cell_guid":"a200880d-c615-4eb3-80fa-d94f91a831ed","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport timm\nimport torch\nimport torch.nn.functional as F\nfrom PIL import Image\nfrom torch.utils.data import DataLoader, Dataset\nimport time\nimport os\nimport csv\nimport logging\nfrom dataclasses import dataclass, field\nfrom typing import List, Dict, Optional, Tuple\nfrom collections import defaultdict\nfrom pathlib import Path\n\nimport torchvision.transforms as T\nfrom torch.amp import autocast\nfrom kornia.contrib import extract_tensor_patches, compute_padding\n\n# Setup logging\nlogging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')\nlogger = logging.getLogger(__name__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T12:26:29.118592Z","iopub.execute_input":"2026-02-15T12:26:29.119399Z","iopub.status.idle":"2026-02-15T12:26:29.124402Z","shell.execute_reply.started":"2026-02-15T12:26:29.119369Z","shell.execute_reply":"2026-02-15T12:26:29.123768Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"# %% [code]\nclass InferenceConfig:\n    \"\"\"Configuration for the inference pipeline.\"\"\"\n    def __init__(\n        self,\n        # Paths\n        model_checkpoint: str = '/kaggle/input/dinov2_patch14_reg4_onlyclassifier_then_all/pytorch/default/3/model_best.pth.tar',\n        species_ids_path: str = '/kaggle/input/plantclef-2026/species_ids.csv',\n        test_images_dir: str = '/kaggle/input/plantclef-2026/PlantCLEF2025_test_images/PlantCLEF2025_test_images/',\n        output_path: str = 'submission.csv',\n        # Model\n        model_name: str = 'vit_base_patch14_reg4_dinov2.lvd142m',\n        # Tiling\n        patch_sizes=None,\n        stride_ratio: float = 0.5,\n        use_pad: bool = True,\n        # Inference\n        batch_size: int = 64,\n        num_workers: int = 4,\n        pin_memory: bool = True,\n        # Prediction filtering\n        min_score: float = 0.05,\n        top_k_tile: int = 3,\n        aggregation: str = 'weighted_vote',\n        final_top_k: int = 50,\n        final_min_score: float = 0.05,\n        # Test-time augmentation\n        use_tta: bool = True,\n        tta_transforms=None,\n        # Logging\n        log_frequency: int = 10,\n    ):\n        self.model_checkpoint = model_checkpoint\n        self.species_ids_path = species_ids_path\n        self.test_images_dir = test_images_dir\n        self.output_path = output_path\n        self.model_name = model_name\n        self.patch_sizes = patch_sizes if patch_sizes is not None else [518]\n        self.stride_ratio = stride_ratio\n        self.use_pad = use_pad\n        self.batch_size = batch_size\n        self.num_workers = num_workers\n        self.pin_memory = pin_memory\n        self.min_score = min_score\n        self.top_k_tile = top_k_tile\n        self.aggregation = aggregation\n        self.final_top_k = final_top_k\n        self.final_min_score = final_min_score\n        self.use_tta = use_tta\n        self.tta_transforms = tta_transforms if tta_transforms is not None else ['none', 'hflip']\n        self.log_frequency = log_frequency\n\n    def __repr__(self):\n        attrs = vars(self)\n        lines = [f\"  {k}={v!r}\" for k, v in attrs.items()]\n        return \"InferenceConfig(\\n\" + \",\\n\".join(lines) + \"\\n)\"\n\n\nclass AverageMeter:\n    \"\"\"Computes and stores the average and current value.\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\nclass ETATimer:\n    \"\"\"Tracks elapsed time and estimates time remaining.\"\"\"\n    def __init__(self, total_steps: int):\n        self.total_steps = total_steps\n        self.start_time = time.time()\n        self.step = 0\n\n    def update(self, step: int):\n        self.step = step\n\n    def eta(self) -> str:\n        elapsed = time.time() - self.start_time\n        if self.step == 0:\n            return \"N/A\"\n        rate = elapsed / self.step\n        remaining = rate * (self.total_steps - self.step)\n        hours, remainder = divmod(int(remaining), 3600)\n        minutes, seconds = divmod(remainder, 60)\n        return f\"{hours:02d}:{minutes:02d}:{seconds:02d}\"\n\n    def elapsed(self) -> str:\n        elapsed = time.time() - self.start_time\n        hours, remainder = divmod(int(elapsed), 3600)\n        minutes, seconds = divmod(remainder, 60)\n        return f\"{hours:02d}:{minutes:02d}:{seconds:02d}\"","metadata":{"_uuid":"81140d29-2323-4545-a43f-3a9ac1ab1a8f","_cell_guid":"05961dbf-3f4d-48fe-a3ef-6a368b11b9c5","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2026-02-15T12:26:29.125708Z","iopub.execute_input":"2026-02-15T12:26:29.126086Z","iopub.status.idle":"2026-02-15T12:26:29.138530Z","shell.execute_reply.started":"2026-02-15T12:26:29.126051Z","shell.execute_reply":"2026-02-15T12:26:29.137825Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patch Creation","metadata":{}},{"cell_type":"code","source":"class PatchDataset(Dataset):\n    \"\"\"Dataset for individual patches from a single image.\"\"\"\n    def __init__(self, patches: torch.Tensor, transform=None):\n        # patches shape: (num_patches, C, H, W)\n        self.patches = patches\n        self.transform = transform\n\n    def __len__(self):\n        return self.patches.size(0)\n\n    def __getitem__(self, idx):\n        patch = self.patches[idx]\n        if self.transform:\n            patch = self.transform(patch)\n        return patch\n\n\nclass TestImageDataset(Dataset):\n    \"\"\"Dataset that loads test images and extracts tiles.\"\"\"\n    SUPPORTED_EXTENSIONS = {'.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.webp'}\n\n    def __init__(self, image_folder: str, patch_size: int = 518,\n                 stride: int = 259, use_pad: bool = True):\n        self.image_folder = Path(image_folder)\n        self.image_paths = sorted([\n            p for p in self.image_folder.iterdir()\n            if p.suffix.lower() in self.SUPPORTED_EXTENSIONS\n        ])\n        self.use_pad = use_pad\n        self.patch_size = patch_size\n        self.stride = stride\n        self.to_tensor = T.ToTensor()\n\n        if len(self.image_paths) == 0:\n            raise ValueError(f\"No images found in {image_folder}\")\n\n        logger.info(f\"Found {len(self.image_paths)} images in {image_folder}\")\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx) -> Tuple[torch.Tensor, str]:\n        image_path = self.image_paths[idx]\n        try:\n            image = Image.open(image_path).convert('RGB')\n        except Exception as e:\n            logger.error(f\"Failed to load image {image_path}: {e}\")\n            # Return a dummy single black patch\n            dummy = torch.zeros(1, 3, self.patch_size, self.patch_size)\n            return dummy, str(image_path)\n\n        image_tensor = self.to_tensor(image).unsqueeze(0)\n        h, w = image_tensor.shape[-2:]\n\n        if self.use_pad:\n            pad = compute_padding(\n                original_size=(h, w),\n                window_size=self.patch_size,\n                stride=self.stride\n            )\n            patches = extract_tensor_patches(\n                image_tensor, self.patch_size, self.stride, padding=pad\n            )\n        else:\n            patches = extract_tensor_patches(\n                image_tensor, self.patch_size, self.stride\n            )\n\n        # patches shape: (1, num_patches, C, H, W) -> (num_patches, C, H, W)\n        patches = patches.squeeze(0)\n        return patches, str(image_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T12:26:29.276161Z","iopub.execute_input":"2026-02-15T12:26:29.276471Z","iopub.status.idle":"2026-02-15T12:26:29.286574Z","shell.execute_reply.started":"2026-02-15T12:26:29.276444Z","shell.execute_reply":"2026-02-15T12:26:29.285764Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# spp predictor","metadata":{}},{"cell_type":"code","source":"class SpeciesPredictor:\n    \"\"\"Handles model loading, inference, and prediction aggregation.\"\"\"\n\n    def __init__(self, config: InferenceConfig):\n        self.config = config\n        self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        logger.info(f\"Using device: {self.device}\")\n\n        # Load species mapping\n        df_species_ids = pd.read_csv(config.species_ids_path)\n        self.class_map = df_species_ids['species_id'].to_dict()\n        self.num_classes = len(df_species_ids)\n        logger.info(f\"Number of species classes: {self.num_classes}\")\n\n        # Load model\n        self.model = self._load_model()\n\n        # Get model data config\n        data_config = timm.data.resolve_model_data_config(self.model)\n        self.model_mean = data_config['mean']\n        self.model_std = data_config['std']\n        self.model_input_size = data_config['input_size'][1]\n\n        # Normalization transform\n        self.normalize = T.Normalize(mean=self.model_mean, std=self.model_std)\n\n    def _load_model(self) -> torch.nn.Module:\n        \"\"\"Load and prepare the model.\"\"\"\n        logger.info(f\"Loading model: {self.config.model_name}\")\n        model = timm.create_model(\n            self.config.model_name,\n            pretrained=False,\n            num_classes=self.num_classes,\n            checkpoint_path=self.config.model_checkpoint\n        )\n        model = model.to(self.device)\n        model.eval()\n        logger.info(\"Model loaded successfully\")\n        return model\n\n    @torch.no_grad()\n    def _predict_patches(self, patches: torch.Tensor) -> Dict[int, List[float]]:\n        \"\"\"\n        Run inference on a set of patches and collect per-species scores.\n\n        Returns:\n            Dictionary mapping species_id -> list of scores across patches.\n        \"\"\"\n        species_scores = defaultdict(list)\n\n        patch_dataset = PatchDataset(patches, transform=self.normalize)\n        patch_loader = DataLoader(\n            patch_dataset,\n            batch_size=self.config.batch_size,\n            shuffle=False,\n            num_workers=0,  # patches are already in memory\n            pin_memory=False\n        )\n\n        for batch_patches in patch_loader:\n            batch_patches = batch_patches.to(self.device, non_blocking=True)\n\n            with autocast('cuda'):\n                outputs = self.model(batch_patches)\n                probabilities = F.softmax(outputs, dim=1)\n\n                top_probs, top_indices = torch.topk(\n                    probabilities, self.config.top_k_tile, dim=1\n                )\n                top_probs = top_probs.cpu().numpy()\n                top_indices = top_indices.cpu().numpy()\n\n            for tile_indices, tile_probs in zip(top_indices, top_probs):\n                for idx, prob in zip(tile_indices, tile_probs):\n                    if prob > self.config.min_score:\n                        species_id = self.class_map[idx]\n                        species_scores[species_id].append(float(prob))\n\n        return species_scores\n\n    def _apply_tta(self, patches: torch.Tensor) -> List[torch.Tensor]:\n        \"\"\"Apply test-time augmentation transforms to patches.\"\"\"\n        augmented = []\n        for transform_name in self.config.tta_transforms:\n            if transform_name == 'none':\n                augmented.append(patches)\n            elif transform_name == 'hflip':\n                augmented.append(torch.flip(patches, dims=[-1]))\n            elif transform_name == 'vflip':\n                augmented.append(torch.flip(patches, dims=[-2]))\n            elif transform_name == 'hflip_vflip':\n                augmented.append(torch.flip(patches, dims=[-1, -2]))\n        return augmented\n\n    def _aggregate_scores(self, species_scores: Dict[int, List[float]]) -> Dict[int, float]:\n        \"\"\"Aggregate per-patch scores into final species scores.\"\"\"\n        aggregated = {}\n\n        for species_id, scores in species_scores.items():\n            if self.config.aggregation == 'max':\n                aggregated[species_id] = max(scores)\n            elif self.config.aggregation == 'mean':\n                aggregated[species_id] = np.mean(scores)\n            elif self.config.aggregation == 'weighted_vote':\n                # Combine frequency (number of patches) with mean confidence\n                count_weight = np.log1p(len(scores))  # log-scaled count\n                mean_score = np.mean(scores)\n                aggregated[species_id] = mean_score * count_weight\n            else:\n                aggregated[species_id] = max(scores)\n\n        return aggregated\n\n    def predict_image(self, patches: torch.Tensor) -> List[int]:\n        \"\"\"\n        Predict species for a single image given its patches.\n\n        Args:\n            patches: Tensor of shape (num_patches, C, H, W)\n\n        Returns:\n            List of predicted species IDs.\n        \"\"\"\n        all_species_scores = defaultdict(list)\n\n        # Apply TTA if enabled\n        if self.config.use_tta:\n            patch_variants = self._apply_tta(patches)\n        else:\n            patch_variants = [patches]\n\n        # Collect scores from all TTA variants\n        for variant_patches in patch_variants:\n            variant_scores = self._predict_patches(variant_patches)\n            for species_id, scores in variant_scores.items():\n                all_species_scores[species_id].extend(scores)\n\n        # Aggregate scores\n        aggregated = self._aggregate_scores(all_species_scores)\n\n        if not aggregated:\n            return []\n\n        # Sort by aggregated score and apply final filtering\n        sorted_species = sorted(aggregated.items(), key=lambda x: x[1], reverse=True)\n\n        # Apply final top-k and minimum score threshold\n        final_predictions = []\n        for species_id, score in sorted_species[:self.config.final_top_k]:\n            if score >= self.config.final_min_score:\n                final_predictions.append(species_id)\n\n        return final_predictions\n\n    def run_inference(self) -> Dict[str, List[int]]:\n        \"\"\"Run inference on all test images.\"\"\"\n        image_predictions = {}\n\n        # Create datasets for each patch size (multi-scale)\n        for patch_size in self.config.patch_sizes:\n            stride = int(patch_size * self.config.stride_ratio)\n            logger.info(f\"Processing with patch_size={patch_size}, stride={stride}\")\n\n            dataset = TestImageDataset(\n                image_folder=self.config.test_images_dir,\n                patch_size=patch_size,\n                stride=stride,\n                use_pad=self.config.use_pad\n            )\n\n            dataloader = DataLoader(\n                dataset,\n                batch_size=1,  # One image at a time (variable patch counts)\n                num_workers=self.config.num_workers,\n                pin_memory=self.config.pin_memory,\n                shuffle=False\n            )\n\n            batch_time = AverageMeter()\n            eta_timer = ETATimer(len(dataloader))\n            end = time.time()\n\n            for batch_idx, (patches, image_path) in enumerate(dataloader):\n                quadrat_id = Path(image_path[0]).stem\n\n                # patches shape from dataloader: (1, num_patches, C, H, W)\n                patches = patches.squeeze(0)\n\n                try:\n                    predictions = self.predict_image(patches)\n                except Exception as e:\n                    logger.error(f\"Error processing {quadrat_id}: {e}\")\n                    predictions = []\n\n                # For multi-scale: merge predictions\n                if quadrat_id in image_predictions:\n                    existing = set(image_predictions[quadrat_id])\n                    existing.update(predictions)\n                    image_predictions[quadrat_id] = list(existing)\n                else:\n                    image_predictions[quadrat_id] = predictions\n\n                # Timing\n                batch_time.update(time.time() - end)\n                end = time.time()\n                eta_timer.update(batch_idx + 1)\n\n                if batch_idx % self.config.log_frequency == 0:\n                    logger.info(\n                        f'[{batch_idx}/{len(dataloader)}] '\n                        f'Time {batch_time.val:.3f}s (avg {batch_time.avg:.3f}s) '\n                        f'ETA {eta_timer.eta()} '\n                        f'Elapsed {eta_timer.elapsed()} '\n                        f'Species found: {len(predictions)}'\n                    )\n\n        logger.info(f\"Inference complete. Processed {len(image_predictions)} images.\")\n        return image_predictions\n\n    @staticmethod\n    def save_submission(predictions: Dict[str, List[int]], output_path: str):\n        \"\"\"Save predictions to submission CSV.\"\"\"\n        df_run = pd.DataFrame(\n            [(qid, species_ids) for qid, species_ids in predictions.items()],\n            columns=['quadrat_id', 'species_ids']\n        )\n        df_run['species_ids'] = df_run['species_ids'].apply(\n            lambda x: str(x) if x else '[]'\n        )\n        df_run.to_csv(output_path, sep=',', index=False, quoting=csv.QUOTE_ALL)\n        logger.info(f\"Submission saved to {output_path}\")\n\n        # Print statistics\n        num_empty = (df_run['species_ids'] == '[]').sum()\n        avg_species = df_run['species_ids'].apply(\n            lambda x: len(eval(x)) if x != '[]' else 0\n        ).mean()\n        logger.info(f\"Statistics: {len(df_run)} quadrats, \"\n                     f\"{num_empty} empty predictions, \"\n                     f\"{avg_species:.1f} avg species per quadrat\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T12:26:29.308952Z","iopub.execute_input":"2026-02-15T12:26:29.309216Z","iopub.status.idle":"2026-02-15T12:26:29.331298Z","shell.execute_reply.started":"2026-02-15T12:26:29.309193Z","shell.execute_reply":"2026-02-15T12:26:29.330568Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"# === Configuration ===\nconfig = InferenceConfig(\n    # Tiling\n    patch_sizes=[518],\n    stride_ratio=0.5,\n    use_pad=True,\n\n    # Inference\n    batch_size=64,\n    num_workers=4,\n\n    # Prediction filtering\n    min_score=0.05,\n    top_k_tile=3,\n    aggregation='weighted_vote',\n    final_top_k=50,\n    final_min_score=0.05,\n\n    # TTA\n    use_tta=True,\n    tta_transforms=['none', 'hflip'],\n\n    # Logging\n    log_frequency=10,\n)\n\n# Print config\nlogger.info(f\"Configuration:\\n{config}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T12:26:29.332666Z","iopub.execute_input":"2026-02-15T12:26:29.332982Z","iopub.status.idle":"2026-02-15T12:26:29.346291Z","shell.execute_reply.started":"2026-02-15T12:26:29.332955Z","shell.execute_reply":"2026-02-15T12:26:29.345520Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Run","metadata":{}},{"cell_type":"code","source":"# === Run Inference ===\npredictor = SpeciesPredictor(config)\npredictions = predictor.run_inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T12:26:29.366102Z","iopub.execute_input":"2026-02-15T12:26:29.366313Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submission\n","metadata":{}},{"cell_type":"code","source":"# === Save Submission ===\npredictor.save_submission(predictions, config.output_path)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Quick sanity check ===\ndf_submission = pd.read_csv(config.output_path)\nprint(f\"\\nSubmission shape: {df_submission.shape}\")\nprint(f\"\\nSample predictions:\")\nprint(df_submission.head(10))\n\n# Distribution of prediction counts\npred_counts = df_submission['species_ids'].apply(lambda x: len(eval(x)))\nprint(f\"\\nPrediction count statistics:\")\nprint(pred_counts.describe())","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}