{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":89850,"databundleVersionId":11256103,"sourceType":"competition"},{"sourceId":228781,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":195042,"modelId":216938}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Introduction\nThe purpose of this notebook is to show how to perform inference on quadrat test images with a tiling-based approach. As you may have have already noticed, predicting the species directly on the entire high-resolution image is a challenging task that often yields suboptimal results. The more intuitive strategy is then split the image into smaller patches, predict the species on each patch and finally aggregate the results.\n\nHere we use a simple method where predictions are thresholded based on species scores and their relative positions among the softmax output. The predictions for each tile are generated using the DINOv2-based model provided in the competition.\nThis notebook can be used as a starting point for further development. Feel free to leave comments on errors or for any improvement.","metadata":{}},{"cell_type":"markdown","source":"### Import libraries","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd\nimport timm \nimport torch\nfrom PIL import Image\nfrom torch.utils.data import DataLoader, Dataset\nimport time\nimport os\nimport torchvision.transforms as T\nfrom torch.amp import autocast\nfrom matplotlib import pyplot as plt\nfrom kornia import tensor_to_image\nfrom kornia.contrib import extract_tensor_patches, compute_padding\nimport csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T16:18:06.743496Z","iopub.execute_input":"2025-03-26T16:18:06.743830Z","iopub.status.idle":"2025-03-26T16:18:06.748420Z","shell.execute_reply.started":"2025-03-26T16:18:06.743797Z","shell.execute_reply":"2025-03-26T16:18:06.747572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AverageMeter:\n    def __init__(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        \n\nclass PatchDataset(Dataset):\n    def __init__(self, patches, transform=None):\n        self.patches = patches.squeeze(0)\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        \n        if self.transform:\n            patch = self.transform(patch)\n        return patch\n\n\nclass TestDataset(Dataset):\n    def __init__(self, image_folder, patch_size=518, stride=259, transform=None, use_pad=False):\n        self.image_folder = image_folder\n        self.image_paths = [os.path.join(image_folder, f) for f in os.listdir(image_folder)]\n        self.transform = transform\n        self.use_pad = use_pad\n        self.patch_size = patch_size\n        self.stride = stride\n        \n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n        image = Image.open(image_path)\n\n        if self.transform:\n            image = self.transform(image).unsqueeze(0)\n        \n        h, w = image.shape[-2:]\n        \n        if self.use_pad:\n            pad = compute_padding(original_size=(h, w), window_size=self.patch_size, stride=self.stride)\n            patches = extract_tensor_patches(image, self.patch_size, self.stride, padding=pad)\n        else:\n            patches = extract_tensor_patches(image, self.patch_size, self.stride)\n\n        return patches, image_path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:27.919915Z","iopub.execute_input":"2025-03-26T14:17:27.920460Z","iopub.status.idle":"2025-03-26T14:17:27.928475Z","shell.execute_reply.started":"2025-03-26T14:17:27.920425Z","shell.execute_reply":"2025-03-26T14:17:27.927644Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_species_ids = pd.read_csv('/kaggle/input/plantclef-2025/species_ids.csv')\n\ndf_metadata = pd.read_csv('/kaggle/input/plantclef-2025/PlantCLEF2024_single_plant_training_metadata.csv', sep=';', dtype={'partner': str})\nclass_map = df_species_ids['species_id'].to_dict() # dictionary to map the species model Id with the species Id\n\ndf_metadata.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:20:30.685466Z","iopub.execute_input":"2025-03-26T14:20:30.685780Z","iopub.status.idle":"2025-03-26T14:20:42.311518Z","shell.execute_reply.started":"2025-03-26T14:20:30.685757Z","shell.execute_reply":"2025-03-26T14:20:42.310675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda')\nmodel = timm.create_model('vit_base_patch14_reg4_dinov2.lvd142m',\n                          pretrained=False,\n                          num_classes=len(df_species_ids),\n                          checkpoint_path='/kaggle/input/dinov2_patch14_reg4_onlyclassifier_then_all/pytorch/default/3/model_best.pth.tar')\nmodel = model.to(device)\nmodel = model.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:40.284927Z","iopub.execute_input":"2025-03-26T14:17:40.285244Z","iopub.status.idle":"2025-03-26T14:17:43.129538Z","shell.execute_reply.started":"2025-03-26T14:17:40.285191Z","shell.execute_reply":"2025-03-26T14:17:43.128570Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_config = timm.data.resolve_model_data_config(model)\nmodel_input_size, model_mean, model_std = data_config['input_size'][1], data_config['mean'], data_config['std']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:43.130480Z","iopub.execute_input":"2025-03-26T14:17:43.130727Z","iopub.status.idle":"2025-03-26T14:17:43.134942Z","shell.execute_reply.started":"2025-03-26T14:17:43.130705Z","shell.execute_reply":"2025-03-26T14:17:43.133972Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Set hyperparameters:\n* batch_size: size of batch of testing images\n* top_k: keep best top_k results for each patch\n* min_score: keep only classes with a score higher than min_score\n* patch_size: size of patches. We will be using the DinoV2 image input size that is 518.\n* stride: overlapping stride. We will be using half of the patch size.","metadata":{}},{"cell_type":"code","source":"batch_size = 64\nmin_score = 0.1\ntop_k_tile = 2\npatch_size = model_input_size\nstride = int(model_input_size / 2)\nuse_pad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:43.135983Z","iopub.execute_input":"2025-03-26T14:17:43.136233Z","iopub.status.idle":"2025-03-26T14:17:43.147088Z","shell.execute_reply.started":"2025-03-26T14:17:43.136191Z","shell.execute_reply":"2025-03-26T14:17:43.146267Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Exemple of tiling on single image. Let's load a random image","metadata":{}},{"cell_type":"code","source":"img = Image.open('/kaggle/input/plantclef-2025/PlantCLEF2025_test_images/PlantCLEF2025_test_images/GUARDEN-CBNMed-30-4-16-3-20240428.jpg')\nplt.imshow(img)\nplt.axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:43.148283Z","iopub.execute_input":"2025-03-26T14:17:43.148615Z","iopub.status.idle":"2025-03-26T14:17:43.766607Z","shell.execute_reply.started":"2025-03-26T14:17:43.148574Z","shell.execute_reply":"2025-03-26T14:17:43.765601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_to_tensor = T.ToTensor()\n\nimage_tensor = image_to_tensor(img).unsqueeze(0)\nh, w = image_tensor.shape[-2:]\n\npad = compute_padding(original_size=(h, w), window_size=patch_size, stride=stride)\n\npatches = extract_tensor_patches(image_tensor, patch_size, stride, padding=pad)\nprint(f\"Shape of image tiles = {patches.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:43.768234Z","iopub.execute_input":"2025-03-26T14:17:43.768490Z","iopub.status.idle":"2025-03-26T14:17:43.844158Z","shell.execute_reply.started":"2025-03-26T14:17:43.768467Z","shell.execute_reply":"2025-03-26T14:17:43.843170Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Let's now visualize the 64 patches extracted with a size of 518 (the same of model input)","metadata":{}},{"cell_type":"code","source":"fig, axs = plt.subplots(8, 8)\naxs = axs.ravel()\n\nfor i in range(len(patches[0])):\n    axs[i].axis(\"off\")\n    axs[i].imshow(tensor_to_image(patches[0][i]))\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:43.845336Z","iopub.execute_input":"2025-03-26T14:17:43.845619Z","iopub.status.idle":"2025-03-26T14:17:47.127629Z","shell.execute_reply.started":"2025-03-26T14:17:43.845596Z","shell.execute_reply":"2025-03-26T14:17:47.126727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example of a single patch\nplt.imshow(tensor_to_image(patches[0][29]))\nplt.axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:17:47.128444Z","iopub.execute_input":"2025-03-26T14:17:47.128674Z","iopub.status.idle":"2025-03-26T14:17:47.274507Z","shell.execute_reply.started":"2025-03-26T14:17:47.128655Z","shell.execute_reply":"2025-03-26T14:17:47.273601Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Run over all the test dataset","metadata":{}},{"cell_type":"code","source":"dataset = TestDataset(image_folder='/kaggle/input/plantclef-2025/PlantCLEF2025_test_images/PlantCLEF2025_test_images/',\n                      patch_size=patch_size,\n                      stride=stride,\n                      use_pad=True,\n                      transform=image_to_tensor)\ndataloader = DataLoader(dataset, batch_size=1, num_workers=4, pin_memory=True)\n\nimage_predictions = {}\n\n# Initialize batch time tracking\nbatch_time = AverageMeter()\nend = time.time()\n\nwith torch.no_grad():\n    for batch_idx, (patches, image_path) in enumerate(dataloader):\n        image_results = {}\n        quadrat_id = os.path.splitext(os.path.basename(image_path[0]))[0]\n        transform_patch = T.Normalize(mean=model_mean, std=model_std)\n        patch_dataset = PatchDataset(patches[0], transform=transform_patch)\n        patch_loader = DataLoader(patch_dataset, batch_size=batch_size, shuffle=False)\n        \n        for batch_patches in patch_loader:\n            batch_patches = batch_patches.to(device)\n            \n            with autocast('cuda'):\n                outputs = model(batch_patches)  # Perform inference on the batch\n                probabilities = torch.nn.functional.softmax(outputs, dim=1)\n        \n                # Get the top-k indices and probabilities\n                top_probs, top_indices = torch.topk(probabilities, top_k_tile)\n                top_probs = top_probs.cpu().numpy()\n                top_indices = top_indices.cpu().numpy()\n                \n                for top_tile_indices, top_tile_probs in zip(top_indices, top_probs):\n                    for top_idx, top_prob in zip(top_tile_indices, top_tile_probs):\n                        species_id = class_map[top_idx]\n                        # Update the results dictionary only if the probability is higher\n                        if top_prob > min_score:\n                            if top_idx not in image_results or image_results[top_idx] < top_prob:\n                                image_results[species_id] = top_prob\n        # store the prediction\n        image_predictions[quadrat_id] = list(image_results.keys())\n        \n        batch_time.update(time.time() - end)\n        end = time.time()\n\n        # Log info at specified frequency\n        if batch_idx % 10 == 0:  # You can set your log frequency here\n            print(f'Predict: [{batch_idx}/{len(dataloader)}] '\n                  f'Time {batch_time.val:.3f} ({batch_time.avg:.3f})')  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:23:31.379676Z","iopub.execute_input":"2025-03-26T14:23:31.379968Z","iopub.status.idle":"2025-03-26T16:13:58.133604Z","shell.execute_reply.started":"2025-03-26T14:23:31.379947Z","shell.execute_reply":"2025-03-26T16:13:58.132415Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Submit prediction","metadata":{"execution":{"iopub.status.busy":"2025-03-26T14:15:44.308841Z","iopub.execute_input":"2025-03-26T14:15:44.309135Z","iopub.status.idle":"2025-03-26T14:15:44.314362Z","shell.execute_reply.started":"2025-03-26T14:15:44.309112Z","shell.execute_reply":"2025-03-26T14:15:44.313293Z"}}},{"cell_type":"code","source":"df_run = pd.DataFrame(list(image_predictions.items()), columns=['quadrat_id', 'species_ids'])\ndf_run['species_ids'] = df_run['species_ids'].apply(str)\ndf_run.to_csv(\"submission.csv\", sep=',', index=False, quoting=csv.QUOTE_ALL)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T16:27:54.439163Z","iopub.execute_input":"2025-03-26T16:27:54.439527Z","iopub.status.idle":"2025-03-26T16:27:54.460491Z","shell.execute_reply.started":"2025-03-26T16:27:54.439494Z","shell.execute_reply":"2025-03-26T16:27:54.459623Z"}},"outputs":[],"execution_count":null}]}