{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":33679,"databundleVersionId":3212216,"sourceType":"competition"},{"sourceId":8050280,"sourceType":"datasetVersion","datasetId":4747361}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Herbarium Dataloading","metadata":{}},{"cell_type":"code","source":"!pip install open_clip_torch","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:00.59896Z","iopub.execute_input":"2024-07-06T22:28:00.599645Z","iopub.status.idle":"2024-07-06T22:28:14.75326Z","shell.execute_reply.started":"2024-07-06T22:28:00.599613Z","shell.execute_reply":"2024-07-06T22:28:14.752297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport open_clip\n\nopen_clip.list_pretrained()","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:14.755483Z","iopub.execute_input":"2024-07-06T22:28:14.755889Z","iopub.status.idle":"2024-07-06T22:28:23.93774Z","shell.execute_reply.started":"2024-07-06T22:28:14.755833Z","shell.execute_reply":"2024-07-06T22:28:23.936786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model, _, preprocess = open_clip.create_model_and_transforms('ViT-B-32', \n                                                             pretrained='laion400m_e32',\n                                                             precision=\"fp16\")","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:23.93934Z","iopub.execute_input":"2024-07-06T22:28:23.939731Z","iopub.status.idle":"2024-07-06T22:28:30.705947Z","shell.execute_reply.started":"2024-07-06T22:28:23.939696Z","shell.execute_reply":"2024-07-06T22:28:30.705149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\ncontext_length = model.context_length\nvocab_size = model.vocab_size\n\nprint(\"Model parameters:\", f\"{np.sum([int(np.prod(p.shape)) for p in model.parameters()]):,}\")\nprint(\"Context length:\", context_length)\nprint(\"Vocab size:\", vocab_size)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:30.708422Z","iopub.execute_input":"2024-07-06T22:28:30.709186Z","iopub.status.idle":"2024-07-06T22:28:30.723387Z","shell.execute_reply.started":"2024-07-06T22:28:30.709152Z","shell.execute_reply":"2024-07-06T22:28:30.722504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocess","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:30.724803Z","iopub.execute_input":"2024-07-06T22:28:30.725041Z","iopub.status.idle":"2024-07-06T22:28:30.734647Z","shell.execute_reply.started":"2024-07-06T22:28:30.72502Z","shell.execute_reply":"2024-07-06T22:28:30.733821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from open_clip import tokenizer\ntokenizer.tokenize(\"Hello World!\")","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:30.736015Z","iopub.execute_input":"2024-07-06T22:28:30.736352Z","iopub.status.idle":"2024-07-06T22:28:30.75744Z","shell.execute_reply.started":"2024-07-06T22:28:30.736321Z","shell.execute_reply":"2024-07-06T22:28:30.75666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from functools import partial\nfrom itertools import islice\nfrom typing import Callable, List, Optional, Sequence, Union\n\nimport torch\nimport torch.nn.functional as F\nfrom torchvision import transforms\nimport tqdm\nfrom skimage import io\n\n\nfrom open_clip import get_tokenizer","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:30.75841Z","iopub.execute_input":"2024-07-06T22:28:30.758673Z","iopub.status.idle":"2024-07-06T22:28:31.358756Z","shell.execute_reply.started":"2024-07-06T22:28:30.758651Z","shell.execute_reply":"2024-07-06T22:28:31.357988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nimport pandas as pd\nimport numpy as np\n\n\nTRAIN_DIR = \"../input/herbarium-2022-fgvc9/train_images/\"\nTEST_DIR = \"../input/herbarium-2022-fgvc9/test_images/\"\n\nwith open(\"../input/herbarium-2022-fgvc9/train_metadata.json\") as json_file:\n    train_meta = json.load(json_file)\n\nimage_ids = [image[\"image_id\"] for image in train_meta[\"images\"]]\nimage_dirs = [TRAIN_DIR + image[\"file_name\"] for image in train_meta[\"images\"]]\n\ncategory_ids = [annot[\"category_id\"] for annot in train_meta[\"annotations\"]]\ngenus_ids = [annot[\"genus_id\"] for annot in train_meta[\"annotations\"]]\n\ncategory_df = pd.DataFrame(train_meta['categories'])\ncategory_df = category_df[['category_id', 'scientificName']]\n\ncategory_dict = category_df.set_index('category_id')['scientificName'].to_dict()\nscientific_names = [category_dict[category_id] for category_id in category_ids]\n\nscientific_names_df = pd.DataFrame(scientific_names, columns=['scientificName'])\n\ntrain_df = pd.DataFrame(data=np.array([image_ids, image_dirs, genus_ids, category_ids]).T, \n                        columns=[\"image_id\", \"directory\", \"genus_id\", \"category_id\"])\n\ntrain = pd.concat([train_df, scientific_names_df], axis=1)\n\ntrain.to_csv(\"train.csv\", index = False)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:31.35985Z","iopub.execute_input":"2024-07-06T22:28:31.3603Z","iopub.status.idle":"2024-07-06T22:28:56.269123Z","shell.execute_reply.started":"2024-07-06T22:28:31.360275Z","shell.execute_reply":"2024-07-06T22:28:56.268285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"categories = pd.read_csv('train.csv')\nprint(f\"Length of Scientific Names : {categories['scientificName'].nunique()}\")\ncategories.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:56.270238Z","iopub.execute_input":"2024-07-06T22:28:56.270549Z","iopub.status.idle":"2024-07-06T22:28:58.325507Z","shell.execute_reply.started":"2024-07-06T22:28:56.270506Z","shell.execute_reply":"2024-07-06T22:28:58.324456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Version_2.1\n# - Converted template into a sequence.\nfrom typing import Callable, List, Optional, Sequence, Union\n\n# Helper function for code tracking\ndef is_sequence(data):\n    if isinstance(data, Sequence):\n        return \"Given data is a sequence.\"\n    else:\n        return \"Given data is not a sequence.\"\n\n    \ntemplate = [lambda c: f\"This is a photo of '{c}'.\"]\n\n\n# Helper function for code tracking\ndef is_sequence_of_callables_or_strings(variable):\n    if not isinstance(variable, Sequence):\n        return False\n    for item in variable:\n        if not (isinstance(item, str) or callable(item)):\n            return False\n    return True\n\nis_sequence_of_callables_or_strings(template)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.330418Z","iopub.execute_input":"2024-07-06T22:28:58.330739Z","iopub.status.idle":"2024-07-06T22:28:58.340735Z","shell.execute_reply.started":"2024-07-06T22:28:58.330713Z","shell.execute_reply":"2024-07-06T22:28:58.339724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Version_2\n# - Takes the classnames from the dataframe and merge the list of labels with template.\n# print(f\"Size of the dataframe: {categories.shape}\")\n# texts = []\n# labels = categories['scientificName'].tolist()\n# for c in labels:\n#     texts.append(template(c))\n    \n# print(f\"Length of the labels: {len(texts)}\")\n\n# Version_3\n# - Convert the labels into a sequence called classnames.\nprint(f\"Size of the dataframe: {categories.shape}\")\n\nclassnames =  categories['scientificName'].unique().tolist()\n\nprint(f\"Number of classnames {len(classnames)}\")\nprint(f\"Is classnames variable a sequence of strings? '{is_sequence(classnames)}'\")\nprint(f\"Data Type of classnames: '{type(classnames)}'\")\nprint(f\"Classnames: {classnames[4:32]}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.342078Z","iopub.execute_input":"2024-07-06T22:28:58.34273Z","iopub.status.idle":"2024-07-06T22:28:58.435582Z","shell.execute_reply.started":"2024-07-06T22:28:58.342696Z","shell.execute_reply":"2024-07-06T22:28:58.434596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas\nimport torch\nimport torchvision\nimport torchvision.datasets as datasets\nfrom torch.utils.data import Dataset, DataLoader\nfrom open_clip import tokenizer\n\nclass HerbariumDataset(Dataset):\n    def __init__(self,data,transform=None):\n        self.annotations = data\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.annotations)\n    \n    def __getitem__(self,index):\n        \n        image_path = self.annotations.iloc[index,1]\n        image = self.transform(Image.open(image_path))\n        y_label = torch.tensor(int(self.annotations.iloc[index,3]))\n        \n        return image,y_label ","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.437008Z","iopub.execute_input":"2024-07-06T22:28:58.437713Z","iopub.status.idle":"2024-07-06T22:28:58.444887Z","shell.execute_reply.started":"2024-07-06T22:28:58.437677Z","shell.execute_reply":"2024-07-06T22:28:58.443855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Version_3\ndef Transform(phase: str):\n    img_size = 224\n    if phase == 'train':\n        return transforms.Compose([\n            transforms.RandomResizedCrop(size=img_size),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomApply([transforms.RandomRotation(degrees=(0, 30))], p=0.5),\n            transforms.ColorJitter(brightness=0.5, contrast=0.5),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  # Example values, modify as neede\n        \n        ])\n    else:\n        return transforms.Compose([\n            transforms.Resize((img_size, img_size)),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  # Example values, modify as needed\n        ])\n\n# Version_2\n\n# def Transform(phase: str):\n#     if phase == 'train':\n#         return A.Compose([\n#             A.RandomResizedCrop(height=224, width=224),\n#             A.HorizontalFlip(p=0.5),\n#             A.ShiftScaleRotate(p=0.5),\n#             A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], max_pixel_value=255.0, p=1.0),\n#             ToTensorV2(),\n#         ])\n#     else:\n#         return A.Compose([\n#             A.Resize(height=224, width=224),\n#             A.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], max_pixel_value=255.0, p=1.0),\n#             ToTensorV2(),\n#         ])","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.445962Z","iopub.execute_input":"2024-07-06T22:28:58.446231Z","iopub.status.idle":"2024-07-06T22:28:58.458412Z","shell.execute_reply.started":"2024-07-06T22:28:58.446209Z","shell.execute_reply":"2024-07-06T22:28:58.457583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.459554Z","iopub.execute_input":"2024-07-06T22:28:58.459884Z","iopub.status.idle":"2024-07-06T22:28:58.470639Z","shell.execute_reply.started":"2024-07-06T22:28:58.459855Z","shell.execute_reply":"2024-07-06T22:28:58.469851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train_dataset = HerbariumDataset(categories,transform=Transform('train'))\nvalid_dataset = HerbariumDataset(categories,preprocess)\n\n#train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=2, pin_memory=True, drop_last=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=batch_size, shuffle=False, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.471821Z","iopub.execute_input":"2024-07-06T22:28:58.472157Z","iopub.status.idle":"2024-07-06T22:28:58.47965Z","shell.execute_reply.started":"2024-07-06T22:28:58.472124Z","shell.execute_reply":"2024-07-06T22:28:58.478893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.480587Z","iopub.execute_input":"2024-07-06T22:28:58.480865Z","iopub.status.idle":"2024-07-06T22:28:58.514085Z","shell.execute_reply.started":"2024-07-06T22:28:58.480842Z","shell.execute_reply":"2024-07-06T22:28:58.513163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Checking the dataset and the dataloader for proper data processing.","metadata":{}},{"cell_type":"code","source":"# Helper functions loaded using ChatGPT.\n\nimport random\nimport numpy as np\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport pandas as pd\n\ndef view_images_from_dataset(dataset: Dataset, seed: int = None, num_samples: int = 5, return_image=False):\n    \"\"\"\n    View images from a Dataset using PIL and check features and number of columns in the dataset.\n\n    Args:\n        dataset (Dataset): The PyTorch dataset.\n        dataframe (pd.DataFrame): The dataframe associated with the dataset.\n        seed (int, optional): The random seed for reproducibility. Defaults to None.\n        num_samples (int, optional): Number of sample images to view. Defaults to 5.\n\n    Returns:\n        None\n    \"\"\"\n    if seed is not None:\n        random.seed(seed)\n        np.random.seed(seed)\n        torch.manual_seed(seed)\n\n    indices = list(range(len(dataset)))\n    random.shuffle(indices)\n    sample_indices = indices[:num_samples]\n    images, labels = [], []\n    for idx in sample_indices:\n        image, label = dataset[idx]\n        print(isinstance(image, torch.Tensor))\n        print(isinstance(label, torch.Tensor))\n        print(f\"Shape of Image: {image.shape}\")\n#         print(f\"value of Image: {image}\")\n        print(f\"Shape of Text: {label.shape}\")\n        if return_image:\n            images.append(image)\n            labels.append(label)\n        else:\n            if isinstance(image, torch.Tensor):\n                image = image.permute(1, 2, 0).numpy()  # Convert from CxHxW to HxWxC for displaying\n            image = Image.fromarray((image * 255).astype(np.uint8))\n            plt.imshow(image)\n            plt.title(f\"Label: {label}\")\n            plt.show()\n    \n    if return_image:\n            return images, labels\n    \n\n# Example usage\n# view_images_from_dataset(my_dataset, my_dataframe, seed=42, num_samples=5)\n\n\ndef view_image_from_dataloader(dataloader: DataLoader, seed: int = None):\n    \"\"\"\n    Load and view an image using next(iter()) from DataLoader.\n\n    Args:\n        dataloader (DataLoader): The PyTorch DataLoader.\n        seed (int, optional): The random seed for reproducibility. Defaults to None.\n\n    Returns:\n        None\n    \"\"\"\n    if seed is not None:\n        random.seed(seed)\n        np.random.seed(seed)\n        torch.manual_seed(seed)\n    \n    data_iter = iter(dataloader)\n    images, labels = next(data_iter)\n    \n    if isinstance(images, torch.Tensor):\n        image = images[0].permute(1, 2, 0).numpy()  # Convert from CxHxW to HxWxC for displaying\n    else:\n        image = images[0]\n    \n    image = Image.fromarray((image * 255).astype(np.uint8))\n    plt.imshow(image)\n    plt.title(f\"Label: {labels[0]}\")\n    plt.show()\n\n# Example usage\n# data_loader = DataLoader(my_dataset, batch_size=4, shuffle=True)\n# view_image_from_dataloader(data_loader, seed=42)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.515414Z","iopub.execute_input":"2024-07-06T22:28:58.515774Z","iopub.status.idle":"2024-07-06T22:28:58.530745Z","shell.execute_reply.started":"2024-07-06T22:28:58.515748Z","shell.execute_reply":"2024-07-06T22:28:58.529891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"view_images_from_dataset(valid_dataset, seed=42, num_samples=5)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:28:58.532055Z","iopub.execute_input":"2024-07-06T22:28:58.532339Z","iopub.status.idle":"2024-07-06T22:29:01.447251Z","shell.execute_reply.started":"2024-07-06T22:28:58.532309Z","shell.execute_reply":"2024-07-06T22:29:01.446263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"view_image_from_dataloader(valid_loader, seed=42)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:01.44875Z","iopub.execute_input":"2024-07-06T22:29:01.449386Z","iopub.status.idle":"2024-07-06T22:29:03.770223Z","shell.execute_reply.started":"2024-07-06T22:29:01.44935Z","shell.execute_reply":"2024-07-06T22:29:03.769305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch in valid_loader:\n    #print(batch)\n    images, targets = batch\n    print(f\"Batch Dimension: {images.shape}\")\n    print(f\"Target Dimension: {targets.shape}\")\n    print(f\"Target: {targets[3]}\")\n    break","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:03.771575Z","iopub.execute_input":"2024-07-06T22:29:03.771878Z","iopub.status.idle":"2024-07-06T22:29:04.953512Z","shell.execute_reply.started":"2024-07-06T22:29:03.771848Z","shell.execute_reply":"2024-07-06T22:29:04.952323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:04.955261Z","iopub.execute_input":"2024-07-06T22:29:04.956236Z","iopub.status.idle":"2024-07-06T22:29:05.211872Z","shell.execute_reply.started":"2024-07-06T22:29:04.956191Z","shell.execute_reply":"2024-07-06T22:29:05.211054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Version 3\n# Sticks to the original build_zero_shot_classifier()\n# This time, rather than transforming the function to suit data, we will process the\n# data to fit the functions.\n\n# NOTE TO SELF: monitor the variables\nimport tqdm\nfrom open_clip import tokenizer\n\n# Helper functions for text processing\ndef process_texts(texts, device):\n    inputs = tokenizer.tokenize(texts).to(device)\n    return inputs\n\n# def encode_texts(model, inputs, normalize=True):\n#     with torch.no_grad():\n#         text_features = model.encode_text(inputs, normalize=False)\n#         if normalize:\n#             text_features = text_features / text_features.norm(dim=-1, keepdim=True)\n#     return text_features\n\ndef encode_texts(model, texts, normalize=True):\n    with torch.no_grad():\n        inputs = tokenizer.tokenize(texts).to(device)\n        text_features = model.encode_text(inputs).float()\n        if normalize==True:\n            text_features /= text_features.norm(dim=-1, keepdim=True)\n    \n    return text_features\n\ndef batched(iterable, n):\n    \"\"\"Batch data into lists of length *n*. The last batch may be shorter.\n    NOTE based on more-itertools impl, to be replaced by python 3.12 itertools.batched impl\n    \"\"\"\n    it = iter(iterable)\n    while True:\n        batch = list(islice(it, n))\n        if not batch:\n            break\n        yield batch\n\n\ndef build_zero_shot_classifier(\n        model,\n        tokenizer,\n#         processor,\n        classnames: Sequence[str],\n        templates: Sequence[Union[Callable, str]],\n        num_classes_per_batch: Optional[int] = 10,\n        device: Union[str, torch.device] = 'cpu',\n        use_tqdm: bool = False,\n):\n    \"\"\" Build zero-shot classifier weights by iterating over class names in batches\n    Args:\n        model: CLIP model instance\n        tokenizer: CLIP tokenizer instance\n        classnames: A sequence of class (label) names\n        templates: A sequence of callables or format() friendly strings to produce templates per class name\n        num_classes_per_batch: The number of classes to batch together in each forward, all if None\n        device: Device to use.\n        use_tqdm: Enable TQDM progress bar.\n    \"\"\"\n    assert isinstance(templates, Sequence) and len(templates) > 0\n    assert isinstance(classnames, Sequence) and len(classnames) > 0\n    use_format = isinstance(templates[0], str)\n    num_templates = len(templates)\n    num_classes = len(classnames)\n    if use_tqdm:\n        num_iter = 1 if num_classes_per_batch is None else ((num_classes - 1) // num_classes_per_batch + 1)\n        iter_wrap = partial(tqdm.tqdm, total=num_iter, unit_scale=num_classes_per_batch)\n    else:\n        iter_wrap = iter\n\n    def _process_batch(batch_classnames):\n        num_batch_classes = len(batch_classnames)\n        texts = [template.format(c) if use_format else template(c) for c in batch_classnames for template in templates]\n        texts = tokenizer.tokenize(texts).to(device)\n        class_embeddings = model.encode_text(texts)\n        class_embeddings = class_embeddings.reshape(num_batch_classes, num_templates, -1).mean(dim=1)\n        class_embeddings = class_embeddings / class_embeddings.norm(dim=1, keepdim=True)\n        class_embeddings = class_embeddings.T\n        return class_embeddings\n\n    with torch.no_grad():\n        if num_classes_per_batch:\n            batched_embeds = [_process_batch(batch) for batch in iter_wrap(batched(classnames, num_classes_per_batch))]\n            zeroshot_weights = torch.cat(batched_embeds, dim=1)\n        else:\n            zeroshot_weights = _process_batch(classnames)\n    return zeroshot_weights\n\n\n# Version_2\n# Had modifications made to extract the class embeddings of all the classes all together.\n# Did not address the zeroshot weights.\n# def batched(iterable, n):\n#     it = iter(iterable)\n#     while True:\n#         batch = list(islice(it, n))\n#         if not batch:\n#             break\n#         yield batch\n\n# def build_zero_shot_classifier(\n#     model,\n#     texts: list,\n#     num_classes: str,\n#     num_classes_per_batch: Optional[int] = 100, #NOTE\n#     device: Union[str, torch.device] = \"cpu\",\n# ):\n#     num_iter = 1 if num_classes_per_batch is None else ((num_classes - 1) // num_classes_per_batch + 1)\n#     iter_wrap = partial(tqdm.tqdm, total=num_iter, unit_scale=num_classes_per_batch)\n    \n#     def _process_batch(batch_classnames):\n#         num_batch_classes = len(batch_classnames)\n#         texts = clip.tokenize(texts).to(device)\n#         class_embeddings = model.encode_text(texts, normalize=True)\n#         class_embeddings = class_embeddings.reshape(num_batch_classes, num_templates, -1).mean(dim=1)\n#         class_embeddings = class_embeddings.T\n#         print(class_embeddings)\n#         return class_embeddings","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:05.213253Z","iopub.execute_input":"2024-07-06T22:29:05.213842Z","iopub.status.idle":"2024-07-06T22:29:05.234043Z","shell.execute_reply.started":"2024-07-06T22:29:05.213808Z","shell.execute_reply":"2024-07-06T22:29:05.233037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"texts = [\"This is a photo of 'Abies amabilis (Douglas ex Loudon) J.Forbes'.\", \"This is a photo of 'Abies balsamea (L.) Mill.'.\", \"This is a photo of 'Abies bracteata (D.Don) Poit.'.\", \"This is a photo of 'Abies concolor (Gordon & Glend.) Lindl. ex Hildebr.'.\", \"This is a photo of 'Abies fraseri (Pursh) Poir.'.\", \"This is a photo of 'Abies grandis (Douglas ex D.Don) Lindl.'.\", \"This is a photo of 'Abies lasiocarpa (Hook.) Nutt.'.\", \"This is a photo of 'Abies magnifica A.Murray bis'.\", \"This is a photo of 'Abies procera Rehder'.\", \"This is a photo of 'Abronia ameliae Lundell'.\"]    \nprint(f\"Length of the labels: {len(texts)}\")\ntokenized_inputs = process_texts(texts, device).to(device)\nprint(type(tokenized_inputs))\nclass_embeddings = encode_texts(model, texts, normalize=True)\nprint(type(class_embeddings))\nprint(f\"Shape of class_embeddings: {class_embeddings.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:05.235387Z","iopub.execute_input":"2024-07-06T22:29:05.235736Z","iopub.status.idle":"2024-07-06T22:29:05.643171Z","shell.execute_reply.started":"2024-07-06T22:29:05.235712Z","shell.execute_reply":"2024-07-06T22:29:05.642268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, label = view_images_from_dataset(valid_dataset, seed=42, num_samples=1, return_image=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:05.644398Z","iopub.execute_input":"2024-07-06T22:29:05.64473Z","iopub.status.idle":"2024-07-06T22:29:06.50017Z","shell.execute_reply.started":"2024-07-06T22:29:05.644702Z","shell.execute_reply":"2024-07-06T22:29:06.499271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:06.501212Z","iopub.execute_input":"2024-07-06T22:29:06.50147Z","iopub.status.idle":"2024-07-06T22:29:06.50839Z","shell.execute_reply.started":"2024-07-06T22:29:06.501448Z","shell.execute_reply":"2024-07-06T22:29:06.507463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(class_embeddings[4:5])","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:06.509506Z","iopub.execute_input":"2024-07-06T22:29:06.509828Z","iopub.status.idle":"2024-07-06T22:29:06.693426Z","shell.execute_reply.started":"2024-07-06T22:29:06.509795Z","shell.execute_reply":"2024-07-06T22:29:06.692377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from open_clip import tokenizer as t\nmodel = model.to(device)\nzeroshot_weights = build_zero_shot_classifier(model=model,tokenizer=t, classnames=classnames, templates=template, device=device, use_tqdm=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:06.694563Z","iopub.execute_input":"2024-07-06T22:29:06.694896Z","iopub.status.idle":"2024-07-06T22:29:31.170277Z","shell.execute_reply.started":"2024-07-06T22:29:06.69487Z","shell.execute_reply":"2024-07-06T22:29:31.16909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(type(zeroshot_weights))\nprint(zeroshot_weights.shape)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:31.17497Z","iopub.execute_input":"2024-07-06T22:29:31.175241Z","iopub.status.idle":"2024-07-06T22:29:31.180026Z","shell.execute_reply.started":"2024-07-06T22:29:31.175218Z","shell.execute_reply":"2024-07-06T22:29:31.179122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Version_3\nimport torch\nfrom contextlib import suppress\nfrom open_clip import tokenizer\n\ndef get_autocast(precision):\n    if precision == 'amp':\n        return torch.cuda.amp.autocast\n    elif precision == 'amp_bfloat16' or precision == 'amp_bf16':\n        # amp_bfloat16 is more stable than amp float16 for clip training\n        return lambda: torch.cuda.amp.autocast(dtype=torch.bfloat16)\n    else:\n        return suppress\n    \ndef get_input_dtype(precision: str):\n    input_dtype = None\n    if precision in ('bf16', 'pure_bf16'):\n        input_dtype = torch.bfloat16\n    elif precision in ('fp16', 'pure_fp16'):\n        input_dtype = torch.float16\n    return input_dtype\n\ndef accuracy(output, target, topk=(1,)):\n#     print(f\"Shape of target: {target.shape}\")\n#     print(f\"Shape of pred before transpose: {output.topk(max(topk), 1, True, True)[1].shape}\")\n    pred = output.topk(max(topk), 1, True, True)[1].t()\n#     print(f\"Shape of pred: {pred.shape}\")\n#     print(f\"Target Processed: {target.view(1, -1).shape}\")\n    correct = pred.eq(target.view(1, -1).expand_as(pred))\n    return [float(correct[:k].reshape(-1).float().sum(0, keepdim=True).cpu().numpy()) for k in topk]\n\n\ndef run(model, classifier, dataloader):\n    autocast = get_autocast(\"amp\")\n    input_dtype = get_input_dtype(\"fp16\")\n\n    with torch.inference_mode():\n        top1, top5, n = 0., 0., 0.\n        for batch in tqdm.tqdm(dataloader, unit_scale=batch_size): # unit_scale = args.batch_size\n            images, targets = batch\n            images = images.to(device=device, dtype=input_dtype)\n            #print(images)\n#             print(images.shape)\n            targets = targets.to(device)\n#             print(targets)\n#             print(f\"Shape of Logits: {targets.shape}\")\n            #print(f\"\\nShape of the Images: {images.shape}\\n\")\n            \n            # Version_3\n            with autocast():\n                # predict\n                output = model(image=images)\n#                 print(f\"Shape of Image Features: {output[0].shape}\")\n                image_features = output['image_features'] if isinstance(output, dict) else output[0]\n                logits = 100. * image_features @ classifier\n#                 print(f\"Shape of Logits: {logits.shape}\")\n           \n            # measure accuracy\n            acc1, acc5 = accuracy(logits, targets, topk=(1,5))\n            top1 += acc1\n            top5 += acc5\n            n += images.size(0)\n            \n#             Version_2.1\n#             for i in range(batch_size):\n#                 image = images[i]\n#                 image = image.unsqueeze(0).to(device, dtype=torch.bfloat16)\n#                 target = targets[i].unsqueeze(0).to(device)\n#                 output = model(image=images)\n#                 image_features = output['image_features'] if isinstance(output, dict) else output[0]\n#                 logits = 100. * image_features @ classifier\n\n#                 # measure accuracy\n#                 acc1, acc5 = accuracy(logits, target, topk=(1, 5))\n#                 top1 += acc1\n#                 top5 += acc5\n#                 n += images.size(0)\n\n    top1 = (top1 / n)\n    top5 = (top5 / n)\n    return top1, top5\n\n\n# Version_2\n# def accuracy(output, target, topk=(1,)):\n#     pred = output.topk(max(topk), 1, True, True)[1].t()\n#     correct = pred.eq(target.view(1, -1).expand_as(pred))\n#     return [float(correct[:k].reshape(-1).float().sum(0, keepdim=True).cpu().numpy()) for k in topk]\n# def run(model, classifier, dataloader):\n#     with torch.no_grad():\n#         top1, top5, n = 0., 0., 0.\n#         for batch in tqdm.tqdm(dataloader, unit_scale=batch_size): # unit_scale = args.batch_size\n#             images, targets = batch\n#             print(f\"Image shape before loading into model: {images.shape}\")\n#             print(f\"Image shape after dimension reduction: {images.shape}\")\n#             images = images.to(device=device, dtype=torch.bfloat16)\n#             target = targets.to(device)\n#             model = model.vision_model\n#             model.to(device)\n#             output = model(pixel_values=images)\n#             image_features = output['last_hidden_state'] if isinstance(output, dict) else output[0]\n#             print(type(classifier))\n#             return 0\n#             logits = 100. * image_features @ classifier\n#             # measure accuracy\n#             acc1, acc5 = accuracy(logits, target, topk=(1, 5))\n#             top1 += acc1\n#             top5 += acc5\n#             n += images.size(0)\n#     top1 = (top1 / n)\n#     top5 = (top5 / n)\n#     return top1, top5\n# model_id = \"openai/clip-vit-base-patch32\"\n# processor = CLIPProcessor.from_pretrained(model_id)\n# model = CLIPModel.from_pretrained(model_id)\n# classifier = build_zero_shot_classifier(\n#     model,\n#     texts,\n#     num_classes=15501,\n#     device= \"cuda\",\n# )\n# results = {}\n# # top1, top5 = run(model, classifier, data['imagenet-v2'].dataloader)\n# top1, top5 = run(model, classifier, valid_loader)\n# results['top1'] = top1\n# results['top5'] = top5","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:31.181297Z","iopub.execute_input":"2024-07-06T22:29:31.181591Z","iopub.status.idle":"2024-07-06T22:29:31.208376Z","shell.execute_reply.started":"2024-07-06T22:29:31.181562Z","shell.execute_reply":"2024-07-06T22:29:31.207584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classifier = zeroshot_weights\nresults = {}\ntop1, top5 = run(model, classifier, valid_loader)\nresults['top1'] = top1\nresults['top5'] = top5\n\n# Specify the file path and name\nfile_path = '/kaggle/working/data.json'\n\n# Save the dictionary to a JSON file\nwith open(file_path, 'w') as json_file:\n    json.dump(results, json_file, indent=4)\n\nprint(f\"Dictionary saved as {file_path}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-06T22:29:31.209309Z","iopub.execute_input":"2024-07-06T22:29:31.209648Z","iopub.status.idle":"2024-07-07T01:21:58.631415Z","shell.execute_reply.started":"2024-07-06T22:29:31.209625Z","shell.execute_reply":"2024-07-07T01:21:58.630045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aaaaaaaaaadssstfdghdfkkkggfdhhkm","metadata":{},"execution_count":null,"outputs":[]}]}