{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"}],"dockerImageVersionId":30616,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **1. Install packages**","metadata":{}},{"cell_type":"code","source":"%%capture\n!pip install torchinfo","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-09T05:18:54.933186Z","iopub.execute_input":"2023-12-09T05:18:54.933553Z","iopub.status.idle":"2023-12-09T05:19:07.561904Z","shell.execute_reply.started":"2023-12-09T05:18:54.93352Z","shell.execute_reply":"2023-12-09T05:19:07.560719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **2. Import libraries**","metadata":{}},{"cell_type":"code","source":"# Data handling\nimport pandas as pd\nimport numpy as np\n\n# Data visualization\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nfrom PIL import Image\nImage.MAX_IMAGE_PIXELS = None\n\n# Preprocessing\nfrom sklearn.model_selection import train_test_split as tts\nfrom sklearn.utils.class_weight import compute_class_weight\n\n# Torch\nimport torch\nfrom torch import nn, optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models import vit_b_16, ViT_B_16_Weights\nfrom torchinfo import summary\n\n# Metrics\nfrom sklearn.metrics import balanced_accuracy_score\nfrom sklearn.metrics import confusion_matrix\n\n# os\nimport os\n\n# Path\nfrom pathlib import Path\n\n# random\nimport random\n\n# OrderedDict\nfrom collections import OrderedDict\n\n# tqdm\nfrom tqdm.auto import tqdm\n\n# warnings\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:07.563926Z","iopub.execute_input":"2023-12-09T05:19:07.564232Z","iopub.status.idle":"2023-12-09T05:19:12.677404Z","shell.execute_reply.started":"2023-12-09T05:19:07.564205Z","shell.execute_reply":"2023-12-09T05:19:12.676554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **3. Load data and EDA**","metadata":{}},{"cell_type":"code","source":"IMAGE_PATH = Path(\"/kaggle/input/UBC-OCEAN/train_thumbnails\")","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:12.67847Z","iopub.execute_input":"2023-12-09T05:19:12.678876Z","iopub.status.idle":"2023-12-09T05:19:12.682845Z","shell.execute_reply.started":"2023-12-09T05:19:12.67885Z","shell.execute_reply":"2023-12-09T05:19:12.681904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_train = pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")\ndata_train.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:12.685167Z","iopub.execute_input":"2023-12-09T05:19:12.685446Z","iopub.status.idle":"2023-12-09T05:19:12.726309Z","shell.execute_reply.started":"2023-12-09T05:19:12.685404Z","shell.execute_reply":"2023-12-09T05:19:12.725457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_train = data_train.loc[data_train[\"is_tma\"] == False]\ndata_train.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:12.727341Z","iopub.execute_input":"2023-12-09T05:19:12.72768Z","iopub.status.idle":"2023-12-09T05:19:12.742774Z","shell.execute_reply.started":"2023-12-09T05:19:12.727647Z","shell.execute_reply":"2023-12-09T05:19:12.74199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_train['image_id'] = data_train['image_id'].apply(lambda x: os.path.join(IMAGE_PATH, f\"{x}_thumbnail.png\"))\ndata_train.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:12.743751Z","iopub.execute_input":"2023-12-09T05:19:12.744067Z","iopub.status.idle":"2023-12-09T05:19:12.758587Z","shell.execute_reply.started":"2023-12-09T05:19:12.74404Z","shell.execute_reply":"2023-12-09T05:19:12.7576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Total Images\nprint(f'Total Images = {len(data_train)}')","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:12.759602Z","iopub.execute_input":"2023-12-09T05:19:12.759874Z","iopub.status.idle":"2023-12-09T05:19:12.76699Z","shell.execute_reply.started":"2023-12-09T05:19:12.75984Z","shell.execute_reply":"2023-12-09T05:19:12.76607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Total classes\nCLASSES = data_train['label'].unique()\nCLASSES = sorted(CLASSES)\nprint(\"==\" * 20)\nprint(\" \" * 10, f\"Total Classes = {len(CLASSES)}\")\nprint(\"==\" * 20)\nprint(CLASSES)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:12.768199Z","iopub.execute_input":"2023-12-09T05:19:12.76847Z","iopub.status.idle":"2023-12-09T05:19:12.781791Z","shell.execute_reply.started":"2023-12-09T05:19:12.768445Z","shell.execute_reply":"2023-12-09T05:19:12.78057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_class = data_train['label'].value_counts().to_frame()\nlabels = df_class.index.to_list()\nvalues = df_class.iloc[:,0]\n\ndef autopct_fun(abs_values):\n    gen = iter(abs_values)\n    return lambda pct: f\"{pct:.1f}%\\n({next(gen)})\"\n\nfig,ax = plt.subplots(figsize = (6,4))\nax.pie(x = values, labels = labels, autopct = autopct_fun(values), shadow = True)\nax.axis(\"equal\")\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:12.782883Z","iopub.execute_input":"2023-12-09T05:19:12.783196Z","iopub.status.idle":"2023-12-09T05:19:13.019016Z","shell.execute_reply.started":"2023-12-09T05:19:12.783172Z","shell.execute_reply":"2023-12-09T05:19:13.01779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **4. Preprocessing**","metadata":{}},{"cell_type":"markdown","source":"We divide our dataframe into 3 data sets:\n- **train** = **70%**\n- **validation** = **15%**\n- **test** = **15%**","metadata":{}},{"cell_type":"code","source":"SEED = 42\n\ndf_train, df_rest = tts(data_train, \n                       test_size = 0.3, \n                       random_state = SEED, \n                       shuffle = True, \n                       stratify = data_train['label'])\n\ndf_val, df_test = tts(df_rest, \n                      test_size = 0.5, \n                      random_state = SEED, \n                      shuffle = True, \n                      stratify = df_rest['label'])","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.02556Z","iopub.execute_input":"2023-12-09T05:19:13.026696Z","iopub.status.idle":"2023-12-09T05:19:13.048222Z","shell.execute_reply.started":"2023-12-09T05:19:13.026643Z","shell.execute_reply":"2023-12-09T05:19:13.047099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_class_train = df_train['label'].value_counts().to_frame()\nlabels_train = df_class_train.index.to_list()\nvalues_train = df_class_train.iloc[:,0]\n\ndf_class_val = df_val['label'].value_counts().to_frame()\nlabels_val = df_class_val.index.to_list()\nvalues_val = df_class_val.iloc[:,0]\n\ndf_class_test = df_test['label'].value_counts().to_frame()\nlabels_test = df_class_test.index.to_list()\nvalues_test = df_class_test.iloc[:,0]\n\ndef autopct_fun(abs_values):\n    gen = iter(abs_values)\n    return lambda pct: f\"{pct:.1f}%\\n({next(gen)})\"\n\nfig,ax = plt.subplots(nrows = 1, ncols = 3, figsize = (12,5))\nax = ax.flat\nax[0].pie(x = values_train, \n          labels = labels_train, \n          autopct = autopct_fun(values_train), \n          shadow = True)\nax[0].axis(\"equal\")\nax[0].set_title(\"Train\", fontsize = 11, fontweight = \"bold\", color = \"black\")\n\nax[1].pie(x = values_val, \n          labels = labels_val, \n          autopct = autopct_fun(values_val), \n          shadow = True)\nax[1].axis(\"equal\")\nax[1].set_title(\"Validation\", fontsize = 11, fontweight = \"bold\", color = \"black\")\n\nax[2].pie(x = values_test, \n          labels = labels_test, \n          autopct = autopct_fun(values_test), \n          shadow = True)\nax[2].axis(\"equal\")\nax[2].set_title(\"Test\", fontsize = 11, fontweight = \"bold\", color = \"black\")\n\nfig.tight_layout()\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.049997Z","iopub.execute_input":"2023-12-09T05:19:13.050923Z","iopub.status.idle":"2023-12-09T05:19:13.735534Z","shell.execute_reply.started":"2023-12-09T05:19:13.050873Z","shell.execute_reply":"2023-12-09T05:19:13.734392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_weight = compute_class_weight('balanced', \n                                    classes = np.unique(np.array(df_train['label'])), \n                                    y = np.array(df_train['label']))\n\nclass_weight = torch.from_numpy(class_weight)\nclass_weight","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.736979Z","iopub.execute_input":"2023-12-09T05:19:13.737329Z","iopub.status.idle":"2023-12-09T05:19:13.840818Z","shell.execute_reply.started":"2023-12-09T05:19:13.737295Z","shell.execute_reply":"2023-12-09T05:19:13.840045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_map = dict(zip(CLASSES, range(len(CLASSES))))\nlabel_map","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.841774Z","iopub.execute_input":"2023-12-09T05:19:13.842038Z","iopub.status.idle":"2023-12-09T05:19:13.84811Z","shell.execute_reply.started":"2023-12-09T05:19:13.842015Z","shell.execute_reply":"2023-12-09T05:19:13.847213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now it's our turn to define our **Dataset** and **DataLoader** for each data set.","metadata":{}},{"cell_type":"markdown","source":"- **Dataset**","metadata":{}},{"cell_type":"code","source":"class CustomImageDataset(Dataset):\n    def __init__(self, df:pd.DataFrame, label_map:dict, image_transforms):\n        self.df = df\n        self.label_map = label_map\n        self.image_transforms = image_transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        df_new = self.df.copy()\n        df_new['label'] = df_new['label'].map(self.label_map)\n        \n        image_path = df_new.iloc[idx, 0]\n        image = Image.open(image_path).convert(\"RGB\")\n        image = self.image_transforms(image)\n        \n        label = df_new.iloc[idx, 1]\n        \n        return image, label","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.849249Z","iopub.execute_input":"2023-12-09T05:19:13.849571Z","iopub.status.idle":"2023-12-09T05:19:13.857831Z","shell.execute_reply.started":"2023-12-09T05:19:13.84954Z","shell.execute_reply":"2023-12-09T05:19:13.856896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We are going to define the **weights** of the **pre-trained model** and the **transformations** that were applied to the **pre-training** images.","metadata":{}},{"cell_type":"code","source":"weights = ViT_B_16_Weights.DEFAULT\nauto_transforms = weights.transforms()\nauto_transforms","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.858916Z","iopub.execute_input":"2023-12-09T05:19:13.859238Z","iopub.status.idle":"2023-12-09T05:19:13.868531Z","shell.execute_reply.started":"2023-12-09T05:19:13.859213Z","shell.execute_reply":"2023-12-09T05:19:13.867664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CustomImageDataset(df_train, label_map, auto_transforms)\nval_dataset = CustomImageDataset(df_val, label_map, auto_transforms)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.869403Z","iopub.execute_input":"2023-12-09T05:19:13.86967Z","iopub.status.idle":"2023-12-09T05:19:13.87721Z","shell.execute_reply.started":"2023-12-09T05:19:13.869647Z","shell.execute_reply":"2023-12-09T05:19:13.876485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- **DataLoader**","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 32\nNUM_WORKERS = os.cpu_count()\n\ntrain_dataloader = DataLoader(dataset = train_dataset, \n                              batch_size = BATCH_SIZE, \n                              shuffle = True, \n                              num_workers = NUM_WORKERS)\n\nval_dataloader = DataLoader(dataset = val_dataset, \n                            batch_size = BATCH_SIZE, \n                            shuffle = True, \n                            num_workers = NUM_WORKERS)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.878178Z","iopub.execute_input":"2023-12-09T05:19:13.878464Z","iopub.status.idle":"2023-12-09T05:19:13.887469Z","shell.execute_reply.started":"2023-12-09T05:19:13.878433Z","shell.execute_reply":"2023-12-09T05:19:13.886598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_images, batch_labels = next(iter(train_dataloader))\n\nbatch_images.shape, batch_labels.shape","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:13.88851Z","iopub.execute_input":"2023-12-09T05:19:13.888972Z","iopub.status.idle":"2023-12-09T05:19:36.312139Z","shell.execute_reply.started":"2023-12-09T05:19:13.888918Z","shell.execute_reply":"2023-12-09T05:19:36.310684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **5. Model**","metadata":{}},{"cell_type":"code","source":"# GPU\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:36.313885Z","iopub.execute_input":"2023-12-09T05:19:36.314284Z","iopub.status.idle":"2023-12-09T05:19:36.381733Z","shell.execute_reply.started":"2023-12-09T05:19:36.314246Z","shell.execute_reply":"2023-12-09T05:19:36.380791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We define the model to use with the pre-trained weights.\nmodel = vit_b_16(weights = weights)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:36.383034Z","iopub.execute_input":"2023-12-09T05:19:36.383339Z","iopub.status.idle":"2023-12-09T05:19:43.404988Z","shell.execute_reply.started":"2023-12-09T05:19:36.383314Z","shell.execute_reply":"2023-12-09T05:19:43.403985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Let's visualize the architecture of our model.\nsummary(model = model, \n        input_size = [48,3,224,224], \n        col_names = [\"input_size\", \"output_size\", \"num_params\", \"trainable\"], \n        col_width = 15, row_settings = [\"var_names\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:43.406334Z","iopub.execute_input":"2023-12-09T05:19:43.40665Z","iopub.status.idle":"2023-12-09T05:19:52.75111Z","shell.execute_reply.started":"2023-12-09T05:19:43.406622Z","shell.execute_reply":"2023-12-09T05:19:52.750215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Because we are going to do **transfer learning**, we need to **freeze the extractor layers**, that means we will freeze the **conv_proj** and **encoder** layers.","metadata":{}},{"cell_type":"code","source":"for param in model.conv_proj.parameters():\n    param.requires_grad = False","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:52.75239Z","iopub.execute_input":"2023-12-09T05:19:52.752748Z","iopub.status.idle":"2023-12-09T05:19:52.757915Z","shell.execute_reply.started":"2023-12-09T05:19:52.752712Z","shell.execute_reply":"2023-12-09T05:19:52.756769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for param in model.encoder.parameters():\n    param.requires_grad = False","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:52.759104Z","iopub.execute_input":"2023-12-09T05:19:52.759447Z","iopub.status.idle":"2023-12-09T05:19:52.768856Z","shell.execute_reply.started":"2023-12-09T05:19:52.759422Z","shell.execute_reply":"2023-12-09T05:19:52.768086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We are going to check if the extractor layers were frozen, viewing the model, in the **Trainable** column it must be in **False** mode.","metadata":{}},{"cell_type":"code","source":"summary(model = model, \n        input_size = [48,3,224,224], \n        col_names = [\"input_size\", \"output_size\", \"num_params\", \"trainable\"], \n        col_width = 15, row_settings = [\"var_names\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:52.769878Z","iopub.execute_input":"2023-12-09T05:19:52.770187Z","iopub.status.idle":"2023-12-09T05:19:53.242252Z","shell.execute_reply.started":"2023-12-09T05:19:52.770164Z","shell.execute_reply":"2023-12-09T05:19:53.241302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Great!!, it turned out well for us, now we just need to modify the size of the **output**, which should be **equal** to the **number of classes** we have.","metadata":{}},{"cell_type":"code","source":"output_shape = len(CLASSES)\n\nmodel.heads = nn.Sequential(OrderedDict([('head', nn.Linear(in_features = 768, out_features = output_shape))]))","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:53.243643Z","iopub.execute_input":"2023-12-09T05:19:53.244017Z","iopub.status.idle":"2023-12-09T05:19:53.250043Z","shell.execute_reply.started":"2023-12-09T05:19:53.243989Z","shell.execute_reply":"2023-12-09T05:19:53.248807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Let's do our last visualization of our model to check if the output is equal to the number of classes.**","metadata":{}},{"cell_type":"code","source":"summary(model = model, \n        input_size = [48,3,224,224], \n        col_names = [\"input_size\", \"output_size\", \"num_params\", \"trainable\"], \n        col_width = 15, \n        row_settings = [\"var_names\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:19:53.251216Z","iopub.execute_input":"2023-12-09T05:19:53.251509Z","iopub.status.idle":"2023-12-09T05:19:53.662363Z","shell.execute_reply.started":"2023-12-09T05:19:53.251466Z","shell.execute_reply":"2023-12-09T05:19:53.661493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we have to define our **loss function** and the **optimizer**.","metadata":{}},{"cell_type":"code","source":"loss_fn = nn.CrossEntropyLoss()\noptimizer = optim.SGD(model.parameters(), lr = 0.01, weight_decay = 0.15)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:57:12.271593Z","iopub.execute_input":"2023-12-09T05:57:12.272007Z","iopub.status.idle":"2023-12-09T05:57:12.278863Z","shell.execute_reply.started":"2023-12-09T05:57:12.271939Z","shell.execute_reply":"2023-12-09T05:57:12.277902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Once this is defined, we are going to create some functions and classes to execute the training, save the best model and make predictions, the name of the functions and class are as follows:\n\n- **EarlyStopping**\n- **train_step**\n- **val_step**\n- **predictions**","metadata":{}},{"cell_type":"code","source":"class EarlyStopping:\n    def __init__(self, patience:int = 5, delta:float = 0.0001, path = \"best_model.pth\"):\n        self.patience = patience\n        self.delta = delta\n        self.path = path\n        self.best_score = None\n        self.counter = 0\n        self.early_stop = False\n        \n    def __call__(self, val_loss, model):\n        if self.best_score is None:\n            self.best_score = val_loss\n            self.save_checkpoint(model)\n            \n        elif val_loss > self.best_score + self.delta:\n            self.counter += 1\n            if self.counter >= self.patience:\n                self.early_stop = True\n                \n        else:\n            self.best_score = val_loss\n            self.save_checkpoint(model)\n            self.counter = 0\n            \n    def save_checkpoint(self, model):\n        torch.save(model.state_dict(), self.path)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:57:16.091541Z","iopub.execute_input":"2023-12-09T05:57:16.092175Z","iopub.status.idle":"2023-12-09T05:57:16.100734Z","shell.execute_reply.started":"2023-12-09T05:57:16.092138Z","shell.execute_reply":"2023-12-09T05:57:16.099692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"early_stopping = EarlyStopping(patience = 10, delta = 0.)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:57:17.477193Z","iopub.execute_input":"2023-12-09T05:57:17.477633Z","iopub.status.idle":"2023-12-09T05:57:17.482591Z","shell.execute_reply.started":"2023-12-09T05:57:17.477592Z","shell.execute_reply":"2023-12-09T05:57:17.481546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_step(model:torch.nn.Module, \n               dataloader:torch.utils.data.DataLoader, \n               loss_fn:torch.nn.Module, \n               optimizer:torch.optim.Optimizer):\n    \n    model.train()\n    \n    train_loss = 0.\n    train_accuracy = 0.\n    \n    for batch,(X,y) in enumerate(dataloader):\n        X,y = X.to(device = device, dtype = torch.float32), y.to(device = device, dtype = torch.long)\n        optimizer.zero_grad()\n        y_pred_logit = model(X)\n        loss = loss_fn(y_pred_logit, y)\n        train_loss += loss.item()\n        \n        loss.backward()\n        optimizer.step()\n        \n        y_pred_prob = y_pred_logit.softmax(dim = 1)\n        y_pred_class = y_pred_prob.argmax(dim = 1)\n        \n        train_accuracy += balanced_accuracy_score(y.cpu().numpy(), y_pred_class.detach().cpu().numpy())\n    \n    train_loss = train_loss / len(dataloader)\n    train_accuracy = train_accuracy / len(dataloader)\n    \n    return train_loss, train_accuracy\n        ","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:57:18.822055Z","iopub.execute_input":"2023-12-09T05:57:18.822426Z","iopub.status.idle":"2023-12-09T05:57:18.831306Z","shell.execute_reply.started":"2023-12-09T05:57:18.822393Z","shell.execute_reply":"2023-12-09T05:57:18.830275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def val_step(model:torch.nn.Module, \n             dataloader:torch.utils.data.DataLoader, \n             loss_fn:torch.nn.Module):\n    \n    model.eval()\n    \n    val_loss = 0.\n    val_accuracy = 0.\n    \n    with torch.inference_mode():\n        for batch,(X,y) in enumerate(dataloader):\n            X,y = X.to(device = device, dtype = torch.float32), y.to(device = device, dtype = torch.long)\n            y_pred_logit = model(X)\n            loss = loss_fn(y_pred_logit, y)\n            val_loss += loss.item()\n            \n            y_pred_prob = y_pred_logit.softmax(dim = 1)\n            y_pred_class = y_pred_prob.argmax(dim = 1)\n            \n            val_accuracy += balanced_accuracy_score(y.cpu().numpy(), y_pred_class.detach().cpu().numpy())\n    \n    val_loss = val_loss / len(dataloader)\n    val_accuracy = val_accuracy / len(dataloader)\n    \n    return val_loss, val_accuracy","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:57:20.281661Z","iopub.execute_input":"2023-12-09T05:57:20.282041Z","iopub.status.idle":"2023-12-09T05:57:20.290324Z","shell.execute_reply.started":"2023-12-09T05:57:20.282007Z","shell.execute_reply":"2023-12-09T05:57:20.289406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model:torch.nn.Module, \n          train_dataloader:torch.utils.data.DataLoader, \n          val_dataloader:torch.utils.data.DataLoader, \n          loss_fn:torch.nn.Module, \n          optimizer:torch.optim.Optimizer, \n          early_stopping, \n          epochs:int = 10):\n    \n    results = {'train_loss':[], 'train_accuracy':[], 'val_loss':[], 'val_accuracy':[]}\n    \n    for epoch in tqdm(range(epochs)):\n        train_loss, train_accuracy = train_step(model = model, \n                                                dataloader = train_dataloader, \n                                                loss_fn = loss_fn, \n                                                optimizer = optimizer)\n        \n        val_loss, val_accuracy = val_step(model = model, \n                                          dataloader = val_dataloader, \n                                          loss_fn = loss_fn)\n        \n        print(f'Epoch: {epoch + 1} | ', \n              f'Train Loss: {train_loss:.4f} | ', \n              f'Train Accuracy: {train_accuracy:.4f} | ', \n              f'Val Loss: {val_loss:.4f} | ', \n              f'Val Accuracy: {val_accuracy:.4f}')\n        \n        early_stopping(val_loss, model)\n        \n        if early_stopping.early_stop == True:\n            print(\"Early Stopping!!\")\n            break\n            \n        results['train_loss'].append(train_loss)\n        results['train_accuracy'].append(train_accuracy)\n        results['val_loss'].append(val_loss)\n        results['val_accuracy'].append(val_accuracy)\n        \n    return results","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:57:21.723791Z","iopub.execute_input":"2023-12-09T05:57:21.724572Z","iopub.status.idle":"2023-12-09T05:57:21.733552Z","shell.execute_reply.started":"2023-12-09T05:57:21.724537Z","shell.execute_reply":"2023-12-09T05:57:21.732562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training!!!\nEPOCHS = 100\n\ntorch.cuda.manual_seed(SEED)\ntorch.manual_seed(SEED)\n\nRESULTS = train(model.to(device), \n                train_dataloader, \n                val_dataloader, \n                loss_fn, \n                optimizer, \n                early_stopping, \n                EPOCHS)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T05:57:23.241337Z","iopub.execute_input":"2023-12-09T05:57:23.241681Z","iopub.status.idle":"2023-12-09T06:30:24.875912Z","shell.execute_reply.started":"2023-12-09T05:57:23.241654Z","shell.execute_reply":"2023-12-09T06:30:24.874754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predictions(test_dataloader:torch.utils.data.DataLoader, best_model:str):\n    \n    loaded_model = vit_b_16()\n    loaded_model.heads = nn.Sequential(OrderedDict([('head', nn.Linear(in_features = 768, \n                                                                       out_features = output_shape))]))\n    \n    checkpoint = torch.load(best_model)\n    \n    loaded_model.load_state_dict(checkpoint)\n    \n    loaded_model.to(device)\n    \n    loaded_model.eval()\n    \n    y_pred_test = []\n    \n    with torch.inference_mode():\n        for X,_ in tqdm(test_dataloader):\n            X = X.to(device, dtype = torch.float32)\n            y_pred_logit = loaded_model(X)\n            y_pred_prob = y_pred_logit.softmax(dim = 1)\n            y_pred_class = y_pred_prob.argmax(dim = 1)\n            y_pred_test.append(y_pred_class.detach().cpu())\n            \n    y_pred_test = torch.cat(y_pred_test)\n    \n    return y_pred_test","metadata":{"execution":{"iopub.status.busy":"2023-12-09T06:30:36.201212Z","iopub.execute_input":"2023-12-09T06:30:36.202204Z","iopub.status.idle":"2023-12-09T06:30:36.211585Z","shell.execute_reply.started":"2023-12-09T06:30:36.202158Z","shell.execute_reply":"2023-12-09T06:30:36.210488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Predictions test**","metadata":{}},{"cell_type":"markdown","source":"Before making the predictions, let's remember that we have to convert it to a **Dataset** and then to a **DataLoader**.","metadata":{}},{"cell_type":"code","source":"test_dataset = CustomImageDataset(df_test, label_map, auto_transforms)\n\ntest_dataloader = DataLoader(dataset = test_dataset, \n                             batch_size = BATCH_SIZE, \n                             shuffle = False, \n                             num_workers = NUM_WORKERS)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T06:30:40.395416Z","iopub.execute_input":"2023-12-09T06:30:40.395842Z","iopub.status.idle":"2023-12-09T06:30:40.401407Z","shell.execute_reply.started":"2023-12-09T06:30:40.39581Z","shell.execute_reply":"2023-12-09T06:30:40.4004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We load the best model and make the predictions.\ny_pred_test = predictions(test_dataloader, \"/kaggle/working/best_model.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-12-09T06:30:42.603264Z","iopub.execute_input":"2023-12-09T06:30:42.603649Z","iopub.status.idle":"2023-12-09T06:30:58.597509Z","shell.execute_reply.started":"2023-12-09T06:30:42.603617Z","shell.execute_reply":"2023-12-09T06:30:58.596275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **6. Metrics**","metadata":{}},{"cell_type":"markdown","source":"- **Balanced Accuracy**","metadata":{}},{"cell_type":"code","source":"acc_test = balanced_accuracy_score(df_test['label'].map(label_map), y_pred_test.numpy())\n\nprint(f'Balanced Accuracy Test = {acc_test:.4f}')","metadata":{"execution":{"iopub.status.busy":"2023-12-09T06:31:03.090171Z","iopub.execute_input":"2023-12-09T06:31:03.090589Z","iopub.status.idle":"2023-12-09T06:31:03.104944Z","shell.execute_reply.started":"2023-12-09T06:31:03.090546Z","shell.execute_reply":"2023-12-09T06:31:03.104013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- **Confusion Matrix**","metadata":{}},{"cell_type":"code","source":"cf_mx_test = confusion_matrix(df_test['label'].map(label_map), y_pred_test.numpy())\n\nfig,ax = plt.subplots(figsize = (7,3.8))\nsns.heatmap(cf_mx_test, \n            cmap = \"Reds\", \n            annot = True, \n            annot_kws = {'fontsize':8, 'fontweight':'bold'}, \n            linewidths = 0.8, \n            cbar = False, \n            square = True, \n            xticklabels = CLASSES, \n            yticklabels = CLASSES, \n            ax = ax)\nax.set_title(\"Confusion Matrix Test\", fontsize = 12, fontweight = \"bold\", color = \"green\")\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-09T06:31:07.8996Z","iopub.execute_input":"2023-12-09T06:31:07.899936Z","iopub.status.idle":"2023-12-09T06:31:08.185287Z","shell.execute_reply.started":"2023-12-09T06:31:07.899911Z","shell.execute_reply":"2023-12-09T06:31:08.184026Z"},"trusted":true},"execution_count":null,"outputs":[]}]}