{"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":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":1669366,"sourceType":"datasetVersion","datasetId":988673},{"sourceId":1685182,"sourceType":"datasetVersion","datasetId":998578},{"sourceId":8047253,"sourceType":"datasetVersion","datasetId":4745259},{"sourceId":8162250,"sourceType":"datasetVersion","datasetId":4829271}],"dockerImageVersionId":30646,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q transformers pytorch-lightning","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys; \npackage_paths = [\n    '../input/image-fmix/FMix-master'\n]\nfor pth in package_paths:\n    sys.path.append(pth)\nfrom fmix import sample_mask, make_low_freq_image, binarise_mask\nimport os\nimport json\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\nfrom PIL import Image\nimport torch \nimport torch.nn as nn\nfrom torch.nn import CrossEntropyLoss\nfrom torch.utils.data import Dataset, DataLoader, SubsetRandomSampler, WeightedRandomSampler, RandomSampler\nimport torchvision.transforms as transforms\nfrom torchvision.transforms import Compose,ToTensor, Resize\nfrom sklearn.model_selection import train_test_split\nfrom transformers import ViTImageProcessor\nimport pytorch_lightning as pl\nfrom transformers import ViTForImageClassification, AdamW\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.callbacks import EarlyStopping,ModelCheckpoint\nfrom imblearn.under_sampling import RandomUnderSampler\nfrom imblearn.over_sampling import SMOTE\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.utils.class_weight import compute_class_weight\nimport seaborn as sns\nimport cv2\nimport albumentations as albu","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-24T08:22:45.390684Z","iopub.execute_input":"2024-04-24T08:22:45.391081Z","iopub.status.idle":"2024-04-24T08:23:17.383064Z","shell.execute_reply.started":"2024-04-24T08:22:45.391052Z","shell.execute_reply":"2024-04-24T08:23:17.381889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = \"../input/cassava-leaf-disease-classification\"\nTRAIN_PATH = \"../input/cassava-leaf-disease-classification/train_images/\"\nTEST_PATH = \"../input/cassava-leaf-disease-classification/test_images/\"\n# MODEL_PATH = \"/kaggle/input/vit-base-patch16-224-in21k/vit-base-patch16-224-in21k\"\nMODEL_PATH = \"/kaggle/input/vit-base-patch16-384-v1/vit-base-patch16-384\"\n\n# model specific global variables\nIMG_SIZE = 384\nBATCH_SIZE = 16\nLR = 1e-4\nGAMMA = 0.7\nMAX_EPOCHS = 10\nCHECKPOINT_PATH = \"\"","metadata":{"execution":{"iopub.status.busy":"2024-04-24T08:26:22.140476Z","iopub.execute_input":"2024-04-24T08:26:22.141477Z","iopub.status.idle":"2024-04-24T08:26:22.148529Z","shell.execute_reply.started":"2024-04-24T08:26:22.141435Z","shell.execute_reply":"2024-04-24T08:26:22.147320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 1) EDA","metadata":{}},{"cell_type":"code","source":"with open(os.path.join(DATA_PATH, \"label_num_to_disease_map.json\")) as file:\n    map_classes = json.loads(file.read())\nlabels_dict = {}\nlabels_dict[0] = map_classes['0']\nlabels_dict[1] = map_classes['1']\nlabels_dict[2] = map_classes['2']\nlabels_dict[3] = map_classes['3']\nlabels_dict[4] = map_classes['4']\nprint(labels_dict)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T08:26:23.636011Z","iopub.execute_input":"2024-04-24T08:26:23.636429Z","iopub.status.idle":"2024-04-24T08:26:23.650937Z","shell.execute_reply.started":"2024-04-24T08:26:23.636398Z","shell.execute_reply":"2024-04-24T08:26:23.649601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_map = {\n 'Cassave Bacterial Blight (CBB)': 0,\n 'Cassava Brown Streak Disease': 1,\n 'Cassava Green Mottle': 2,\n 'Cassava Mosaic Disease (CMD)': 3,\n 'Healthy': 5,\n}\n\nid_to_class = {\n     0: \"Cassava Bacterial Blight (CBB)\",\n     1: \"Cassava Brown Streak Disease (CBSD)\",\n     2: \"Cassava Green Mottle (CGM)\",\n     3: \"Cassava Mosaic Disease (CMD)\",\n     4: \"Healthy\"\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-24T08:26:26.277820Z","iopub.execute_input":"2024-04-24T08:26:26.278626Z","iopub.status.idle":"2024-04-24T08:26:26.284990Z","shell.execute_reply.started":"2024-04-24T08:26:26.278590Z","shell.execute_reply":"2024-04-24T08:26:26.283526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(DATA_PATH, \"train.csv\"))\ndf_temp =  pd.read_csv(os.path.join(DATA_PATH, \"train.csv\"))\ndf_temp[\"class_name\"] = df_temp[\"label\"].astype(str).map(map_classes)\n\ndf_temp","metadata":{"execution":{"iopub.status.busy":"2024-04-24T08:26:27.433489Z","iopub.execute_input":"2024-04-24T08:26:27.433965Z","iopub.status.idle":"2024-04-24T08:26:27.534214Z","shell.execute_reply.started":"2024-04-24T08:26:27.433927Z","shell.execute_reply":"2024-04-24T08:26:27.532870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tabulate import tabulate\n#number of pictures for each disease\nlabel_counts = df_temp.groupby('class_name')['image_id'].count()\ncounts_df = label_counts.reset_index()\ncounts_df.columns = [\"Category\", \"No. of Plants\"]\ncounts_df = counts_df.sort_values(by=\"No. of Plants\", ascending=False)\ncounts_df = counts_df.reset_index(drop=True)\nprint(tabulate(counts_df, headers='keys', tablefmt='pretty'))","metadata":{"execution":{"iopub.status.busy":"2024-04-24T08:38:08.029969Z","iopub.execute_input":"2024-04-24T08:38:08.030445Z","iopub.status.idle":"2024-04-24T08:38:08.048857Z","shell.execute_reply.started":"2024-04-24T08:38:08.030376Z","shell.execute_reply":"2024-04-24T08:38:08.047617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Image Exploration","metadata":{}},{"cell_type":"code","source":"def plot_image_examples(class_label_num, n_rows=3, n_cols=2):\n    \"\"\"Plots a grid with examples of images for the specified class.\"\"\"\n\n    # Gets the label of the class.\n    class_label = labels_dict[class_label_num]\n\n    # Filter the images by label.\n    label_df = df_temp[df_temp.label == class_label_num]\n\n    # Random indices to plot.\n    rand_idx = np.random.randint(0, len(label_df), n_rows * n_cols)\n\n    fig, axs = plt.subplots(n_rows, n_cols, figsize=(10, 10))\n\n    for row in range(n_rows):\n        for col in range(n_cols):\n            idx = rand_idx[row * n_cols + col]\n            img_path = os.path.join(TRAIN_PATH, label_df.image_id.values[idx])\n            img = Image.open(img_path)\n            axs[row, col].imshow(img)\n            axs[row, col].axis('off')\n\n            # Get image size\n            img_width, img_height = img.size\n            img_size_text = f'Size: {img_width}x{img_height}'\n\n            # Add text annotation for image size\n            axs[row, col].text(0, 0, img_size_text, color='white', backgroundcolor='black', fontsize=8,\n                               verticalalignment='top')\n\n            axs[row, col].set_title(label_df.class_name.values[idx])\n\n    plt.suptitle(class_label)\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-04-24T07:46:05.949290Z","iopub.execute_input":"2024-04-24T07:46:05.949756Z","iopub.status.idle":"2024-04-24T07:46:05.963158Z","shell.execute_reply.started":"2024-04-24T07:46:05.949719Z","shell.execute_reply":"2024-04-24T07:46:05.961773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 2) Data Preprocessing","metadata":{}},{"cell_type":"code","source":"feature_extractor = ViTImageProcessor.from_pretrained(MODEL_PATH)\nprint(feature_extractor)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:20.055771Z","iopub.execute_input":"2024-04-22T10:07:20.056522Z","iopub.status.idle":"2024-04-22T10:07:20.065664Z","shell.execute_reply.started":"2024-04-22T10:07:20.056493Z","shell.execute_reply":"2024-04-22T10:07:20.064619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = albu.Compose([\n    albu.RandomResizedCrop(height=IMG_SIZE, width=IMG_SIZE),\n    albu.Transpose(p=0.5),\n    albu.HorizontalFlip(p=0.5),\n    albu.VerticalFlip(p=0.5),\n    albu.HueSaturationValue(hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.5),\n    albu.RandomBrightnessContrast(p=0.5),\n    albu.ShiftScaleRotate(p=0.5),\n    albu.CoarseDropout(p=0.5),\n    albu.Cutout(p=0.5),\n    albu.Normalize(    \n        mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n    ToTensorV2(p=1.0),\n], p=1.)\n\nvalid_transform = albu.Compose([\n    albu.CenterCrop(height=IMG_SIZE, width=IMG_SIZE, p=1.),\n    albu.Resize(height=IMG_SIZE, width=IMG_SIZE),\n    albu.Normalize(\n        mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n    ToTensorV2(p=1.0),\n], p=1.)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:22.908431Z","iopub.execute_input":"2024-04-22T10:07:22.908796Z","iopub.status.idle":"2024-04-22T10:07:22.920368Z","shell.execute_reply.started":"2024-04-22T10:07:22.908767Z","shell.execute_reply":"2024-04-22T10:07:22.919397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_augmentation(image_id,transform):\n    plt.figure(figsize=(16,4))\n    \n    img=cv2.imread(os.path.join(TRAIN_IMAGES_DIR,image_id))\n    img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)\n    \n    plt.subplot(1,3,1)\n    plt.imshow(img)\n    plt.axis(\"off\")\n    \n    plt.subplot(1,3,2)\n    x=transform(image=img)[\"image\"]\n    plt.imshow(x)\n    plt.axis(\"off\")\n    \n    plt.subplot(1,3,3)\n    x=transform(image=img)[\"image\"]\n    plt.imshow(x)\n    plt.axis(\"off\")\n    \n    plt.show()\n    \ndef visualize(images, transform):\n    \"\"\"\n    Plot images and their transformations\n    \"\"\"\n    fig = plt.figure(figsize=(32, 16))\n    \n    for i, im in enumerate(images):\n        ax = fig.add_subplot(2, 5, i + 1, xticks=[], yticks=[])\n        plt.imshow(im)\n        \n    for i, im in enumerate(images):\n        ax = fig.add_subplot(2, 5, i + 6, xticks=[], yticks=[])\n        plt.imshow(transform(image=im)['image'])","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:24.502172Z","iopub.execute_input":"2024-04-22T10:07:24.502563Z","iopub.status.idle":"2024-04-22T10:07:24.512406Z","shell.execute_reply.started":"2024-04-22T10:07:24.502521Z","shell.execute_reply":"2024-04-22T10:07:24.511405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rand_bbox(size, lam):\n    W = size[0]\n    H = size[1]\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = np.int(W * cut_rat)\n    cut_h = np.int(H * cut_rat)\n\n    # uniform\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    bbx1 = np.clip(cx - cut_w // 2, 0, W)\n    bby1 = np.clip(cy - cut_h // 2, 0, H)\n    bbx2 = np.clip(cx + cut_w // 2, 0, W)\n    bby2 = np.clip(cy + cut_h // 2, 0, H)\n    return bbx1, bby1, bbx2, bby2\n# Custom Dataset\nclass CassavaDataset(Dataset):\n    def __init__(self, df, data_path=DATA_PATH,\n                 mode=\"train\",\n                 transform=None,\n                 output_label=True, \n                 one_hot_label=False,\n                 do_fmix=False, \n                 fmix_params={\n                     'alpha': 1., \n                     'decay_power': 3., \n                     'shape': (IMG_SIZE,IMG_SIZE),\n                     'max_soft': True, \n                     'reformulate': False\n                 },\n                 do_cutmix=False,\n                 cutmix_params={\n                     'alpha': 1,\n                 }\n                ):\n        self.df = df\n       \n        self.data_path = data_path\n        self.transform = transform\n        self.mode = mode\n        self.do_fmix = do_fmix\n        self.fmix_params = fmix_params\n        self.do_cutmix = do_cutmix\n        self.cutmix_params = cutmix_params\n        \n        self.output_label = output_label\n        self.one_hot_label = one_hot_label\n        self.data_dir = \"train_images\" if mode == \"train\" else \"test_images\"\n        \n        if output_label == True:\n            self.labels = self.df['label'].values\n\n            if one_hot_label is True:\n                self.labels = np.eye(self.df['label'].max()+1)[self.labels]\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        image_name = self.df.loc[idx, 'image_id']\n        label = self.df.loc[idx, 'label']\n        image_path = os.path.join(self.data_path, self.data_dir, image_name)\n\n        img=cv2.imread(image_path,cv2.IMREAD_COLOR)\n        img=cv2.cvtColor(img,cv2.COLOR_BGR2RGB)\n\n        if self.transform is not None:\n            img = self.transform(image=img)['image']\n            \n        if self.do_fmix and np.random.uniform(0., 1., size=1)[0] > 0.5:\n            with torch.no_grad():\n                lam = np.clip(np.random.beta(self.fmix_params['alpha'], self.fmix_params['alpha']),0.6,0.7)\n                \n                # Make mask, get mean / std\n                mask = make_low_freq_image(self.fmix_params['decay_power'], self.fmix_params['shape'])\n                mask = binarise_mask(mask, lam, self.fmix_params['shape'], self.fmix_params['max_soft'])\n                fmix_ix = np.random.choice(self.df.index, size=1)[0]\n                image_id = self.df.iloc[fmix_ix]['image_id']\n                path_of_image = os.path.join(self.data_path, self.data_dir, image_id)\n                img_x=cv2.imread(path_of_image,cv2.IMREAD_COLOR)\n                img_x=cv2.cvtColor(img_x,cv2.COLOR_BGR2RGB)\n#                 fmix_img  = get_img(\"{}/{}\".format(self.data_root, self.df.iloc[fmix_ix]['image_id']))\n\n                if self.transform:\n                    fmix_img = self.transform(image=img_x)['image']\n\n                mask_torch = torch.from_numpy(mask)\n                \n                # mix image\n                img = mask_torch*img+(1.-mask_torch)*fmix_img\n\n                rate = mask.sum()/IMG_SIZE/IMG_SIZE\n                label = rate*label + (1.-rate)*self.labels[fmix_ix]\n                \n        if self.do_cutmix and np.random.uniform(0., 1., size=1)[0] > 0.5:\n            with torch.no_grad():\n                cmix_ix = np.random.choice(self.df_data.index, size=1)[0]\n                image_id_c = self.df.iloc[fmix_ix]['image_id']\n                path_of_image_c = os.path.join(self.data_path, self.data_dir, image_id_c)\n                img_c=cv2.imread(path_of_image_c,cv2.IMREAD_COLOR)\n                img_c=cv2.cvtColor(img_c,cv2.COLOR_BGR2RGB)\n#                 cmix_img  = get_img(\"{}/{}\".format(self.data_root, self.df.iloc[cmix_ix]['image_id']))\n                if self.transforms:\n                    cmix_img = self.transform(image=img_c)['image']\n                    \n                lam = np.clip(np.random.beta(self.cutmix_params['alpha'], self.cutmix_params['alpha']),0.3,0.4)\n                bbx1, bby1, bbx2, bby2 = rand_bbox((IMG_SIZE, IMG_SIZE), lam)\n\n                img[:, bbx1:bbx2, bby1:bby2] = cmix_img[:, bbx1:bbx2, bby1:bby2]\n\n                rate = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (IMG_SIZE * IMG_SIZE))\n                label = rate*label + (1.-rate)*self.labels[cmix_ix]\n\n        return img, label","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:25.839674Z","iopub.execute_input":"2024-04-22T10:07:25.840022Z","iopub.status.idle":"2024-04-22T10:07:25.864535Z","shell.execute_reply.started":"2024-04-22T10:07:25.839996Z","shell.execute_reply":"2024-04-22T10:07:25.863405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# #Dataset \n# # train_dataset = CassavaDataset(df, transform = train_transform)\n\n# #SubsetRandomSampler\n\n# train_size = len(df)\n# dataset_indices = list(range(train_size))\n# np.random.shuffle(dataset_indices)\n# train_split_idx = int(np.floor(0.10 * train_size))\n# train_idx, val_idx = dataset_indices[train_split_idx:], dataset_indices[:train_split_idx]\n\n# train_ = df.loc[train_idx]\n# valid_ = df.loc[val_idx]\n\n# train_ds = CassavaDataset(train_,transform = train_transform, output_label=True, one_hot_label=False, do_fmix=0, do_cutmix=0)\n# valid_ds = CassavaDataset(valid_,transform = valid_transform, output_label=True)\n\n# # class_weights = compute_class_weight(class_weight='balanced', classes=np.unique(df['label'].values), y=df['label'].values)\n# # weights = [class_weights[label] for label in df['label'].values]\n\n# train_sampler = SubsetRandomSampler(train_idx)\n# # train_sampler = WeightedRandomSampler(weights, len(weights))\n# val_sampler = SubsetRandomSampler(val_idx)\n\n# #Pass samplers to the dataloader\n# train_loader = DataLoader(dataset=train_ds,batch_size = BATCH_SIZE, sampler=train_sampler, num_workers=4, drop_last = True)\n# val_loader = DataLoader(dataset=valid_ds, batch_size = BATCH_SIZE, sampler = val_sampler, num_workers=4, drop_last = True)\n# # val_loader.dataset.transforms = valid_transform\n# # test_loader = val_loader","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:27.282271Z","iopub.execute_input":"2024-04-22T10:07:27.282651Z","iopub.status.idle":"2024-04-22T10:07:27.287961Z","shell.execute_reply.started":"2024-04-22T10:07:27.282625Z","shell.execute_reply":"2024-04-22T10:07:27.287020Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_dataloader(folds,trn_idx, val_idx):\n      \n    train_ = folds.loc[trn_idx].reset_index(drop=True)\n    valid_ = folds.loc[val_idx].reset_index(drop=True)\n    \n    train_ds = CassavaDataset(train_,transform = train_transform, output_label=True, one_hot_label=False, do_fmix=0, do_cutmix=0)\n    valid_ds = CassavaDataset(valid_,transform = valid_transform, output_label=True)\n    class_weights = compute_class_weight(class_weight='balanced', classes=np.unique(train_['label'].values), y=train_['label'].values)\n    weights = [class_weights[label] for label in train_['label'].values]\n\n# train_sampler = SubsetRandomSampler(train_idx)\n    train_sampler = WeightedRandomSampler(weights, len(weights))\n\n#     train_sampler = RandomSampler(train_)\n    val_sampler =  RandomSampler(valid_)\n\n    #Pass samplers to the dataloader\n    train_loader = DataLoader(dataset=train_ds,batch_size = BATCH_SIZE, sampler=train_sampler, num_workers=4, drop_last = True)\n    val_loader = DataLoader(dataset=valid_ds, batch_size = BATCH_SIZE, sampler = val_sampler, num_workers=4, drop_last = True)\n    \n    return train_loader, val_loader","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:28.716804Z","iopub.execute_input":"2024-04-22T10:07:28.717546Z","iopub.status.idle":"2024-04-22T10:07:28.724593Z","shell.execute_reply.started":"2024-04-22T10:07:28.717514Z","shell.execute_reply":"2024-04-22T10:07:28.723422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_image(img_dict):\n    image_tensor = img_dict[0]\n    target = img_dict[1]\n    plt.figure(figsize=(10, 10))\n    image = image_tensor.permute(1, 2, 0) \n    plt.imshow(image)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:30.449720Z","iopub.execute_input":"2024-04-22T10:07:30.450467Z","iopub.status.idle":"2024-04-22T10:07:30.455620Z","shell.execute_reply.started":"2024-04-22T10:07:30.450437Z","shell.execute_reply":"2024-04-22T10:07:30.454439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_image(train_ds[5])","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:32.013580Z","iopub.execute_input":"2024-04-22T10:07:32.013956Z","iopub.status.idle":"2024-04-22T10:07:32.018576Z","shell.execute_reply.started":"2024-04-22T10:07:32.013928Z","shell.execute_reply":"2024-04-22T10:07:32.017446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\") if torch.cuda.is_available() else torch.device(\"cpu\")\nprint(\"Device:\", device)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:33.569165Z","iopub.execute_input":"2024-04-22T10:07:33.569893Z","iopub.status.idle":"2024-04-22T10:07:33.575099Z","shell.execute_reply.started":"2024-04-22T10:07:33.569861Z","shell.execute_reply":"2024-04-22T10:07:33.574047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 3) Defining ViT Model","metadata":{}},{"cell_type":"code","source":"class ViTLightningModule(pl.LightningModule):\n    def __init__(self, num_labels=5):\n        super(ViTLightningModule, self).__init__()\n        self.vit = ViTForImageClassification.from_pretrained(MODEL_PATH,\n                                                             ignore_mismatched_sizes=True,\n                                                              num_labels=5,\n                                                              id2label=id_to_class,\n                                                              label2id=class_map)\n\n\n    def forward(self, pixel_values):\n        outputs = self.vit(pixel_values=pixel_values)\n        return outputs.logits\n        \n    def common_step(self, batch, batch_idx):\n        pixel_values, labels = batch\n        logits = self(pixel_values)\n        \n        criterion = nn.CrossEntropyLoss()\n        loss = criterion(logits, labels)\n        predictions = logits.argmax(-1)\n        correct = (predictions == labels).sum().item()\n        accuracy = correct/pixel_values.shape[0]\n\n        return loss, accuracy\n      \n    def training_step(self, batch, batch_idx):\n        loss, accuracy = self.common_step(batch, batch_idx)     \n        # logs metrics for each training_step,\n        # and the average across the epoch\n        self.log(\"training_loss\", loss)\n        self.log(\"training_accuracy\", accuracy)\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        loss, accuracy = self.common_step(batch, batch_idx)     \n        self.log(\"validation_loss\", loss, on_epoch=True)\n        self.log(\"validation_accuracy\", accuracy, on_epoch=True)\n        return loss\n\n    def test_step(self, batch, batch_idx):\n        loss, accuracy = self.common_step(batch, batch_idx)     \n\n        return loss\n\n    def configure_optimizers(self):\n        # We could make the optimizer more fancy by adding a scheduler and specifying which parameters do\n        # not require weight_decay but just using AdamW out-of-the-box works fine\n        return AdamW(self.parameters(), LR)\n\n    def train_dataloader(self):\n        return train_loader\n\n    def val_dataloader(self):\n        return val_loader\n\n#     def test_dataloader(self):\n#         return test_loader","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:36.183881Z","iopub.execute_input":"2024-04-22T10:07:36.184619Z","iopub.status.idle":"2024-04-22T10:07:36.196320Z","shell.execute_reply.started":"2024-04-22T10:07:36.184589Z","shell.execute_reply":"2024-04-22T10:07:36.195368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"early_stop_callback = EarlyStopping(\n    monitor='val_loss',\n    patience=20,\n    strict=False,\n    verbose=True,\n    mode='min'\n)\n\n\nclass ValidationMetricsCallback(pl.Callback):\n    def __init__(self, fold_index):\n        super().__init__()\n        self.best_epoch = 0\n        self.best_accuracy = 0.0\n        self.fold_index = fold_index\n        \n    def on_validation_end(self, trainer, pl_module):\n        val_loss = trainer.callback_metrics[\"validation_loss\"]\n        val_accuracy = trainer.callback_metrics[\"validation_accuracy\"]\n        print(f'Validation Loss: {val_loss:.4f}, Accuracy: {val_accuracy:.4f}')\n        torch.save(pl_module.state_dict(), f'vit_base_patch_16_384_fold_{self.fold_index}_final_epoch_{self.best_epoch}')","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:08:12.206596Z","iopub.execute_input":"2024-04-22T10:08:12.207009Z","iopub.status.idle":"2024-04-22T10:08:12.215133Z","shell.execute_reply.started":"2024-04-22T10:08:12.206974Z","shell.execute_reply":"2024-04-22T10:08:12.214083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\n\nfolds = df.copy()\nFold = StratifiedKFold(n_splits=5, shuffle=True, random_state=4012)\n\nfor n, (train_index, val_index) in enumerate(Fold.split(folds, folds[\"label\"])):\n    if n != 4:\n         continue\n    train_loader, val_loader = prepare_dataloader(folds,train_index, val_index)\n    batch = next(iter(train_loader))\n    print(batch[0].shape)\n    model = ViTLightningModule()\n    validation_callback = ValidationMetricsCallback(fold_index=n)\n#     early_stop_callback = EarlyStopping(monitor='val_loss', patience=20, strict=False, verbose=True, mode='min')\n\n    trainer = Trainer(callbacks=[validation_callback], max_epochs=MAX_EPOCHS)\n    trainer.fit(model, train_loader, val_loader)","metadata":{"execution":{"iopub.status.busy":"2024-04-22T11:41:13.163461Z","iopub.status.idle":"2024-04-22T11:41:13.163803Z","shell.execute_reply.started":"2024-04-22T11:41:13.163641Z","shell.execute_reply":"2024-04-22T11:41:13.163658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pred_class(img):\n    # transform images\n    img_tens = valid_transform(image=img)['image']\n    img_im = img_tens.unsqueeze(0).cuda() \n    uinput = Variable(img_im)\n    uinput = uinput.to(device)\n    out = model(uinput)\n    # convert image to numpy format in cpu and snatching max prediction score class index\n    index = out.data.cpu().numpy().argmax()    \n    return index","metadata":{"execution":{"iopub.status.busy":"2024-04-22T10:07:41.796671Z","iopub.status.idle":"2024-04-22T10:07:41.797184Z","shell.execute_reply.started":"2024-04-22T10:07:41.796934Z","shell.execute_reply":"2024-04-22T10:07:41.796959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.autograd import Variable\nmodel.eval()\nmodel.cuda()\n\nsubmission = pd.DataFrame(columns=['image_id', 'label'])\n\nfor filename in os.listdir(TEST_PATH):\n    image_path = os.path.join(TEST_PATH, filename)\n    x=cv2.imread(image_path,cv2.IMREAD_COLOR)\n    x=cv2.cvtColor(x,cv2.COLOR_BGR2RGB)\n\n    index = pred_class(x)\n    pred =  index\n    \n    submission = pd.concat([submission, pd.DataFrame({'image_id': [filename], 'label': [pred]})], ignore_index=True)\n    \nsubmission.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}