{"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":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Attention UNet using TensorFlow for SenNet and HOA","metadata":{}},{"cell_type":"markdown","source":"## 1. Imports","metadata":{}},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')\n\nimport os\nimport numpy as np\n\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\n\nimport tensorflow as tf\nfrom tensorflow.keras.layers import Input, Conv2D, MaxPooling2D, UpSampling2D, concatenate, Conv2DTranspose\nfrom tensorflow.keras.layers import multiply, add, Activation, BatchNormalization\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.utils import plot_model\n\nfrom sklearn.model_selection import train_test_split","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-11-14T16:04:27.388688Z","iopub.execute_input":"2023-11-14T16:04:27.389832Z","iopub.status.idle":"2023-11-14T16:04:27.397195Z","shell.execute_reply.started":"2023-11-14T16:04:27.389792Z","shell.execute_reply":"2023-11-14T16:04:27.395861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Checking if GPU is Available","metadata":{}},{"cell_type":"code","source":"# Check if GPU is available and output the device name\ngpu_devices = tf.config.experimental.list_physical_devices('GPU')\nif gpu_devices:\n    print(\"GPU is available:\", gpu_devices)\n    for device in gpu_devices:\n        tf.config.experimental.set_memory_growth(device, True)\nelse:\n    print(\"GPU is not available, using CPU instead.\")","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:04:31.185775Z","iopub.execute_input":"2023-11-14T16:04:31.186174Z","iopub.status.idle":"2023-11-14T16:04:31.605536Z","shell.execute_reply.started":"2023-11-14T16:04:31.186143Z","shell.execute_reply":"2023-11-14T16:04:31.604302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Visualizing Sample Images","metadata":{}},{"cell_type":"code","source":"# Set the base path\nbase_path = '/kaggle/input/blood-vessel-segmentation/train'  \n\n# Replace 'kidney_1_dense' with the dataset you want to explore\ndataset = 'kidney_1_voi'\n\n# Paths to images and labels\nimages_path = os.path.join(base_path, dataset, 'images')\nlabels_path = os.path.join(base_path, dataset, 'labels')\n\n# List the files in the directories\nimage_files = sorted([os.path.join(images_path, f) for f in os.listdir(images_path) if f.endswith('.tif')])\nlabel_files = sorted([os.path.join(labels_path, f) for f in os.listdir(labels_path) if f.endswith('.tif')])\n\n# Function to display a set of images\ndef show_images(images, titles=None, cmap='gray'):\n    n = len(images)\n    fig, axes = plt.subplots(1, n, figsize=(20, 10))\n    if not isinstance(axes, np.ndarray):\n        axes = [axes]\n    for idx, ax in enumerate(axes):\n        ax.imshow(images[idx], cmap=cmap)\n        if titles:\n            ax.set_title(titles[idx])\n        ax.axis('off')\n    plt.tight_layout()\n    plt.show()\n\n# Load and display the first image and its mask\nfirst_image = tiff.imread(image_files[0])\nfirst_label = tiff.imread(label_files[0])\n\nshow_images([first_image, first_label], titles=['First Image', 'First Label'])","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:04:36.533854Z","iopub.execute_input":"2023-11-14T16:04:36.534731Z","iopub.status.idle":"2023-11-14T16:04:38.540211Z","shell.execute_reply.started":"2023-11-14T16:04:36.534696Z","shell.execute_reply":"2023-11-14T16:04:38.539153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Basic statistics about the images\nimage_shapes = [tiff.imread(file).shape for file in image_files]\n\nprint(f\"Number of images: {len(image_files)}\")\nprint(f\"Image shapes: {set(image_shapes)}\")","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:04:43.722073Z","iopub.execute_input":"2023-11-14T16:04:43.722472Z","iopub.status.idle":"2023-11-14T16:06:21.248013Z","shell.execute_reply.started":"2023-11-14T16:04:43.722439Z","shell.execute_reply":"2023-11-14T16:06:21.247032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# It's useful to see the distribution of pixel values\npixel_values = first_image.flatten()\nplt.hist(pixel_values, bins=50, color='blue', alpha=0.7)\nplt.title('Pixel Value Distribution')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:06:21.250258Z","iopub.execute_input":"2023-11-14T16:06:21.250682Z","iopub.status.idle":"2023-11-14T16:06:21.599590Z","shell.execute_reply.started":"2023-11-14T16:06:21.250647Z","shell.execute_reply":"2023-11-14T16:06:21.598374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Preprocessing","metadata":{}},{"cell_type":"code","source":"def preprocess_image(path):\n    # Load the image using tifffile\n    image = tiff.imread(path)\n    \n    # If the image has more than one channel, extract just one channel\n    if image.ndim > 2 and image.shape[2] > 1:\n        image = image[..., 0]\n    \n    # Normalize the image to [0, 1] range\n    image = image / 255.0\n    \n    # Convert image to a TensorFlow tensor\n    image_tensor = tf.convert_to_tensor(image, dtype=tf.float32)\n    \n    # Add a channel dimension if it does not exist\n    if image_tensor.ndim == 2:\n        image_tensor = image_tensor[..., tf.newaxis]\n    \n    # Ensure image tensor is 3D at this point\n    if image_tensor.ndim != 3:\n        raise ValueError('Image tensor must be 3 dimensions [height, width, channels]')\n    \n    # Resize the image to the desired size\n    image_tensor = tf.image.resize(image_tensor, [256, 256])\n    \n    return image_tensor","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:06:21.601207Z","iopub.execute_input":"2023-11-14T16:06:21.601946Z","iopub.status.idle":"2023-11-14T16:06:21.609229Z","shell.execute_reply.started":"2023-11-14T16:06:21.601907Z","shell.execute_reply":"2023-11-14T16:06:21.608199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_mask(path):\n    # Load the mask using tifffile\n    mask = tiff.imread(path)\n    \n    # If the mask has more than one channel, extract just one channel\n    if mask.ndim > 2 and mask.shape[2] > 1:\n        mask = mask[..., 0]\n    \n    # Normalize the mask to be in [0, 1]\n    mask = mask / 255.0 if mask.max() > 1 else mask\n    \n    # Convert mask to a TensorFlow tensor\n    mask_tensor = tf.convert_to_tensor(mask, dtype=tf.float32)\n    \n    # Add a channel dimension if it does not exist\n    if mask_tensor.ndim == 2:\n        mask_tensor = mask_tensor[..., tf.newaxis]\n    \n    # Ensure mask tensor is 3D at this point\n    if mask_tensor.ndim != 3:\n        raise ValueError('Mask tensor must be 3 dimensions [height, width, channels]')\n    \n    # Resize the mask to the desired size\n    mask_tensor = tf.image.resize(mask_tensor, [256, 256], method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)\n    \n    # The resize operation could push the values away from 0 and 1, we threshold to ensure it's a proper mask\n    mask_tensor = tf.where(mask_tensor > 0.5, 1, 0)\n    \n    return mask_tensor","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:06:21.611638Z","iopub.execute_input":"2023-11-14T16:06:21.611961Z","iopub.status.idle":"2023-11-14T16:06:21.623756Z","shell.execute_reply.started":"2023-11-14T16:06:21.611925Z","shell.execute_reply":"2023-11-14T16:06:21.622793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. Data Split","metadata":{}},{"cell_type":"code","source":"# Subset 20% of the dataset for quick experiments\nsubset_size = int(0.2 * len(image_files))\nimage_files_subset = image_files[:subset_size]\nlabel_files_subset = label_files[:subset_size]\n\n# Preprocess and load images into memory (This might take a lot of RAM, be careful with large datasets)\nimages = np.array([preprocess_image(f) for f in image_files_subset])\nmasks = np.array([preprocess_mask(f) for f in label_files_subset])\n\n# Split into train and validation sets\nX_train, X_val, y_train, y_val = train_test_split(images, masks, test_size=0.1, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:06:21.624968Z","iopub.execute_input":"2023-11-14T16:06:21.625243Z","iopub.status.idle":"2023-11-14T16:06:41.586668Z","shell.execute_reply.started":"2023-11-14T16:06:21.625219Z","shell.execute_reply":"2023-11-14T16:06:41.585519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6. Building The Attention UNet Model Architechture","metadata":{}},{"cell_type":"code","source":"def attention_gate(inp_1, inp_2, n_intermediate_channels):\n    inp_1_conv = Conv2D(n_intermediate_channels, (1, 1), padding='same')(inp_1)\n    inp_2_conv = Conv2D(n_intermediate_channels, (1, 1), padding='same')(inp_2)\n    f = add([inp_1_conv, inp_2_conv])\n    f = Activation('relu')(f)\n    g = Conv2D(1, (1, 1), padding='same')(f)\n    gate = Activation('sigmoid')(g)\n\n    return multiply([inp_2, gate])\n\ndef conv_block(input_tensor, num_filters):\n    encoder = Conv2D(num_filters, (3, 3), padding='same')(input_tensor)\n    encoder = Activation('relu')(encoder)\n    encoder = BatchNormalization()(encoder)\n    encoder = Conv2D(num_filters, (3, 3), padding='same')(encoder)\n    encoder = Activation('relu')(encoder)\n    encoder = BatchNormalization()(encoder)\n    return encoder\n\ndef encoder_block(input_tensor, num_filters):\n    encoder = conv_block(input_tensor, num_filters)\n    encoder_pool = MaxPooling2D((2, 2), strides=(2, 2))(encoder)\n    return encoder_pool, encoder\n\ndef decoder_block(input_tensor, concat_tensor, num_filters):\n    decoder = Conv2DTranspose(num_filters, (2, 2), strides=(2, 2), padding='same')(input_tensor)\n    decoder = concatenate([decoder, concat_tensor], axis=-1)\n    decoder = conv_block(decoder, num_filters)\n    return decoder\n\ndef get_attention_unet(input_shape, num_filters_start=16, num_classes=1):\n    inputs = Input(input_shape)\n\n    # Downsampling through the model\n    encoder_pool0, encoder0 = encoder_block(inputs, num_filters_start)\n    encoder_pool1, encoder1 = encoder_block(encoder_pool0, num_filters_start*2)\n    encoder_pool2, encoder2 = encoder_block(encoder_pool1, num_filters_start*4)\n    encoder_pool3, encoder3 = encoder_block(encoder_pool2, num_filters_start*8)\n\n    center = conv_block(encoder_pool3, num_filters_start*16)\n\n    # Upsampling and establishing the skip connections\n    decoder3 = decoder_block(center, encoder3, num_filters_start*8)\n    attn3 = attention_gate(encoder3, decoder3, num_filters_start*8)\n    decoder2 = decoder_block(attn3, encoder2, num_filters_start*4)\n    attn2 = attention_gate(encoder2, decoder2, num_filters_start*4)\n    decoder1 = decoder_block(attn2, encoder1, num_filters_start*2)\n    attn1 = attention_gate(encoder1, decoder1, num_filters_start*2)\n    decoder0 = decoder_block(attn1, encoder0, num_filters_start)\n\n    # Output\n    outputs = Conv2D(num_classes, (1, 1), activation='sigmoid')(decoder0)\n\n    model = Model(inputs=[inputs], outputs=[outputs])\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:06:41.588433Z","iopub.execute_input":"2023-11-14T16:06:41.588743Z","iopub.status.idle":"2023-11-14T16:06:41.605433Z","shell.execute_reply.started":"2023-11-14T16:06:41.588709Z","shell.execute_reply":"2023-11-14T16:06:41.604405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 7. Visualizing Attention UNet Model Summary and Architechture","metadata":{}},{"cell_type":"code","source":"# Define model parameters\ninput_shape = (256, 256, 1)  # or your preferred dimensions\nnum_filters_start = 32  # number of filters in the first layer of U-Net\nnum_classes = 1  # binary segmentation\n\n# Create a new model instance\nattention_unet_model = get_attention_unet(input_shape, num_filters_start, num_classes)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-11-14T16:06:41.606739Z","iopub.execute_input":"2023-11-14T16:06:41.607083Z","iopub.status.idle":"2023-11-14T16:06:42.690833Z","shell.execute_reply.started":"2023-11-14T16:06:41.607050Z","shell.execute_reply":"2023-11-14T16:06:42.689675Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from keras import backend as K\nfrom keras.losses import binary_crossentropy\nimport tensorflow as tf\n\ndef dice_coef(y_true, y_pred, smooth=1):\n    y_true_f = K.flatten(y_true)\n    y_pred_f = K.flatten(y_pred)\n    intersection = K.sum(y_true_f * y_pred_f)\n    return (2. * intersection + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth)\n\ndef iou_coef(y_true, y_pred, smooth=1):\n  intersection = K.sum(K.abs(y_true * y_pred), axis=[1,2,3])\n  union = K.sum(y_true,[1,2,3])+K.sum(y_pred,[1,2,3])-intersection\n  iou = K.mean((intersection + smooth) / (union + smooth), axis=0)\n  return iou\n\ndef dice_loss(y_true, y_pred):\n    smooth = 1.\n    y_true_f = K.flatten(y_true)\n    y_pred_f = K.flatten(y_pred)\n    intersection = y_true_f * y_pred_f\n    score = (2. * K.sum(intersection) + smooth) / (K.sum(y_true_f) + K.sum(y_pred_f) + smooth)\n    return 1. - score\n\ndef bce_dice_loss(y_true, y_pred):\n    return binary_crossentropy(tf.cast(y_true, tf.float32), y_pred) + 0.5 * dice_loss(tf.cast(y_true, tf.float32), y_pred)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:10:02.722024Z","iopub.execute_input":"2023-11-14T16:10:02.723067Z","iopub.status.idle":"2023-11-14T16:10:02.734048Z","shell.execute_reply.started":"2023-11-14T16:10:02.723025Z","shell.execute_reply":"2023-11-14T16:10:02.732945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compile the model\nattention_unet_model.compile(optimizer='adam', loss=bce_dice_loss,metrics=[dice_coef,iou_coef])\n\n# Summary of the model\nattention_unet_model.summary()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-11-14T16:10:23.224019Z","iopub.execute_input":"2023-11-14T16:10:23.224695Z","iopub.status.idle":"2023-11-14T16:10:23.447822Z","shell.execute_reply.started":"2023-11-14T16:10:23.224662Z","shell.execute_reply":"2023-11-14T16:10:23.446839Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the model\nplot_model(attention_unet_model, to_file='model.png', show_shapes=True, show_layer_names=True)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-11-14T16:10:31.347244Z","iopub.execute_input":"2023-11-14T16:10:31.347646Z","iopub.status.idle":"2023-11-14T16:10:32.806058Z","shell.execute_reply.started":"2023-11-14T16:10:31.347616Z","shell.execute_reply":"2023-11-14T16:10:32.804975Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 8. Training the Attention UNet Model","metadata":{}},{"cell_type":"code","source":"# Training the model (make sure 'train_images', 'train_masks', 'val_images', 'val_masks' are loaded and preprocessed)\nresults = attention_unet_model.fit(X_train, \n                                   y_train, \n                                   batch_size=32, \n                                   epochs=10, \n                                   validation_data=(X_val, y_val))","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:10:39.466700Z","iopub.execute_input":"2023-11-14T16:10:39.467549Z","iopub.status.idle":"2023-11-14T16:11:25.544404Z","shell.execute_reply.started":"2023-11-14T16:10:39.467518Z","shell.execute_reply":"2023-11-14T16:11:25.543511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 9. Evaluating the Model","metadata":{}},{"cell_type":"code","source":"# Evaluate the model\nval_loss, val_dice_coef, val_iou_coef = attention_unet_model.evaluate(X_val, y_val)\nprint(f\"Validation Accuracy: {val_acc}\")","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:13:02.288061Z","iopub.execute_input":"2023-11-14T16:13:02.289024Z","iopub.status.idle":"2023-11-14T16:13:02.470428Z","shell.execute_reply.started":"2023-11-14T16:13:02.288990Z","shell.execute_reply":"2023-11-14T16:13:02.469353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\n\n# Assuming X_val and y_val are your validation images and masks\nnum_samples = 5  # Choose the number of samples you want to display\nsample_indices = random.sample(range(len(X_val)), num_samples)\n\nsample_images = X_val[sample_indices]\nsample_true_masks = y_val[sample_indices]\n","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:13:12.875863Z","iopub.execute_input":"2023-11-14T16:13:12.876345Z","iopub.status.idle":"2023-11-14T16:13:12.882457Z","shell.execute_reply.started":"2023-11-14T16:13:12.876303Z","shell.execute_reply":"2023-11-14T16:13:12.881350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_pred_masks = attention_unet_model.predict(sample_images)\n\n# Thresholding example (adjust threshold as needed)\nsample_pred_masks = (sample_pred_masks > 0.5).astype(np.uint8)","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:14:44.028461Z","iopub.execute_input":"2023-11-14T16:14:44.028838Z","iopub.status.idle":"2023-11-14T16:14:44.112693Z","shell.execute_reply.started":"2023-11-14T16:14:44.028813Z","shell.execute_reply":"2023-11-14T16:14:44.111679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(num_samples):\n    plt.figure(figsize=(12, 5))\n\n    # Display original image\n    plt.subplot(1, 3, 1)\n    plt.imshow(sample_images[i], cmap='gray')\n    plt.title('Original Image')\n    plt.axis('off')\n\n    # Display true mask\n    plt.subplot(1, 3, 2)\n    plt.imshow(sample_true_masks[i], cmap='gray')\n    plt.title('True Mask')\n    plt.axis('off')\n\n    # Display predicted mask\n    plt.subplot(1, 3, 3)\n    plt.imshow(sample_pred_masks[i], cmap='gray')\n    plt.title('Predicted Mask')\n    plt.axis('off')\n\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-11-14T16:14:47.450868Z","iopub.execute_input":"2023-11-14T16:14:47.451218Z","iopub.status.idle":"2023-11-14T16:14:49.533049Z","shell.execute_reply.started":"2023-11-14T16:14:47.451190Z","shell.execute_reply":"2023-11-14T16:14:49.532003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 11. Future Directions","metadata":{}},{"cell_type":"markdown","source":"\nThis is just the model architechture, try using callbacks to control overfitting, maybe use Keras Tuner for Hyperparameter Tuning, or try denoising the image before any other preprocessing. This will improve the model performance. Youc an also try using transfer learning.\n\nIf you are interested in exploring other Biomedical Segmentation models then checkout the following two starter notebooks:\n\n* [SegNet using TensorFlow Starter for SenNet + HOA](https://www.kaggle.com/code/salmankhaliq22/sennet-hoa-segnet-tensorflow-starter)\n* [UNet using TensorFlow Starter for SenNet + HOA](https://www.kaggle.com/code/salmankhaliq22/unet-tensorflow-starter-sennet-hoa/notebook)","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}