{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":206671331,"sourceType":"kernelVersion"}],"dockerImageVersionId":30580,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# CZII UNET w/Callback, Max Epochs","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/code/stpeteishii/czii-prepare-image-dataset-for-unet","metadata":{"papermill":{"duration":0.00738,"end_time":"2023-10-31T14:35:49.360958","exception":false,"start_time":"2023-10-31T14:35:49.353578","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"An effective method to find the optimal number of epochs.Save the predicted mask images generated every 3 epoch and try arranging them later. You can see the best number of epochs at a glance. If the epoch number is too large, the mask image will be lost.","metadata":{"papermill":{"duration":0.007482,"end_time":"2023-10-31T14:35:49.37588","exception":false,"start_time":"2023-10-31T14:35:49.368398","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install imantics --quiet","metadata":{"papermill":{"duration":13.989536,"end_time":"2023-10-31T14:36:03.372848","exception":false,"start_time":"2023-10-31T14:35:49.383312","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nimport tensorflow as tf\nimport json\nimport os\nimport imantics\nfrom PIL import Image\nfrom skimage.transform import resize\nimport random\nfrom sklearn.model_selection import train_test_split\n%matplotlib inline","metadata":{"papermill":{"duration":9.744181,"end_time":"2023-10-31T14:36:13.125092","exception":false,"start_time":"2023-10-31T14:36:03.380911","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"base_dir = '/kaggle/input/czii-prepare-image-dataset-for-unet'\nimages_dir = f'{base_dir}/image' \nmasks_dir = f'{base_dir}/mask' \ntimages_dir = f'{base_dir}/timage' ","metadata":{"papermill":{"duration":0.014689,"end_time":"2023-10-31T14:36:13.147877","exception":false,"start_time":"2023-10-31T14:36:13.133188","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.listdir(images_dir)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images_listdir = os.listdir(images_dir)\nrandom_images = np.random.choice(images_listdir, size = 9, replace = False)","metadata":{"papermill":{"duration":0.022329,"end_time":"2023-10-31T14:36:13.178043","exception":false,"start_time":"2023-10-31T14:36:13.155714","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_size=512","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_image(path):\n    img = cv2.imread(path)\n    #img = cv2.imread(path,cv2.IMREAD_ANYDEPTH)\n    #img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img2 = cv2.resize(img, (image_size, image_size))\n    #print(img.shape,img2.shape)\n    return img2","metadata":{"papermill":{"duration":0.014979,"end_time":"2023-10-31T14:36:13.222538","exception":false,"start_time":"2023-10-31T14:36:13.207559","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Input images ","metadata":{"papermill":{"duration":0.007468,"end_time":"2023-10-31T14:36:13.237532","exception":false,"start_time":"2023-10-31T14:36:13.230064","status":"completed"},"tags":[]}},{"cell_type":"code","source":"rows = 3\ncols = 3\nfig, ax = plt.subplots(rows, cols, figsize = (12,12))\nfor i, ax in enumerate(ax.flat):\n    if i < len(random_images):\n        img = read_image(f\"{images_dir}/{random_images[i]}\")\n        ax.set_title(f\"{random_images[i]}\")\n        ax.imshow(img)\n        ax.axis('off')","metadata":{"papermill":{"duration":1.589493,"end_time":"2023-10-31T14:36:14.834556","exception":false,"start_time":"2023-10-31T14:36:13.245063","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Ground truth masks","metadata":{"papermill":{"duration":0.01276,"end_time":"2023-10-31T14:36:14.860722","exception":false,"start_time":"2023-10-31T14:36:14.847962","status":"completed"},"tags":[]}},{"cell_type":"code","source":"fig, ax = plt.subplots(rows, cols, figsize = (12,12))\nfor i, ax in enumerate(ax.flat):\n    if i < len(random_images):\n        file = random_images[i]\n        img = read_image(f\"{masks_dir}/{file}\")       \n        #img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        ax.set_title(f\"{random_images[i]}\")\n        ax.imshow(img)\n        ax.axis('off')","metadata":{"papermill":{"duration":1.425215,"end_time":"2023-10-31T14:36:16.298726","exception":false,"start_time":"2023-10-31T14:36:14.873511","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#MASKS=np.zeros((1,image_size, image_size, 1), dtype=bool) \nMASKS=np.zeros((1,image_size,image_size,3),dtype=np.uint8)\nIMAGES=np.zeros((1,image_size,image_size,3),dtype=np.uint8)\n\nfor j,file in enumerate(images_listdir[0:81]): ##the smaller, the faster\n    #print(j)\n    image = read_image(f\"{images_dir}/{file}\")\n    image_ex = np.expand_dims(image, axis=0)\n    IMAGES = np.vstack([IMAGES, image_ex])\n    \n    file2=file\n    mask = read_image(f\"{masks_dir}/{file2}\")\n    #mask = cv2.cvtColor(mask, cv2.COLOR_BGR2GRAY) #####\n    mask_ex = np.expand_dims(mask, axis=0)    \n    MASKS = np.vstack([MASKS, mask_ex])\n","metadata":{"papermill":{"duration":1.769615,"end_time":"2023-10-31T14:36:18.112401","exception":false,"start_time":"2023-10-31T14:36:16.342786","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images=np.array(IMAGES)[1:81]\nmasks=np.array(MASKS)[1:81]\nprint(images.shape,masks.shape)","metadata":{"papermill":{"duration":0.040793,"end_time":"2023-10-31T14:36:18.168251","exception":false,"start_time":"2023-10-31T14:36:18.127458","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images_train, images_test, masks_train, masks_test = train_test_split(\n    images, masks, test_size=0.3, random_state=42)","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.038918,"end_time":"2023-10-31T14:36:18.221744","exception":false,"start_time":"2023-10-31T14:36:18.182826","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(images_train), len(masks_train))\nprint(len(images_test), len(masks_test))","metadata":{"papermill":{"duration":0.022288,"end_time":"2023-10-31T14:36:18.259006","exception":false,"start_time":"2023-10-31T14:36:18.236718","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# U-Net","metadata":{"papermill":{"duration":0.014232,"end_time":"2023-10-31T14:36:18.287583","exception":false,"start_time":"2023-10-31T14:36:18.273351","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def conv_block(input, num_filters):\n    conv = tf.keras.layers.Conv2D(num_filters, 3, padding=\"same\")(input)\n    conv = tf.keras.layers.BatchNormalization()(conv)\n    conv = tf.keras.layers.Activation(\"relu\")(conv)\n    conv = tf.keras.layers.Conv2D(num_filters, 3, padding=\"same\")(conv)\n    conv = tf.keras.layers.BatchNormalization()(conv)\n    conv = tf.keras.layers.Activation(\"relu\")(conv)\n    return conv\n\ndef encoder_block(input, num_filters):\n    skip = conv_block(input, num_filters)\n    pool = tf.keras.layers.MaxPool2D((2,2))(skip)\n    return skip, pool\n\ndef decoder_block(input, skip, num_filters):\n    up_conv = tf.keras.layers.Conv2DTranspose(num_filters, (2,2), strides=2, padding=\"same\")(input)\n    conv = tf.keras.layers.Concatenate()([up_conv, skip])\n    conv = conv_block(conv, num_filters)\n    return conv\n\ndef Unet(input_shape):\n    inputs = tf.keras.layers.Input(input_shape)\n\n    skip1, pool1 = encoder_block(inputs, 64)\n    skip2, pool2 = encoder_block(pool1, 128)\n    skip3, pool3 = encoder_block(pool2, 256)\n    skip4, pool4 = encoder_block(pool3, 512)\n\n    bridge = conv_block(pool4, 1024)\n\n    decode1 = decoder_block(bridge, skip4, 512)\n    decode2 = decoder_block(decode1, skip3, 256)\n    decode3 = decoder_block(decode2, skip2, 128)\n    decode4 = decoder_block(decode3, skip1, 64)\n\n    outputs = tf.keras.layers.Conv2D(3, 1, padding=\"same\", activation=\"sigmoid\")(decode4) #####\n\n    model = tf.keras.models.Model(inputs, outputs, name=\"U-Net\")\n    return model\n\nunet_model = Unet((512,512,3)) # Unet((512,512,3)) # actual size(1041,1511,3)\nunet_model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])\nunet_model.summary()","metadata":{"_kg_hide-output":true,"papermill":{"duration":3.768253,"end_time":"2023-10-31T14:36:22.070257","exception":false,"start_time":"2023-10-31T14:36:18.302004","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Callback","metadata":{"papermill":{"duration":0.028918,"end_time":"2023-10-31T14:36:22.128723","exception":false,"start_time":"2023-10-31T14:36:22.099805","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#!rm -rf output_masks","metadata":{"papermill":{"duration":0.035785,"end_time":"2023-10-31T14:36:22.193477","exception":false,"start_time":"2023-10-31T14:36:22.157692","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow import keras\n\nclass SaveMaskImagesCallback(keras.callbacks.Callback):\n    def __init__(self, output_dir, images, masks, save_every_n_epochs=1):\n        super(SaveMaskImagesCallback, self).__init__()\n        self.output_dir = output_dir\n        self.images = images\n        self.masks = masks\n        self.save_every_n_epochs = save_every_n_epochs\n\n    def on_epoch_end(self, epoch, logs=None):\n        if (epoch + 1) % self.save_every_n_epochs == 0:\n            predictions = self.model.predict(self.images)\n            for i, pred_mask in enumerate(predictions):\n                output_path = os.path.join(self.output_dir, f\"mask_{i:02d}_ep_{epoch + 1:03d}.png\")\n                pred_mask = (pred_mask * 255).astype(np.uint8)\n                keras.preprocessing.image.save_img(output_path, pred_mask)\n\noutput_directory = \"output_masks\"  \nos.makedirs(output_directory, exist_ok=True)\nmask_callback = SaveMaskImagesCallback(output_directory, images_train, masks_train)","metadata":{"papermill":{"duration":0.039875,"end_time":"2023-10-31T14:36:22.262429","exception":false,"start_time":"2023-10-31T14:36:22.222554","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Train with 40 epochs","metadata":{"papermill":{"duration":0.028873,"end_time":"2023-10-31T14:36:22.320261","exception":false,"start_time":"2023-10-31T14:36:22.291388","status":"completed"},"tags":[]}},{"cell_type":"code","source":"unet_result = unet_model.fit(\n    images_train, masks_train, \n    validation_split=0.2, batch_size=2, epochs=40,\n    callbacks=[mask_callback]\n)","metadata":{"papermill":{"duration":380.649972,"end_time":"2023-10-31T14:42:42.999007","exception":false,"start_time":"2023-10-31T14:36:22.349035","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Predict","metadata":{"papermill":{"duration":0.103907,"end_time":"2023-10-31T14:42:43.416004","exception":false,"start_time":"2023-10-31T14:42:43.312097","status":"completed"},"tags":[]}},{"cell_type":"code","source":"unet_predict = unet_model.predict(images_test)","metadata":{"papermill":{"duration":12.243025,"end_time":"2023-10-31T14:42:55.762948","exception":false,"start_time":"2023-10-31T14:42:43.519923","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_result(idx, og, unet, target, p):\n\n    print(unet.flatten().min(),unet.flatten().max())\n    \n    fig, axs = plt.subplots(1, 3, figsize=(12,12))\n    axs[0].set_title(\"Original \"+str(idx))\n    axs[0].imshow(og)\n    axs[0].axis('off')\n    \n    axs[1].set_title(\"U-Net: p>\"+str(p))\n    axs[1].imshow(unet*255)\n    axs[1].axis('off')\n    \n    axs[2].set_title(\"Ground Truth\")\n    axs[2].imshow(target)\n    axs[2].axis('off')\n\n    plt.show()","metadata":{"papermill":{"duration":0.114414,"end_time":"2023-10-31T14:42:55.982901","exception":false,"start_time":"2023-10-31T14:42:55.868487","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"r1,r2,r3,r4=0.7,0.8,0.9,0.99","metadata":{"papermill":{"duration":0.112198,"end_time":"2023-10-31T14:42:56.200267","exception":false,"start_time":"2023-10-31T14:42:56.088069","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unet_predict1 = (unet_predict > r1).astype(np.uint8)\nunet_predict2 = (unet_predict > r2).astype(np.uint8)\nunet_predict3 = (unet_predict > r3).astype(np.uint8)\nunet_predict4 = (unet_predict > r4).astype(np.uint8)","metadata":{"papermill":{"duration":0.134428,"end_time":"2023-10-31T14:42:56.440017","exception":false,"start_time":"2023-10-31T14:42:56.305589","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"show_test_idx = random.sample(range(len(unet_predict)), 3)\nfor idx in show_test_idx: \n    show_result(idx, images_test[idx], unet_predict1[idx], masks_test[idx], r1)\n    show_result(idx, images_test[idx], unet_predict2[idx], masks_test[idx], r2)\n    show_result(idx, images_test[idx], unet_predict3[idx], masks_test[idx], r3)\n    show_result(idx, images_test[idx], unet_predict4[idx], masks_test[idx], r4)\n    print()","metadata":{"papermill":{"duration":5.183271,"end_time":"2023-10-31T14:43:01.729239","exception":false,"start_time":"2023-10-31T14:42:56.545968","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.125499,"end_time":"2023-10-31T14:43:01.9836","exception":false,"start_time":"2023-10-31T14:43:01.858101","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}