{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport csv\nimport shutil\nimport cv2\nimport pandas as pd\nimport tensorflow as tf\nimport numpy as np\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom matplotlib import pyplot as plt\nfrom scipy.linalg import fractional_matrix_power\nfrom scipy.signal import wiener\nfrom PIL import Image\nfrom tensorflow.keras.regularizers import l1, l2, L1L2\nfrom tensorflow.keras.layers import GlobalAveragePooling2D, Reshape, Dense, Input\n","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:57:01.803898Z","iopub.execute_input":"2023-07-28T14:57:01.804266Z","iopub.status.idle":"2023-07-28T14:57:01.811158Z","shell.execute_reply.started":"2023-07-28T14:57:01.804235Z","shell.execute_reply":"2023-07-28T14:57:01.809804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Image preprocessing","metadata":{}},{"cell_type":"code","source":"# CLAHE on green channel of RGB image \n\ndef enhance_contrast(image): \n\n    image = np.uint8(image)\n    r, g, b = cv2.split(image)\n    clahe = cv2.createCLAHE(clipLimit=2, tileGridSize=(8,8))\n    g = wiener(g, 2, 0.01) # weiner \n    g = np.uint8(g)\n    \n    g_enhanced = clahe.apply(g)\n    \n    image =  np.stack([r,g_enhanced,b],axis=-1)\n\n    image = cv2.medianBlur(image, 3)\n    image = cv2.resize(image, (224, 224))\n    \n    return image\n    '''\n    # Split the image into RGB channels\n    red_channel, green_channel, blue_channel = cv2.split(image)\n    # Apply CLAHE to the green channel only\n    clahe = cv2.createCLAHE(clipLimit=2, tileGridSize=(8,8))\n    green_channel = wiener(green_channel, 2, 0.01) # weiner filter\n    green_enhanced = clahe.apply(green_channel)\n\n    # Merge the enhanced green channel with the original red and blue channels\n    enhanced_image = cv2.merge([red_channel, green_enhanced, blue_channel])\n    #enhanced_image = cv2.blur(image, (3,3)) # median filter\n    return enhanced_image\n    '''\n# crop image \ndef crop_image_from_gray(img,tol=7):\n    if img.ndim ==1:\n        return img\n    if img.ndim ==2:\n        mask = img>tol\n        return img[np.ix_(mask.any(1),mask.any(0))]\n    elif img.ndim==3:\n        gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        mask = gray_img>tol\n        \n        check_shape = img[:,:,0][np.ix_(mask.any(1),mask.any(0))].shape[0]\n        if (check_shape == 0): # image is too dark so that we crop out everything,\n            return img # return original image\n        else:\n            img1=img[:,:,0][np.ix_(mask.any(1),mask.any(0))]\n            img2=img[:,:,1][np.ix_(mask.any(1),mask.any(0))]\n            img3=img[:,:,2][np.ix_(mask.any(1),mask.any(0))]\n    #         print(img1.shape,img2.shape,img3.shape)\n            img = np.stack([img1,img2,img3],axis=-1)\n    #         print(img.shape)\n        return img\n    \ndef preprocess_image(image):\n    image = crop_image_from_gray(image)\n    image = enhance_contrast(image)\n    #image = adjust_brightness(image)\n    return image\n# check cropping\nax = plt.subplot(1,3,1)\nimg = '/kaggle/input/diabetic-retinopathy-resized/resized_train/resized_train/10003_left.jpeg'\nimg = cv2.imread(img)\nimage = img[:,:,1]\nplt.imshow(image, cmap ='gray')\nimg = preprocess_image(img)\nax = plt.subplot(1,3,2)\nplt.imshow(img)","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:57:01.817284Z","iopub.execute_input":"2023-07-28T14:57:01.818078Z","iopub.status.idle":"2023-07-28T14:57:02.369779Z","shell.execute_reply.started":"2023-07-28T14:57:01.818046Z","shell.execute_reply":"2023-07-28T14:57:02.368855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Select images and split into training, validation and test directory\n","metadata":{}},{"cell_type":"code","source":"labels = ['label0', 'label1', 'label2', 'label3', 'label4']\n\nimg_dir = '/kaggle/input/diabeticretinopathycropped/diabetic retinopathy/diabetic retinopathy cropped'\nlabel_csv = '/kaggle/input/diabeticretinopathycropped/diabetic retinopathy/trainLabels.csv'\n\ntrain_dir = 'train'\nvalid_dir = 'tmp_valid'\ntest_dir = 'tmp_test'\n\n# clear and create new directory\nif os.path.isdir(train_dir):\n    shutil.rmtree(train_dir)\nif os.path.isdir(valid_dir):\n    shutil.rmtree(valid_dir)\nif os.path.isdir(test_dir):\n    shutil.rmtree(test_dir)\n    \nos.makedirs(train_dir)\nos.makedirs(valid_dir)\nos.makedirs(test_dir)\n\nfor label in labels:\n    sub_dir = os.path.join(train_dir, label)\n    if not os.path.exists(sub_dir):\n        os.makedirs(sub_dir)\n    sub_dir = os.path.join(valid_dir, label)\n    if not os.path.exists(sub_dir):\n        os.makedirs(sub_dir)\n    sub_dir = os.path.join(test_dir, label)\n    if not os.path.exists(sub_dir):\n        os.makedirs(sub_dir)\n\n# open csv file and read to dict img - label\nimg_labels = {}\nwith open(label_csv, 'r') as f:\n    reader = csv.reader(f)\n    next(reader) # skip header row\n    for row in reader:\n        img_name = row[0] + '.jpeg'\n        label = int(row[1])\n        img_labels[img_name] = label\n\n# randomly select one fifth of the images and write their rows to a new CSV file\nimg_number = int(len(list(img_labels.keys()))/10)################# change to /3, /2 or /1 for bigger dataset \n\nselected_imgs = random.sample(list(img_labels.keys()), img_number)\nrandom.shuffle(selected_imgs)\n# split the selected imgs to train, test and validation\ntrain_ratio = 0.7 \nvalid_ratio = 0.2 \ntest_ratio = 0.1\n\ntrain_split = int(len(selected_imgs) * train_ratio)\nvalid_split = int(len(selected_imgs) * valid_ratio)\n\n# Split the images into train and test sets\ntrain_imgs = selected_imgs[:train_split]\nvalid_imgs = selected_imgs[train_split:train_split+valid_split]\ntest_imgs = selected_imgs[train_split+valid_split:]\n\nfor img_name in train_imgs:\n    label = img_labels[img_name]\n    src_path = os.path.join(img_dir, img_name)\n    dst_path = os.path.join(train_dir, labels[label], img_name)\n    shutil.copy(src_path, dst_path)\nfor img_name in valid_imgs:\n    label = img_labels[img_name]\n    src_path = os.path.join(img_dir, img_name)\n    dst_path = os.path.join(valid_dir, labels[label], img_name)\n    shutil.copy(src_path, dst_path)\nfor img_name in test_imgs:\n    label = img_labels[img_name]\n    src_path = os.path.join(img_dir, img_name)\n    dst_path = os.path.join(test_dir, labels[label], img_name)\n    shutil.copy(src_path, dst_path)","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:57:02.372251Z","iopub.execute_input":"2023-07-28T14:57:02.372936Z","iopub.status.idle":"2023-07-28T14:57:24.222280Z","shell.execute_reply.started":"2023-07-28T14:57:02.372902Z","shell.execute_reply":"2023-07-28T14:57:24.221106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport matplotlib.pyplot as plt\ndir = train_dir\nimages = []\nnum = 4\n\nfor label in labels:\n    img_dir = os.path.join(dir, label)\n    image_paths = os.listdir(img_dir)\n    image_paths = list(map(lambda img: os.path.join(img_dir, img), image_paths))\n    images.extend(random.sample(image_paths, num))\nplt.figure(figsize=(20, 20))\nimage = cv2.imread(images[0])\nprint(images[0])\nfor i in range(5 * num):\n    image = plt.imread(images[i])\n    image = preprocess_image(image)\n    ax = plt.subplot(5, num, i + 1)\n    plt.imshow(image)\n    plt.title(os.path.basename(images[i]))\n    plt.axis(\"off\")\n'''\nimg_dir = os.path.join(dir, labels[2])\nimage_paths = os.listdir(img_dir)\nimages = random.sample(list(map(lambda img: os.path.join(img_dir, img), image_paths)), 2)\nplt.figure(figsize=(20, 20))\nfor i in range(2):\n    image = plt.imread(images[i])\n    image = preprocess_image(image)\n    ax = plt.subplot(1, num, i + 1)\n    height, width, channel = image.shape\n    plt.imshow(image)\n    plt.title(os.path.basename(images[i]) + ' ' + str(height) + '*' + str(width))\n    plt.axis(\"off\")\n'''","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:57:24.223776Z","iopub.execute_input":"2023-07-28T14:57:24.224123Z","iopub.status.idle":"2023-07-28T14:57:27.716907Z","shell.execute_reply.started":"2023-07-28T14:57:24.224089Z","shell.execute_reply":"2023-07-28T14:57:27.714562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\ni = random.randint(0, 10)\n# Reading an image in default mode\nimage = cv2.imread(images[0])\nimage = crop_image_from_gray(image)\nB, G, R = cv2.split(image)\nclahe = cv2.createCLAHE(clipLimit=4, tileGridSize=(8,8))\nG_enhanced = clahe.apply(G)\nG_wiener = wiener(G, 2, 0.01)\nG_wiener = clahe.apply(np.uint8(G_wiener))\n# Display the individual channels using plt.imshow()\nplt.figure(figsize=(20, 20))\nax = plt.subplot(3, 2, 1)\nplt.imshow(G_enhanced, cmap='gray')\nplt.title('G_enhanced Channel')\nax = plt.subplot(3, 2, 2)\nplt.hist(G_enhanced.flatten(),256,[0,256],color='r')\n\nax = plt.subplot(3, 2, 3)\nplt.imshow(G_wiener, cmap='gray')\nplt.title('G_enhanced wiener Channel')\nax = plt.subplot(3, 2, 4)\nplt.hist(G_wiener.flatten(),256,[0,256],color='r')\n\nax = plt.subplot(3, 2, 5)\nplt.imshow(G, cmap='gray')\nplt.title('Green Channel')\nax = plt.subplot(3, 2, 6)\nplt.hist(G.flatten(),256,[0,256],color='r')","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:57:27.719292Z","iopub.execute_input":"2023-07-28T14:57:27.720195Z","iopub.status.idle":"2023-07-28T14:57:31.275340Z","shell.execute_reply.started":"2023-07-28T14:57:27.720158Z","shell.execute_reply":"2023-07-28T14:57:31.274534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Apply preprocessing and augmentation on the dataset\n","metadata":{}},{"cell_type":"code","source":"# apply augmentation and CLAHE on training set -> move to augmentation set\n\n\n# Specify the paths to the train folder and the augmented output folder\ntrain_dir = 'train'\naug_dir = 'aug'\n\nif os.path.isdir(aug_dir):\n    shutil.rmtree(aug_dir)\n\nos.makedirs(aug_dir)\n\nfor label in labels:\n    os.makedirs(os.path.join(aug_dir, label))\n    \n# create data frame\ntrain_paths = []\nlabel_image = []\nfor label in labels:\n    for image_path in [os.path.join(train_dir, label, image) for image in os.listdir(os.path.join(train_dir, label))]:\n        train_paths.append(image_path)\n        label_image.append(label)\n        \nFseries=pd.Series(train_paths, name='filepaths')\nLseries=pd.Series(label_image, name='labels')\ntrain_df = pd.concat([Fseries, Lseries], axis = 1)\nprint(train_df.head())\nprint('Total number of images before undersampling: {num}'.format(num = len(train_df)))\nprint(train_df['labels'].value_counts())\n\n# undersampling by dividing to the average of 5 images\nmax_value = 0\nfor value in train_df['labels'].value_counts():\n    max_value += value\nmax_value = int(max_value/5)\n\n# limit the number of image to max_value\nsample_list = []\ngroups=train_df.groupby('labels')\nfor label in train_df['labels'].unique():                 \n    group=groups.get_group(label)\n    sample_count=len(group)    \n    if sample_count > max_value:\n        samples=group.sample(max_value, replace=False, weights=None, random_state=123, axis=0).reset_index(drop=True)\n    else:\n        samples=group.sample(frac=1.0, replace=False, random_state=123, axis=0).reset_index(drop=True)\n    sample_list.append(samples)\ntrain_df=pd.concat(sample_list, axis=0).reset_index(drop=True)\nprint('Total number of images after undersampling: {num}'.format(num = len(train_df)))\nprint(train_df['labels'].value_counts())    \n\n# create aug images to balance after undersampling\ntarget = max_value\ngen = ImageDataGenerator(horizontal_flip=True, vertical_flip=True, rotation_range=20, preprocessing_function=preprocess_image)\ngroups = train_df.groupby('labels') # group by class\n\nfor label in train_df['labels'].unique():  # for every class               \n    group = groups.get_group(label)  # a dataframe holding only rows with the specified label \n    sample_count = len(group)   # determine how many samples there are in this class  \n    if sample_count < target: # if the class has less than target number of images\n        aug_img_count = 0\n        delta = target - sample_count  # number of augmented images to create\n        target_dir = os.path.join(aug_dir, label)  # define where to write the images    \n        aug_gen = gen.flow_from_dataframe(group,  x_col='filepaths', y_col=None, target_size=(224, 224), class_mode=None, batch_size=1,\n                                         shuffle=False, save_to_dir=target_dir, save_prefix='aug',save_format='jpeg')\n        while aug_img_count<delta:\n            images=next(aug_gen)   \n            aug_img_count += len(images)\n            \n# copy image (with CLAHE only)\ngen = ImageDataGenerator(preprocessing_function=preprocess_image)\nfor label in train_df['labels'].unique():\n    group = groups.get_group(label)\n    target_dir=os.path.join(aug_dir, label)\n    aug_gen=gen.flow_from_dataframe(group,  x_col='filepaths', y_col=None, target_size=(224, 224), class_mode=None, batch_size=1,\n                                    shuffle=False, save_to_dir=target_dir, save_prefix='aug',save_format='jpeg')\n    num_images = len(group)\n    for _ in range(num_images):\n        images = next(aug_gen)\n\nprint('Number of images in each label after augmentation: ')\nfor label in labels:\n    path = os.path.join(aug_dir, label)\n    print(label + ': {}'.format(len(os.listdir(path))) + ' images')","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:57:31.276856Z","iopub.execute_input":"2023-07-28T14:57:31.277497Z","iopub.status.idle":"2023-07-28T14:58:32.174952Z","shell.execute_reply.started":"2023-07-28T14:57:31.277445Z","shell.execute_reply":"2023-07-28T14:58:32.174026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# apply CLAHE on validation and test dataset\ngen = ImageDataGenerator(preprocessing_function=preprocess_image)\n\nnew_train_dir = 'not_aug'\nnew_val_dir = 'valid'\nnew_test_dir = 'test'\n\nif os.path.isdir(new_val_dir):\n    shutil.rmtree(new_val_dir)\nif os.path.isdir(new_test_dir):\n    shutil.rmtree(new_test_dir)\nif os.path.isdir(new_train_dir):\n    shutil.rmtree(new_train_dir)\nos.makedirs(new_val_dir)\nos.makedirs(new_test_dir)\nos.makedirs(new_train_dir)\nfor label in labels:\n    os.makedirs(os.path.join(new_val_dir, label))\n    os.makedirs(os.path.join(new_test_dir, label))\n    os.makedirs(os.path.join(new_train_dir, label))\n\n# create df from folder to apply clahe\ntrain_paths = []\nlabel_image = []\nfor label in labels:\n    for image_path in [os.path.join(train_dir, label, image) for image in os.listdir(os.path.join(train_dir, label))]:\n        train_paths.append(image_path)\n        label_image.append(label)\n        \nFseries=pd.Series(train_paths, name='filepaths')\nLseries=pd.Series(label_image, name='labels')\ntest_df = pd.concat([Fseries, Lseries], axis = 1)\n\n\nvalid_paths = []\nlabel_image = []\nfor label in labels:\n    for image_path in [os.path.join(valid_dir, label, image) for image in os.listdir(os.path.join(valid_dir, label))]:\n        valid_paths.append(image_path)\n        label_image.append(label)\n        \nFseries=pd.Series(valid_paths, name='filepaths')\nLseries=pd.Series(label_image, name='labels')\nvalid_df = pd.concat([Fseries, Lseries], axis = 1)\n\ntest_paths = []\nlabel_image = []\nfor label in labels:\n    for image_path in [os.path.join(test_dir, label, image) for image in os.listdir(os.path.join(test_dir, label))]:\n        test_paths.append(image_path)\n        label_image.append(label)\n        \nFseries=pd.Series(test_paths, name='filepaths')\nLseries=pd.Series(label_image, name='labels')\ntest_df = pd.concat([Fseries, Lseries], axis = 1)\n\ngroups = train_df.groupby('labels')\nfor label in test_df['labels'].unique():\n    group = groups.get_group(label)\n    target_dir=os.path.join(new_train_dir, label)\n    aug_gen=gen.flow_from_dataframe(group,  x_col='filepaths', y_col=None, target_size=(224, 224), class_mode=None, batch_size=1,\n                                    shuffle=False, save_to_dir=target_dir, save_prefix='aug',save_format='jpeg')\n    num_images = len(group)\n    for _ in range(num_images):\n        images = next(aug_gen)\n        \ngroups = valid_df.groupby('labels')\nfor label in valid_df['labels'].unique():\n    group = groups.get_group(label)\n    target_dir=os.path.join(new_val_dir, label)\n    aug_gen=gen.flow_from_dataframe(group,  x_col='filepaths', y_col=None, target_size=(224, 224), class_mode=None, batch_size=1,\n                                    shuffle=False, save_to_dir=target_dir, save_prefix='aug',save_format='jpeg')\n    num_images = len(group)\n    print(len(group))\n    for _ in range(num_images):\n        images = next(aug_gen)\n\ngroups = test_df.groupby('labels')\nfor label in test_df['labels'].unique():\n    group = groups.get_group(label)\n    target_dir=os.path.join(new_test_dir, label)\n    aug_gen=gen.flow_from_dataframe(group,  x_col='filepaths', y_col=None, target_size=(224, 224), class_mode=None, batch_size=1,\n                                    shuffle=False, save_to_dir=target_dir, save_prefix='aug',save_format='jpeg')\n    num_images = len(group)\n    for _ in range(num_images):\n        images = next(aug_gen)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-28T15:01:34.438869Z","iopub.execute_input":"2023-07-28T15:01:34.439264Z","iopub.status.idle":"2023-07-28T15:02:19.094559Z","shell.execute_reply.started":"2023-07-28T15:01:34.439233Z","shell.execute_reply":"2023-07-28T15:02:19.093523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport PIL\nimport PIL.Image\nimport tensorflow as tf\nimport pathlib\naug_dir = pathlib.Path('aug').with_suffix('')\nnew_train_dir = pathlib.Path('not_aug').with_suffix('')\nnew_test_dir = pathlib.Path('test').with_suffix('')\nnew_val_dir = pathlib.Path('valid').with_suffix('')\n\nbatch_size = 32 ################\nimg_height = 224 ################\nimg_width = 224 ################\nnum_classes = 5\n\ntrain_ds = tf.keras.utils.image_dataset_from_directory(\n  new_train_dir,# change to new_train_dir for non-augmentation\n  shuffle=True,\n  image_size=(img_height, img_width),\n  batch_size=batch_size)\n\n\nvalidation_ds = tf.keras.utils.image_dataset_from_directory(\n  new_val_dir,\n  shuffle=False,\n  image_size=(img_height, img_width),\n  batch_size=batch_size)\n\ntest_ds = tf.keras.utils.image_dataset_from_directory(\n    new_test_dir,\n    shuffle=False,\n    image_size=(img_height, img_width),\n    batch_size=batch_size)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-28T15:03:04.528937Z","iopub.execute_input":"2023-07-28T15:03:04.529300Z","iopub.status.idle":"2023-07-28T15:03:04.749246Z","shell.execute_reply.started":"2023-07-28T15:03:04.529268Z","shell.execute_reply":"2023-07-28T15:03:04.748359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# improve performance\nAUTOTUNE = tf.data.AUTOTUNE\n\ntrain_ds = train_ds.prefetch(buffer_size=AUTOTUNE)\nvalidation_ds = validation_ds.prefetch(buffer_size=AUTOTUNE)\ntest_ds = test_ds.prefetch(buffer_size=AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2023-07-28T15:03:08.917382Z","iopub.execute_input":"2023-07-28T15:03:08.918256Z","iopub.status.idle":"2023-07-28T15:03:08.927481Z","shell.execute_reply.started":"2023-07-28T15:03:08.918211Z","shell.execute_reply":"2023-07-28T15:03:08.926520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# Convolutional Block Attention Module (CBAM)","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.layers import GlobalAveragePooling2D, GlobalMaxPooling2D, Reshape, Dense, Input\nfrom tensorflow.keras.layers import Activation, Concatenate, Conv2D, Multiply\n\n\ndef channel_attention_module(x, ratio=8):\n    batch, _, _, channel = x.shape\n\n    ## Shared layers\n    l1 = Dense(channel//ratio, activation=\"relu\", use_bias=False)\n    l2 = Dense(channel, use_bias=False)\n\n    ## Global Average Pooling\n    x1 = GlobalAveragePooling2D()(x)\n    x1 = l1(x1)\n    x1 = l2(x1)\n\n    ## Global Max Pooling\n    x2 = GlobalMaxPooling2D()(x)\n    x2 = l1(x2)\n    x2 = l2(x2)\n\n    ## Add both the features and pass through sigmoid\n    feats = x1 + x2\n    feats = Activation(\"sigmoid\")(feats)\n    feats = Multiply()([x, feats])\n\n    return feats\n\ndef spatial_attention_module(x):\n    ## Average Pooling\n    x1 = tf.reduce_mean(x, axis=-1)\n    x1 = tf.expand_dims(x1, axis=-1)\n\n    ## Max Pooling\n    x2 = tf.reduce_max(x, axis=-1)\n    x2 = tf.expand_dims(x2, axis=-1)\n\n    ## Concatenat both the features\n    feats = Concatenate()([x1, x2])\n    ## Conv layer\n    feats = Conv2D(1, kernel_size=7, padding=\"same\", activation=\"sigmoid\")(feats)\n    feats = Multiply()([x, feats])\n\n    return feats\n\ndef cbam(x):\n    x = channel_attention_module(x)\n    x = spatial_attention_module(x)\n    return x\n\ndef SE(inputs, ratio=8):\n    b, _, _, c = inputs.shape\n    x = GlobalAveragePooling2D()(inputs)\n    x = Dense(c//ratio, activation=\"relu\", use_bias=False)(x)\n    x = Dense(c, activation=\"sigmoid\", use_bias=False)(x)\n    x = inputs * x\n    return x\n\nif __name__ == \"__main__\":\n    inputs = Input(shape=(128, 128, 32))\n    y = cbam(inputs)\n    print(y.shape)\n    inputs = Input(shape=(128, 128, 32))\n    y = SE(inputs)\n    print(y.shape)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.349132Z","iopub.status.idle":"2023-07-28T14:58:32.352291Z","shell.execute_reply.started":"2023-07-28T14:58:32.352045Z","shell.execute_reply":"2023-07-28T14:58:32.352070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transfer learning\nimg_shape = (img_height, img_width, 3,)\n#preprocess_input = tf.keras.applications.mobilenet_v2.preprocess_input\n#base_model = tf.keras.applications.MobileNetV2(input_shape=img_shape,\n                                               #include_top=False,\n                                               #weights='imagenet')\nbase_model = tf.keras.applications.EfficientNetB0(include_top=False, weights='imagenet', input_shape = img_shape)###############\nbase_model.trainable = False\nglobal_max_layer = tf.keras.layers.GlobalAveragePooling2D()\n#prediction_layer = tf.keras.layers.Dense(num_classes, activation='softmax')\n\n\nRegularizer = L1L2(l1=1e-4, l2=1e-3)\n# build\ninputs = tf.keras.Input(shape=img_shape)\nx = base_model(inputs, training=False)\nx = cbam(x)\n#x = SE(x)\nx = global_max_layer(x)\nx = tf.keras.layers.Dropout(0.5)(x)\nx = tf.keras.layers.Dense(320, activation='softmax')(x)\nx = tf.keras.layers.Dropout(0.25)(x)\noutputs = tf.keras.layers.Dense(5, activation='softmax',kernel_regularizer=Regularizer)(x)\n\n\n#outputs = tf.keras.layers.Dense(5, activation='softmax')(x)\nmodel = tf.keras.Model(inputs, outputs)\n# compile\n\nbase_learning_rate = 0.0001 ###################\nmodel.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=base_learning_rate), loss = 'sparse_categorical_crossentropy', metrics=['accuracy'])\n\n","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.353660Z","iopub.status.idle":"2023-07-28T14:58:32.354399Z","shell.execute_reply.started":"2023-07-28T14:58:32.354160Z","shell.execute_reply":"2023-07-28T14:58:32.354182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.summary()","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.355761Z","iopub.status.idle":"2023-07-28T14:58:32.362855Z","shell.execute_reply.started":"2023-07-28T14:58:32.362568Z","shell.execute_reply":"2023-07-28T14:58:32.362593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from keras.callbacks import EarlyStopping\n\n# Define the EarlyStopping callback\nearly_stopping = EarlyStopping(monitor='val_loss', patience=10, mode='min', verbose=1)\n\ninitial_epochs = 20\nhistory = model.fit(train_ds,\n                    epochs=initial_epochs,\n                    validation_data=validation_ds, callbacks=[early_stopping])","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.364198Z","iopub.status.idle":"2023-07-28T14:58:32.364956Z","shell.execute_reply.started":"2023-07-28T14:58:32.364718Z","shell.execute_reply":"2023-07-28T14:58:32.364740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Evaluate the model on the test dataset\nloss, accuracy = model.evaluate(test_ds)\n\n# Print the evaluation results\nprint('Test Loss:', loss)\nprint('Test Accuracy:', accuracy)","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.366283Z","iopub.status.idle":"2023-07-28T14:58:32.367040Z","shell.execute_reply.started":"2023-07-28T14:58:32.366801Z","shell.execute_reply":"2023-07-28T14:58:32.366832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_model.trainable = True\nmodel.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=base_learning_rate/10), loss = 'sparse_categorical_crossentropy', metrics=['accuracy'])\n\n\n# Fine-tune from this layer onwards\nfine_tune_at = 0\n\n# Freeze all the layers before the `fine_tune_at` layer\nfor layer in base_model.layers[:fine_tune_at]:\n    layer.trainable = False\n\nfine_tune_epochs = 20\ntotal_epochs =  initial_epochs + fine_tune_epochs\n\nhistory_fine = model.fit(train_ds,\n                         epochs=total_epochs,\n                         initial_epoch=history.epoch[-1],\n                         validation_data=validation_ds, callbacks=[early_stopping])\n","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.368353Z","iopub.status.idle":"2023-07-28T14:58:32.369088Z","shell.execute_reply.started":"2023-07-28T14:58:32.368854Z","shell.execute_reply":"2023-07-28T14:58:32.368876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Confusion matrix, f1-score","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix, classification_report\n\n# Evaluate the model on the test dataset\nloss, accuracy = model.evaluate(test_ds)\n\n# Print the evaluation results\nprint('Test Loss:', loss)\nprint('Test Accuracy:', accuracy)\n\ny_pred = model.predict(test_ds)\ny_pred = np.argmax(y_pred, axis=1)\ny_test = tf.concat([y for x, y in test_ds], axis=0).numpy()\nprint(y_test)\nprint(y_pred)\n# Print f1, precision, and recall scores\nprint('overall precision, recall and f1:')\nprint(precision_score(y_test, y_pred , average=\"macro\"))\nprint(recall_score(y_test, y_pred , average=\"macro\"))\nprint(f1_score(y_test, y_pred , average=\"macro\"))\nprint('report f1 on each class: ')\nprint(classification_report(y_test, y_pred))\n","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.370398Z","iopub.status.idle":"2023-07-28T14:58:32.371163Z","shell.execute_reply.started":"2023-07-28T14:58:32.370927Z","shell.execute_reply":"2023-07-28T14:58:32.370957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Accuracy, Loss Graph\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nacc = history.history['accuracy']\nval_acc = history.history['val_accuracy']\n\nloss = history.history['loss']\nval_loss = history.history['val_loss']\n\nplt.figure(figsize=(8, 8))\nplt.subplot(2, 1, 1)\nplt.plot(acc, label='Training Accuracy')\nplt.plot(val_acc, label='Validation Accuracy')\nplt.legend(loc='lower right')\nplt.ylabel('Accuracy')\nplt.ylim([min(plt.ylim()),1])\nplt.title('Training and Validation Accuracy')\n\nplt.subplot(2, 1, 2)\nplt.plot(loss, label='Training Loss')\nplt.plot(val_loss, label='Validation Loss')\nplt.legend(loc='upper right')\nplt.ylabel('Cross Entropy')\nplt.ylim([0,4])\nplt.title('Training and Validation Loss')\nplt.xlabel('epoch')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-28T14:58:32.379217Z","iopub.status.idle":"2023-07-28T14:58:32.379968Z","shell.execute_reply.started":"2023-07-28T14:58:32.379730Z","shell.execute_reply":"2023-07-28T14:58:32.379751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Other dataset (apply preprocessing + move to test ds)","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n","metadata":{},"execution_count":null,"outputs":[]}]}