{
  "id": 619283,
  "title": "Baseline segmentation",
  "url": "/competitions/vesuvius-challenge-surface-detection/discussion/619283",
  "author_name": "Diwakar30",
  "post_date": "2025-11-13T20:26:39.560000",
  "votes": 1,
  "comment_count": 0,
  "views": 0,
  "content": "<h1>Vesuvius Challenge - Phase 1 Baseline</h1>\n<h1>2D U-Net segmentation with TensorFlow</h1>\n<p>import os\nimport numpy as np\nimport tensorflow as tf\nfrom PIL import Image\nimport matplotlib.pyplot as plt</p>\n<h1>==========================</h1>\n<h1>Paths</h1>\n<h1>==========================</h1>\n<p>train_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_images\"\ntrain_labels_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_labels\"\ntest_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/test_images\"</p>\n<p>train_image_files = sorted(tf.io.gfile.glob(os.path.join(train_images_dir, \"<em>.tif\")))\ntrain_label_files = sorted(tf.io.gfile.glob(os.path.join(train_labels_dir, \"</em>.tif\")))</p>\n<p>print(\"Total training samples:\", len(train_image_files))</p>\n<h1>==========================</h1>\n<h1>Data loader</h1>\n<h1>==========================</h1>\n<p>def _decode_tiff_pil(file_path, is_mask=False):\n    path_str = file_path.numpy().decode(\"utf-8\")\n    img = Image.open(path_str)\n    if is_mask:\n        img = img.convert(\"L\")\n        arr = np.array(img, dtype=np.float32)\n        arr = np.expand_dims(arr, axis=-1)\n        return arr\n    else:\n        img = img.convert(\"RGB\")\n        return np.array(img, dtype=np.float32)</p>\n<p>def load_image_pair(image_path, label_path):\n    [image] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=False),\n        inp=[image_path],\n        Tout=[tf.float32]\n    )\n    [label] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=True),\n        inp=[label_path],\n        Tout=[tf.float32]\n    )\n    image.set_shape([None, None, 3])\n    label.set_shape([None, None, 1])\n    return image, label</p>\n<p>def preprocess(image, label, size=(256,256)):\n    image = tf.image.resize(image, size, method=\"bilinear\") / 255.0\n    label = tf.image.resize(label, size, method=\"nearest\")\n    return image, tf.cast(label, tf.int32)</p>\n<p>train_ds = tf.data.Dataset.from_tensor_slices((train_image_files, train_label_files))\ntrain_ds = train_ds.map(load_image_pair, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.map(lambda x, y: preprocess(x, y, size=(256,256)), num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.shuffle(100).batch(4).prefetch(tf.data.AUTOTUNE)</p>\n<h1>==========================</h1>\n<h1>Simple 2D U-Net</h1>\n<h1>==========================</h1>\n<p>def conv_block(x, filters):\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    return x</p>\n<p>def build_unet(input_shape=(256,256,3), num_classes=3):\n    inputs = tf.keras.Input(shape=input_shape)</p>\n<pre><code>\nc1 = conv_block(inputs, 32)\np1 = tf.keras.layers.MaxPool2D()(c1)\nc2 = conv_block(p1, 64)\np2 = tf.keras.layers.MaxPool2D()(c2)\nc3 = conv_block(p2, 128)\np3 = tf.keras.layers.MaxPool2D()(c3)\nc4 = conv_block(p3, 256)\np4 = tf.keras.layers.MaxPool2D()(c4)\n\n\nbn = conv_block(p4, 512)\n\n\nu1 = tf.keras.layers.UpSampling2D()(bn)\nu1 = tf.keras.layers.Concatenate()([u1, c4])\nc5 = conv_block(u1, 256)\n\nu2 = tf.keras.layers.UpSampling2D()(c5)\nu2 = tf.keras.layers.Concatenate()([u2, c3])\nc6 = conv_block(u2, 128)\n\nu3 = tf.keras.layers.UpSampling2D()(c6)\nu3 = tf.keras.layers.Concatenate()([u3, c2])\nc7 = conv_block(u3, 64)\n\nu4 = tf.keras.layers.UpSampling2D()(c7)\nu4 = tf.keras.layers.Concatenate()([u4, c1])\nc8 = conv_block(u4, 32)\n\noutputs = tf.keras.layers.Conv2D(num_classes, 1, activation=)(c8)\nreturn tf.keras.Model(inputs, outputs)\n</code></pre>\n<p>model = build_unet()\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(1e-4),\n    loss=tf.keras.losses.SparseCategoricalCrossentropy(),\n    metrics=[\"accuracy\"]\n)\nmodel.summary()</p>\n<h1>==========================</h1>\n<h1>Train model</h1>\n<h1>==========================</h1>\n<p>history = model.fit(\n    train_ds,\n    epochs=5  # increase later\n)</p>\n<h1>==========================</h1>\n<h1>Visualize predictions</h1>\n<h1>==========================</h1>\n<p>def visualize(image, mask, pred_mask=None):\n    plt.figure(figsize=(12,4))\n    plt.subplot(1,3,1)\n    plt.imshow(image)\n    plt.title(\"Image\")\n    plt.subplot(1,3,2)\n    plt.imshow(mask.squeeze(), cmap='gray')\n    plt.title(\"Mask\")\n    if pred_mask is not None:\n        plt.subplot(1,3,3)\n        plt.imshow(pred_mask.squeeze(), cmap='gray')\n        plt.title(\"Prediction\")\n    plt.show()</p>\n<p>for img, lbl in train_ds.take(1):\n    preds = model.predict(img)\n    pred_classes = tf.argmax(preds, axis=-1)[…, tf.newaxis]\n    for i in range(2):\n        visualize(img[i].numpy(), lbl[i].numpy(), pred_classes[i].numpy())</p>\n<h1>==========================</h1>\n<h1>Save predictions for submission</h1>\n<h1>==========================</h1>\n<p>import tifffile</p>\n<p>test_image_files = sorted(tf.io.gfile.glob(os.path.join(test_images_dir, \"*.tif\")))</p>\n<p>output_dir = \"/kaggle/working/predictions\"\nos.makedirs(output_dir, exist_ok=True)</p>\n<p>for path in test_image_files[:5]:  # change to all images for final submission\n    img = preprocess(_decode_tiff_pil(path, is_mask=False), np.zeros((1,1,1)), size=(256,256))[0]\n    img_batch = tf.expand_dims(img, 0)\n    pred = model.predict(img_batch)\n    pred_class = tf.argmax(pred, axis=-1).numpy()[0].astype(np.uint8)\n    tifffile.imwrite(os.path.join(output_dir, os.path.basename(path)), pred_class)</p>",
  "messages": [
    {
      "id": 3322911,
      "postDate": "2025-11-13T20:26:39.560Z",
      "content": "<h1>Vesuvius Challenge - Phase 1 Baseline</h1>\n<h1>2D U-Net segmentation with TensorFlow</h1>\n<p>import os\nimport numpy as np\nimport tensorflow as tf\nfrom PIL import Image\nimport matplotlib.pyplot as plt</p>\n<h1>==========================</h1>\n<h1>Paths</h1>\n<h1>==========================</h1>\n<p>train_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_images\"\ntrain_labels_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_labels\"\ntest_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/test_images\"</p>\n<p>train_image_files = sorted(tf.io.gfile.glob(os.path.join(train_images_dir, \"<em>.tif\")))\ntrain_label_files = sorted(tf.io.gfile.glob(os.path.join(train_labels_dir, \"</em>.tif\")))</p>\n<p>print(\"Total training samples:\", len(train_image_files))</p>\n<h1>==========================</h1>\n<h1>Data loader</h1>\n<h1>==========================</h1>\n<p>def _decode_tiff_pil(file_path, is_mask=False):\n    path_str = file_path.numpy().decode(\"utf-8\")\n    img = Image.open(path_str)\n    if is_mask:\n        img = img.convert(\"L\")\n        arr = np.array(img, dtype=np.float32)\n        arr = np.expand_dims(arr, axis=-1)\n        return arr\n    else:\n        img = img.convert(\"RGB\")\n        return np.array(img, dtype=np.float32)</p>\n<p>def load_image_pair(image_path, label_path):\n    [image] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=False),\n        inp=[image_path],\n        Tout=[tf.float32]\n    )\n    [label] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=True),\n        inp=[label_path],\n        Tout=[tf.float32]\n    )\n    image.set_shape([None, None, 3])\n    label.set_shape([None, None, 1])\n    return image, label</p>\n<p>def preprocess(image, label, size=(256,256)):\n    image = tf.image.resize(image, size, method=\"bilinear\") / 255.0\n    label = tf.image.resize(label, size, method=\"nearest\")\n    return image, tf.cast(label, tf.int32)</p>\n<p>train_ds = tf.data.Dataset.from_tensor_slices((train_image_files, train_label_files))\ntrain_ds = train_ds.map(load_image_pair, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.map(lambda x, y: preprocess(x, y, size=(256,256)), num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.shuffle(100).batch(4).prefetch(tf.data.AUTOTUNE)</p>\n<h1>==========================</h1>\n<h1>Simple 2D U-Net</h1>\n<h1>==========================</h1>\n<p>def conv_block(x, filters):\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    return x</p>\n<p>def build_unet(input_shape=(256,256,3), num_classes=3):\n    inputs = tf.keras.Input(shape=input_shape)</p>\n<pre><code>\nc1 = conv_block(inputs, 32)\np1 = tf.keras.layers.MaxPool2D()(c1)\nc2 = conv_block(p1, 64)\np2 = tf.keras.layers.MaxPool2D()(c2)\nc3 = conv_block(p2, 128)\np3 = tf.keras.layers.MaxPool2D()(c3)\nc4 = conv_block(p3, 256)\np4 = tf.keras.layers.MaxPool2D()(c4)\n\n\nbn = conv_block(p4, 512)\n\n\nu1 = tf.keras.layers.UpSampling2D()(bn)\nu1 = tf.keras.layers.Concatenate()([u1, c4])\nc5 = conv_block(u1, 256)\n\nu2 = tf.keras.layers.UpSampling2D()(c5)\nu2 = tf.keras.layers.Concatenate()([u2, c3])\nc6 = conv_block(u2, 128)\n\nu3 = tf.keras.layers.UpSampling2D()(c6)\nu3 = tf.keras.layers.Concatenate()([u3, c2])\nc7 = conv_block(u3, 64)\n\nu4 = tf.keras.layers.UpSampling2D()(c7)\nu4 = tf.keras.layers.Concatenate()([u4, c1])\nc8 = conv_block(u4, 32)\n\noutputs = tf.keras.layers.Conv2D(num_classes, 1, activation=)(c8)\nreturn tf.keras.Model(inputs, outputs)\n</code></pre>\n<p>model = build_unet()\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(1e-4),\n    loss=tf.keras.losses.SparseCategoricalCrossentropy(),\n    metrics=[\"accuracy\"]\n)\nmodel.summary()</p>\n<h1>==========================</h1>\n<h1>Train model</h1>\n<h1>==========================</h1>\n<p>history = model.fit(\n    train_ds,\n    epochs=5  # increase later\n)</p>\n<h1>==========================</h1>\n<h1>Visualize predictions</h1>\n<h1>==========================</h1>\n<p>def visualize(image, mask, pred_mask=None):\n    plt.figure(figsize=(12,4))\n    plt.subplot(1,3,1)\n    plt.imshow(image)\n    plt.title(\"Image\")\n    plt.subplot(1,3,2)\n    plt.imshow(mask.squeeze(), cmap='gray')\n    plt.title(\"Mask\")\n    if pred_mask is not None:\n        plt.subplot(1,3,3)\n        plt.imshow(pred_mask.squeeze(), cmap='gray')\n        plt.title(\"Prediction\")\n    plt.show()</p>\n<p>for img, lbl in train_ds.take(1):\n    preds = model.predict(img)\n    pred_classes = tf.argmax(preds, axis=-1)[…, tf.newaxis]\n    for i in range(2):\n        visualize(img[i].numpy(), lbl[i].numpy(), pred_classes[i].numpy())</p>\n<h1>==========================</h1>\n<h1>Save predictions for submission</h1>\n<h1>==========================</h1>\n<p>import tifffile</p>\n<p>test_image_files = sorted(tf.io.gfile.glob(os.path.join(test_images_dir, \"*.tif\")))</p>\n<p>output_dir = \"/kaggle/working/predictions\"\nos.makedirs(output_dir, exist_ok=True)</p>\n<p>for path in test_image_files[:5]:  # change to all images for final submission\n    img = preprocess(_decode_tiff_pil(path, is_mask=False), np.zeros((1,1,1)), size=(256,256))[0]\n    img_batch = tf.expand_dims(img, 0)\n    pred = model.predict(img_batch)\n    pred_class = tf.argmax(pred, axis=-1).numpy()[0].astype(np.uint8)\n    tifffile.imwrite(os.path.join(output_dir, os.path.basename(path)), pred_class)</p>",
      "rawMarkdown": "# Vesuvius Challenge - Phase 1 Baseline\n# 2D U-Net segmentation with TensorFlow\n\nimport os\nimport numpy as np\nimport tensorflow as tf\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# ==========================\n# Paths\n# ==========================\ntrain_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_images\"\ntrain_labels_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_labels\"\ntest_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/test_images\"\n\ntrain_image_files = sorted(tf.io.gfile.glob(os.path.join(train_images_dir, \"*.tif\")))\ntrain_label_files = sorted(tf.io.gfile.glob(os.path.join(train_labels_dir, \"*.tif\")))\n\nprint(\"Total training samples:\", len(train_image_files))\n\n# ==========================\n# Data loader\n# ==========================\ndef _decode_tiff_pil(file_path, is_mask=False):\n    path_str = file_path.numpy().decode(\"utf-8\")\n    img = Image.open(path_str)\n    if is_mask:\n        img = img.convert(\"L\")\n        arr = np.array(img, dtype=np.float32)\n        arr = np.expand_dims(arr, axis=-1)\n        return arr\n    else:\n        img = img.convert(\"RGB\")\n        return np.array(img, dtype=np.float32)\n\ndef load_image_pair(image_path, label_path):\n    [image] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=False),\n        inp=[image_path],\n        Tout=[tf.float32]\n    )\n    [label] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=True),\n        inp=[label_path],\n        Tout=[tf.float32]\n    )\n    image.set_shape([None, None, 3])\n    label.set_shape([None, None, 1])\n    return image, label\n\ndef preprocess(image, label, size=(256,256)):\n    image = tf.image.resize(image, size, method=\"bilinear\") / 255.0\n    label = tf.image.resize(label, size, method=\"nearest\")\n    return image, tf.cast(label, tf.int32)\n\ntrain_ds = tf.data.Dataset.from_tensor_slices((train_image_files, train_label_files))\ntrain_ds = train_ds.map(load_image_pair, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.map(lambda x, y: preprocess(x, y, size=(256,256)), num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.shuffle(100).batch(4).prefetch(tf.data.AUTOTUNE)\n\n# ==========================\n# Simple 2D U-Net\n# ==========================\ndef conv_block(x, filters):\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    return x\n\ndef build_unet(input_shape=(256,256,3), num_classes=3):\n    inputs = tf.keras.Input(shape=input_shape)\n    \n    # Encoder\n    c1 = conv_block(inputs, 32)\n    p1 = tf.keras.layers.MaxPool2D()(c1)\n    c2 = conv_block(p1, 64)\n    p2 = tf.keras.layers.MaxPool2D()(c2)\n    c3 = conv_block(p2, 128)\n    p3 = tf.keras.layers.MaxPool2D()(c3)\n    c4 = conv_block(p3, 256)\n    p4 = tf.keras.layers.MaxPool2D()(c4)\n\n    # Bottleneck\n    bn = conv_block(p4, 512)\n\n    # Decoder\n    u1 = tf.keras.layers.UpSampling2D()(bn)\n    u1 = tf.keras.layers.Concatenate()([u1, c4])\n    c5 = conv_block(u1, 256)\n\n    u2 = tf.keras.layers.UpSampling2D()(c5)\n    u2 = tf.keras.layers.Concatenate()([u2, c3])\n    c6 = conv_block(u2, 128)\n\n    u3 = tf.keras.layers.UpSampling2D()(c6)\n    u3 = tf.keras.layers.Concatenate()([u3, c2])\n    c7 = conv_block(u3, 64)\n\n    u4 = tf.keras.layers.UpSampling2D()(c7)\n    u4 = tf.keras.layers.Concatenate()([u4, c1])\n    c8 = conv_block(u4, 32)\n\n    outputs = tf.keras.layers.Conv2D(num_classes, 1, activation=\"softmax\")(c8)\n    return tf.keras.Model(inputs, outputs)\n\nmodel = build_unet()\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(1e-4),\n    loss=tf.keras.losses.SparseCategoricalCrossentropy(),\n    metrics=[\"accuracy\"]\n)\nmodel.summary()\n\n# ==========================\n# Train model\n# ==========================\nhistory = model.fit(\n    train_ds,\n    epochs=5  # increase later\n)\n\n# ==========================\n# Visualize predictions\n# ==========================\ndef visualize(image, mask, pred_mask=None):\n    plt.figure(figsize=(12,4))\n    plt.subplot(1,3,1)\n    plt.imshow(image)\n    plt.title(\"Image\")\n    plt.subplot(1,3,2)\n    plt.imshow(mask.squeeze(), cmap='gray')\n    plt.title(\"Mask\")\n    if pred_mask is not None:\n        plt.subplot(1,3,3)\n        plt.imshow(pred_mask.squeeze(), cmap='gray')\n        plt.title(\"Prediction\")\n    plt.show()\n\nfor img, lbl in train_ds.take(1):\n    preds = model.predict(img)\n    pred_classes = tf.argmax(preds, axis=-1)[..., tf.newaxis]\n    for i in range(2):\n        visualize(img[i].numpy(), lbl[i].numpy(), pred_classes[i].numpy())\n\n# ==========================\n# Save predictions for submission\n# ==========================\nimport tifffile\n\ntest_image_files = sorted(tf.io.gfile.glob(os.path.join(test_images_dir, \"*.tif\")))\n\noutput_dir = \"/kaggle/working/predictions\"\nos.makedirs(output_dir, exist_ok=True)\n\nfor path in test_image_files[:5]:  # change to all images for final submission\n    img = preprocess(_decode_tiff_pil(path, is_mask=False), np.zeros((1,1,1)), size=(256,256))[0]\n    img_batch = tf.expand_dims(img, 0)\n    pred = model.predict(img_batch)\n    pred_class = tf.argmax(pred, axis=-1).numpy()[0].astype(np.uint8)\n    tifffile.imwrite(os.path.join(output_dir, os.path.basename(path)), pred_class)\n",
      "votes": 1
    }
  ],
  "comments": [],
  "raw_markdown_by_id": {
    "3322911": "# Vesuvius Challenge - Phase 1 Baseline\n# 2D U-Net segmentation with TensorFlow\n\nimport os\nimport numpy as np\nimport tensorflow as tf\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# ==========================\n# Paths\n# ==========================\ntrain_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_images\"\ntrain_labels_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_labels\"\ntest_images_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/test_images\"\n\ntrain_image_files = sorted(tf.io.gfile.glob(os.path.join(train_images_dir, \"*.tif\")))\ntrain_label_files = sorted(tf.io.gfile.glob(os.path.join(train_labels_dir, \"*.tif\")))\n\nprint(\"Total training samples:\", len(train_image_files))\n\n# ==========================\n# Data loader\n# ==========================\ndef _decode_tiff_pil(file_path, is_mask=False):\n    path_str = file_path.numpy().decode(\"utf-8\")\n    img = Image.open(path_str)\n    if is_mask:\n        img = img.convert(\"L\")\n        arr = np.array(img, dtype=np.float32)\n        arr = np.expand_dims(arr, axis=-1)\n        return arr\n    else:\n        img = img.convert(\"RGB\")\n        return np.array(img, dtype=np.float32)\n\ndef load_image_pair(image_path, label_path):\n    [image] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=False),\n        inp=[image_path],\n        Tout=[tf.float32]\n    )\n    [label] = tf.py_function(\n        func=lambda x: _decode_tiff_pil(x, is_mask=True),\n        inp=[label_path],\n        Tout=[tf.float32]\n    )\n    image.set_shape([None, None, 3])\n    label.set_shape([None, None, 1])\n    return image, label\n\ndef preprocess(image, label, size=(256,256)):\n    image = tf.image.resize(image, size, method=\"bilinear\") / 255.0\n    label = tf.image.resize(label, size, method=\"nearest\")\n    return image, tf.cast(label, tf.int32)\n\ntrain_ds = tf.data.Dataset.from_tensor_slices((train_image_files, train_label_files))\ntrain_ds = train_ds.map(load_image_pair, num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.map(lambda x, y: preprocess(x, y, size=(256,256)), num_parallel_calls=tf.data.AUTOTUNE)\ntrain_ds = train_ds.shuffle(100).batch(4).prefetch(tf.data.AUTOTUNE)\n\n# ==========================\n# Simple 2D U-Net\n# ==========================\ndef conv_block(x, filters):\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    x = tf.keras.layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n    x = tf.keras.layers.BatchNormalization()(x)\n    return x\n\ndef build_unet(input_shape=(256,256,3), num_classes=3):\n    inputs = tf.keras.Input(shape=input_shape)\n    \n    # Encoder\n    c1 = conv_block(inputs, 32)\n    p1 = tf.keras.layers.MaxPool2D()(c1)\n    c2 = conv_block(p1, 64)\n    p2 = tf.keras.layers.MaxPool2D()(c2)\n    c3 = conv_block(p2, 128)\n    p3 = tf.keras.layers.MaxPool2D()(c3)\n    c4 = conv_block(p3, 256)\n    p4 = tf.keras.layers.MaxPool2D()(c4)\n\n    # Bottleneck\n    bn = conv_block(p4, 512)\n\n    # Decoder\n    u1 = tf.keras.layers.UpSampling2D()(bn)\n    u1 = tf.keras.layers.Concatenate()([u1, c4])\n    c5 = conv_block(u1, 256)\n\n    u2 = tf.keras.layers.UpSampling2D()(c5)\n    u2 = tf.keras.layers.Concatenate()([u2, c3])\n    c6 = conv_block(u2, 128)\n\n    u3 = tf.keras.layers.UpSampling2D()(c6)\n    u3 = tf.keras.layers.Concatenate()([u3, c2])\n    c7 = conv_block(u3, 64)\n\n    u4 = tf.keras.layers.UpSampling2D()(c7)\n    u4 = tf.keras.layers.Concatenate()([u4, c1])\n    c8 = conv_block(u4, 32)\n\n    outputs = tf.keras.layers.Conv2D(num_classes, 1, activation=\"softmax\")(c8)\n    return tf.keras.Model(inputs, outputs)\n\nmodel = build_unet()\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(1e-4),\n    loss=tf.keras.losses.SparseCategoricalCrossentropy(),\n    metrics=[\"accuracy\"]\n)\nmodel.summary()\n\n# ==========================\n# Train model\n# ==========================\nhistory = model.fit(\n    train_ds,\n    epochs=5  # increase later\n)\n\n# ==========================\n# Visualize predictions\n# ==========================\ndef visualize(image, mask, pred_mask=None):\n    plt.figure(figsize=(12,4))\n    plt.subplot(1,3,1)\n    plt.imshow(image)\n    plt.title(\"Image\")\n    plt.subplot(1,3,2)\n    plt.imshow(mask.squeeze(), cmap='gray')\n    plt.title(\"Mask\")\n    if pred_mask is not None:\n        plt.subplot(1,3,3)\n        plt.imshow(pred_mask.squeeze(), cmap='gray')\n        plt.title(\"Prediction\")\n    plt.show()\n\nfor img, lbl in train_ds.take(1):\n    preds = model.predict(img)\n    pred_classes = tf.argmax(preds, axis=-1)[..., tf.newaxis]\n    for i in range(2):\n        visualize(img[i].numpy(), lbl[i].numpy(), pred_classes[i].numpy())\n\n# ==========================\n# Save predictions for submission\n# ==========================\nimport tifffile\n\ntest_image_files = sorted(tf.io.gfile.glob(os.path.join(test_images_dir, \"*.tif\")))\n\noutput_dir = \"/kaggle/working/predictions\"\nos.makedirs(output_dir, exist_ok=True)\n\nfor path in test_image_files[:5]:  # change to all images for final submission\n    img = preprocess(_decode_tiff_pil(path, is_mask=False), np.zeros((1,1,1)), size=(256,256))[0]\n    img_batch = tf.expand_dims(img, 0)\n    pred = model.predict(img_batch)\n    pred_class = tf.argmax(pred, axis=-1).numpy()[0].astype(np.uint8)\n    tifffile.imwrite(os.path.join(output_dir, os.path.basename(path)), pred_class)\n"
  }
}