{"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":9988,"databundleVersionId":868324,"sourceType":"competition"}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nimport cv2\nimport tensorflow as tf","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:04.107955Z","iopub.execute_input":"2023-12-28T15:07:04.108899Z","iopub.status.idle":"2023-12-28T15:07:15.606469Z","shell.execute_reply.started":"2023-12-28T15:07:04.108865Z","shell.execute_reply":"2023-12-28T15:07:15.605519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras.layers import Conv2D, BatchNormalization, Activation, MaxPool2D, Conv2DTranspose, Concatenate, Input, Lambda, Dropout, MaxPooling2D\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.applications import ResNet50\n#from tensorflow.keras.layers import Concatenate as concatenate\nconcatenate = tf.concat\n\ndef conv_block(input, num_filters):\n    x = Conv2D(num_filters, 3, padding=\"same\")(input)\n    x = BatchNormalization()(x)\n    x = Activation(\"relu\")(x)\n\n    x = Conv2D(num_filters, 3, padding=\"same\")(x)\n    x = BatchNormalization()(x)\n    x = Activation(\"relu\")(x)\n\n    return x\n\ndef decoder_block(input, skip_features, num_filters):\n    x = Conv2DTranspose(num_filters, (2, 2), strides=2, padding=\"same\")(input)\n    x = Concatenate()([x, skip_features])\n    x = conv_block(x, num_filters)\n    return x\n\ndef build_resnet50_unet(input_shape):\n    \"\"\" Input \"\"\"\n    inputs = Input(input_shape)\n\n    \"\"\" Pre-trained ResNet50 Model \"\"\"\n    resnet50 = ResNet50(include_top=False, weights=\"imagenet\", input_tensor=inputs)\n    #resnet50.trainable = False\n\n    \"\"\" Encoder \"\"\"\n    s1 = resnet50.get_layer(\"input_1\").output           ## (512 x 512)\n    s2 = resnet50.get_layer(\"conv1_relu\").output        ## (256 x 256)\n    s3 = resnet50.get_layer(\"conv2_block3_out\").output  ## (128 x 128)\n    s4 = resnet50.get_layer(\"conv3_block4_out\").output  ## (64 x 64)\n\n    \"\"\" Bridge \"\"\"\n    b1 = resnet50.get_layer(\"conv4_block6_out\").output  ## (32 x 32)\n\n    \"\"\" Decoder \"\"\"\n    d1 = decoder_block(b1, s4, 512)                     ## (64 x 64)\n    d2 = decoder_block(d1, s3, 256)                     ## (128 x 128)\n    d3 = decoder_block(d2, s2, 128)                     ## (256 x 256)\n    d4 = decoder_block(d3, s1, 64)                      ## (512 x 512)\n\n    \"\"\" Output \"\"\"\n    outputs = tf.nn.softmax(Conv2D(2, 1, padding=\"same\", activation=\"sigmoid\")(d4))\n\n    model = Model(inputs, outputs, name=\"ResNet50_U-Net\")\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:15.608635Z","iopub.execute_input":"2023-12-28T15:07:15.609648Z","iopub.status.idle":"2023-12-28T15:07:15.720488Z","shell.execute_reply.started":"2023-12-28T15:07:15.609609Z","shell.execute_reply":"2023-12-28T15:07:15.719694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nimport os\nimport tensorflow as tf\nfrom tensorflow import keras\n\n# set the random seed:\nSEED = 42\nrandom.seed(SEED)\n\nTRAIN_DIR = '/kaggle/input/airbus-ship-detection/train_v2/'\nTEST_DIR = '/kaggle/input/airbus-ship-detection/test_v2/'\n\nCORRUPTED_TRAIN_IMAGE_IDS = ['6384c3e78.jpg']","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:15.721499Z","iopub.execute_input":"2023-12-28T15:07:15.72178Z","iopub.status.idle":"2023-12-28T15:07:15.726708Z","shell.execute_reply.started":"2023-12-28T15:07:15.721755Z","shell.execute_reply":"2023-12-28T15:07:15.725817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_mask = pd.read_csv('/kaggle/input/airbus-ship-detection/' + 'train_ship_segmentations_v2.csv', dtype = 'string', index_col = 'ImageId')\ndf_mask","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:15.727945Z","iopub.execute_input":"2023-12-28T15:07:15.728545Z","iopub.status.idle":"2023-12-28T15:07:17.152278Z","shell.execute_reply.started":"2023-12-28T15:07:15.728512Z","shell.execute_reply":"2023-12-28T15:07:17.15134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_mask = pd.DataFrame((df_mask['EncodedPixels'].fillna('') + ' ').groupby('ImageId').sum().str[:-1]).drop(CORRUPTED_TRAIN_IMAGE_IDS).reset_index()\ndf_mask = df_mask.sample(len(df_mask), random_state = 42)\nindex = ['train'] * int(len(df_mask) * 0.8) + ['val'] * int(len(df_mask) * 0.1)\nindex += ['test'] * (len(df_mask) - len(index))\ndf_mask.index = index\ndf_mask","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:17.155109Z","iopub.execute_input":"2023-12-28T15:07:17.1554Z","iopub.status.idle":"2023-12-28T15:07:17.849881Z","shell.execute_reply.started":"2023-12-28T15:07:17.155374Z","shell.execute_reply":"2023-12-28T15:07:17.848974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_filtered_mask = df_mask[df_mask['EncodedPixels'] != '']\ndf_filtered_mask","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:17.851064Z","iopub.execute_input":"2023-12-28T15:07:17.851364Z","iopub.status.idle":"2023-12-28T15:07:17.898148Z","shell.execute_reply.started":"2023-12-28T15:07:17.851341Z","shell.execute_reply":"2023-12-28T15:07:17.897271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = tf.data.Dataset.from_tensor_slices(df_filtered_mask.loc['train'])\nds_valid = tf.data.Dataset.from_tensor_slices(df_filtered_mask.loc['val'])\nds_test = tf.data.Dataset.from_tensor_slices(df_mask.loc['test'])","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:17.899128Z","iopub.execute_input":"2023-12-28T15:07:17.899385Z","iopub.status.idle":"2023-12-28T15:07:20.770526Z","shell.execute_reply.started":"2023-12-28T15:07:17.899361Z","shell.execute_reply":"2023-12-28T15:07:20.769516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def decode_rle(rle, shape = (768, 768)):\n    shape = tf.convert_to_tensor(shape, tf.int64)\n    rle = tf.strings.to_number(tf.strings.split(rle), tf.int64)\n    starts = rle[::2] - 1\n    lens = rle[1::2]\n    ones_len = tf.reduce_sum(lens)\n    ones = tf.ones([ones_len], tf.uint8)\n    # Make scattering indices\n    r = tf.range(ones_len)\n    lens_cum = tf.math.cumsum(lens)\n    s = tf.searchsorted(lens_cum, r, 'right')\n    idx = r + tf.gather(starts - tf.pad(lens_cum[:-1], [(1, 0)]), s)\n    return tf.transpose(tf.reshape(\n        tf.scatter_nd(tf.expand_dims(idx, 1), ones, [tf.math.reduce_prod(shape)]),\n        shape))","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:20.77165Z","iopub.execute_input":"2023-12-28T15:07:20.771969Z","iopub.status.idle":"2023-12-28T15:07:20.779809Z","shell.execute_reply.started":"2023-12-28T15:07:20.771941Z","shell.execute_reply":"2023-12-28T15:07:20.778838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"VALIDATION_LENGTH = len(df_filtered_mask.loc['val'])\nTEST_LENGTH = len(df_filtered_mask.loc['test'])\nTRAIN_LENGTH = len(df_filtered_mask.loc['train'])\nBATCH_SIZE = 8#16\nBUFFER_SIZE = 1000\n_SIZE_VAL = 128 * 2\nIMG_SIZE = (_SIZE_VAL, _SIZE_VAL)\nNUM_CLASSES = 2","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:20.780973Z","iopub.execute_input":"2023-12-28T15:07:20.781237Z","iopub.status.idle":"2023-12-28T15:07:20.805572Z","shell.execute_reply.started":"2023-12-28T15:07:20.781216Z","shell.execute_reply":"2023-12-28T15:07:20.804721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_train_image2(tensor) -> tuple:\n    img_id = tensor[0]\n    encoded_pixels = tensor[1]\n    img = tf.image.convert_image_dtype(\n        tf.io.decode_jpeg(\n            tf.io.read_file(TRAIN_DIR + img_id)\n        ),\n        tf.float32)\n    img = tf.image.resize(img, IMG_SIZE)\n    mask = decode_rle(encoded_pixels)\n    mask = tf.image.resize(mask[..., None], IMG_SIZE, method='nearest')[..., 0]\n    weights = tf.gather([0.05, 0.95], indices = tf.cast(mask, tf.int32), name = 'cast_sample_weights')\n    mask = tf.cast(tf.one_hot(mask, NUM_CLASSES), dtype=tf.float32)\n    \n    return img, mask, weights\n\n#ds_train = ds_train.map(lambda x: tf.py_function(load_train_image2, [x], [tf.float32, tf.float32]), num_parallel_calls=tf.data.AUTOTUNE)\n#ds_valid = ds_valid.map(lambda x: tf.py_function(load_train_image2, [x], [tf.float32, tf.float32]), num_parallel_calls=tf.data.AUTOTUNE)\n#ds_test = ds_test.map(lambda x: tf.py_function(load_train_image2, [x], [tf.float32, tf.float32]), num_parallel_calls=tf.data.AUTOTUNE)\nds_train = ds_train.map(load_train_image2, num_parallel_calls=tf.data.AUTOTUNE)\nds_valid = ds_valid.map(load_train_image2, num_parallel_calls=tf.data.AUTOTUNE)\nds_test = ds_test.map(load_train_image2, num_parallel_calls=tf.data.AUTOTUNE)\n\nclass Augment(tf.keras.layers.Layer):\n    def __init__(self, seed = SEED):\n        super().__init__()\n\n        self.rng = np.random.default_rng(seed)\n        self.rand_flip_imgs = keras.layers.RandomFlip(mode = \"horizontal\", seed = seed)\n        self.rand_flip_masks = keras.layers.RandomFlip(mode = \"horizontal\", seed = seed)\n    \n    def call(self, imgs, masks, weights):\n        imgs = self.rand_flip_imgs(imgs)\n        masks = self.rand_flip_masks(masks)\n\n        k = self.rng.choice(4)\n        imgs = tf.image.rot90(imgs, k)\n        masks = tf.image.rot90(masks, k)\n\n        return imgs, masks, weights\n\ntrain_batches = (\n    ds_train\n    .repeat()\n    .map(Augment())\n    .batch(BATCH_SIZE))\n\nvalidation_batches = ds_valid.batch(BATCH_SIZE)\n\ntest_batches = ds_test.batch(BATCH_SIZE)","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:20.806547Z","iopub.execute_input":"2023-12-28T15:07:20.8068Z","iopub.status.idle":"2023-12-28T15:07:21.927271Z","shell.execute_reply.started":"2023-12-28T15:07:20.806777Z","shell.execute_reply":"2023-12-28T15:07:21.926258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import keras.backend as K\nimport tensorflow_addons as tfa\n\nclass UNetModel:\n    def __init__(self, input_shape=(128, 128, 3)):\n        self._model = self._build_model(input_shape)\n\n    @property\n    def model(self) -> tf.keras.Model:\n        return self._model\n    \n    def _build_model(self, input_shape, num_classes=NUM_CLASSES) -> tf.keras.Model:\n        inputs = tf.keras.layers.Input(shape=input_shape)\n        \n        filters_list = [16, 32, 64]\n\n        # apply Encoder\n        encoder_outputs = self._encoder(input_shape, filters_list)(inputs)\n        print(f'Encoder output tensors: {encoder_outputs}')\n\n        # apply Decoder and establishing the skip connections\n        x = self._decoder(encoder_outputs, filters_list[::-1])\n\n        # This is the last layers of the model\n        last = self._conv_blocks(num_classes, size=1)(x)\n        outputs = tf.keras.activations.softmax(last)\n\n        return tf.keras.Model(inputs=inputs, outputs=outputs)\n    \n    def _encoder(self, input_shape, filters_list):\n        inputs = tf.keras.layers.Input(shape=input_shape)\n        outputs = []\n\n        model = tf.keras.Sequential()\n        x = model(inputs)\n\n        for filters in filters_list:\n            x = self._conv_blocks(filters=filters, size=3, apply_instance_norm=True)(x)\n            x = self._conv_blocks(filters=filters, size=1, apply_instance_norm=True)(x)\n            outputs.append(x)\n            x = tf.keras.layers.MaxPool2D(pool_size=(2, 2))(x)\n\n        output = self._conv_blocks(filters=128, size=3, apply_batch_norm=True, apply_dropout=False)(x)\n        outputs.append(output)\n\n        # Create the feature extraction model\n        encoder = tf.keras.Model(inputs=inputs, outputs=outputs, name=\"encoder\")\n        encoder.trainable = True\n        return encoder\n    \n    def _decoder(self, encoder_outputs, filters_list):     \n        x = encoder_outputs[-1]\n        for filters, skip, apply_dropout in zip(filters_list, encoder_outputs[-2::-1], [False] * 4):\n            x = self._upsample_block(filters, 3)(x)\n            x = tf.keras.layers.Concatenate()([x, skip])\n            x = self._conv_blocks(filters, size=3, apply_batch_norm=True, apply_dropout=apply_dropout)(x)\n            x = self._conv_blocks(filters, size=1, apply_batch_norm=True)(x)\n        return x\n    \n    def _conv_blocks(self, filters, size, apply_batch_norm=False, apply_instance_norm=False, apply_dropout=False):\n        \"\"\"Downsamples an input. Conv2D => Batchnorm => Dropout => LeakyReLU\n            :param:\n                filters: number of filters\n                size: filter size\n                apply_dropout: If True, adds the dropout layer\n            :return: Downsample Sequential Model\n        \"\"\"\n        initializer = tf.random_normal_initializer(0., 0.02)\n        result = tf.keras.Sequential()\n        result.add(\n          tf.keras.layers.Conv2D(filters, size, strides=1,\n                                 padding='same', use_bias=False,\n                                 kernel_initializer=initializer,))\n        if apply_batch_norm:\n            result.add(tf.keras.layers.BatchNormalization())\n        if apply_instance_norm:\n            result.add(tfa.layers.InstanceNormalization())\n        result.add(tf.keras.layers.Activation(tfa.activations.mish))\n        if apply_dropout:\n            result.add(tf.keras.layers.Dropout(0.55))\n        return result\n    \n    def _upsample_block(self, filters, size, apply_dropout=False):\n        \"\"\"Upsamples an input. Conv2DTranspose => Batchnorm => Dropout => LeakyReLU\n            :param:\n                filters: number of filters\n                size: filter size\n                apply_dropout: If True, adds the dropout layer\n            :return: Upsample Sequential Model\n        \"\"\"\n        initializer = tf.random_normal_initializer(0., 0.02)\n        result = tf.keras.Sequential()\n        result.add(\n          tf.keras.layers.Conv2DTranspose(filters, size, strides=2,\n                                          padding='same',\n                                          kernel_initializer=initializer,\n                                          use_bias=False))\n        result.add(tf.keras.layers.BatchNormalization())\n        if apply_dropout:\n            result.add(tf.keras.layers.Dropout(0.1))\n        result.add(tf.keras.layers.Activation(tfa.activations.mish))\n        return result\n    \n\ndef dice(targets, inputs, smooth=1e-6):\n    #axis = [1,2,3]\n    #intersection = K.sum(targets * inputs, axis=axis)\n    #intersection = tf.math.reduce_sum(targets * inputs)\n    #dice = (2 * intersection + smooth) / (K.sum(targets, axis=axis) + K.sum(inputs, axis=axis) + smooth)\n    #return dice\n    return (2 * tf.math.reduce_sum(targets * inputs)) / (tf.math.reduce_sum(targets) + tf.math.reduce_sum(inputs))\n\ndef bce_loss(targets, inputs, smooth=1e-6):\n    axis = [1,2,3]\n    loss = K.sum(targets * tf.math.log(inputs + smooth) + (1 - targets) * tf.math.log(1 - inputs + smooth), axis=axis)\n    return - loss\n\ndef bce_dice_loss(targets, inputs):\n    return bce_loss(targets, inputs) - tf.math.log(dice(targets, inputs))","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:21.928599Z","iopub.execute_input":"2023-12-28T15:07:21.928889Z","iopub.status.idle":"2023-12-28T15:07:22.262113Z","shell.execute_reply.started":"2023-12-28T15:07:21.928864Z","shell.execute_reply":"2023-12-28T15:07:22.26121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class IoU(tf.keras.metrics.Metric):\n    def __init__(self, num_classes: int, target_class_ids: list, sparse_y_true: bool, sparse_y_pred: bool,\n                 axis: int = -1, name=None, dtype=None):\n        super(IoU, self).__init__(name=name, dtype=dtype)\n        self.num_classes = num_classes\n        self.target_class_ids = target_class_ids\n        self.sparse_y_true = sparse_y_true\n        self.sparse_y_pred = sparse_y_pred\n        self.axis = axis\n\n        # Variable to accumulate the predictions in the confusion matrix.\n        self.total_cm = self.add_weight(\n            'total_confusion_matrix',\n            shape=(num_classes, num_classes),\n            initializer='zeros')\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        \"\"\"Accumulates the confusion matrix statistics.\n        Args:\n          y_true: The ground truth values.\n          y_pred: The predicted values.\n          sample_weight: Optional weighting of each example. Defaults to 1. Can be a\n            `Tensor` whose rank is either 0, or the same rank as `y_true`, and must\n            be broadcastable to `y_true`.\n        Returns:\n          Update op.\n        \"\"\"\n        \n        y_true = tf.reshape(y_true, [-1] + list(y_pred.shape[1:]))\n        \n        if not self.sparse_y_true:\n            y_true = tf.argmax(y_true, axis=self.axis)\n        if not self.sparse_y_pred:\n            y_pred = tf.argmax(y_pred, axis=self.axis)\n            \n        y_true = tf.cast(y_true, self._dtype)\n        y_pred = tf.cast(y_pred, self._dtype)\n\n        # Flatten the input if its rank > 1.\n        if y_pred.shape.ndims > 1:\n            y_pred = tf.reshape(y_pred, [-1])\n\n        if y_true.shape.ndims > 1:\n            y_true = tf.reshape(y_true, [-1])\n\n        if sample_weight is not None:\n            sample_weight = tf.reshape(sample_weight, [-1, 128, 128])\n            sample_weight = tf.cast(sample_weight, self._dtype)\n            if sample_weight.shape.ndims > 1:\n                sample_weight = tf.reshape(sample_weight, [-1])\n\n        # Accumulate the prediction to current confusion matrix.\n        current_cm = tf.math.confusion_matrix(y_true, y_pred, self.num_classes, weights=sample_weight, dtype=self._dtype)\n        return self.total_cm.assign_add(current_cm)\n    \n    def reset_state(self):\n        tf.keras.backend.set_value(\n            self.total_cm, np.zeros((self.num_classes, self.num_classes))\n        )\n    \n    def result(self):\n        \"\"\"Compute the intersection-over-union via the confusion matrix.\"\"\"\n        sum_over_row = tf.cast(\n            tf.reduce_sum(self.total_cm, axis=0), dtype=self._dtype)\n        sum_over_col = tf.cast(\n            tf.reduce_sum(self.total_cm, axis=1), dtype=self._dtype)\n        true_positives = tf.cast(\n            tf.linalg.tensor_diag_part(self.total_cm), dtype=self._dtype)\n\n        # sum_over_row + sum_over_col = 2 * true_positives + false_positives + false_negatives.\n        denominator = sum_over_row + sum_over_col - true_positives\n\n        # Only keep the target classes\n        true_positives = tf.gather(true_positives, self.target_class_ids)\n        denominator = tf.gather(denominator, self.target_class_ids)\n\n        # If the denominator is 0, we need to ignore the class.\n        num_valid_entries = tf.reduce_sum(\n            tf.cast(tf.not_equal(denominator, 0), dtype=self._dtype))\n\n        iou = tf.math.divide_no_nan(true_positives, denominator)\n\n        return tf.math.divide_no_nan(\n            tf.reduce_sum(iou, name='mean_iou'), num_valid_entries)\n    \n    def get_config(self):\n        config = {\n            \"num_classes\": self.num_classes,\n            \"target_class_ids\": self.target_class_ids,\n            \"sparse_y_true\": self.sparse_y_true,\n            \"sparse_y_pred\": self.sparse_y_pred,\n            \"axis\": self.axis,\n        }\n        base_config = super().get_config()\n        return dict(list(base_config.items()) + list(config.items()))","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:22.26358Z","iopub.execute_input":"2023-12-28T15:07:22.263961Z","iopub.status.idle":"2023-12-28T15:07:22.281839Z","shell.execute_reply.started":"2023-12-28T15:07:22.263927Z","shell.execute_reply":"2023-12-28T15:07:22.280974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"STEPS_PER_EPOCH = TRAIN_LENGTH // BATCH_SIZE\n\"\"\"\noptimizer = tfa.optimizers.RectifiedAdam(\n    learning_rate=0.005,\n    total_steps=EPOCHS * STEPS_PER_EPOCH,\n    warmup_proportion=0.3,\n    min_lr=0.00001,\n)\noptimizer = tfa.optimizers.Lookahead(optimizer)\n\"\"\"\nloss = tf.keras.losses.CategoricalCrossentropy()\nmIoU = IoU(num_classes=2, target_class_ids=[0, 1], sparse_y_true=False, sparse_y_pred=False, name='mean-IoU')\n\n#model = UNetModel(IMG_SIZE + (3,)).model\nmodel = build_resnet50_unet(IMG_SIZE + (3, ))\n\nmodel.compile(optimizer='adam', \n              loss=loss, # bce_dice_loss,\n              metrics=[mIoU, dice],)\n\ntrainable_params = np.sum([np.prod(v.get_shape().as_list()) for v in model.trainable_variables])\nprint(f'Trainable params: {trainable_params}')\n\ntf.keras.utils.plot_model(model, show_shapes=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:22.323527Z","iopub.execute_input":"2023-12-28T15:07:22.3238Z","iopub.status.idle":"2023-12-28T15:07:27.800392Z","shell.execute_reply.started":"2023-12-28T15:07:22.323776Z","shell.execute_reply":"2023-12-28T15:07:27.799226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_filepath = '/kaggle/working/checkpoints/model-checkpoint'\nsave_callback = keras.callbacks.ModelCheckpoint(\n    filepath=checkpoint_filepath,\n    monitor='val_mean-IoU',\n    mode='max',\n    save_best_only=True\n)\n\nmodel_history = model.fit(train_batches,\n                          epochs = 3,\n                          steps_per_epoch = STEPS_PER_EPOCH,\n                          validation_data = validation_batches,\n                          callbacks = [save_callback])","metadata":{"execution":{"iopub.status.busy":"2023-12-28T15:07:27.801759Z","iopub.execute_input":"2023-12-28T15:07:27.802247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(image):\n    image = np.expand_dims(image, axis=0)\n    pred_mask = model.predict(image)[0].argmax(axis=-1)\n    return pred_mask","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N = 50\n\nf,ax = plt.subplots(N, 3, figsize=(10, 4 * N))\ni = 0\nfor image, mask in ds_test.take(N):\n    mask = mask.numpy().argmax(axis=-1)\n    ax[i, 0].imshow(image)\n    ax[i, 0].set_title('image')\n    ax[i, 1].imshow(mask)\n    ax[i, 1].set_title('true mask')\n\n    pred_mask = predict(image)\n    ax[i, 2].imshow(pred_mask)\n    ax[i, 2].set_title('predicted mask')\n    i += 1\n\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}