{"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"}],"dockerImageVersionId":30665,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Import Necessary Libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport os.path as osp\nimport copy\n\nimport time\nimport numpy as np\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\n\nimport torch\nfrom torch.utils.data.dataset import Dataset\nfrom torch import nn\nfrom torch.utils.data.dataloader import DataLoader\nfrom torchvision import transforms\nfrom torchvision import models\nfrom torchvision.utils import make_grid\nfrom torchinfo import summary\nimport torchvision.transforms.functional as F\n\nimport albumentations \nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom sklearn.metrics import precision_score, recall_score, f1_score\nfrom sklearn.metrics import classification_report, confusion_matrix, ConfusionMatrixDisplay\n\nimport random\nimport json\n\nimport polars as pl\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.axes_grid1 import ImageGrid\nimport matplotlib.patches as mpatches\nimport seaborn as sns\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nplt.style.use('ggplot')\npl.Config().set_tbl_rows(50)\npl.Config().set_tbl_cols(-1)\npl.Config().set_fmt_str_lengths(100)\n\ndef seed_everything(seed):\n    np.random.seed(seed) \n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    random.seed(seed)\n\nseed_everything(42)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-09T10:42:53.607714Z","iopub.execute_input":"2024-04-09T10:42:53.608767Z","iopub.status.idle":"2024-04-09T10:43:01.598274Z","shell.execute_reply.started":"2024-04-09T10:42:53.608714Z","shell.execute_reply":"2024-04-09T10:43:01.597257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.1 Read Data and LabelMap","metadata":{}},{"cell_type":"code","source":"ROOT_DIR = '../input/cassava-leaf-disease-classification'\nTRAIN_IMAGES_DIR = osp.join(ROOT_DIR, 'train_images')\ntrain_df = pl.read_csv(os.path.join(ROOT_DIR, 'train.csv'))\n\n# Create image paths\ntrain_df = train_df.with_columns(\n    path=pl.concat_str(pl.lit(f'{TRAIN_IMAGES_DIR}/'), pl.col('image_id'))\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:43:01.599988Z","iopub.execute_input":"2024-04-09T10:43:01.600430Z","iopub.status.idle":"2024-04-09T10:43:01.770339Z","shell.execute_reply.started":"2024-04-09T10:43:01.600404Z","shell.execute_reply":"2024-04-09T10:43:01.769558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"../input/cassava-leaf-disease-classification/label_num_to_disease_map.json\", \"r\") as f:\n    labelmap = json.load(f)\n    labelmap = {int(k): v for k, v in labelmap.items()}\n\nprint(json.dumps(labelmap, indent=4))","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:43:01.771345Z","iopub.execute_input":"2024-04-09T10:43:01.771593Z","iopub.status.idle":"2024-04-09T10:43:01.780694Z","shell.execute_reply.started":"2024-04-09T10:43:01.771572Z","shell.execute_reply":"2024-04-09T10:43:01.779784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.2 Checking Existence","metadata":{}},{"cell_type":"code","source":"train_df = train_df.with_columns(\n    is_exists=pl.col('path').map_elements(lambda x: osp.exists(x))\n)\n\nn_not_exists = train_df.filter(pl.col('is_exists') == False).shape[0]\nassert n_not_exists == 0, print(f'There are {n_not_exists} non-exists files')\n\nn_images = train_df.filter(pl.col('is_exists') == True).shape[0]\nprint(f'Total images: {n_images}')\n\ntrain_df.head(20)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:43:01.782969Z","iopub.execute_input":"2024-04-09T10:43:01.783300Z","iopub.status.idle":"2024-04-09T10:43:27.485413Z","shell.execute_reply.started":"2024-04-09T10:43:01.783276Z","shell.execute_reply":"2024-04-09T10:43:27.484540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.3 Train-val-test split","metadata":{}},{"cell_type":"code","source":"val_size = 0.2\ntest_size = 0.2\n\nval_split = n_images - int(n_images * (test_size + val_size))\ntest_split = n_images - int(n_images * test_size)\n\nval_df = train_df[val_split:test_split]\ntest_df = train_df[test_split:]\ntrain_df = train_df[:val_split]\n\nn_train = len(train_df)\nn_val = len(val_df)\nn_test = len(test_df)\n\nprint('Splitted dataset:')\nprint(f'\\t- Training set: {n_train}')\nprint(f'\\t- Validation set: {n_val}')\nprint(f'\\t- Testing set: {n_test}')","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:43:27.486848Z","iopub.execute_input":"2024-04-09T10:43:27.487202Z","iopub.status.idle":"2024-04-09T10:43:27.497835Z","shell.execute_reply.started":"2024-04-09T10:43:27.487172Z","shell.execute_reply":"2024-04-09T10:43:27.496893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.3 Number of samples by class","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(10,6))\nax = sns.countplot(x=\"label\", data=train_df.to_pandas())\nax = ax.bar_label(ax.containers[0])\n\nplt.title('Number of samples by class')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:43:27.499105Z","iopub.execute_input":"2024-04-09T10:43:27.499398Z","iopub.status.idle":"2024-04-09T10:43:27.891486Z","shell.execute_reply.started":"2024-04-09T10:43:27.499375Z","shell.execute_reply":"2024-04-09T10:43:27.890607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.4 Color Histogram Statistics","metadata":{}},{"cell_type":"code","source":"def stats_color_hist(df: pl.DataFrame, nb_bins=256):\n    means = []\n    stds = []\n    count_r = np.zeros(nb_bins)\n    count_g = np.zeros(nb_bins)\n    count_b = np.zeros(nb_bins)\n    \n    paths = df['path']\n    n = len(paths)\n    \n    for i in tqdm(range(n)):\n        img = cv2.imread(paths[i])\n        means.append(img.mean())\n        stds.append(img.std())\n        hist_r = np.histogram(img[:, :, 2], bins=nb_bins, range=[0, 255])\n        hist_g = np.histogram(img[:, :, 1], bins=nb_bins, range=[0, 255])\n        hist_b = np.histogram(img[:, :, 0], bins=nb_bins, range=[0, 255])\n        count_r += hist_r[0]\n        count_g += hist_g[0]\n        count_b += hist_b[0]\n        \n    means = np.array(means)\n    stds = np.array(stds)\n\n    count_r = count_r / n\n    count_g = count_g / n\n    count_b = count_b / n\n\n    bins = hist_r[1]\n    histogram = {\n        'r': count_r,\n        'g': count_g,\n        'b': count_b\n    }\n    return histogram, bins\n","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:43:27.892717Z","iopub.execute_input":"2024-04-09T10:43:27.893131Z","iopub.status.idle":"2024-04-09T10:43:27.903357Z","shell.execute_reply.started":"2024-04-09T10:43:27.893097Z","shell.execute_reply":"2024-04-09T10:43:27.902300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train-set histogram\nhist_train, bins_train = stats_color_hist(train_df)\n\n# Val-set histogram\nhist_val, bins_val = stats_color_hist(val_df)\n\n# Test-set histogram\nhist_test, bins_test = stats_color_hist(test_df)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:43:27.904565Z","iopub.execute_input":"2024-04-09T10:43:27.904885Z","iopub.status.idle":"2024-04-09T10:58:38.894692Z","shell.execute_reply.started":"2024-04-09T10:43:27.904856Z","shell.execute_reply":"2024-04-09T10:58:38.893798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Overall color histogram of each set","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(3, 1, figsize=(10, 10), dpi=300)\naxes[0].set_title('Training set: Histogram')\naxes[0].bar(bins_train[:-1], hist_train['r'], color='r', alpha=0.33)\naxes[0].bar(bins_train[:-1], hist_train['g'], color='g', alpha=0.33)\naxes[0].bar(bins_train[:-1], hist_train['b'], color='b', alpha=0.33)\n\naxes[1].set_title('Validation set: Histogram')\naxes[1].bar(bins_val[:-1], hist_val['r'], color='r', alpha=0.33)\naxes[1].bar(bins_val[:-1], hist_val['g'], color='g', alpha=0.33)\naxes[1].bar(bins_val[:-1], hist_val['b'], color='b', alpha=0.33)\n\naxes[2].set_title('Testing set: Histogram')\naxes[2].bar(bins_test[:-1], hist_test['r'], color='r', alpha=0.33)\naxes[2].bar(bins_test[:-1], hist_test['g'], color='g', alpha=0.33)\naxes[2].bar(bins_test[:-1], hist_test['b'], color='b', alpha=0.33)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:58:38.895893Z","iopub.execute_input":"2024-04-09T10:58:38.896212Z","iopub.status.idle":"2024-04-09T10:58:43.080445Z","shell.execute_reply.started":"2024-04-09T10:58:38.896184Z","shell.execute_reply":"2024-04-09T10:58:43.079575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Closer look of color histogram in the range 0 to 5000 of frequency","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(3, 1, figsize=(10, 10), dpi=300)\naxes[0].set_title('Training set: Histogram')\naxes[0].bar(bins_train[:-1], hist_train['r'], color='r', alpha=0.33)\naxes[0].bar(bins_train[:-1], hist_train['g'], color='g', alpha=0.33)\naxes[0].bar(bins_train[:-1], hist_train['b'], color='b', alpha=0.33)\naxes[0].set_ylim(0, 5000)\n\naxes[1].set_title('Validation set: Histogram')\naxes[1].bar(bins_val[:-1], hist_val['r'], color='r', alpha=0.33)\naxes[1].bar(bins_val[:-1], hist_val['g'], color='g', alpha=0.33)\naxes[1].bar(bins_val[:-1], hist_val['b'], color='b', alpha=0.33)\naxes[1].set_ylim(0, 5000)\n\naxes[2].set_title('Testing set: Histogram')\naxes[2].bar(bins_test[:-1], hist_test['r'], color='r', alpha=0.33)\naxes[2].bar(bins_test[:-1], hist_test['g'], color='g', alpha=0.33)\naxes[2].bar(bins_test[:-1], hist_test['b'], color='b', alpha=0.33)\naxes[2].set_ylim(0, 5000)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:58:43.085280Z","iopub.execute_input":"2024-04-09T10:58:43.085977Z","iopub.status.idle":"2024-04-09T10:58:48.023363Z","shell.execute_reply.started":"2024-04-09T10:58:43.085940Z","shell.execute_reply":"2024-04-09T10:58:48.022462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"$\\Rightarrow$ Color histograms of three channels of training, validation and testing set are <u>*similar*</u>.","metadata":{}},{"cell_type":"markdown","source":"## 2.5 Image shapes","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.6 Visualize sample images","metadata":{}},{"cell_type":"code","source":"def plot_images(row, col):\n    # Create a figure with the desired size\n    fig = plt.figure(figsize=(5 * col, 5 * row))\n\n    # Create a grid of subplots\n    grid = ImageGrid(fig, 111, nrows_ncols=(row, col), axes_pad=0.5)\n\n    # Sample some data from the train_df dataframe\n    sample = train_df.sample(row * col)\n    paths = sample['path'].to_list()\n    labels = sample['label'].to_list()\n\n    for i, ax in enumerate(grid):\n        img = plt.imread(paths[i])\n        ax.imshow(img)\n        ax.set_title(labelmap[labels[i]])\n        ax.axis('off')\n\n    plt.show()\n\nplot_images(5, 6)","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:58:48.024381Z","iopub.execute_input":"2024-04-09T10:58:48.024635Z","iopub.status.idle":"2024-04-09T10:58:55.616198Z","shell.execute_reply.started":"2024-04-09T10:58:48.024613Z","shell.execute_reply":"2024-04-09T10:58:55.615151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Image Augmentation","metadata":{}},{"cell_type":"markdown","source":"## 3.1 Original images","metadata":{}},{"cell_type":"code","source":"N = 3\nfig = plt.figure(figsize=(5, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 1), axes_pad=0.5)\nsample = train_df.sample(N)\npaths = sample['path'].to_list()\nfor i, ax in enumerate(grid):\n    img = plt.imread(paths[i])\n    ax.imshow(img)\n    ax.set_title(f'Original [{i}]')\n    ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:58:55.617483Z","iopub.execute_input":"2024-04-09T10:58:55.617845Z","iopub.status.idle":"2024-04-09T10:58:56.714777Z","shell.execute_reply.started":"2024-04-09T10:58:55.617815Z","shell.execute_reply":"2024-04-09T10:58:56.713772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.2 Vertical Flip","metadata":{}},{"cell_type":"code","source":"vertical_transform = albumentations.VerticalFlip(p=1)\n\nfig = plt.figure(figsize=(5 * 2, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 2), axes_pad=0.5)\nfor i, ax in enumerate(grid):\n    # Random new sample\n    if i % 2 == 0:\n        org_img = plt.imread(paths[i // 2])\n        img = org_img.copy()\n        ax.imshow(img)\n        ax.set_title(f'Original [{i // 2}]')\n    else:\n        img = vertical_transform(image=img)['image']\n        ax.imshow(img)\n        ax.set_title(f'Vertical Flip [{i // 2}]')\n    \n    ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:58:56.716039Z","iopub.execute_input":"2024-04-09T10:58:56.716330Z","iopub.status.idle":"2024-04-09T10:58:58.616799Z","shell.execute_reply.started":"2024-04-09T10:58:56.716305Z","shell.execute_reply":"2024-04-09T10:58:58.615760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.3 Horizontal Flip","metadata":{}},{"cell_type":"code","source":"horizontal_transform = albumentations.HorizontalFlip(p=1)\n\nfig = plt.figure(figsize=(5 * 2, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 2), axes_pad=0.5)\nfor i, ax in enumerate(grid):\n    # Random new sample\n    if i % 2 == 0:\n        org_img = plt.imread(paths[i // 2])\n        img = org_img.copy()\n        ax.imshow(img)\n        ax.set_title(f'Original [{i // 2}]')\n    else:\n        img = horizontal_transform(image=img)['image']\n        ax.imshow(img)\n        ax.set_title(f'Horizontal Flip [{i // 2}]')\n    \n    ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:58:58.618160Z","iopub.execute_input":"2024-04-09T10:58:58.618512Z","iopub.status.idle":"2024-04-09T10:59:00.465328Z","shell.execute_reply.started":"2024-04-09T10:58:58.618484Z","shell.execute_reply":"2024-04-09T10:59:00.464368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.4 Random Resized Crop","metadata":{}},{"cell_type":"code","source":"random_crop_transform = albumentations.RandomResizedCrop(width=512, height=512)\n\nfig = plt.figure(figsize=(5 * 2, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 2), axes_pad=0.5)\nfor i, ax in enumerate(grid):\n    # Random new sample\n    if i % 2 == 0:\n        org_img = plt.imread(paths[i // 2])\n        img = org_img.copy()\n        ax.imshow(img)\n        ax.set_title(f'Original [{i // 2}]')\n    else:\n        img = random_crop_transform(image=img)['image']\n        ax.imshow(img)\n        ax.set_title(f'Random Resized Crop [{i // 2}]')\n    \n    ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2024-04-09T10:59:00.466591Z","iopub.execute_input":"2024-04-09T10:59:00.466989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.5 Transpose","metadata":{}},{"cell_type":"code","source":"transpose_transform = albumentations.Transpose(p=1)\n\nfig = plt.figure(figsize=(5 * 2, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 2), axes_pad=0.5)\nfor i, ax in enumerate(grid):\n    # Random new sample\n    if i % 2 == 0:\n        org_img = plt.imread(paths[i // 2])\n        img = org_img.copy()\n        ax.imshow(img)\n        ax.set_title(f'Original [{i // 2}]')\n    else:\n        img = transpose_transform(image=img)['image']\n        ax.imshow(img)\n        ax.set_title(f'Tranpose [{i // 2}]')\n    \n    ax.axis('off')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.6 Shift Scale Rotate","metadata":{}},{"cell_type":"code","source":"shift_scale_rotate_transform = albumentations.ShiftScaleRotate(p=1)\n\nfig = plt.figure(figsize=(5 * 2, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 2), axes_pad=0.5)\nfor i, ax in enumerate(grid):\n    # Random new sample\n    if i % 2 == 0:\n        org_img = plt.imread(paths[i // 2])\n        img = org_img.copy()\n        ax.imshow(img)\n        ax.set_title(f'Original [{i // 2}]')\n    else:\n        img = shift_scale_rotate_transform(image=img)['image']\n        ax.imshow(img)\n        ax.set_title(f'Shift Scale Rotate [{i // 2}]')\n    \n    ax.axis('off')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.7 Hue Saturation Value Transform","metadata":{}},{"cell_type":"code","source":"hue_saturation_value_transform = albumentations.HueSaturationValue(\n    hue_shift_limit=20,\n    sat_shift_limit=50,\n    val_shift_limit=20,\n    p=1)\n\nfig = plt.figure(figsize=(5 * 2, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 2), axes_pad=0.5)\nfor i, ax in enumerate(grid):\n    # Random new sample\n    if i % 2 == 0:\n        org_img = plt.imread(paths[i // 2])\n        img = org_img.copy()\n        ax.imshow(img)\n        ax.set_title(f'Original [{i // 2}]')\n    else:\n        img = hue_saturation_value_transform(image=img)['image']\n        ax.imshow(img)\n        ax.set_title(f'HSV Transform [{i // 2}]')\n    \n    ax.axis('off')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.8 Random Brightness Contrast","metadata":{}},{"cell_type":"code","source":"random_brightness_contrast_transform = albumentations.RandomBrightnessContrast(\n    brightness_limit=(-0.1, 0.1), \n    contrast_limit=(-0.1, 0.1), \n    p=1)\n\nfig = plt.figure(figsize=(5 * 2, 5 * N))\ngrid = ImageGrid(fig, 111, nrows_ncols=(N, 2), axes_pad=0.5)\nfor i, ax in enumerate(grid):\n    # Random new sample\n    if i % 2 == 0:\n        org_img = plt.imread(paths[i // 2])\n        img = org_img.copy()\n        ax.imshow(img)\n        ax.set_title(f'Original [{i // 2}]')\n    else:\n        img = random_brightness_contrast_transform(image=img)['image']\n        ax.imshow(img)\n        ax.set_title(f'Random Brightness Contrast [{i // 2}]')\n    \n    ax.axis('off')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Modeling","metadata":{}},{"cell_type":"code","source":"WIDTH = 512\nHEIGHT = 512\nNUM_CLASSES = 5\nBATCH_SIZE = 32\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f'Device: {DEVICE}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = albumentations.Compose([\n    \n    albumentations.RandomResizedCrop(WIDTH, HEIGHT),\n    albumentations.HorizontalFlip(p=0.5),\n    albumentations.Transpose(p=0.5),\n    albumentations.VerticalFlip(p=0.5),\n    albumentations.ShiftScaleRotate(p=0.5),\n    albumentations.HueSaturationValue(\n                hue_shift_limit=0.2, \n                sat_shift_limit=0.2, \n                val_shift_limit=0.2, \n                p=0.5\n            ),\n    albumentations.RandomBrightnessContrast(\n                brightness_limit=(-0.1, 0.1), \n                contrast_limit=(-0.1, 0.1), \n                p=0.5),\n    albumentations.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n    \n])\n\ntest_transforms = albumentations.Compose([\n    albumentations.CenterCrop(WIDTH, HEIGHT, p=1.0),\n    albumentations.Resize(WIDTH, HEIGHT),\n    albumentations.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2(),\n])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    \n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        \n        img = Image.open(self.df['path'][index])\n        img = np.array(img)\n        label = torch.tensor(self.df['label'][index], dtype=torch.long)\n        \n        if self.transform:\n            return self.transform(image=img)['image'], label \n        else:\n            return img, label ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CassavaDataset(train_df, transform=train_transforms)\nval_dataset = CassavaDataset(val_df, transform=test_transforms)\ntest_dataset = CassavaDataset(test_df, transform=test_transforms)\n\ntrain_dl = DataLoader(train_dataset, BATCH_SIZE, shuffle=True)\nval_dl = DataLoader(val_dataset, BATCH_SIZE, shuffle=True)\ntest_dl = DataLoader(test_dataset, BATCH_SIZE, shuffle=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_resnet_model():\n    \n    model = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n    \n    for params in model.parameters():\n        params.requires_grad = False\n        \n    in_feat = model.fc.in_features\n        \n    model.fc = nn.Sequential(\n          nn.Linear(in_feat, 256),\n          nn.ReLU(),\n          nn.Dropout(p=0.3),\n          nn.Linear(256, NUM_CLASSES))\n    \n    model = model.to(DEVICE)\n    \n    return model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Model summary","metadata":{}},{"cell_type":"code","source":"model = get_resnet_model()\n\nsummary(model, input_size=(BATCH_SIZE, 3, WIDTH, HEIGHT))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, num_epochs, train_dl, valid_dl):\n    \n    loss_hist_train = [0] * num_epochs\n    accuracy_hist_train = [0] * num_epochs\n    loss_hist_valid = [0] * num_epochs\n    accuracy_hist_valid = [0] * num_epochs\n    \n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n    min_valid_loss = np.inf\n    \n    for epoch in range(num_epochs):\n        \n        model.train()\n        \n        batch_num = 0\n        \n        for x_batch, y_batch in tqdm(train_dl):\n            \n            x_batch = x_batch.to(DEVICE)\n            y_batch = y_batch.to(DEVICE)\n            \n            batch_num += 1\n            #if (batch_num % 100 == 0):\n                #print(f'Batch number: {batch_num}')\n            \n            pred = model(x_batch)\n            loss = loss_fn(pred, y_batch)\n            loss.backward()\n            optimizer.step()\n            optimizer.zero_grad()\n            \n            loss_hist_train[epoch] += loss.item() * y_batch.size(0)\n            is_correct = (torch.argmax(pred, dim=1) == y_batch).float()\n            accuracy_hist_train[epoch] += is_correct.sum().item()\n        \n        \n        loss_hist_train[epoch] /= len(train_dl.dataset)\n        accuracy_hist_train[epoch] /= len(train_dl.dataset)\n        \n        scheduler.step()\n        \n        model.eval()\n        \n        with torch.no_grad():\n            \n            for x_batch, y_batch in valid_dl:\n                \n                x_batch = x_batch.to(DEVICE)\n                y_batch = y_batch.to(DEVICE)\n                \n                pred = model(x_batch)\n                loss = loss_fn(pred, y_batch)\n                loss_hist_valid[epoch] += loss.item() * y_batch.size(0)\n                is_correct = (torch.argmax(pred, dim=1) == y_batch).float()\n                accuracy_hist_valid[epoch] += is_correct.sum().item()\n                \n        loss_hist_valid[epoch] /= len(valid_dl.dataset)\n        accuracy_hist_valid[epoch] /= len(valid_dl.dataset)\n        \n        if accuracy_hist_valid[epoch] > best_acc:\n            best_acc = accuracy_hist_valid[epoch]\n            best_model_wts = copy.deepcopy(model.state_dict())\n        \n        print(f'Epoch {epoch+1}:   Train accuracy: {accuracy_hist_train[epoch]:.4f}    Validation accuracy: {accuracy_hist_valid[epoch]:.4f} ')\n    \n    \n        if loss_hist_valid[epoch] < min_valid_loss:\n            counter = 0\n        else:\n            counter += 1\n    \n        if counter >= patience:\n            break\n    \n    \n    model.load_state_dict(best_model_wts)\n    \n    history = {}\n    history['loss_hist_train'] = loss_hist_train\n    history['loss_hist_valid'] = loss_hist_valid\n    history['accuracy_hist_train'] = accuracy_hist_train\n    history['accuracy_hist_valid'] = accuracy_hist_valid\n    \n    return model, history","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 10\npatience = 3\nloss_fn = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=1, eta_min=1e-6, last_epoch=-1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Train model","metadata":{"execution":{"iopub.status.busy":"2024-03-30T10:31:28.218982Z","iopub.execute_input":"2024-03-30T10:31:28.219955Z","iopub.status.idle":"2024-03-30T10:31:28.225385Z","shell.execute_reply.started":"2024-03-30T10:31:28.219913Z","shell.execute_reply":"2024-03-30T10:31:28.224174Z"}}},{"cell_type":"code","source":"best_model, hist = train(model, num_epochs, train_dl, val_dl)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5. Evaluation","metadata":{}},{"cell_type":"code","source":"label_list = []\nprediction_list = []\n\nwith torch.no_grad():\n    for image, label in tqdm(test_dl):\n        \n        image = image.to(DEVICE)\n        logits = best_model(image)\n        probs = torch.nn.functional.softmax(logits, dim=1).detach().cpu().numpy()\n        prediction = np.argmax(probs, axis=1)\n        label_list += label.numpy().tolist()\n        prediction_list += prediction.tolist()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(classification_report(label_list, prediction_list))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cm = confusion_matrix(label_list, prediction_list)\nplt.figure(figsize=(6, 6))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Greens', cbar=False, linewidth=1, linecolor='white')\nplt.xlabel('Predicted labels')\nplt.ylabel('True labels')\nplt.title('Confusion Matrix')\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}