{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91448,"databundleVersionId":11249847,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🎯 Few-shot fungi classification\nThis notebook provides a simple baselines for the few-shot classification of fungi. Both baselines require you to pre-extracted features using any encoder, e.g., BioCLIP, and DINOv2. For the classification, we provide:\n1. **Centroid-based classifier** computes the class prototype by averaging the features of the support set. The class prototype is then used to classify the query set.\n2. **Nearest neighbour classifier** classifies the query set by finding the nearest neighbour in the support set.","metadata":{}},{"cell_type":"markdown","source":"## ⚙️ Experiment Setup\nInstall and import for the required libraries. \nSetting up pathing etc.\n\nWe also set path to dataset (which are downloaded automatically).","metadata":{}},{"cell_type":"code","source":"###Install required libraries:\n!pip install git+https://github.com/mlfoundations/open_clip.git\n!pip install faiss-gpu -qq","metadata":{"execution":{"iopub.status.busy":"2025-03-10T21:15:10.983887Z","iopub.execute_input":"2025-03-10T21:15:10.984265Z","iopub.status.idle":"2025-03-10T21:15:29.226959Z","shell.execute_reply.started":"2025-03-10T21:15:10.984236Z","shell.execute_reply":"2025-03-10T21:15:29.225799Z"},"tags":["formatted"],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport yaml\nfrom pathlib import Path\nfrom types import SimpleNamespace\nimport argparse\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport faiss\n\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom torchvision import transforms as tfms\nimport torchvision.transforms as T\nimport open_clip\n\nfrom typing import Sequence, Tuple, Any, Dict, List, Optional, Union\nimport importlib","metadata":{"execution":{"iopub.status.busy":"2025-03-10T21:15:29.228281Z","iopub.execute_input":"2025-03-10T21:15:29.228607Z","iopub.status.idle":"2025-03-10T21:15:42.173144Z","shell.execute_reply.started":"2025-03-10T21:15:29.228582Z","shell.execute_reply":"2025-03-10T21:15:42.172090Z"},"tags":["formatted"],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# path to fungitatsic dataset\ndata_path = '/kaggle/input/fungi-clef-2025/'\n","metadata":{"execution":{"iopub.status.busy":"2025-03-10T21:15:42.174968Z","iopub.execute_input":"2025-03-10T21:15:42.175241Z","iopub.status.idle":"2025-03-10T21:15:42.179343Z","shell.execute_reply.started":"2025-03-10T21:15:42.175217Z","shell.execute_reply":"2025-03-10T21:15:42.178225Z"},"tags":["formatted"],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔬 1. Define BioCLIP feature extractor\n\n### BioCLIP: A Deep Learning Model for Biological Image Feature Extraction\n\n🔗 [BioCLIP](https://huggingface.co/imageomics/bioclip) is a deep learning model designed for extracting high-quality feature embeddings from biological images. It is built on the **Contrastive Language-Image Pretraining (CLIP) framework**, specifically tailored for biological applications. The model is hosted on the **Hugging Face Hub**, making it easily accessible for researchers and developers working with biological image data.\n\n#### Purpose and Applications\n\nBioCLIP is designed to bridge the gap between visual data and semantic understanding in biological research. Its main applications include:\n\n- **Automated Species Identification**: Extracting meaningful features from images of fungi, plants, or animals to assist in species recognition.\n- **Biodiversity Monitoring**: Enabling large-scale ecological studies by processing images from camera traps or field surveys.\n- **Medical and Microscopy Analysis**: Supporting tasks like cell classification, pathology slide analysis, and disease detection.\n- **Bioinformatics and AI-powered Conservation**: Facilitating the development of models that leverage image-based insights for conservation efforts.\n\n#### Key Features\n\n- **Pretrained on Biological Datasets**: Unlike general-purpose CLIP models, BioCLIP has been fine-tuned on biological image datasets, improving its performance in domain-specific tasks.\n- **Feature Extraction for Downstream Tasks**: Generates robust, high-dimensional embeddings that can be used for classification, clustering, or retrieval.\n- **L2-Normalized Embeddings**: Ensures that extracted feature vectors have unit length, making them suitable for similarity comparisons.\n- **Seamless Integration with PyTorch**: Designed to work within PyTorch-based workflows for machine learning and deep learning applications.\n\n#### Code and Implementation\n\nThe implementation of BioCLIP for feature extraction is available in the following repository:\n\nThis model can be integrated into Python workflows using `open_clip`, allowing users to load the model, preprocess images, and extract embeddings efficiently.","metadata":{}},{"cell_type":"code","source":"class BioCLIP(torch.nn.Module):\n\n    def __init__(self, device):\n        \"\"\"\n        Initialize the BioCLIP feature extractor.\n        \"\"\"\n        super(BioCLIP, self).__init__()\n        self.device = device\n        self.model = None\n        self.processor = None\n        self.size = (224, 224)\n\n    def load(self):\n        \"\"\"\n        Load the BioCLIP model and its associated image processor.\n        \n        The model is loaded from the Hugging Face Hub and moved to the specified device.\n        \"\"\"\n        self.model, _, self.processor = open_clip.create_model_and_transforms('hf-hub:imageomics/bioclip')\n        self.model.to(self.device)\n\n    def extract_features(self, image):\n        \"\"\"\n        Extract normalized feature embeddings from a given image.\n\n        Args:\n            image (PIL.Image.Image): The input image from which to extract features.\n\n        Returns:\n            torch.Tensor: A normalized feature embedding vector for the input image.\n\n        Raises:\n            ValueError: If the model has not been loaded prior to calling this method.\n        \"\"\"\n        if self.model is None:\n            raise ValueError('Model not loaded')\n        image = image.resize(self.size, Image.BICUBIC)\n        image_tensor_proc = self.processor(image)[None]\n        features = self.model.encode_image(image_tensor_proc.to(self.device))\n        return self.normalize_embedding(features)\n\n    @staticmethod\n    def normalize_embedding(embs):\n        \"\"\"\n        Normalize the embedding vectors to have unit length.\n\n        Args:\n            embs (torch.Tensor): The raw embedding vectors.\n\n        Returns:\n            torch.Tensor: L2-normalized embedding vectors.\n        \"\"\"\n        return torch.nn.functional.normalize(embs.float(), dim=1, p=2)\n","metadata":{"execution":{"iopub.status.busy":"2025-03-10T21:15:42.180889Z","iopub.execute_input":"2025-03-10T21:15:42.181229Z","iopub.status.idle":"2025-03-10T21:15:42.208347Z","shell.execute_reply.started":"2025-03-10T21:15:42.181205Z","shell.execute_reply":"2025-03-10T21:15:42.207337Z"},"tags":["formatted"],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 📙 2. Define dataset class for FungiTastic\n\n`FungiTastic` is a PyTorch dataset class designed for handling the **Few-Shot subset** of the Danish Fungi dataset, with support for **train, validation, and test splits**. The class allows for efficient data handling, including image loading, embedding management, and transformation application.\n\n#### Key Features:\n- **Dataset Initialization**: Loads metadata (e.g., file paths, class IDs) from CSV files. Supports splits (`train`, `val`, `test`) and applies optional image transformations.\n- **Embedding Handling**: Adds precomputed embeddings to the dataset and allows retrieval of embeddings for specific classes.\n- **Data Access**: Implements the `__getitem__()` method to fetch images, class IDs, and file paths. It also returns embeddings if available.\n- **Class Mapping**: Maps class IDs to species names and allows retrieving the category of samples by ID.\n- **Sample Visualization**: Provides the `show_sample()` method to display images along with their class labels.\n\n#### Functions:\n- **`add_embeddings(embeddings)`**: Merges new embeddings with the dataset.\n- **`get_embeddings_for_class(id)`**: Retrieves embeddings for a given class ID.\n- **`_get_df(data_path, split)`**: Loads the dataset metadata into a pandas DataFrame.\n- **`__getitem__(idx)`**: Fetches a sample by index, including image, category ID, and embedding (if available).\n- **`__len__()`**: Returns the total number of samples in the dataset.\n- **`get_class_id(idx)`**: Retrieves the class ID for a specific sample.\n- **`show_sample(idx)`**: Displays an image with its associated class label.\n","metadata":{}},{"cell_type":"code","source":"class FungiTastic(torch.nn.Module):\n    \"\"\"\n    Dataset class for the FewShot subset of the Danish Fungi dataset (size 300, closed-set).\n\n    This dataset loader supports training, validation, and testing splits, and provides\n    convenient access to images, class IDs, and file paths. It also supports optional\n    image transformations.\n    \"\"\"\n\n    SPLIT2STR = {'train': 'Train', 'val': 'Val', 'test': 'Test'}\n\n    def __init__(self, root: str, split: str = 'val', transform=None):\n        \"\"\"\n        Initializes the FungiTastic dataset.\n\n        Args:\n            root (str): The root directory of the dataset.\n            split (str, optional): The dataset split to use. Must be one of {'train', 'val', 'test'}.\n                Defaults to 'val'.\n            transform (callable, optional): Optional transform to be applied on a sample.\n        \"\"\"\n        super().__init__()\n        self.split = split\n        self.transform = transform\n        self.df = self._get_df(root, split)\n\n        assert \"image_path\" in self.df\n        if self.split != 'test':\n            assert \"category_id\" in self.df\n            self.n_classes = len(self.df['category_id'].unique())\n            self.category_id2label = {\n                k: v[0] for k, v in self.df.groupby('category_id')['species'].unique().to_dict().items()\n            }\n            self.label2category_id = {\n                v: k for k, v in self.category_id2label.items()\n            }\n\n    def add_embeddings(self, embeddings: pd.DataFrame):\n        \"\"\"\n        Updates the dataset instance with new embeddings.\n\n        Args:\n            embeddings (pd.DataFrame): A DataFrame containing an 'embedding' column.\n                                       It must align with `self.df` in terms of indexing.\n        \"\"\"\n        assert isinstance(embeddings, pd.DataFrame), \"Embeddings must be a pandas DataFrame.\"\n        assert \"embedding\" in embeddings.columns, \"Embeddings DataFrame must have an 'embedding' column.\"\n        assert len(embeddings) == len(self.df), \"Embeddings must match dataset length.\"\n\n        self.df = pd.merge(self.df, embeddings, on=\"filename\", how=\"inner\")\n\n    def get_embeddings_for_class(self, id):\n        # return the embeddings for class class_idx\n        class_idxs = self.df[self.df['category_id'] == id].index\n        return self.df.iloc[class_idxs]['embedding']\n    \n    @staticmethod\n    def _get_df(data_path: str, split: str) -> pd.DataFrame:\n        \"\"\"\n        Loads the dataset metadata as a pandas DataFrame.\n\n        Args:\n            data_path (str): The root directory where the dataset is stored.\n            split (str): The dataset split to load. Must be one of {'train', 'val', 'test'}.\n\n        Returns:\n            pd.DataFrame: A DataFrame containing metadata and file paths for the split.\n        \"\"\"\n        df_path = os.path.join(\n            data_path,\n            \"metadata\",\n            \"FungiTastic-FewShot\",\n            f\"FungiTastic-FewShot-{FungiTastic.SPLIT2STR[split]}.csv\"\n        )\n        df = pd.read_csv(df_path)\n        df[\"image_path\"] = df.filename.apply(\n            lambda x: os.path.join(data_path, \"FungiTastic-FewShot\", split, '300p', x)\n        )\n        return df\n\n    def __getitem__(self, idx: int):\n        \"\"\"\n        Retrieves a single data sample by index.\n    \n        Args:\n            idx (int): Index of the sample to retrieve.\n            ret_image (bool, optional): Whether to explicitly return the image. Defaults to False.\n    \n        Returns:\n            tuple:\n                - If embeddings exist: (image?, embedding, category_id, file_path)\n                - If no embeddings: (image, category_id, file_path) (original version)\n        \"\"\"\n        file_path = self.df[\"image_path\"].iloc[idx].replace('FungiTastic-FewShot', 'images/FungiTastic-FewShot')\n    \n        if self.split != 'test':\n            category_id = self.df[\"category_id\"].iloc[idx]\n        else:\n            category_id = None\n    \n        image = Image.open(file_path)\n    \n        if self.transform:\n            image = self.transform(image)\n    \n        # Check if embeddings exist\n        if \"embedding\" in self.df.columns:\n            emb = torch.tensor(self.df.iloc[idx]['embedding'], dtype=torch.float32).squeeze()\n        else:\n            emb = None  # No embeddings available\n    \n\n        return image, category_id, file_path, emb\n\n\n    def __len__(self):\n        \"\"\"\n        Returns the number of samples in the dataset.\n        \"\"\"\n        return len(self.df)\n\n    def get_class_id(self, idx: int) -> int:\n        \"\"\"\n        Returns the class ID of a specific sample.\n        \"\"\"\n        return self.df[\"category_id\"].iloc[idx]\n\n    def show_sample(self, idx: int) -> None:\n        \"\"\"\n        Displays a sample image along with its class name and index.\n        \"\"\"\n        image, category_id, _, _ = self.__getitem__(idx)\n        class_name = self.category_id2label[category_id]\n\n        plt.imshow(image)\n        plt.title(f\"Class: {class_name}; id: {idx}\")\n        plt.axis('off')\n        plt.show()\n\n    def get_category_idxs(self, category_id: int) -> List[int]:\n        \"\"\"\n        Retrieves all indexes for a given category ID.\n        \"\"\"\n        return self.df[self.df.category_id == category_id].index.tolist()","metadata":{"execution":{"iopub.status.busy":"2025-03-10T21:15:42.209424Z","iopub.execute_input":"2025-03-10T21:15:42.209817Z","iopub.status.idle":"2025-03-10T21:15:42.232637Z","shell.execute_reply.started":"2025-03-10T21:15:42.209780Z","shell.execute_reply":"2025-03-10T21:15:42.231714Z"},"tags":["formatted"],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Load the datasets\n\ntrain_dataset = FungiTastic(root=data_path, split='train', transform=None)\ntest_dataset = FungiTastic(root=data_path, split='test', transform=None)\n\ntrain_dataset.df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:15:42.233681Z","iopub.execute_input":"2025-03-10T21:15:42.234064Z","iopub.status.idle":"2025-03-10T21:15:42.561805Z","shell.execute_reply.started":"2025-03-10T21:15:42.234028Z","shell.execute_reply":"2025-03-10T21:15:42.560867Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Visualize few samples\n\ntrain_dataset.show_sample(1), train_dataset.show_sample(500), train_dataset.show_sample(1000) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:15:42.562731Z","iopub.execute_input":"2025-03-10T21:15:42.563130Z","iopub.status.idle":"2025-03-10T21:15:43.299071Z","shell.execute_reply.started":"2025-03-10T21:15:42.563097Z","shell.execute_reply":"2025-03-10T21:15:43.298091Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ⏳ 3. Precompute embeddings using BioCLIP\n\nThis code snippet demonstrates how to precompute image embeddings using the **BioCLIP** model. It loads the BioCLIP model, processes a given dataset, and stores the generated embeddings.\n\n#### Steps:\n1. **Initialize the Model**: The code sets up the **BioCLIP** model on the available device (GPU if available, else CPU). The model is then loaded and set to evaluation mode (`model.eval()`).\n2. **Generate Embeddings**: \n   - The `generate_embeddings()` function loops through all the samples in the dataset.\n   - For each image, it extracts the feature embeddings using `model.extract_features()`.\n   - The extracted embeddings are stored in a pandas DataFrame with the image filename.\n3. **Add Embeddings to Dataset**: The generated embeddings are then added to the respective dataset (`test_dataset` and `train_dataset`) using the `add_embeddings()` method.\n\n#### Key Components:\n- **BioCLIP model**: Used to extract embeddings for each image in the dataset.\n- **Data Handling**: `im_names` stores image filenames, and `embs` stores the corresponding embeddings.\n- **Embedding Storage**: Embeddings are stored in a pandas DataFrame, which is later merged with the dataset.\n\nThis process enables efficient embedding-based learning and retrieval tasks on the dataset.\n","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = BioCLIP(device=device)\nmodel.load()\nmodel.eval()\n\ndef generate_embeddings(dataset):\n\n    idxs = np.arange(len(dataset))\n    im_names, embs = [], []\n    for idx in tqdm(idxs):\n        im, label, file_path, _ = dataset[idx]\n\n        with torch.no_grad():\n            feat = model.extract_features(im)\n            \n        # feat_quant = model.quantize_normalized_embedding(feat)\n\n        im_names.append(os.path.basename(file_path))\n        embs.append(feat.detach().cpu().numpy())\n\n    embeddings = pd.DataFrame({'filename': im_names, 'embedding': embs})\n\n    return embeddings","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:15:43.301127Z","iopub.execute_input":"2025-03-10T21:15:43.301368Z","iopub.status.idle":"2025-03-10T21:15:49.230614Z","shell.execute_reply.started":"2025-03-10T21:15:43.301349Z","shell.execute_reply":"2025-03-10T21:15:49.229848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"embeddings = generate_embeddings(test_dataset)\ntest_dataset.add_embeddings(embeddings)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:15:49.231699Z","iopub.execute_input":"2025-03-10T21:15:49.231981Z","iopub.status.idle":"2025-03-10T21:16:33.136872Z","shell.execute_reply.started":"2025-03-10T21:15:49.231947Z","shell.execute_reply":"2025-03-10T21:16:33.135832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"embeddings = generate_embeddings(train_dataset)\ntrain_dataset.add_embeddings(embeddings)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:16:33.137959Z","iopub.execute_input":"2025-03-10T21:16:33.138247Z","iopub.status.idle":"2025-03-10T21:19:25.785845Z","shell.execute_reply.started":"2025-03-10T21:16:33.138217Z","shell.execute_reply":"2025-03-10T21:19:25.784598Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🏃🏻‍♂4. Run the few-shot classifier on the precomputed features","metadata":{}},{"cell_type":"markdown","source":"### 4.1 Classification with kNN\n\nThis code defines a **Nearest Neighbor (NN) Classifier** that uses precomputed embeddings to make predictions by comparing them to class prototypes. It employs **cosine similarity** to match test embeddings with training embeddings, making it suitable for few-shot and embedding-based classification tasks.\n\n#### Key Steps:\n\n1. **Classifier Initialization**:\n   - The `NNClassifier` class is initialized with the training embeddings and sets up a similarity search index using **FAISS** (Facebook AI Similarity Search). \n   - The `build_index()` method creates an index for efficient similarity searches based on the cosine similarity between embeddings.\n\n2. **Making Predictions**:\n   - The `make_prediction()` method computes the similarity between each test embedding and the stored class prototypes. It returns the predicted class labels (`cls`) and confidence scores (`conf`) based on the similarity values.\n   - **FAISS** is used to search the nearest prototypes for each test embedding.\n\n3. **Generating Predictions for the Test Dataset**:\n   - After initializing the classifier, predictions are made for the **test dataset** using the `make_prediction()` method.\n   - Predicted labels are added to the test dataset (`test_dataset.df[\"preds\"]`).\n\n4. **Submission Generation**:\n   - The predictions are grouped by `observationID` and formatted as a space-separated string of class IDs.\n   - The predictions are saved to a CSV file for submission.\n\n#### Key Components:\n- **FAISS**: Efficiently computes nearest neighbor search using cosine similarity.\n- **Embedding-based Classification**: The model leverages precomputed embeddings for each sample and matches them to class prototypes.\n- **Prediction Confidence**: The similarity values serve as confidence scores for each prediction.\n\nThis approach is effective for embedding-based classification tasks where models make predictions based on the similarity of embeddings to predefined class prototypes.\n","metadata":{}},{"cell_type":"code","source":"class NNClassifier():\n    def __init__(self, train_embeddings, device='cuda'):\n        \"\"\"\n        :param cfg: config object, namespace\n        :param train_embeddings: list of C torch arrays of shape [N_C, D] where N_C is the number of training samples\n        of class C and D is the dimensionality of the embeddings\n        \"\"\"\n        \n        self.device = device\n        self.index, self.idx2cls = self.build_index(train_embeddings)\n\n    def make_prediction(self, embeddings, plot_sim_hist=False, ret_probs=False):\n        \"\"\"\n        :param embeddings: torch.Tensor of shape (batch_size, n_channels, height, width)\n        :return: probabilities of shape (batch_size, n_classes) computed based on\n        the similarity of the embeddings to the class prototypes\n        \"\"\"\n        # compute the similarity of each embedding to each prototype\n        # embeddings - [N, D], class_prototypes - [C, D]\n        similarities, indices = self.index.search(embeddings, 1)\n        # get the classes for the indices\n        cls = self.idx2cls[indices.squeeze()]\n        # get the confidence of the prediction\n        conf = similarities\n        \n        return cls, conf\n\n    def build_index(self, train_dataset):\n        idx2cls = train_dataset.df.category_id.values\n        # concatenate the embeddings\n        embs = np.array(train_dataset.df.embedding.values.tolist(), dtype=np.float32).squeeze()\n        # build the index for cosine similarity search\n        index = faiss.IndexFlatIP(embs.shape[1])\n        index.add(embs)\n        return index, idx2cls","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:19:25.786830Z","iopub.execute_input":"2025-03-10T21:19:25.787115Z","iopub.status.idle":"2025-03-10T21:19:25.793266Z","shell.execute_reply.started":"2025-03-10T21:19:25.787088Z","shell.execute_reply":"2025-03-10T21:19:25.792354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"classifier = NNClassifier(train_dataset, device='cpu')\n\ncls, conf = classifier.make_prediction(np.array(test_dataset.df.embedding.values.tolist(), dtype=np.float32).squeeze())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:19:25.794132Z","iopub.execute_input":"2025-03-10T21:19:25.794396Z","iopub.status.idle":"2025-03-10T21:19:25.916738Z","shell.execute_reply.started":"2025-03-10T21:19:25.794375Z","shell.execute_reply":"2025-03-10T21:19:25.915716Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset.df[\"preds\"] = cls\nsubmission = (\n    test_dataset.df.groupby(\"observationID\")[\"preds\"]\n    .apply(lambda x: \" \".join(map(str, sorted(set(x)))))  # Sort and remove duplicates\n    .reset_index()\n)\n\nsubmission = submission.rename(columns={\"observationID\": \"observationId\", \"preds\": \"predictions\"})\nsubmission = submission.drop_duplicates(subset=\"observationId\")\nsubmission.to_csv(\"baseline-submission-with-nn.csv\", index=None)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:19:25.917662Z","iopub.execute_input":"2025-03-10T21:19:25.917914Z","iopub.status.idle":"2025-03-10T21:19:26.199588Z","shell.execute_reply.started":"2025-03-10T21:19:25.917892Z","shell.execute_reply":"2025-03-10T21:19:26.198853Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4.2 Classification with Prototypes\n\nThe `PrototypeClassifier` is a PyTorch model designed for **embedding-based classification** using **prototype-based classification**. It learns the class prototypes by averaging the embeddings of each class and uses cosine similarity to classify new embeddings based on their proximity to the prototypes.\n\n#### Key Steps:\n\n1. **Initialization**:\n   - The model is initialized with a **train dataset** and computes class prototypes using embeddings of each class.\n   - The `class_prototypes` are stored as **non-trainable parameters** (`requires_grad=False`), representing the average embeddings of each class.\n\n2. **Embedding Retrieval**:\n   - The `_get_classifier_embeddings()` method retrieves the embeddings for each class in the training dataset.\n   - If a class has no embeddings, a zero vector is used as a placeholder for that class.\n\n3. **Prototype Computation**:\n   - The `get_prototypes()` method computes the **prototype** for each class by averaging its embeddings. These prototypes serve as the central representation of each class.\n\n4. **Prediction**:\n   - The `make_prediction()` method calculates the cosine similarity between each input embedding and all class prototypes.\n   - The predicted class (`cls`) is the one with the highest similarity, and the prediction confidence (`conf`) is derived from the maximum softmax probability over the similarities.\n\n#### Key Components:\n- **Class Prototypes**: Each class is represented by its average embedding, which acts as the class prototype.\n- **Cosine Similarity**: The model uses cosine similarity to compare input embeddings with class prototypes and determine the predicted class.\n- **Prediction Confidence**: The confidence is computed using the softmax of the cosine similarity values, providing a measure of certainty for each prediction.\n\nThis classifier is well-suited for tasks where class prototypes (average embeddings) can effectively represent the classes, and it works well with embedding-based models for few-shot or metric learning tasks.\n","metadata":{"execution":{"iopub.status.busy":"2025-03-10T15:10:15.574562Z","iopub.execute_input":"2025-03-10T15:10:15.574896Z","iopub.status.idle":"2025-03-10T15:10:15.594254Z","shell.execute_reply.started":"2025-03-10T15:10:15.574828Z","shell.execute_reply":"2025-03-10T15:10:15.593200Z"}}},{"cell_type":"code","source":"class PrototypeClassifier(torch.nn.Module):\n    def __init__(self, train_dataset, device='cuda'):\n        super().__init__()\n        self.device = device\n        self.train_dataset = train_dataset\n\n        class_embeddings, _ = self._get_classifier_embeddings(train_dataset)\n        self.class_prototypes = torch.nn.Parameter(self.get_prototypes(class_embeddings), requires_grad=False)\n\n    def _get_classifier_embeddings(self, dataset_train):\n        class_embeddings = []\n        empty_classes = []\n        n_classes = min(torch.inf, dataset_train.n_classes)\n        for cls in range(n_classes):\n            cls_embs = dataset_train.get_embeddings_for_class(cls)\n            if len(cls_embs) == 0:\n                # if no embeddings for class, use zeros\n                empty_classes.append(cls)\n                class_embeddings.append(torch.zeros(1, dataset_train.emb_dim))\n            else:\n                class_embeddings.append(torch.tensor(np.vstack(cls_embs.values)))\n        return class_embeddings, empty_classes\n\n    def get_prototypes(self, embeddings):\n        return torch.stack([class_embs.mean(dim=0) for class_embs in embeddings])\n\n    def make_prediction(self, embeddings):\n        similarities = torch.nn.functional.cosine_similarity(embeddings, self.class_prototypes, dim=-1)\n        cls = torch.argmax(similarities, dim=1)\n        conf = torch.nn.functional.softmax(similarities, dim=1).max(dim=1).values\n        return cls, conf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:19:26.200364Z","iopub.execute_input":"2025-03-10T21:19:26.200593Z","iopub.status.idle":"2025-03-10T21:19:26.208086Z","shell.execute_reply.started":"2025-03-10T21:19:26.200573Z","shell.execute_reply":"2025-03-10T21:19:26.207224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"classifier = PrototypeClassifier(train_dataset, device='cpu')\ncls, conf = classifier.make_prediction(torch.tensor(np.array(test_dataset.df.embedding.values.tolist(), dtype=np.float32)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:19:26.209082Z","iopub.execute_input":"2025-03-10T21:19:26.209423Z","iopub.status.idle":"2025-03-10T21:19:45.089333Z","shell.execute_reply.started":"2025-03-10T21:19:26.209388Z","shell.execute_reply":"2025-03-10T21:19:45.086910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset.df[\"preds\"] = cls\nsubmission = (\n    test_dataset.df.groupby(\"observationID\")[\"preds\"]\n    .apply(lambda x: \" \".join(map(str, sorted(set(x)))))  # Sort and remove duplicates\n    .reset_index()\n)\n\nsubmission = submission.rename(columns={\"observationID\": \"observationId\", \"preds\": \"predictions\"})\nsubmission = submission.drop_duplicates(subset=\"observationId\")\nsubmission.to_csv(\"baseline-submission-with-prototypes.csv\", index=None)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-10T21:19:45.090339Z","iopub.execute_input":"2025-03-10T21:19:45.090652Z","iopub.status.idle":"2025-03-10T21:19:45.126922Z","shell.execute_reply.started":"2025-03-10T21:19:45.090620Z","shell.execute_reply":"2025-03-10T21:19:45.126105Z"}},"outputs":[],"execution_count":null}]}