{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.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":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14266465,"sourceType":"datasetVersion","datasetId":8751895},{"sourceId":290917305,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# About\n\n- Load `tfrecord` and build `tf.data` dataloader.\n    - **Note**: This can be replaced with `torch.utils.data.DataLoader`. \n- Use `medicai` for volume transformation and **3D** model, i.e. [`SegFormer3D`](https://arxiv.org/abs/2404.10156) - written in **Keras 3**.\n    - Will use `torch` backend.\n- Train the model with pure **PyTorch** custom training pipeline. [Learn](https://keras.io/guides/writing_a_custom_training_loop_in_torch/) more about it.","metadata":{}},{"cell_type":"code","source":"from IPython.display import clear_output\n\nvar=\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\"\n!pip install \\\n    \"$var\"/keras_nightly-3.12.0.dev2025100703-py3-none-any.whl \\\n    --no-index \\\n    --find-links \"$var\"\n\n!pip install git+https://github.com/innat/medic-ai.git -q\nclear_output()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:04:38.384925Z","iopub.execute_input":"2026-02-02T16:04:38.385680Z","iopub.status.idle":"2026-02-02T16:04:55.752643Z","shell.execute_reply.started":"2026-02-02T16:04:38.385650Z","shell.execute_reply":"2026-02-02T16:04:55.751790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\n\nimport os, warnings\nos.environ[\"KERAS_BACKEND\"] = \"torch\"\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:04:55.754535Z","iopub.execute_input":"2026-02-02T16:04:55.754871Z","iopub.status.idle":"2026-02-02T16:04:56.041298Z","shell.execute_reply.started":"2026-02-02T16:04:55.754842Z","shell.execute_reply":"2026-02-02T16:04:56.040543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras\nfrom keras import ops\nfrom keras.optimizers import SGD, AdamW, Muon\nfrom keras.optimizers.schedules import CosineDecay, PolynomialDecay\n\nimport torch\nimport tensorflow as tf\n\n# keras.mixed_precision.set_global_policy(\"mixed_float16\")\nkeras.version(), torch.__version__, keras.config.backend()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:04:56.042225Z","iopub.execute_input":"2026-02-02T16:04:56.042726Z","iopub.status.idle":"2026-02-02T16:05:15.595276Z","shell.execute_reply.started":"2026-02-02T16:04:56.042699Z","shell.execute_reply":"2026-02-02T16:05:15.594690Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# mainly for 3D or 2D models, transformation, loss, metrics etc\nimport medicai\nfrom medicai.transforms import (\n    Compose,\n    NormalizeIntensity,\n    ScaleIntensityRange,\n    Resize,\n    RandShiftIntensity,\n    RandRotate90,\n    RandRotate,\n    RandFlip,\n    RandCutOut,\n    RandSpatialCrop\n)\nfrom medicai.layers import ResizingND\nfrom medicai.models import (\n    UNet, SegFormer, TransUNet, SwinUNETR, UPerNet, ConvNeXtV2Tiny, UNETRPlusPlus\n)\nfrom medicai.losses import (\n    SparseDiceCELoss, SparseTverskyLoss, SparseCenterlineDiceLoss\n)\nfrom medicai.metrics import SparseDiceMetric\nfrom medicai.callbacks import SlidingWindowInferenceCallback\nfrom medicai.utils import SlidingWindowInference\nfrom medicai.utils import soft_skeletonize","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:15.596103Z","iopub.execute_input":"2026-02-02T16:05:15.596602Z","iopub.status.idle":"2026-02-02T16:05:15.641247Z","shell.execute_reply.started":"2026-02-02T16:05:15.596578Z","shell.execute_reply":"2026-02-02T16:05:15.640631Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loader","metadata":{}},{"cell_type":"code","source":"input_shape=(128, 128, 128)\nbatch_size=1\nnum_classes=3\n\n# Each tfrecord contains 6 samples, total 786 samples.\nnum_samples = 780\nepochs = 50","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:15.643550Z","iopub.execute_input":"2026-02-02T16:05:15.643783Z","iopub.status.idle":"2026-02-02T16:05:15.671047Z","shell.execute_reply.started":"2026-02-02T16:05:15.643762Z","shell.execute_reply":"2026-02-02T16:05:15.670383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        ## Geometric transformation\n        RandSpatialCrop(\n            keys=[\"image\", \"label\"],\n            roi_size=input_shape,\n            random_center=True,\n            random_size=False,\n            invalid_label=2,         \n            min_valid_ratio=0.5,     \n            max_attempts=10\n        ),\n        RandFlip(keys=[\"image\", \"label\"], spatial_axis=[0], prob=0.5),\n        RandFlip(keys=[\"image\", \"label\"], spatial_axis=[1], prob=0.5),\n        RandFlip(keys=[\"image\", \"label\"], spatial_axis=[2], prob=0.5),\n        RandRotate90(\n            keys=[\"image\", \"label\"], \n            prob=0.4, \n            max_k=3, \n            spatial_axes=(0, 1)\n        ),\n\n        ## Z-score norm\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n\n        ## Intensiry transformation\n        RandShiftIntensity(\n            keys=[\"image\"], offsets=0.10, prob=0.5\n        ),\n        \n        ## Spatial transformation \n        RandCutOut(\n            keys=[\"image\", \"label\"],\n            invalid_label=2, \n            mask_size=[\n                input_shape[1]//4,\n                input_shape[2]//4\n            ],\n            fill_mode=\"constant\", # \"constant\", \"gaussian\"\n            cutout_mode='volume', # \"slice\", \"volume\"\n            prob=0.2,\n            num_cuts=2,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n\n\ndef val_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:15.671928Z","iopub.execute_input":"2026-02-02T16:05:15.672218Z","iopub.status.idle":"2026-02-02T16:05:15.683362Z","shell.execute_reply.started":"2026-02-02T16:05:15.672196Z","shell.execute_reply":"2026-02-02T16:05:15.682690Z"}},"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,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:15.684390Z","iopub.execute_input":"2026-02-02T16:05:15.684662Z","iopub.status.idle":"2026-02-02T16:05:15.696703Z","shell.execute_reply.started":"2026-02-02T16:05:15.684632Z","shell.execute_reply":"2026-02-02T16:05:15.695992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_inputs(image, label):\n    # Add channel dimension\n    image = image[..., None] # (D, H, W, 1)\n    label = label[..., None] # (D, H, W, 1)\n\n    # Convert to float32\n    image = tf.cast(image, tf.float32)\n    label = tf.cast(label, tf.float32)\n    return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:15.697629Z","iopub.execute_input":"2026-02-02T16:05:15.697982Z","iopub.status.idle":"2026-02-02T16:05:15.708898Z","shell.execute_reply.started":"2026-02-02T16:05:15.697953Z","shell.execute_reply":"2026-02-02T16:05:15.708289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def tfrecord_loader(tfrecord_pattern, batch_size=1, shuffle=True):\n    dataset = tf.data.TFRecordDataset(\n        tf.io.gfile.glob(tfrecord_pattern)\n    )\n    dataset = dataset.shuffle(buffer_size=100) if shuffle else dataset \n    dataset = dataset.map(\n        parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE\n    )\n    dataset = dataset.map(\n        prepare_inputs,\n        num_parallel_calls=tf.data.AUTOTUNE\n    )\n    if shuffle:\n        dataset = dataset.map(\n            train_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    else:\n        dataset = dataset.map(\n            val_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    dataset = dataset.batch(batch_size, drop_remainder=shuffle)\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    return dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:15.709658Z","iopub.execute_input":"2026-02-02T16:05:15.709936Z","iopub.status.idle":"2026-02-02T16:05:15.720254Z","shell.execute_reply.started":"2026-02-02T16:05:15.709917Z","shell.execute_reply":"2026-02-02T16:05:15.719636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_tfrec = sorted(\n    glob.glob(\"/kaggle/input/vesuvius-tfrecords/*.tfrec\"),\n    key=lambda x: int(x.split(\"_\")[-1].replace(\".tfrec\", \"\"))\n)\n\nval_idx = -1\nval_patterns = [all_tfrec[val_idx]]\ntrain_patterns = [\n    f for i, f in enumerate(all_tfrec) if i != len(all_tfrec) + val_idx\n]\n\ntrain_ds = tfrecord_loader(\n    train_patterns, batch_size=batch_size, shuffle=True\n)\nval_ds = tfrecord_loader(\n    val_patterns, batch_size=1, shuffle=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:15.721038Z","iopub.execute_input":"2026-02-02T16:05:15.721308Z","iopub.status.idle":"2026-02-02T16:05:19.382518Z","shell.execute_reply.started":"2026-02-02T16:05:15.721282Z","shell.execute_reply":"2026-02-02T16:05:19.381868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_ds))\nx.shape, y.shape ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:19.383480Z","iopub.execute_input":"2026-02-02T16:05:19.383803Z","iopub.status.idle":"2026-02-02T16:05:24.115433Z","shell.execute_reply.started":"2026-02-02T16:05:19.383778Z","shell.execute_reply":"2026-02-02T16:05:24.114681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_sample(x, y, sample_idx=0, max_slices=16):\n    img = np.squeeze(x[sample_idx])  # (D, H, W)\n    mask = np.squeeze(y[sample_idx])  # (D, H, W)\n    D = img.shape[0]\n\n    # Decide which slices to plot\n    step = max(1, D // max_slices)\n    slices = range(0, D, step)\n\n    n_slices = len(slices)\n    fig, axes = plt.subplots(2, n_slices, figsize=(3*n_slices, 6))\n\n    for i, s in enumerate(slices):\n        axes[0, i].imshow(img[s], cmap='gray')\n        axes[0, i].set_title(f\"Slice {s}\")\n        axes[0, i].axis('off')\n\n        axes[1, i].imshow(mask[s], cmap='gray')\n        axes[1, i].set_title(f\"Mask {s}\")\n        axes[1, i].axis('off')\n\n    plt.suptitle(f\"Sample {sample_idx}\")\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:24.116409Z","iopub.execute_input":"2026-02-02T16:05:24.116737Z","iopub.status.idle":"2026-02-02T16:05:24.122978Z","shell.execute_reply.started":"2026-02-02T16:05:24.116714Z","shell.execute_reply":"2026-02-02T16:05:24.122154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_planes(image, mask, alpha=0.4):\n    # Central slices\n    d, h, w = image.shape\n    axial_img    = image[d // 2]\n    coronal_img  = image[:, h // 2, :]\n    sagittal_img = image[:, :, w // 2]\n\n    axial_msk    = mask[d // 2]\n    coronal_msk  = mask[:, h // 2, :]\n    sagittal_msk = mask[:, :, w // 2]\n\n    slices_img = [axial_img, coronal_img, sagittal_img]\n    slices_msk = [axial_msk, coronal_msk, sagittal_msk]\n    \n    titles = [\"Axial (XY plane)\", \"Coronal (XZ plane)\", \"Sagittal (YZ plane)\"]\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n\n    for i, ax in enumerate(axes):\n        ax.imshow(slices_img[i], cmap=\"gray\")\n\n        # overlay jet only where mask > 0\n        m = slices_msk[i]\n        if m.max() > 0:\n            ax.imshow(m, cmap=\"jet\", alpha=alpha)\n\n        ax.set_title(titles[i])\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-02-02T16:05:24.123838Z","iopub.execute_input":"2026-02-02T16:05:24.124093Z","iopub.status.idle":"2026-02-02T16:05:24.135339Z","shell.execute_reply.started":"2026-02-02T16:05:24.124073Z","shell.execute_reply":"2026-02-02T16:05:24.134629Z"}},"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":"2026-02-02T16:05:24.138613Z","iopub.execute_input":"2026-02-02T16:05:24.139314Z","iopub.status.idle":"2026-02-02T16:05:25.029059Z","shell.execute_reply.started":"2026-02-02T16:05:24.139280Z","shell.execute_reply":"2026-02-02T16:05:25.028319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_planes(\n    np.squeeze(x[0]), # picking one sample\n    np.squeeze(y[0])  # picking one sample\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:25.030311Z","iopub.execute_input":"2026-02-02T16:05:25.030639Z","iopub.status.idle":"2026-02-02T16:05:25.978307Z","shell.execute_reply.started":"2026-02-02T16:05:25.030596Z","shell.execute_reply":"2026-02-02T16:05:25.977164Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"# medicai.models.list_models()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:25.979409Z","iopub.execute_input":"2026-02-02T16:05:25.979736Z","iopub.status.idle":"2026-02-02T16:05:25.983324Z","shell.execute_reply.started":"2026-02-02T16:05:25.979712Z","shell.execute_reply":"2026-02-02T16:05:25.982540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = SegFormer(\n    input_shape=input_shape + (1,),\n    encoder_name='mit_b0',\n    classifier_activation='softmax',\n    num_classes=num_classes,\n    dropout=0.2,\n)\nmodel.count_params() / 1e6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:25.984412Z","iopub.execute_input":"2026-02-02T16:05:25.984978Z","iopub.status.idle":"2026-02-02T16:05:27.319844Z","shell.execute_reply.started":"2026-02-02T16:05:25.984956Z","shell.execute_reply":"2026-02-02T16:05:27.319236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ALERT: This attributes only available in medicai (not in core keras)\ntry:\n    print(model.instance_describe())\nexcept AttributeError:\n    pass","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.320720Z","iopub.execute_input":"2026-02-02T16:05:27.321039Z","iopub.status.idle":"2026-02-02T16:05:27.325155Z","shell.execute_reply.started":"2026-02-02T16:05:27.321018Z","shell.execute_reply":"2026-02-02T16:05:27.324616Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Pipeline","metadata":{}},{"cell_type":"code","source":"steps_per_epoch = num_samples // batch_size\ntotal_steps = steps_per_epoch * epochs\nwarmup_steps = int(total_steps * 0.05)\ndecay_steps = max(1, total_steps - warmup_steps)\nlr_schedule = CosineDecay(\n    initial_learning_rate=1e-6,\n    decay_steps=decay_steps,\n    warmup_target=min(3e-4, 1e-4 * (batch_size / 2)),\n    warmup_steps=warmup_steps,\n    alpha=0.1,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.326071Z","iopub.execute_input":"2026-02-02T16:05:27.326294Z","iopub.status.idle":"2026-02-02T16:05:27.360105Z","shell.execute_reply.started":"2026-02-02T16:05:27.326275Z","shell.execute_reply":"2026-02-02T16:05:27.359494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define optomizer, loss, metrics\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n)\n\ndice_ce_loss_fn = SparseDiceCELoss(\n    from_logits=False, \n    num_classes=num_classes,\n    ignore_class_ids=2,\n)\ncldice_loss_fn = SparseCenterlineDiceLoss(\n    from_logits=False, \n    num_classes=num_classes,\n    target_class_ids=1,\n    ignore_class_ids=2,\n    iters=1, # ideal to set 20-50 - computationally expensive\n)\ncombined_loss_fn = lambda y_true, y_pred: (\n    dice_ce_loss_fn(y_true, y_pred) + cldice_loss_fn(y_true, y_pred)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.360896Z","iopub.execute_input":"2026-02-02T16:05:27.361144Z","iopub.status.idle":"2026-02-02T16:05:27.375883Z","shell.execute_reply.started":"2026-02-02T16:05:27.361124Z","shell.execute_reply":"2026-02-02T16:05:27.375286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define sliding-window-inferencer for validation\nswi = SlidingWindowInference(\n    model,\n    num_classes=num_classes,\n    roi_size=input_shape,\n    sw_batch_size=1,\n    overlap=0.5,\n    mode='gaussian',\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.376677Z","iopub.execute_input":"2026-02-02T16:05:27.376935Z","iopub.status.idle":"2026-02-02T16:05:27.386273Z","shell.execute_reply.started":"2026-02-02T16:05:27.376906Z","shell.execute_reply":"2026-02-02T16:05:27.385538Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, dataloader, metrics):\n    loop = tqdm(dataloader, desc=\"Training\", leave=False)\n    \n    for imgs, labels in loop:\n        # forward pass\n        outputs = model(imgs)\n        loss = combined_loss_fn(labels, outputs)\n\n        # backward pass\n        model.zero_grad()\n        trainable_weights = [v for v in model.trainable_weights]\n\n        # call torch.Tensor.backward() on the loss to compute gradients\n        loss.backward()\n        gradients = [v.value.grad for v in trainable_weights]\n\n        # update weights\n        with torch.no_grad():\n            optim.apply(gradients, trainable_weights)\n\n        # update training metric\n        metrics.update_state(\n            ops.convert_to_tensor(labels), \n            ops.convert_to_tensor(outputs)\n        )\n        \n        # Update tqdm\n        loss_score = ops.convert_to_numpy(loss)\n        metrics_score = ops.convert_to_numpy(metrics.result())\n        loop.set_postfix(\n            loss=loss_score,\n            dice=metrics_score,\n        )\n\n    return loss, metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.387774Z","iopub.execute_input":"2026-02-02T16:05:27.388028Z","iopub.status.idle":"2026-02-02T16:05:27.396857Z","shell.execute_reply.started":"2026-02-02T16:05:27.388010Z","shell.execute_reply":"2026-02-02T16:05:27.396166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, dataloader, metrics):\n    for sample in dataloader:\n        x, y = sample\n        output = swi(x)\n        y = ops.convert_to_tensor(y)\n        output = ops.convert_to_tensor(output)\n        metrics.update_state(y, output)\n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.398035Z","iopub.execute_input":"2026-02-02T16:05:27.398305Z","iopub.status.idle":"2026-02-02T16:05:27.408194Z","shell.execute_reply.started":"2026-02-02T16:05:27.398286Z","shell.execute_reply":"2026-02-02T16:05:27.407570Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_training(train_loader, val_loader, model, epochs=20):\n    # metrics for train\n    train_metrics = SparseDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        ignore_class_ids=2,\n        name='dice'\n    )\n\n    # metrics for validation\n    val_metrics = SparseDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        ignore_class_ids=2,\n        name='val_dice'\n    )\n    \n    # Initialize best validation dice score\n    best_val_dice = 0.0\n\n    for epoch in range(epochs):\n        print(f'Epoch {epoch+1}/{epochs}')\n        \n        # Training\n        loss, train_metrics = train_one_epoch(\n            model, train_loader, train_metrics,\n        )\n        # display training logs at the end of epoch\n        train_metrics_score = ops.convert_to_numpy(train_metrics.result())\n        loss_score = ops.convert_to_numpy(loss)\n        \n        # reset training metrics at the end of each epoch\n        train_metrics.reset_state()\n\n        # Validation [at every 5 epoch]\n        if (epoch + 1) % 2 == 0:\n            val_metrics = validate(model, val_loader, val_metrics)\n            val_metrics_score = ops.convert_to_numpy(\n                val_metrics.result()\n            )\n            val_metrics.reset_state()\n            print(\n                f'Training - Loss: {loss_score:.4f}, Dice: {train_metrics_score:.4f}'\n                f'\\nValidation - Dice: {val_metrics_score:.4f}\\n'\n            )\n\n            # Save best model weights\n            if val_metrics_score > best_val_dice:\n                best_val_dice = val_metrics_score\n                model.save_weights('model.weights.h5')\n                # torch.save(model.state_dict(), 'model.pth') # OK too.\n                print(\n                    f'Dice score improved: {best_val_dice}. Model saved.'\n                )\n        else:\n            print(\n                f'Training - Loss: {loss_score:.4f}, Dice: {train_metrics_score:.4f}\\n'\n            )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.409144Z","iopub.execute_input":"2026-02-02T16:05:27.409406Z","iopub.status.idle":"2026-02-02T16:05:27.421844Z","shell.execute_reply.started":"2026-02-02T16:05:27.409382Z","shell.execute_reply":"2026-02-02T16:05:27.421193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_training(\n    train_ds, val_ds, model, epochs=epochs\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T16:05:27.422484Z","iopub.execute_input":"2026-02-02T16:05:27.422701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_ds))\nx.shape, y.shape","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_pred = swi(x)\ny_pred.shape","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"segment = y_pred.argmax(-1).astype(np.uint8)\nsegment.shape, np.unique(segment)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_sample(\n    x, segment, sample_idx=0, max_slices=4\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}