{"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-01T15:45:36.646158Z","iopub.execute_input":"2024-04-01T15:45:36.646755Z","iopub.status.idle":"2024-04-01T15:45:36.661945Z","shell.execute_reply.started":"2024-04-01T15:45:36.646715Z","shell.execute_reply":"2024-04-01T15:45:36.660797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Exploratory Data Analysis","metadata":{}},{"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-01T15:16:38.133243Z","iopub.execute_input":"2024-04-01T15:16:38.133725Z","iopub.status.idle":"2024-04-01T15:16:38.293420Z","shell.execute_reply.started":"2024-04-01T15:16:38.133697Z","shell.execute_reply":"2024-04-01T15:16:38.292653Z"},"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-01T15:16:38.294378Z","iopub.execute_input":"2024-04-01T15:16:38.294643Z","iopub.status.idle":"2024-04-01T15:16:38.303593Z","shell.execute_reply.started":"2024-04-01T15:16:38.294621Z","shell.execute_reply":"2024-04-01T15:16:38.302630Z"},"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-01T15:16:42.958328Z","iopub.execute_input":"2024-04-01T15:16:42.959181Z","iopub.status.idle":"2024-04-01T15:17:24.428493Z","shell.execute_reply.started":"2024-04-01T15:16:42.959143Z","shell.execute_reply":"2024-04-01T15:17:24.427523Z"},"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-01T15:17:24.430452Z","iopub.execute_input":"2024-04-01T15:17:24.430776Z","iopub.status.idle":"2024-04-01T15:17:24.439412Z","shell.execute_reply.started":"2024-04-01T15:17:24.430752Z","shell.execute_reply":"2024-04-01T15:17:24.438495Z"},"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-03-30T07:46:30.830893Z","iopub.execute_input":"2024-03-30T07:46:30.831621Z","iopub.status.idle":"2024-03-30T07:46:31.129710Z","shell.execute_reply.started":"2024-03-30T07:46:30.831586Z","shell.execute_reply":"2024-03-30T07:46:31.128710Z"},"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-03-30T07:51:59.332573Z","iopub.execute_input":"2024-03-30T07:51:59.333052Z","iopub.status.idle":"2024-03-30T07:51:59.342903Z","shell.execute_reply.started":"2024-03-30T07:51:59.333004Z","shell.execute_reply":"2024-03-30T07:51:59.341758Z"},"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-03-30T08:21:23.991762Z","iopub.execute_input":"2024-03-30T08:21:23.992180Z","iopub.status.idle":"2024-03-30T08:35:28.995907Z","shell.execute_reply.started":"2024-03-30T08:21:23.992149Z","shell.execute_reply":"2024-03-30T08:35:28.994965Z"},"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-03-30T08:46:18.561604Z","iopub.execute_input":"2024-03-30T08:46:18.562469Z","iopub.status.idle":"2024-03-30T08:46:22.836665Z","shell.execute_reply.started":"2024-03-30T08:46:18.562434Z","shell.execute_reply":"2024-03-30T08:46:22.835756Z"},"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-03-30T08:47:37.661766Z","iopub.execute_input":"2024-03-30T08:47:37.662412Z","iopub.status.idle":"2024-03-30T08:47:42.188217Z","shell.execute_reply.started":"2024-03-30T08:47:37.662378Z","shell.execute_reply":"2024-03-30T08:47:42.187269Z"},"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-03-19T12:10:51.727110Z","iopub.execute_input":"2024-03-19T12:10:51.727375Z","iopub.status.idle":"2024-03-19T12:10:59.033874Z","shell.execute_reply.started":"2024-03-19T12:10:51.727352Z","shell.execute_reply":"2024-03-19T12:10:59.032076Z"},"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-03-30T09:34:00.319367Z","iopub.execute_input":"2024-03-30T09:34:00.320398Z","iopub.status.idle":"2024-03-30T09:34:01.340401Z","shell.execute_reply.started":"2024-03-30T09:34:00.320362Z","shell.execute_reply":"2024-03-30T09:34:01.339431Z"},"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-03-30T09:35:55.619350Z","iopub.execute_input":"2024-03-30T09:35:55.620422Z","iopub.status.idle":"2024-03-30T09:35:57.364550Z","shell.execute_reply.started":"2024-03-30T09:35:55.620387Z","shell.execute_reply":"2024-03-30T09:35:57.363496Z"},"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-03-30T09:37:01.929377Z","iopub.execute_input":"2024-03-30T09:37:01.929799Z","iopub.status.idle":"2024-03-30T09:37:03.645656Z","shell.execute_reply.started":"2024-03-30T09:37:01.929766Z","shell.execute_reply":"2024-03-30T09:37:03.644649Z"},"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-03-30T09:38:36.731322Z","iopub.execute_input":"2024-03-30T09:38:36.731716Z","iopub.status.idle":"2024-03-30T09:38:38.391680Z","shell.execute_reply.started":"2024-03-30T09:38:36.731675Z","shell.execute_reply":"2024-03-30T09:38:38.390207Z"},"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":{"execution":{"iopub.status.busy":"2024-03-30T09:41:23.956523Z","iopub.execute_input":"2024-03-30T09:41:23.957235Z","iopub.status.idle":"2024-03-30T09:41:25.659715Z","shell.execute_reply.started":"2024-03-30T09:41:23.957204Z","shell.execute_reply":"2024-03-30T09:41:25.658737Z"},"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":{"execution":{"iopub.status.busy":"2024-03-30T09:43:12.159455Z","iopub.execute_input":"2024-03-30T09:43:12.160173Z","iopub.status.idle":"2024-03-30T09:43:14.069464Z","shell.execute_reply.started":"2024-03-30T09:43:12.160138Z","shell.execute_reply":"2024-03-30T09:43:14.068524Z"},"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":{"execution":{"iopub.status.busy":"2024-03-30T09:44:34.327491Z","iopub.execute_input":"2024-03-30T09:44:34.327904Z","iopub.status.idle":"2024-03-30T09:44:36.079926Z","shell.execute_reply.started":"2024-03-30T09:44:34.327872Z","shell.execute_reply":"2024-03-30T09:44:36.078877Z"},"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":{"execution":{"iopub.status.busy":"2024-03-30T09:45:56.589888Z","iopub.execute_input":"2024-03-30T09:45:56.590879Z","iopub.status.idle":"2024-03-30T09:45:58.212077Z","shell.execute_reply.started":"2024-03-30T09:45:56.590843Z","shell.execute_reply":"2024-03-30T09:45:58.210904Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:17:53.505438Z","iopub.execute_input":"2024-04-01T15:17:53.505816Z","iopub.status.idle":"2024-04-01T15:17:53.535250Z","shell.execute_reply.started":"2024-04-01T15:17:53.505789Z","shell.execute_reply":"2024-04-01T15:17:53.534042Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:17:55.538883Z","iopub.execute_input":"2024-04-01T15:17:55.539238Z","iopub.status.idle":"2024-04-01T15:17:55.548122Z","shell.execute_reply.started":"2024-04-01T15:17:55.539208Z","shell.execute_reply":"2024-04-01T15:17:55.547168Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:17:56.442008Z","iopub.execute_input":"2024-04-01T15:17:56.442439Z","iopub.status.idle":"2024-04-01T15:17:56.449453Z","shell.execute_reply.started":"2024-04-01T15:17:56.442404Z","shell.execute_reply":"2024-04-01T15:17:56.448526Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:17:57.980319Z","iopub.execute_input":"2024-04-01T15:17:57.980687Z","iopub.status.idle":"2024-04-01T15:17:57.986889Z","shell.execute_reply.started":"2024-04-01T15:17:57.980654Z","shell.execute_reply":"2024-04-01T15:17:57.985855Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:18:09.020519Z","iopub.execute_input":"2024-04-01T15:18:09.020861Z","iopub.status.idle":"2024-04-01T15:18:09.026930Z","shell.execute_reply.started":"2024-04-01T15:18:09.020836Z","shell.execute_reply":"2024-04-01T15:18:09.026005Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:18:21.571252Z","iopub.execute_input":"2024-04-01T15:18:21.571644Z","iopub.status.idle":"2024-04-01T15:18:24.383210Z","shell.execute_reply.started":"2024-04-01T15:18:21.571603Z","shell.execute_reply":"2024-04-01T15:18:24.382264Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:18:28.813885Z","iopub.execute_input":"2024-04-01T15:18:28.814718Z","iopub.status.idle":"2024-04-01T15:18:28.829150Z","shell.execute_reply.started":"2024-04-01T15:18:28.814684Z","shell.execute_reply":"2024-04-01T15:18:28.828236Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:18:54.824874Z","iopub.execute_input":"2024-04-01T15:18:54.825265Z","iopub.status.idle":"2024-04-01T15:18:54.831926Z","shell.execute_reply.started":"2024-04-01T15:18:54.825234Z","shell.execute_reply":"2024-04-01T15:18:54.831018Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T15:18:57.215400Z","iopub.execute_input":"2024-04-01T15:18:57.215773Z","iopub.status.idle":"2024-04-01T15:38:45.664262Z","shell.execute_reply.started":"2024-04-01T15:18:57.215746Z","shell.execute_reply":"2024-04-01T15:38:45.663304Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T16:11:25.045234Z","iopub.execute_input":"2024-04-01T16:11:25.046069Z","iopub.status.idle":"2024-04-01T16:12:39.954045Z","shell.execute_reply.started":"2024-04-01T16:11:25.046035Z","shell.execute_reply":"2024-04-01T16:12:39.953139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(classification_report(label_list, prediction_list))","metadata":{"execution":{"iopub.status.busy":"2024-04-01T16:16:13.168332Z","iopub.execute_input":"2024-04-01T16:16:13.168995Z","iopub.status.idle":"2024-04-01T16:16:13.192578Z","shell.execute_reply.started":"2024-04-01T16:16:13.168958Z","shell.execute_reply":"2024-04-01T16:16:13.191626Z"},"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":{"execution":{"iopub.status.busy":"2024-04-01T16:17:16.344354Z","iopub.execute_input":"2024-04-01T16:17:16.344714Z","iopub.status.idle":"2024-04-01T16:17:16.610733Z","shell.execute_reply.started":"2024-04-01T16:17:16.344688Z","shell.execute_reply":"2024-04-01T16:17:16.609789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}