{"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":"!git clone https://github.com/jakeret/unet\n!pip install ../working/unet/","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-11-19T17:52:03.769281Z","iopub.execute_input":"2022-11-19T17:52:03.769765Z","iopub.status.idle":"2022-11-19T17:52:18.796655Z","shell.execute_reply.started":"2022-11-19T17:52:03.76973Z","shell.execute_reply":"2022-11-19T17:52:18.795253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport os\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm import tqdm\nfrom sklearn.utils import shuffle\nfrom sklearn.utils import class_weight\nfrom sklearn.preprocessing import minmax_scale\nimport random\nimport cv2\nfrom imgaug import augmenters as iaa\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers\nfrom tensorflow.keras import Input\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import Dense, Flatten, Dropout, Activation\nfrom tensorflow.keras.layers import BatchNormalization, GlobalAveragePooling2D\nfrom tensorflow.keras.callbacks import ModelCheckpoint, ReduceLROnPlateau, EarlyStopping\n\nfrom unet import utils\nfrom unet.datasets import circles\nimport unet","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:18.801014Z","iopub.execute_input":"2022-11-19T17:52:18.801415Z","iopub.status.idle":"2022-11-19T17:52:18.813573Z","shell.execute_reply.started":"2022-11-19T17:52:18.801376Z","shell.execute_reply":"2022-11-19T17:52:18.811333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_folder = '../input/cassava-leaf-disease-classification/train_images/'\nsamples_df = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\nsamples_df[\"label\"] = samples_df[\"label\"].astype(\"str\")\nsamples_df.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:18.81472Z","iopub.execute_input":"2022-11-19T17:52:18.815058Z","iopub.status.idle":"2022-11-19T17:52:18.885536Z","shell.execute_reply.started":"2022-11-19T17:52:18.815029Z","shell.execute_reply":"2022-11-19T17:52:18.884446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"samples_df = samples_df.query(\"label=='4'\")","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:18.88914Z","iopub.execute_input":"2022-11-19T17:52:18.889782Z","iopub.status.idle":"2022-11-19T17:52:18.902424Z","shell.execute_reply.started":"2022-11-19T17:52:18.88974Z","shell.execute_reply":"2022-11-19T17:52:18.901174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_percentage = 0.8\ntraining_item_count = int(len(samples_df)*training_percentage)\nvalidation_item_count = len(samples_df)-int(len(samples_df)*training_percentage)\ntraining_df = samples_df[:training_item_count]\nvalidation_df = samples_df[training_item_count:]","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:18.904498Z","iopub.execute_input":"2022-11-19T17:52:18.905009Z","iopub.status.idle":"2022-11-19T17:52:18.912672Z","shell.execute_reply.started":"2022-11-19T17:52:18.904964Z","shell.execute_reply":"2022-11-19T17:52:18.911272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_ECI_band(img):\n    '''\n    Return the ECI band calculated from an RGB image between 0 and 255\n    using the formula below:    \n    ECI = (red_channel-1)^2 + green_channel^2/0.16\n    '''\n    img = img/255.\n    img = cv2.GaussianBlur(img,(35,35),0)\n    ECI_band = np.power(img[:,:,0]-1,2) + np.power(img[:,:,1],2)/0.16\n    normalized_ECI_band = (ECI_band/ECI_band.max()*255).astype(np.uint8)\n    return normalized_ECI_band\n\n\ndef get_CIVE_band(img):\n    '''\n    Return the CIVE band calculated from an RGB image between 0 and 255\n    using the formula below:\n    CIVE = 0.441*red_channel - 0.881*green_channel + 0.385*blue_channel + 18.787\n    '''\n    img = cv2.GaussianBlur(img,(35,35),0)\n    CIVE_band = 0.441*img[:,:,0] - 0.881*img[:,:,1] + 0.385*img[:,:,2] + 18.787\n    normalized_CIVE_band = (((CIVE_band+abs(CIVE_band.min()))/CIVE_band.max())).astype(np.uint8)\n    return normalized_CIVE_band\n\n\ndef apply_ECI_mask(img, vegetation_index_band):\n    '''\n    Apply a binary mask on an image and return the masked image\n    '''\n    ret, otsu = cv2.threshold(vegetation_index_band,0,255,cv2.THRESH_BINARY+cv2.THRESH_OTSU)\n    masked_img = cv2.bitwise_and(img,img,mask = otsu)\n    return masked_img\n\n\ndef apply_CIVE_mask(img, vegetation_index_band):\n    '''\n    Apply a binary mask on an image and return the masked image\n    '''\n    ret, otsu = cv2.threshold(vegetation_index_band,0,255,cv2.THRESH_BINARY_INV+cv2.THRESH_OTSU)\n    masked_img = cv2.bitwise_and(img,img,mask = otsu)\n    return masked_img","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:18.9147Z","iopub.execute_input":"2022-11-19T17:52:18.915141Z","iopub.status.idle":"2022-11-19T17:52:18.929103Z","shell.execute_reply.started":"2022-11-19T17:52:18.915108Z","shell.execute_reply":"2022-11-19T17:52:18.927736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 20))\nitems = 6\nfor idx, image_id in enumerate(samples_df.image_id[:items]):\n    img_path = training_folder+image_id\n    img = np.array(Image.open(img_path))\n    ax = plt.subplot(items, 3, idx*3 + 1)\n    ax.set_title(\"original\")\n    plt.imshow(img)\n    \n    ECI_band =  get_ECI_band(img)\n    ax = plt.subplot(items, 3, idx*3 + 2)\n    ax.set_title(\"ECI band\")\n    plt.imshow(ECI_band)\n    \n    masked_img = apply_ECI_mask(img, ECI_band)\n    ax = plt.subplot(items, 3, idx*3 + 3)\n    ax.set_title(\"ECI+Otsu\")\n    plt.imshow(masked_img)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:18.930475Z","iopub.execute_input":"2022-11-19T17:52:18.930814Z","iopub.status.idle":"2022-11-19T17:52:23.346387Z","shell.execute_reply.started":"2022-11-19T17:52:18.930783Z","shell.execute_reply":"2022-11-19T17:52:23.34525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 20))\nitems = 6\nfor idx, image_id in enumerate(samples_df.image_id[:items]):\n    img_path = training_folder+image_id\n    img = np.array(Image.open(img_path))\n    ax = plt.subplot(items, 3, idx*3 + 1)\n    ax.set_title(\"original\")\n    plt.imshow(img)\n    \n    CIVE_band = get_CIVE_band(img)\n    ax = plt.subplot(items, 3, idx*3 + 2)\n    ax.set_title(\"CIVE band\")\n    plt.imshow(CIVE_band)\n    \n    masked_img = apply_CIVE_mask(img, CIVE_band)\n    ax = plt.subplot(items, 3, idx*3 + 3)\n    ax.set_title(\"CIVE+Otsu\")\n    plt.imshow(masked_img)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:23.347959Z","iopub.execute_input":"2022-11-19T17:52:23.348726Z","iopub.status.idle":"2022-11-19T17:52:27.586554Z","shell.execute_reply.started":"2022-11-19T17:52:23.348686Z","shell.execute_reply":"2022-11-19T17:52:27.5853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 20))\nitems = 6\nfor idx, image_id in enumerate(samples_df.image_id[10:10+items]):\n    img_path = training_folder+image_id\n    img = np.array(Image.open(img_path))\n    ax = plt.subplot(items, 3, idx*3 + 1)\n    ax.set_title(\"original\")\n    plt.imshow(img)\n    \n    ECI_band =  get_ECI_band(img)\n    CIVE_band = get_CIVE_band(img)\n\n    masked_img = apply_ECI_mask(img, ECI_band)\n    ax = plt.subplot(items, 3, idx*3 + 2)\n    ax.set_title(\"ECI+Otsu\")\n    plt.imshow(masked_img)\n    \n    masked_img = apply_CIVE_mask(img, CIVE_band)\n    ax = plt.subplot(items, 3, idx*3 + 3)\n    ax.set_title(\"CIVE+Otsu\")\n    plt.imshow(masked_img)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:27.58806Z","iopub.execute_input":"2022-11-19T17:52:27.588442Z","iopub.status.idle":"2022-11-19T17:52:32.090659Z","shell.execute_reply.started":"2022-11-19T17:52:27.588409Z","shell.execute_reply":"2022-11-19T17:52:32.089378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_input_image(img_path, crop_size=512):\n    '''\n    Randomly select a 512x512 pixel area from the full image\n    '''\n    img = Image.open(img_path)\n    img_height, img_width = img.size\n    img = np.array(img)\n    y = random.randint(0,img_height-crop_size)\n    x = random.randint(0,img_width-crop_size)\n    cropped_img = img[x:x+crop_size , y:y+crop_size,:]\n    \n    return cropped_img\n\n\ndef get_groundtruth_mask(img, crop_size=512):\n    '''\n    Generate the groundtruth mask using the CIVE band\n    '''\n    img = cv2.GaussianBlur(img,(35,35),0)\n    cive_band = 0.441*img[:,:,0] - 0.881*img[:,:,1] + 0.385*img[:,:,2] + 18.787\n    normalized_cive_band = (((cive_band+abs(cive_band.min()))/cive_band.max())).astype(np.uint8)\n    ret, otsu_mask = cv2.threshold(normalized_cive_band,0,1,cv2.THRESH_BINARY_INV+cv2.THRESH_OTSU)\n    \n    veg_mask = otsu_mask.astype(np.uint8)\n    soil_mask = np.where(otsu_mask==1, 0, 1).astype(np.uint8)\n    masks = np.transpose(np.array([soil_mask, veg_mask]),(1,2,0))\n    \n    return masks","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:32.094954Z","iopub.execute_input":"2022-11-19T17:52:32.096295Z","iopub.status.idle":"2022-11-19T17:52:32.106312Z","shell.execute_reply.started":"2022-11-19T17:52:32.096248Z","shell.execute_reply":"2022-11-19T17:52:32.104901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"augmentation = iaa.Sequential([\n    iaa.Multiply((0.5, 1.5)),\n    iaa.Affine(scale={\"x\": (1, 1.2), \"y\":(1, 1.2)}),\n    iaa.Sometimes(0.4,\n                 iaa.GaussianBlur(sigma=(0,2))),\n    iaa.Sometimes(0.3,\n                 iaa.Grayscale(alpha=(0.0, 1.0))),\n    iaa.Sometimes(0.3,\n                 iaa.SigmoidContrast(gain=(3, 6), cutoff=(0.4, 0.6))),\n    iaa.Sometimes(0.3,\n        iaa.CoarseDropout((0.0, 0.05), size_percent=(0.05, 0.6), per_channel=0.5))\n])","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:32.108031Z","iopub.execute_input":"2022-11-19T17:52:32.10858Z","iopub.status.idle":"2022-11-19T17:52:32.126123Z","shell.execute_reply.started":"2022-11-19T17:52:32.108518Z","shell.execute_reply":"2022-11-19T17:52:32.124539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = [get_input_image(training_folder+image_filename) for image_filename in samples_df[:10].image_id.values]\nlabels = samples_df[:10].label.values\naugmented_images = augmentation.augment_images(images=images)\n\nsample_number = len(augmented_images)\nfig = plt.figure(figsize = (20,sample_number))\nfor i in range(0,sample_number):\n    ax = fig.add_subplot(2, 5, i+1)\n    ax.imshow(augmented_images[i])\n    ax.set_title(str(labels[i]))\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:32.127847Z","iopub.execute_input":"2022-11-19T17:52:32.128239Z","iopub.status.idle":"2022-11-19T17:52:34.363456Z","shell.execute_reply.started":"2022-11-19T17:52:32.128178Z","shell.execute_reply":"2022-11-19T17:52:34.361607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def custom_generator(image_path_list, folder, batch_size=16, training_mode=True):\n    \n    while True:\n        for start in range(0, len(image_path_list), batch_size):\n            X_batch = []\n            Y_batch = []\n            end = min(start + batch_size, training_item_count)\n            \n            image_list = [get_input_image(folder+\"/\"+image_path) for image_path in image_path_list[start:end]]\n            groundtruth_image_list = [get_groundtruth_mask(image) for image in image_list]\n            \n            #only apply augmentation during training\n            if training_mode:\n                image_list = augmentation.augment_images(images=image_list)\n\n            X_batch = np.array(image_list)/255.\n            Y_batch = np.array(groundtruth_image_list)\n\n            yield X_batch, Y_batch","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:34.364809Z","iopub.execute_input":"2022-11-19T17:52:34.365149Z","iopub.status.idle":"2022-11-19T17:52:34.374413Z","shell.execute_reply.started":"2022-11-19T17:52:34.36512Z","shell.execute_reply":"2022-11-19T17:52:34.37326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 8\ncallbacks = [ReduceLROnPlateau(monitor='val_loss', patience=1, verbose=1, factor=0.5),\n             EarlyStopping(monitor='val_loss', patience=3),\n             ModelCheckpoint(filepath='best_model.h5', monitor='val_loss', save_best_only=True)]","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:34.376415Z","iopub.execute_input":"2022-11-19T17:52:34.376961Z","iopub.status.idle":"2022-11-19T17:52:34.389077Z","shell.execute_reply.started":"2022-11-19T17:52:34.376905Z","shell.execute_reply":"2022-11-19T17:52:34.387713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unet_model = unet.build_model(512,\n                              channels=3,\n                              num_classes=2,\n                              layer_depth=4,\n                              filters_root=64,\n                              padding=\"same\")\n\nunet.finalize_model(unet_model,\n                    loss=tf.keras.losses.BinaryCrossentropy(),\n                    metrics=[tf.keras.metrics.BinaryAccuracy()],\n                    auc=False,\n                    learning_rate=1e-4)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T17:52:34.390942Z","iopub.execute_input":"2022-11-19T17:52:34.391399Z","iopub.status.idle":"2022-11-19T17:52:34.830553Z","shell.execute_reply.started":"2022-11-19T17:52:34.391365Z","shell.execute_reply":"2022-11-19T17:52:34.829465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = unet_model.fit_generator(custom_generator(training_df[\"image_id\"], training_folder, batch_size=batch_size, training_mode=True),\n                  steps_per_epoch = int(len(training_df)/batch_size),\n                  epochs = 20, \n                  validation_data=custom_generator(validation_df[\"image_id\"], training_folder, batch_size=batch_size),\n                  validation_steps=int(len(validation_df)/batch_size),\n                  callbacks=callbacks)","metadata":{"execution":{"iopub.status.busy":"2022-11-19T18:34:05.580335Z","iopub.execute_input":"2022-11-19T18:34:05.581187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unet_model.load_weights(\"best_model.h5\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cat ../input/cassava-leaf-disease-classification/label_num_to_disease_map.json","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_folder = '../input/cassava-leaf-disease-classification/train_images/'\ndisease_df = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\ndisease_df = disease_df.query(\"label!=4\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 20))\nitems = 5\nfor idx, image_id in enumerate(validation_df.image_id[:items]):\n    img_path = training_folder+image_id\n    img = get_input_image(img_path)\n    ax = plt.subplot(items, 3, idx*3 + 1)\n    ax.set_title(\"original\")\n    plt.imshow(img)\n    \n    CIVE_band = get_CIVE_band(img)\n    masked_img = apply_CIVE_mask(img, CIVE_band)\n    ax = plt.subplot(items, 3, idx*3 + 2)\n    ax.set_title(\"CIVE+Otsu\")\n    plt.imshow(masked_img)\n\n    mask = unet_model.predict(np.array([img]))\n    soil_mask = np.where(np.transpose(mask[0],(2,0,1))[0]<=0.5,1,0).astype(np.uint8)\n    masked_img = cv2.bitwise_and(img,img,mask = soil_mask)\n    ax = plt.subplot(items, 3, idx*3 + 3)\n    ax.set_title(\"Unet\")\n    plt.imshow(masked_img)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 20))\nitems = 5\nfor idx, image_id in enumerate(disease_df.image_id[10:10+items]):\n    img_path = training_folder+image_id\n    img = get_input_image(img_path)\n    ax = plt.subplot(items, 3, idx*3 + 1)\n    ax.set_title(\"original\")\n    plt.imshow(img)\n    \n    CIVE_band = get_CIVE_band(img)\n    masked_img = apply_CIVE_mask(img, CIVE_band)\n    ax = plt.subplot(items, 3, idx*3 + 2)\n    ax.set_title(\"CIVE+Otsu\")\n    plt.imshow(masked_img)\n\n    mask = unet_model.predict(np.array([img]))\n    soil_mask = np.where(np.transpose(mask[0],(2,0,1))[0]<=0.5,1,0).astype(np.uint8)\n    masked_img = cv2.bitwise_and(img,img,mask = soil_mask)\n    ax = plt.subplot(items, 3, idx*3 + 3)\n    ax.set_title(\"Unet\")\n    plt.imshow(masked_img)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}