{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"},{"sourceId":13753886,"sourceType":"datasetVersion","datasetId":8751895}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# About\n\n- Build `2.5D` input data.\n- Build `tf.data` API with **TFRecord**.\n- Build **2D** model from [`medicai`](https://github.com/innat/medic-ai), i.e. [**TransUNet**.](https://github.com/innat/medic-ai/blob/main/medicai/models/transunet/README.md)\n- **Multi-GPU** (or **TPU**) setup with `jax` backend.\n- Train model with **Keras** built-in training API.","metadata":{}},{"cell_type":"markdown","source":"## Installation and Import","metadata":{}},{"cell_type":"code","source":"# The `medicai` is medical-based 2D and 3D ML library. \n# We'll use it for segmentaiton model, 3D volume transformation, etc.\n!pip install git+https://github.com/innat/medic-ai.git -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:08:11.246443Z","iopub.execute_input":"2025-11-20T11:08:11.246681Z","iopub.status.idle":"2025-11-20T11:08:20.941104Z","shell.execute_reply.started":"2025-11-20T11:08:11.246657Z","shell.execute_reply":"2025-11-20T11:08:20.940407Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, warnings\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\nwarnings.filterwarnings('ignore')\n\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:08:20.942169Z","iopub.execute_input":"2025-11-20T11:08:20.942473Z","iopub.status.idle":"2025-11-20T11:08:21.212084Z","shell.execute_reply.started":"2025-11-20T11:08:20.942448Z","shell.execute_reply":"2025-11-20T11:08:21.211505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras\nfrom keras import ops\n\nimport jax\nimport tensorflow as tf\n\nkeras.version(), jax.__version__, keras.config.backend()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:08:21.213609Z","iopub.execute_input":"2025-11-20T11:08:21.214075Z","iopub.status.idle":"2025-11-20T11:08:36.209358Z","shell.execute_reply.started":"2025-11-20T11:08:21.214047Z","shell.execute_reply":"2025-11-20T11:08:36.208708Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Multi-GPU Setup\n\n**Note**: Same setup would also work on **TPU-VM**","metadata":{}},{"cell_type":"code","source":"# due to distributed training only\nkeras.config.disable_flash_attention()\n\n# reproducibility\nkeras.utils.set_random_seed(101)\n\n# distributed config\ndevices = keras.distribution.list_devices()\ndata_parallel = keras.distribution.DataParallel(devices=devices)\nkeras.distribution.set_distribution(data_parallel)\ntotal_device = len(devices)\n\nprint(f'detected devices: {devices}')\nprint(f'total device: {total_device}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:08:36.210005Z","iopub.execute_input":"2025-11-20T11:08:36.210441Z","iopub.status.idle":"2025-11-20T11:08:36.917273Z","shell.execute_reply.started":"2025-11-20T11:08:36.210415Z","shell.execute_reply":"2025-11-20T11:08:36.916489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_tfrecord_fn(example):\n    feature_description = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"label\": tf.io.FixedLenFeature([], tf.string),\n        \"image_shape\": tf.io.FixedLenFeature([3], tf.int64),\n        \"label_shape\": tf.io.FixedLenFeature([3], tf.int64),\n    }\n    parsed_example = tf.io.parse_single_example(example, feature_description)\n    image = tf.io.decode_raw(parsed_example[\"image\"], tf.uint8)\n    label = tf.io.decode_raw(parsed_example[\"label\"], tf.uint8)\n    image_shape = tf.cast(parsed_example[\"image_shape\"], tf.int64)\n    label_shape = tf.cast(parsed_example[\"label_shape\"], tf.int64)\n    image = tf.reshape(image, image_shape)\n    label = tf.reshape(label, label_shape)\n    return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:08:50.115801Z","iopub.execute_input":"2025-11-20T11:08:50.116104Z","iopub.status.idle":"2025-11-20T11:08:50.121325Z","shell.execute_reply.started":"2025-11-20T11:08:50.116082Z","shell.execute_reply":"2025-11-20T11:08:50.120649Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"input_shape = 320\nbatch_size = 6 * total_device\nslices_radius = 16 # input_channel = 2*slices_radius + 1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:08:55.548145Z","iopub.execute_input":"2025-11-20T11:08:55.548873Z","iopub.status.idle":"2025-11-20T11:08:55.552365Z","shell.execute_reply.started":"2025-11-20T11:08:55.548847Z","shell.execute_reply":"2025-11-20T11:08:55.551554Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prepare 2.5D Input\n\n**Key idea:** Use multiple slices as context to predict a single slice’s mask.\n\n1. **Input preparation**\n    - For each training sample, select a **center slice** from the `3D` volume.\n    - Stack **neighboring slices** around the center to form a `2.5D` input:\n    - `Input shape: (H, W, 2*slices_radius + 1)`\n2. **Label selection**\n    - Use only the **center slice’s** label as the target:","metadata":{}},{"cell_type":"code","source":"def process_inputs(image, label):\n    # cast to float32\n    image = tf.cast(image, tf.float32)\n    label = tf.cast(label, tf.float32)\n\n    # resize image with linaer interpolation\n    image = tf.image.resize(\n        images=image, \n        size=[input_shape, input_shape], \n        method='bilinear'\n    )\n    # normalize to [0, 1]\n    image = image / 255.\n\n    # resize label / mask with nearest interpolation\n    label = tf.image.resize(\n        images=label, \n        size=[input_shape, input_shape], \n        method='nearest'\n    )\n\n    return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:08:59.467339Z","iopub.execute_input":"2025-11-20T11:08:59.468070Z","iopub.status.idle":"2025-11-20T11:08:59.472527Z","shell.execute_reply.started":"2025-11-20T11:08:59.468039Z","shell.execute_reply":"2025-11-20T11:08:59.471810Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_training_inputs(image, label, k=2):\n    z = tf.shape(image)[0]\n\n    # valid center slice\n    z_center = tf.random.uniform(\n        shape=[], \n        minval=k, \n        maxval=z-k, \n        dtype=tf.int32\n    )\n\n    # gather neighbor slices: [z-k, ..., z, ..., z+k]\n    idxs = tf.range(z_center - k, z_center + k + 1)\n    image_25d = tf.gather(image, idxs, axis=0)\n    image_25d = tf.transpose(image_25d, [1, 2, 0])\n\n    # central slice (Y, X)\n    label_2d = label[z_center]   \n    label_2d = (label_2d == 1)\n    label_2d = label_2d[..., None]\n\n    image_25d, label_2d = process_inputs(\n        image_25d, label_2d\n    )\n    return image_25d, label_2d\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:00.659296Z","iopub.execute_input":"2025-11-20T11:09:00.659991Z","iopub.status.idle":"2025-11-20T11:09:00.664885Z","shell.execute_reply.started":"2025-11-20T11:09:00.659968Z","shell.execute_reply":"2025-11-20T11:09:00.664057Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_validation_inputs(image, label, k=2):\n    z = tf.shape(image)[0]\n    idxs = tf.range(k, z - k)\n\n    def create_slices(z_center):\n        window_idxs = tf.range(z_center - k, z_center + k + 1)\n        image_25d = tf.gather(image, window_idxs, axis=0)\n        image_25d = tf.transpose(image_25d, [1, 2, 0]) \n\n        # central slice (Y, X)\n        label_2d = label[z_center]\n        label_2d = (label_2d == 1)\n        label_2d = label_2d[..., None]\n\n        image_25d, label_2d = process_inputs(\n            image_25d, label_2d\n        )\n        return image_25d, label_2d\n\n    return tf.data.Dataset.from_tensor_slices(idxs).map(\n        create_slices, num_parallel_calls=tf.data.AUTOTUNE\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:01.722877Z","iopub.execute_input":"2025-11-20T11:09:01.723155Z","iopub.status.idle":"2025-11-20T11:09:01.729103Z","shell.execute_reply.started":"2025-11-20T11:09:01.723135Z","shell.execute_reply":"2025-11-20T11:09:01.728454Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Augmentation and Data Loader","metadata":{}},{"cell_type":"code","source":"aug_layers = [\n    keras.layers.RandomFlip(\"horizontal_and_vertical\"),\n    keras.layers.RandomBrightness(factor=0.2, value_range=(0, 1))\n]\n\ndef augment_data(x, y):\n    c = tf.shape(x)[-1]\n    z = tf.concat([x, y], axis=-1)\n    \n    # apply augmentations\n    for layer in aug_layers:\n        z = layer(z)\n\n    # split back\n    x = z[..., :c]\n    y = z[..., c:]\n\n    # ensure mask is binary again\n    y = tf.cast(tf.round(y), tf.float32)\n\n    return x, y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:02.764435Z","iopub.execute_input":"2025-11-20T11:09:02.764957Z","iopub.status.idle":"2025-11-20T11:09:03.823595Z","shell.execute_reply.started":"2025-11-20T11:09:02.764931Z","shell.execute_reply":"2025-11-20T11:09:03.822984Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_tfrecord_dataset(\n    tfrecord_pattern, batch_size=1, training=True\n):\n    dataset = tf.data.TFRecordDataset(\n        tf.io.gfile.glob(tfrecord_pattern)\n    )\n    dataset = dataset.shuffle(buffer_size=100) if training else dataset \n    dataset = dataset.map(\n        parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE\n    )\n\n    if training:\n        dataset = dataset.map(\n            lambda x, y: prepare_training_inputs(x, y, k=slices_radius),\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n        dataset = dataset.map(\n            augment_data, num_parallel_calls=tf.data.AUTOTUNE\n        )\n    else:\n        dataset = dataset.flat_map(\n            lambda x, y: prepare_validation_inputs(x, y, k=slices_radius)\n        )\n    \n    dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n    return dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:03.824826Z","iopub.execute_input":"2025-11-20T11:09:03.825760Z","iopub.status.idle":"2025-11-20T11:09:03.830871Z","shell.execute_reply.started":"2025-11-20T11:09:03.825739Z","shell.execute_reply":"2025-11-20T11:09:03.830006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_pattern = \"/kaggle/input/vesuvius-tfrecords/training_shard_[0-7].tfrec\"\nval_pattern   = \"/kaggle/input/vesuvius-tfrecords/training_shard_8.tfrec\"\n\ntrain_ds = load_tfrecord_dataset(\n    train_pattern, batch_size=batch_size, training=True\n)\nval_ds = load_tfrecord_dataset(\n    val_pattern, batch_size=batch_size, training=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:04.481658Z","iopub.execute_input":"2025-11-20T11:09:04.481948Z","iopub.status.idle":"2025-11-20T11:09:07.059453Z","shell.execute_reply.started":"2025-11-20T11:09:04.481927Z","shell.execute_reply":"2025-11-20T11:09:07.058558Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**viz**","metadata":{}},{"cell_type":"code","source":"x, y = next(iter(train_ds))\nx.shape, y.shape, np.unique(y.numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T04:47:22.224013Z","iopub.execute_input":"2025-11-20T04:47:22.224292Z","iopub.status.idle":"2025-11-20T04:47:32.349705Z","shell.execute_reply.started":"2025-11-20T04:47:22.224272Z","shell.execute_reply":"2025-11-20T04:47:32.349062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_ds))\nx.shape, y.shape, np.unique(y.numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T06:49:32.526513Z","iopub.execute_input":"2025-11-20T06:49:32.527329Z","iopub.status.idle":"2025-11-20T06:49:34.744754Z","shell.execute_reply.started":"2025-11-20T06:49:32.527299Z","shell.execute_reply":"2025-11-20T06:49:34.743937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_sample(x, y, sample_idx=0, max_slices=16):\n    img = x[sample_idx]\n    mask = np.squeeze(y[sample_idx])\n    H, W, C = img.shape\n\n    # select slices/channels to display\n    step = max(1, C // max_slices)\n    channels = range(0, C, step)\n\n    n_slices = len(channels)\n    fig, axes = plt.subplots(1, n_slices, figsize=(3*n_slices, 3))\n\n    for i, c in enumerate(channels):\n        axes[i].imshow(img[..., c], cmap='gray')\n        axes[i].imshow(mask, cmap='jet', alpha=0.3)\n        axes[i].set_title(f\"Ch {c}\")\n        axes[i].axis('off')\n\n    plt.suptitle(f\"2.5D sample (k={C//2})\")\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:08.243836Z","iopub.execute_input":"2025-11-20T11:09:08.244431Z","iopub.status.idle":"2025-11-20T11:09:08.250135Z","shell.execute_reply.started":"2025-11-20T11:09:08.244406Z","shell.execute_reply":"2025-11-20T11:09:08.249213Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_sample(\n    x, y, sample_idx=0, max_slices=4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T06:49:40.489366Z","iopub.execute_input":"2025-11-20T06:49:40.489753Z","iopub.status.idle":"2025-11-20T06:49:41.185889Z","shell.execute_reply.started":"2025-11-20T06:49:40.489727Z","shell.execute_reply":"2025-11-20T06:49:41.184797Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"import medicai\nfrom medicai.models import TransUNet, UNet\nfrom medicai.losses import BinaryDiceCELoss\nfrom medicai.metrics import BinaryDiceMetric","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:10.076666Z","iopub.execute_input":"2025-11-20T11:09:10.076956Z","iopub.status.idle":"2025-11-20T11:09:10.109076Z","shell.execute_reply.started":"2025-11-20T11:09:10.076935Z","shell.execute_reply":"2025-11-20T11:09:10.108530Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# medicai.models.list_models()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:10.843038Z","iopub.execute_input":"2025-11-20T11:09:10.843540Z","iopub.status.idle":"2025-11-20T11:09:10.847247Z","shell.execute_reply.started":"2025-11-20T11:09:10.843514Z","shell.execute_reply":"2025-11-20T11:09:10.846345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = 1\nmodel = UNet(\n    input_shape=(\n        input_shape, input_shape, slices_radius*2 + 1\n    ),\n    encoder_name='resnet18',\n    encoder_depth=4,\n    classifier_activation='sigmoid',\n    num_classes=num_classes,\n)\nmodel.count_params() / 1e6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:11.777926Z","iopub.execute_input":"2025-11-20T11:09:11.778247Z","iopub.status.idle":"2025-11-20T11:09:16.954961Z","shell.execute_reply.started":"2025-11-20T11:09:11.778222Z","shell.execute_reply":"2025-11-20T11:09:16.954145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define optomizer, loss, metrics\noptim = keras.optimizers.AdamW(\n    learning_rate=1e-4,\n    weight_decay=1e-5,\n)\n\nloss_fn = BinaryDiceCELoss(\n    from_logits=False, \n    num_classes=num_classes\n)\n\nmetrics = [\n    BinaryDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        name='dice'\n    ),\n]\n\nmodel.compile(\n    optimizer=optim,\n    loss=loss_fn,\n    metrics=metrics\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:21.022602Z","iopub.execute_input":"2025-11-20T11:09:21.023098Z","iopub.status.idle":"2025-11-20T11:09:21.051546Z","shell.execute_reply.started":"2025-11-20T11:09:21.023073Z","shell.execute_reply":"2025-11-20T11:09:21.050913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_checkpoint_callback = keras.callbacks.ModelCheckpoint(\n    \"model.weights.h5\",\n    monitor=\"val_dice\",\n    verbose=0,\n    save_best_only=True,\n    save_weights_only=True,\n    mode=\"max\",\n    save_freq=\"epoch\",\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:09:25.448017Z","iopub.execute_input":"2025-11-20T11:09:25.448599Z","iopub.status.idle":"2025-11-20T11:09:25.452797Z","shell.execute_reply.started":"2025-11-20T11:09:25.448571Z","shell.execute_reply":"2025-11-20T11:09:25.451881Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_all_original_labels(tfrecord_pattern):\n    raw_dataset = tf.data.TFRecordDataset(tf.io.gfile.glob(tfrecord_pattern))\n    # Load all original 3D volumes\n    volumes_info = []\n    for _, label in raw_dataset.map(parse_tfrecord_fn).as_numpy_iterator():\n        # Ensure mask is binary (0/1)\n        label = (label == 1).astype(np.float32)[..., None]\n        volumes_info.append({\n            'label': label, \n            'depth': label.shape[0],\n            'height': label.shape[1],\n            'width': label.shape[2],\n            \"channel\": label.shape[3]\n        })\n    return volumes_info","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:22:29.277242Z","iopub.execute_input":"2025-11-20T11:22:29.277603Z","iopub.status.idle":"2025-11-20T11:22:29.282881Z","shell.execute_reply.started":"2025-11-20T11:22:29.277581Z","shell.execute_reply":"2025-11-20T11:22:29.282247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_original_volume(\n    predicted_slices,\n    original_volume_info,\n    slices_radius,\n    metrics\n):\n    k = slices_radius\n    slice_counts = [info['depth'] - 2 * k for info in original_volume_info]\n\n    target_maps = []\n    prediction_maps = []\n    volume_metric_scores = []\n\n    start_idx = 0\n\n    for i, info in enumerate(original_volume_info):\n        slice_count = slice_counts[i]\n        end_idx = start_idx + slice_count\n        volume_preds = predicted_slices[start_idx:end_idx]\n\n        orig_depth, orig_height, orig_width, channel = (\n            info['depth'], info['height'], info['width'], info['channel']\n        )\n\n        predicted_volume = np.zeros(\n            (orig_depth, orig_height, orig_width, channel), dtype=np.float32\n        )\n        predicted_volume[k : orig_depth - k] = volume_preds\n\n        original_label = info['label'][None, ...]\n        predicted_volume = predicted_volume[None, ...]\n\n        metrics.reset_state()\n        metrics.update_state(\n            original_label,\n            predicted_volume\n        )\n        volume_metric_scores.append(\n            metrics.result()\n        )\n        prediction_maps.append(\n            ops.convert_to_numpy(predicted_volume)\n        )\n        target_maps.append(\n            ops.convert_to_numpy(original_label)\n        )\n        start_idx = end_idx \n\n    avg_volume_metric = np.mean(volume_metric_scores)\n\n    return {\n        'avg_volume_metric': avg_volume_metric,\n        'per_volume_scores': volume_metric_scores,\n        \"prediction_maps\": prediction_maps,\n        \"target_maps\": target_maps\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:22:32.218824Z","iopub.execute_input":"2025-11-20T11:22:32.219406Z","iopub.status.idle":"2025-11-20T11:22:32.226270Z","shell.execute_reply.started":"2025-11-20T11:22:32.219382Z","shell.execute_reply":"2025-11-20T11:22:32.225455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_volumes_info = load_all_original_labels(val_pattern)\nval_metrics = BinaryDiceMetric(\n    from_logits=False, \n    num_classes=num_classes, \n    name='dice'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:22:34.252986Z","iopub.execute_input":"2025-11-20T11:22:34.253614Z","iopub.status.idle":"2025-11-20T11:22:35.099958Z","shell.execute_reply.started":"2025-11-20T11:22:34.253589Z","shell.execute_reply":"2025-11-20T11:22:35.099117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VolumeMetricCallback(keras.callbacks.Callback):\n    def __init__(\n        self, \n        val_dataloader, \n        slices_radius, \n        val_volumes_info, \n        val_metrics,\n        model_save_path,\n    ):\n        super().__init__()\n        self.val_dataloader = val_dataloader\n        self.slices_radius = slices_radius\n        self.val_volumes_info = val_volumes_info\n        self.val_metrics = val_metrics\n        self.model_save_path = model_save_path\n        self.best_val_score = 0.0\n\n    def on_epoch_end(self, epoch, logs=None):\n        predicted_slices = self.model.predict(\n            self.val_dataloader\n        )\n        predicted_slices = (predicted_slices > 0.5).astype(np.float32)\n        results = evaluate_original_volume(\n            predicted_slices=predicted_slices,\n            original_volume_info=self.val_volumes_info,\n            slices_radius=self.slices_radius,\n            metrics=self.val_metrics\n        )\n\n        avg_dice = results['avg_volume_metric']\n        print(\n            f\"\\nAverage Volume Metric ({self.val_metrics.name}) after epoch {epoch + 1}: {avg_dice:.4f}\"\n        )\n\n        if avg_dice > best_val_dice:\n            best_val_dice = avg_dice\n            model.save_weights('model.weights.h5')\n            print('Validation score improved. Model saved.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:22:35.101049Z","iopub.execute_input":"2025-11-20T11:22:35.101379Z","iopub.status.idle":"2025-11-20T11:22:35.107125Z","shell.execute_reply.started":"2025-11-20T11:22:35.101360Z","shell.execute_reply":"2025-11-20T11:22:35.106324Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_volumes_info = load_all_original_labels(val_pattern)\nvolume_metric_callback = VolumeMetricCallback(\n    val_dataloader=val_ds,\n    slices_radius=slices_radius,\n    val_volumes_info=val_volumes_info,\n    val_metrics=val_metrics\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:22:36.505694Z","iopub.execute_input":"2025-11-20T11:22:36.506320Z","iopub.status.idle":"2025-11-20T11:22:37.338461Z","shell.execute_reply.started":"2025-11-20T11:22:36.506291Z","shell.execute_reply":"2025-11-20T11:22:37.337778Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.fit(\n    train_ds,\n    epochs=20,\n    callbacks=[\n        volume_metric_callback\n    ],\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:22:37.508781Z","iopub.execute_input":"2025-11-20T11:22:37.509378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_weights(\n    \"model.weights.h5\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-17T00:25:02.335232Z","iopub.execute_input":"2025-11-17T00:25:02.335935Z","iopub.status.idle":"2025-11-17T00:25:06.755960Z","shell.execute_reply.started":"2025-11-17T00:25:02.335907Z","shell.execute_reply":"2025-11-17T00:25:06.755322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_ds))\nx.shape, y.shape, np.unique(y.numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T04:30:11.465005Z","iopub.execute_input":"2025-11-20T04:30:11.465299Z","iopub.status.idle":"2025-11-20T04:30:16.332243Z","shell.execute_reply.started":"2025-11-20T04:30:11.465277Z","shell.execute_reply":"2025-11-20T04:30:16.331537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_pred = model.predict(x)\ny_pred.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-17T00:25:31.156819Z","iopub.execute_input":"2025-11-17T00:25:31.157635Z","iopub.status.idle":"2025-11-17T00:25:38.911705Z","shell.execute_reply.started":"2025-11-17T00:25:31.157599Z","shell.execute_reply":"2025-11-17T00:25:38.910898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"segment = (y_pred > 0.35).astype(np.uint8)\nsegment.shape, np.unique(segment)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-17T00:40:17.501903Z","iopub.execute_input":"2025-11-17T00:40:17.502619Z","iopub.status.idle":"2025-11-17T00:40:17.517395Z","shell.execute_reply.started":"2025-11-17T00:40:17.502594Z","shell.execute_reply":"2025-11-17T00:40:17.516677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Image+GT**","metadata":{}},{"cell_type":"code","source":"plot_sample(\n    x, y, sample_idx=0, max_slices=4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-17T00:40:04.459908Z","iopub.execute_input":"2025-11-17T00:40:04.460133Z","iopub.status.idle":"2025-11-17T00:40:04.964733Z","shell.execute_reply.started":"2025-11-17T00:40:04.460118Z","shell.execute_reply":"2025-11-17T00:40:04.963880Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Image+Prediction**","metadata":{}},{"cell_type":"code","source":"plot_sample(\n    x, segment, sample_idx=0, max_slices=4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-17T00:40:19.421100Z","iopub.execute_input":"2025-11-17T00:40:19.421362Z","iopub.status.idle":"2025-11-17T00:40:19.939600Z","shell.execute_reply.started":"2025-11-17T00:40:19.421345Z","shell.execute_reply":"2025-11-17T00:40:19.938830Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}