{"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":51753,"databundleVersionId":5692552,"sourceType":"competition"}],"dockerImageVersionId":30636,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# imports\nimport os\nimport random\nfrom path import Path\n\n# Data analysis and manipulation\nimport numpy as np\nimport pandas as pd\n\nfrom tqdm import tqdm\n\n# Data visualization\nfrom matplotlib import animation\nimport matplotlib.pyplot as plt\nfrom IPython import display\nimport seaborn as sns\nimport plotly\n\n# ML, DL & Modelling\n# from sklearn import \nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras import layers\nfrom tensorflow.keras import callbacks\nprint(f\"TensorFlow version: {tf.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:39:32.343179Z","iopub.execute_input":"2024-01-31T14:39:32.343988Z","iopub.status.idle":"2024-01-31T14:39:45.416119Z","shell.execute_reply.started":"2024-01-31T14:39:32.343954Z","shell.execute_reply":"2024-01-31T14:39:45.415147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# Loading a fraction of the dataset\n\nBASE_DIR = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\n\n\n# Building list of record_ids\nrecord_ids_train = os.listdir(BASE_DIR + '/train/')\nrecord_ids_val = os.listdir(BASE_DIR + '/validation/')\nrecord_ids_test = os.listdir(BASE_DIR + '/test/')\n\n# print(record_ids_train)\n\nprint (f\"Number of training record ids: {len(record_ids_train)}\")\nprint(f\"Number of validation validation ids: {len(record_ids_val)}\")\nprint(f\"Number of test test ids: {len(record_ids_test)}\")\n\n# sample_record_ids_train = record_ids_train\n\n# Randomly selecting observations\nnumber_observations =5_000\nsample_record_ids_train = random.sample(record_ids_train, number_observations) #record_ids\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:39:49.42786Z","iopub.execute_input":"2024-01-31T14:39:49.429014Z","iopub.status.idle":"2024-01-31T14:39:49.836012Z","shell.execute_reply.started":"2024-01-31T14:39:49.428977Z","shell.execute_reply":"2024-01-31T14:39:49.835065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Class Distribution **As we can see there are 11270 negative examples and 9259 positive examples. This shows there are slightly more examples with no contrails however the split seems pretty close to 50/50.","metadata":{}},{"cell_type":"code","source":"# Initialize the count for positive examples\npositive_ids = 0\n\n# Iterate through training examples and count positive instances\nfor record in tqdm(record_ids_train):\n    pixel_masks = np.load(f\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/{record}/human_pixel_masks.npy\")\n    if len(np.unique(pixel_masks)) > 1:\n        positive_ids += 1\n\n# Calculate the number of negative examples\nnegative_ids = len(record_ids_train) - positive_ids\n\n# Display the count of positive and negative examples\nprint(f\"Number of positive examples: {positive_ids}, Number of negative examples: {negative_ids}\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Understanding Spectral Bands**\nIn the dataset, there are 9 bands for each example, and each band represents a series of images captured at different wavelengths of light. For instance, band 08 consists of a series of images highlighting a specific wavelength, while band 16 contains images capturing a different wavelength. It is important to note that each band contains images of the same thing just at different wavelengths so each band will look slightly different and contain more information than others because contrails look different at different wavelengths\n\nThese individual bands contain images taken at 10-minute intervals. There are 8 images in total for each band, taken across 80 minutes with a 10-minute interval between each image. When we examine band 8, for example, we are looking at a sequence of images, each taken 10 minutes after the previous one. This time series of images provides valuable information about the expansion and shape changes of contrails over time.\n\nNow let's discuss the segmentation masks. There are two files: human_pixel_masks and human_individual_masks. The human_individual_masks represent labels generated by multiple labellers. These labels are then compared and combined to create a final ground truth, which can be found in the human_pixel_masks file. This mask corresponds to the 5th image in the bands. The purpose of this approach is to enhance the accuracy of the ground truth. By involving multiple labellers and aggregating their findings, we can achieve a more precise and reliable ground truth. To draw an analogy, just as a group of doctors can collectively provide a more accurate assessment of a disease, multiple labellers collaborating on the labels can produce a more accurate ground truth.","metadata":{}},{"cell_type":"code","source":"image = np.load(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/1000216489776414077/band_08.npy\")\nprint(f\"band 08 shape {image.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:40:03.635405Z","iopub.execute_input":"2024-01-31T14:40:03.636165Z","iopub.status.idle":"2024-01-31T14:40:03.713125Z","shell.execute_reply.started":"2024-01-31T14:40:03.636134Z","shell.execute_reply":"2024-01-31T14:40:03.712212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plotting So We Can Visualize Our Understanding\nAs we can see in the plot as you go across from the left to the right you are seeing the different spectral bands we are looking at, and as you go down you see the image at different time steps (10 minutes apart). You can see how the image looks slightly different across time and bands. The human pixel mask (ground truth) corresponds to the image at time step 5. Hopefully this image can make it easier to understand what is going on, and please feel free to try this out on other examples too!","metadata":{}},{"cell_type":"code","source":"bands = ['08', '09', '10', '11', '12', '13', '14', '15', '16']\n\ndef plot_example():\n    \n    \"\"\" Args: example_id(str): The id of the example i.e. '1000216489776414077' split_dir(str): The split directoryu i.e. 'test', 'train', 'val'\n    \"\"\"\n    \nfig, axs = plt.subplots(8, len(bands), figsize=(16, 16))\n\nfor j, band in enumerate(bands):\n    img = np.load(BASE_DIR + f\"/train/1000603527582775543/band_{band}.npy\")\n    for i in range(8):\n        axs[i, j].imshow(img[..., i]) \n        axs[i, j].set_title(f\"Band {band}\\nTime Step {i+1}\") \n\nplt.tight_layout()  \nplt.show()\nplot_example()","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:40:08.030704Z","iopub.execute_input":"2024-01-31T14:40:08.031447Z","iopub.status.idle":"2024-01-31T14:40:19.352398Z","shell.execute_reply.started":"2024-01-31T14:40:08.031414Z","shell.execute_reply":"2024-01-31T14:40:19.351178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Checking out the Masks**\n\nAs we can see the human individual masks contains 4 (256, 256, 1) masks. These 4 masks corresponding to the masks created by 4 labellers. Although there are 4 in this example other examples can have a varying amount. The human pixel masks are the ground truth and there is only one mask of shape (256, 256, 1). This should be expected because the input images are of shape 256 x 256 so we should expect the output shape to be the same.","metadata":{}},{"cell_type":"code","source":"# Load human individual masks and human pixel masks from specified paths\nindividual_masks_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/1000216489776414077/human_individual_masks.npy\"\npixel_masks_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/1000216489776414077/human_pixel_masks.npy\"\n\nindividual_masks = np.load(individual_masks_path)\npixel_ground_truth = np.load(pixel_masks_path)\n\n# Display the shapes of the loaded masks\nprint(f\"The shape of human individual masks are {individual_masks.shape}\")\nprint(f\"The shape of human pixel masks are {pixel_ground_truth.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:40:26.552323Z","iopub.execute_input":"2024-01-31T14:40:26.553093Z","iopub.status.idle":"2024-01-31T14:40:26.576234Z","shell.execute_reply.started":"2024-01-31T14:40:26.553062Z","shell.execute_reply":"2024-01-31T14:40:26.575384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_masks(example_id, split_dir):\n    \"\"\"\n    Args:\n        example_id (str): The id of the example, e.g., '1000216489776414077'\n        split_dir (str): The split directory, e.g., 'train', 'test', 'val' \n    \"\"\"\n    individual_masks_path = Path(BASE_DIR) / split_dir / example_id / 'human_individual_masks.npy'\n    pixel_ground_truth_path = Path(BASE_DIR) / split_dir / example_id / 'human_pixel_masks.npy'\n\n    individual_masks = np.load(individual_masks_path)\n    pixel_ground_truth = np.load(pixel_ground_truth_path)\n    \n    fig, axs = plt.subplots(1, len(individual_masks[0, 0, 0]) + 1, figsize=(2 * (len(individual_masks[0, 0, 0]) + 1), 16))\n    \n    for i in range(len(individual_masks[0, 0, 0])):\n        axs[i].imshow(individual_masks[..., i] ,cmap='viridis')\n        axs[i].set_title(f\"Labeller {i + 1}\")\n        \n    axs[i + 1].imshow(pixel_ground_truth,)\n    axs[i + 1].set_title(\"Aggregated/Ground Truth\")\n    \n    plt.tight_layout()\n    plt.show()\n\n# Example usage\nplot_masks('1000603527582775543', 'train')","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:40:29.864131Z","iopub.execute_input":"2024-01-31T14:40:29.86452Z","iopub.status.idle":"2024-01-31T14:40:30.649115Z","shell.execute_reply.started":"2024-01-31T14:40:29.864489Z","shell.execute_reply":"2024-01-31T14:40:30.64823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**False Color Images**\n\nWe know that the image labellers are putting their final annotations on the image at step 5 (the fifth element) in each band. However, the image that the labellers end up annotating is not found in the spectral bands given to use but rather a false color image. A false color image is a image that can be generated from the spectral bands given to us and it is meant to make contrails appear dark relative to their surroundings making them easier to detect. This false color image is what the labelers ended up using to derive their ground truths.\n\nThe false color images are generated using the **ash color scheme**. Here is a brief overview of the color scheme \n\nThe ash color scheme is a false color representation commonly used for visualizing volcanic ash plumes or volcanic cloud observations. It employes a three-channel color scheme (red, green, and blue) to highlight specific features of interest:\n\n* **Red Channel**: Represents the temperature or thermal information of the volcanic plume, with warmer regions shown in red/orange and cooler regions in shades of blue/green.\n\n* **Green Channel**: Indicates the particular size or density of volcanic ash particles in the plume. Darker green indicates denser or larger particles, while lighter green represents finer particles.\n\n* **Blue Channel**: Provides additional information about the plume, such as its height or altitude. Higher blue values suggest greater plume attitude, while darker shades imply lower altitudes.\n\nBy combining these color channels, the ash color scheme enhances the visibility of different properties within the volcanic plume, facilitating analysis and interpretation of volcanic cloud data. \n\nAlthough it is commonly used to study volcanic activity it is also useful for identifying contrails","metadata":{}},{"cell_type":"code","source":"# Defining normalization function\n\ndef normalize_range(data, bounds):\n    \"\"\"Maps data to the range [0,1]\"\"\"\n    normalized_data = (data - bounds[0])/(bounds[1] - bounds[0])\n    return normalized_data","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:40:41.58739Z","iopub.execute_input":"2024-01-31T14:40:41.587739Z","iopub.status.idle":"2024-01-31T14:40:41.592661Z","shell.execute_reply.started":"2024-01-31T14:40:41.587713Z","shell.execute_reply":"2024-01-31T14:40:41.591648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\n# Loop through the list of record IDs to select specific bands and build X_init\nselected_bands = ['band_11.npy', 'band_14.npy', 'band_15.npy']\ntarget_suffix = 'human_pixel_masks.npy'\nN_TIMES_BEFORE = 4\n\nrecords_list = []\ntargets_list = []\nskipped_observations = 0\n\nfor record_id in sample_record_ids_train:\n    \n    # Build target paths\n    target_path = os.path.join(BASE_DIR, 'train', record_id, target_suffix)\n\n    try:\n        target = np.load(open(target_path, 'rb'))\n    except FileNotFoundError:\n        print(f\"Warning: Target file not found for record_id {record_id}. Skipping.\")\n        skipped_observations += 1\n        continue\n\n#     # Build target paths\n#     target_path = os.path.join(BASE_DIR, 'train', record_id, target_suffix)\n#     target = np.load(open(target_path, 'rb'))\n\n    # Skip observations with no contrails\n    if target.sum() == 0:\n        continue\n    else:\n        # Build band paths\n        band1_path = os.path.join(BASE_DIR, 'train', record_id, selected_bands[0])\n        band2_path = os.path.join(BASE_DIR, 'train', record_id, selected_bands[1])\n        band3_path = os.path.join(BASE_DIR, 'train', record_id, selected_bands[2])\n\n        # Load each band\n        band1 = np.load(open(band1_path, 'rb'))[:, :, N_TIMES_BEFORE]\n        band2 = np.load(open(band2_path, 'rb'))[:, :, N_TIMES_BEFORE]\n        band3 = np.load(open(band3_path, 'rb'))[:, :, N_TIMES_BEFORE]\n\n        # Normalize each band with its relevant bounds\n        # Define bounds for each band\n        _T11_BOUNDS = (243, 303)\n        _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n        _TDIFF_BOUNDS = (-4, 2)\n\n        # Apply normalization functions\n        normalized_r = normalize_range(band3 - band2, _TDIFF_BOUNDS)\n        normalized_g = normalize_range(band2 - band1, _CLOUD_TOP_TDIFF_BOUNDS)\n        normalized_b = normalize_range(band2, _T11_BOUNDS)\n\n        # Build a single record from all bands\n        record = np.clip(np.stack([normalized_r, normalized_g, normalized_b], axis=2), 0, 1)\n        \n        # Append target and observation lists\n        targets_list.append(target)\n        records_list.append(record)\n\nX_init = np.stack(records_list, axis=0)\ny_init = np.stack(targets_list, axis=0).astype(float)\n\nprint(f\"We skipped {skipped_observations} observations with no contrails.\")","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:40:44.791004Z","iopub.execute_input":"2024-01-31T14:40:44.791815Z","iopub.status.idle":"2024-01-31T14:46:44.367576Z","shell.execute_reply.started":"2024-01-31T14:40:44.791781Z","shell.execute_reply":"2024-01-31T14:46:44.366611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Checking all shapes\nsample_ids_array = np.array(sample_record_ids_train)\nprint(\"X_init shape is:\", X_init.shape)\nprint(\"y_init shape is:\", y_init.shape)\nprint(\"sample_ids_array shape is:\", sample_ids_array.shape)\n\ntest_img = X_init[28]\ntarget_img= y_init[28]\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:47:39.684384Z","iopub.execute_input":"2024-01-31T14:47:39.685262Z","iopub.status.idle":"2024-01-31T14:47:39.693238Z","shell.execute_reply.started":"2024-01-31T14:47:39.685226Z","shell.execute_reply":"2024-01-31T14:47:39.692107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Showing image test\nplt.figure(figsize=(20, 8))\nax = plt.subplot(1, 3, 1)\nax.imshow(test_img);\nax.set_title('False color image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(target_img, interpolation='none');\nax.set_title('Ground truth contrail mask')\n\nax = plt.subplot(1, 3, 3)\nax.imshow(test_img)\nax.imshow(target_img, cmap='Reds', alpha=.3, interpolation='none')\nax.set_title('Contrail mask on false color image');\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:47:42.614619Z","iopub.execute_input":"2024-01-31T14:47:42.615453Z","iopub.status.idle":"2024-01-31T14:47:43.547673Z","shell.execute_reply.started":"2024-01-31T14:47:42.615423Z","shell.execute_reply":"2024-01-31T14:47:43.546699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Defining proba_to_pixel to transform predicted probas into categories for the dice metric\n\n# def proba_to_pixel(y):\n#     return tf.where(y > 0.5, tf.ones_like(y),tf.zeros_like(y))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss function for model","metadata":{}},{"cell_type":"code","source":"# Defining proba_to_pixel to transform predicted probas into categories for the dice metric\n\ndef proba_to_pixel(y):\n    return tf.where(y > 0.5, tf.ones_like(y),tf.zeros_like(y))\n\n# Defining the dice metric to be used in model compilation\n\ndef dice_metric(y_true, y_pred):\n        #converting y_pred probas to pixels (between 0 and 1)\n        y_pred = proba_to_pixel(y_pred)\n        y_true = proba_to_pixel(y_true)\n\n        # Define epsilon to prevent division by zero\n        smooth = 1e-5 \n\n        # Calculate the sum of y_true and y_pred for each class\n        y_true_sum = tf.reduce_sum(y_true)\n        y_pred_sum = tf.reduce_sum(y_pred)\n\n        # Calculate the intersection and union of y_true and y_pred\n        intersection = tf.reduce_sum(y_true * y_pred)\n        union = y_true_sum + y_pred_sum\n\n        # Calculate the Dice coefficient for each class\n        dice = (2. * intersection + smooth) / (union + smooth)\n\n        return dice\n    \n# Defining the dice_closs as loss function used for the model \n\ndef dice_loss(y_true, y_pred):\n        \n    # No need to convert y_pred to pixels for the loss \n    # Define epsilon to prevent division by zero\n    smooth = 1e-5 \n\n    # Calculate the sum of y_true and y_pred for each class\n    y_true_sum = tf.reduce_sum(y_true)\n    y_pred_sum = tf.reduce_sum(y_pred)\n\n    # Calculate the intersection and union of y_true and y_pred\n    intersection = tf.reduce_sum(y_true * y_pred)\n    union = y_true_sum + y_pred_sum\n\n    # Calculate the Dice coefficient for each class\n    dice = (2. * intersection + smooth) / (union + smooth)\n\n    return 1 - dice","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:47:51.252533Z","iopub.execute_input":"2024-01-31T14:47:51.253385Z","iopub.status.idle":"2024-01-31T14:47:51.261713Z","shell.execute_reply.started":"2024-01-31T14:47:51.25335Z","shell.execute_reply":"2024-01-31T14:47:51.260726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Testing the dice loss function \n\npredictions_test = np.array([0.03, 0.9, 0.8])\ntrue_set = np.array([1.0, 1.0, 0.0])\n\ndice_loss(true_set, predictions_test)\n\n\n    ","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:47:55.068697Z","iopub.execute_input":"2024-01-31T14:47:55.069063Z","iopub.status.idle":"2024-01-31T14:47:55.995494Z","shell.execute_reply.started":"2024-01-31T14:47:55.069035Z","shell.execute_reply":"2024-01-31T14:47:55.99454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Doing this process to get the average for 1K observations\n# 1-Retrieving the number of [1] in one 'human_pixel_masks.npy)\nfrom statistics import mean\n\ndef proportion_1_pixels(y):\n    \n    values, counts = np.unique(y, return_counts=True)\n    proportion_0_pixels = counts[0]/counts.sum()\n    proportion_1_pixels = 1 - proportion_0_pixels\n    \n    return proportion_1_pixels\n\np = proportion_1_pixels(y_init)","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:47:59.234254Z","iopub.execute_input":"2024-01-31T14:47:59.234656Z","iopub.status.idle":"2024-01-31T14:48:03.849617Z","shell.execute_reply.started":"2024-01-31T14:47:59.234626Z","shell.execute_reply":"2024-01-31T14:48:03.848792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def baseline_predict_flexible(y, p): \n    values = [0,1]\n    probas = [1-p, p]\n    sample_size = np.prod([y.shape[i] for i in range(0, len(y.shape))])\n#     y.shape[0] * y.shape[1]\n    baseline_pred = np.random.choice(values, size=sample_size, p=probas)\n    baseline_pred = baseline_pred.reshape(y.shape)\n    \n    return baseline_pred\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:48:10.848287Z","iopub.execute_input":"2024-01-31T14:48:10.848994Z","iopub.status.idle":"2024-01-31T14:48:10.854497Z","shell.execute_reply.started":"2024-01-31T14:48:10.848959Z","shell.execute_reply":"2024-01-31T14:48:10.853524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred_baseline = baseline_predict_flexible(y_init, p=p).astype(float)\nprint(y_pred_baseline.shape)","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:48:14.086814Z","iopub.execute_input":"2024-01-31T14:48:14.087685Z","iopub.status.idle":"2024-01-31T14:48:17.815182Z","shell.execute_reply.started":"2024-01-31T14:48:14.087652Z","shell.execute_reply":"2024-01-31T14:48:17.814266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unique_values_baseline, counts_baseline = np.unique(y_pred_baseline, return_counts=True)\nprint(unique_values_baseline)\nprint(counts_baseline)\nprint(counts_baseline.sum())","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:48:20.325373Z","iopub.execute_input":"2024-01-31T14:48:20.326216Z","iopub.status.idle":"2024-01-31T14:48:25.009293Z","shell.execute_reply.started":"2024-01-31T14:48:20.326186Z","shell.execute_reply":"2024-01-31T14:48:25.008378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Computing the dice metric for y_pred_baseline\n#tf.convert_to_tensor\n\ndice_init = dice_loss(y_init, y_pred_baseline)\ndice_init","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:48:34.421303Z","iopub.execute_input":"2024-01-31T14:48:34.421898Z","iopub.status.idle":"2024-01-31T14:48:43.217817Z","shell.execute_reply.started":"2024-01-31T14:48:34.421857Z","shell.execute_reply":"2024-01-31T14:48:43.216738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Building a first simple CNN model\nOur first model architecture will have 2 building blocks:\n\nDownsampling path / Decoder = convolutional layers extracting features from the image while reducing it size\nUpsampling / Decoder = expanding the size of the image using Transpose convolution to reach an output (the mask) with same size as input image","metadata":{}},{"cell_type":"code","source":"# Build the autoencoder model\nmodel = Sequential()\n\n# Encoder\nmodel.add(layers.Conv2D(16, (3, 3), input_shape=(256, 256, 3), padding='same', activation='relu'))\nmodel.add(layers.MaxPool2D(pool_size=(2, 2)))\n\nmodel.add(layers.Conv2D(32, (2, 2), padding='same', activation='relu'))\nmodel.add(layers.MaxPool2D(pool_size=(2, 2)))\n\n# Decoder\nmodel.add(layers.Conv2DTranspose(32, (2, 2), padding='same', activation='relu', strides=(2, 2)))\nmodel.add(layers.Conv2DTranspose(1, (2, 2), padding='same', activation='sigmoid', strides=(2, 2)))\n\n# Display model summary\nmodel.summary()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# Defining an optimizer with specific parameters\noptimizer = tf.keras.optimizers.Adam(\n    learning_rate=0.03,\n    beta_1=0.9,\n    beta_2=0.999,\n    epsilon=1e-07,\n    amsgrad=False,\n    weight_decay=None,\n    clipnorm=None,\n    clipvalue=None,\n    global_clipnorm=None,\n    use_ema=False,\n    ema_momentum=0.99,\n    ema_overwrite_frequency=None,\n    jit_compile=True,\n    name='Adam'\n)\n\n# Compiling the model with the Adam optimizer, dice loss function, and dice metric\nmodel.compile(optimizer=optimizer,\n              loss=dice_loss,\n              metrics=dice_metric)\n\n# Fitting the model to the training data with early stopping\nearly_stopping = callbacks.EarlyStopping(patience=30)\nhistory_base_model = model.fit(X_init, y_init,\n                               batch_size=8, \n                               epochs=30,\n                               validation_split=0.3,\n                               callbacks=[early_stopping],\n                               verbose=1)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# history_base_model.__dict__","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_history(history, title='', axs=None, exp_name=\"\"):\n    if axs is not None:\n        ax1, ax2 = axs\n    else:\n        f, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))\n    \n    if len(exp_name) > 0 and exp_name[0] != '_':\n        exp_name = '_' + exp_name\n    ax1.plot(history.history['loss'], label = 'train' + exp_name)\n    ax1.plot(history.history['val_loss'], label = 'val' + exp_name)\n#     ax1.set_ylim(0., 2.2)\n    ax1.set_title('loss')\n    ax1.legend()\n\n    ax2.plot(history.history['dice_metric'], label='train dice metric'  + exp_name)\n    ax2.plot(history.history['val_dice_metric'], label='val dice metric'  + exp_name)\n#     ax2.set_ylim(0.25, 1.)\n    ax2.set_title('Dice metric')\n    ax2.legend()\n    return (ax1, ax2)","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:48:52.988394Z","iopub.execute_input":"2024-01-31T14:48:52.988763Z","iopub.status.idle":"2024-01-31T14:48:52.996622Z","shell.execute_reply.started":"2024-01-31T14:48:52.988736Z","shell.execute_reply":"2024-01-31T14:48:52.99554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(history_base_model)\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\n\ndef dice_loss(y_true, y_pred):\n    smooth = 1e-5\n    \n    # Flatten the tensors\n    y_true_flat = tf.keras.layers.Flatten()(y_true)\n    y_pred_flat = tf.keras.layers.Flatten()(y_pred)\n\n    intersection = tf.reduce_sum(y_true_flat * y_pred_flat)\n    union = tf.reduce_sum(y_true_flat) + tf.reduce_sum(y_pred_flat)\n\n    dice = (2.0 * intersection + smooth) / (union + smooth)\n\n    return 1.0 - dice\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T15:25:24.86993Z","iopub.execute_input":"2024-01-31T15:25:24.87029Z","iopub.status.idle":"2024-01-31T15:25:24.876796Z","shell.execute_reply.started":"2024-01-31T15:25:24.870264Z","shell.execute_reply":"2024-01-31T15:25:24.875769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Advanced U-net model","metadata":{}},{"cell_type":"code","source":"\n\nimport numpy as np\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models, callbacks\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom sklearn.model_selection import train_test_split\n\n# Assuming X_init and y_init are defined\nX_train, X_val, y_train, y_val = train_test_split(X_init, y_init, test_size=0.3, random_state=42)\n\n\n# Assuming img_size_target and number_channels_target are defined\nimg_size_target = X_init.shape[1]\nnumber_channels_target = X_init.shape[-1]\nstart_neurons = 4\n\n# Define the U-Net model\ndef unet_model(input_shape=(256, 256, 3), start_neurons=4, dropout_rate=0.1):\n    inputs = tf.keras.Input(shape=input_shape)\n    \n    # Encoder\n    conv1 = layers.Conv2D(start_neurons, 3, activation='relu', padding='same')(inputs)\n    conv1 = layers.Conv2D(start_neurons, 3, activation='relu', padding='same')(conv1)\n    pool1 = layers.MaxPooling2D(pool_size=(2, 2))(conv1)\n    pool1 = layers.Dropout(dropout_rate)(pool1)\n\n    conv2 = layers.Conv2D(start_neurons * 2, 3, activation='relu', padding='same')(pool1)\n    conv2 = layers.Conv2D(start_neurons * 2, 3, activation='relu', padding='same')(conv2)\n    pool2 = layers.MaxPooling2D(pool_size=(2, 2))(conv2)\n    pool2 = layers.Dropout(dropout_rate)(pool2)\n\n    conv3 = layers.Conv2D(start_neurons * 4, 3, activation='relu', padding='same')(pool2)\n    conv3 = layers.Conv2D(start_neurons * 4, 3, activation='relu', padding='same')(conv3)\n    pool3 = layers.MaxPooling2D(pool_size=(2, 2))(conv3)\n    pool3 = layers.Dropout(dropout_rate)(pool3)\n\n    # Bottleneck\n    conv4 = layers.Conv2D(start_neurons * 8, 3, activation='relu', padding='same')(pool3)\n    conv4 = layers.Conv2D(start_neurons * 8, 3, activation='relu', padding='same')(conv4)\n\n    # Decoder\n    up5 = layers.Conv2DTranspose(start_neurons * 4, (2, 2), strides=(2, 2), padding='same')(conv4)\n    up5 = layers.concatenate([up5, conv3], axis=-1)\n    up5 = layers.Dropout(dropout_rate)(up5)\n    conv5 = layers.Conv2D(start_neurons * 4, 3, activation='relu', padding='same')(up5)\n    conv5 = layers.Conv2D(start_neurons * 4, 3, activation='relu', padding='same')(conv5)\n\n    up6 = layers.Conv2DTranspose(start_neurons * 2, (2, 2), strides=(2, 2), padding='same')(conv5)\n    up6 = layers.concatenate([up6, conv2], axis=-1)\n    up6 = layers.Dropout(dropout_rate)(up6)\n    conv6 = layers.Conv2D(start_neurons * 2, 3, activation='relu', padding='same')(up6)\n    conv6 = layers.Conv2D(start_neurons * 2, 3, activation='relu', padding='same')(conv6)\n\n    up7 = layers.Conv2DTranspose(start_neurons, (2, 2), strides=(2, 2), padding='same')(conv6)\n    up7 = layers.concatenate([up7, conv1], axis=-1)\n    up7 = layers.Dropout(dropout_rate)(up7)\n    conv7 = layers.Conv2D(start_neurons, 3, activation='relu', padding='same')(up7)\n    conv7 = layers.Conv2D(start_neurons, 3, activation='relu', padding='same')(conv7)\n\n    # Output layer\n    output = layers.Conv2D(1, 1, activation='sigmoid')(conv7)\n\n    model = models.Model(inputs=inputs, outputs=output)\n    return model\n\n# Create the U-Net model\nunet_model = unet_model(input_shape=(img_size_target, img_size_target, number_channels_target), start_neurons=start_neurons, dropout_rate=0.2)\n\n# Compile the model\noptimizer = tf.keras.optimizers.Adam(learning_rate=0.001)\nunet_model.compile(optimizer=optimizer, loss='binary_crossentropy', metrics=['accuracy'])\n\n# Define callbacks\nes = callbacks.EarlyStopping(patience=30)\nlrp = callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.2, patience=5, min_lr=0.001)\n\n# Data augmentation\ndatagen = ImageDataGenerator(\n    rotation_range=20,\n    width_shift_range=0.2,\n    height_shift_range=0.2,\n    shear_range=0.2,\n    zoom_range=0.2,\n    horizontal_flip=True,\n    fill_mode='nearest'\n)\n\n# Fit the model using data augmentation\nhistory_unet_model = unet_model.fit(\n    datagen.flow(X_train, y_train, batch_size=8),\n    epochs=50,\n    validation_data=(X_val, y_val),\n    callbacks=[es, lrp],\n    verbose=1\n)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T15:35:59.255037Z","iopub.execute_input":"2024-01-31T15:35:59.255825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Access loss values from the training history\ntrain_loss = history_unet_model.history['loss']\nval_loss = history_unet_model.history['val_loss']\n\n# Plotting the loss\nepochs = range(1, len(train_loss) + 1)\nplt.plot(epochs, train_loss, label='Training Loss')\nplt.plot(epochs, val_loss, label='Validation Loss')\nplt.title('Training and Validation Loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-01-31T15:17:31.52624Z","iopub.execute_input":"2024-01-31T15:17:31.531157Z","iopub.status.idle":"2024-01-31T15:17:31.814305Z","shell.execute_reply.started":"2024-01-31T15:17:31.531119Z","shell.execute_reply":"2024-01-31T15:17:31.813564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# img_size_target = X_init.shape[1]\n# number_channels_target = X_init.shape[-1]\n# start_neurons = 4\n\n# # Define the input layer\n# input_layer = layers.Input((img_size_target, img_size_target, number_channels_target))\n\n# # Build the model\n# output_layer = model(input_layer, start_neurons)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# unet model with Keras Functional API\nunet_model = tf.keras.Model(input_layer, output_layer, name=\"U-Net\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# optimizer_2 = tf.keras.optimizers.legacy.Adam(\n#     learning_rate=0.01,\n# )","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time \n# 2. Compile\n# unet_model.compile(optimizer=optimizer_2,\n#                   loss=dice_loss,\n#                   metrics=dice_metric)\nunet_model.compile(optimizer='adam', loss=dice_loss, metrics=[dice_metric])\n\n\n# 3. Fit \nes = callbacks.EarlyStopping(patience=30)\nlrp = callbacks.ReduceLROnPlateau(monitor='val_loss',\n                                  factor=0.2,\n                                  patience=5,\n                                  min_lr=0.001)\n\nhistory_unet_model = unet_model.fit(X_init,\n                                    y_init,\n                                    batch_size=8,\n                                    epochs=50,\n                                    validation_split=0.3,\n                                    callbacks=[es, lrp],\n                                    verbose=1)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(history_unet_model)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-01-31T14:57:20.765897Z","iopub.execute_input":"2024-01-31T14:57:20.766289Z","iopub.status.idle":"2024-01-31T14:57:21.356096Z","shell.execute_reply.started":"2024-01-31T14:57:20.766248Z","shell.execute_reply":"2024-01-31T14:57:21.354809Z"},"trusted":true},"execution_count":null,"outputs":[]}]}