{"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":"tpuV5e8","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14750380,"sourceType":"datasetVersion","datasetId":9427246},{"sourceId":14771190,"sourceType":"datasetVersion","datasetId":9441760}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# About\n\n- 200 images generated by inference for each epoch based on the Innat program\n- [Innat - [Train] Vesuvius Surface 3D Detection on TPU](https://www.kaggle.com/code/ipythonx/train-vesuvius-surface-3d-detection-on-tpu)\n- The training was carried out on a single RTX-3090Ti graphics accelerator for 30 hours.\n","metadata":{}},{"cell_type":"code","source":"import os\nimport subprocess\n\nis_kaggle = os.environ.get(\"KAGGLE_KERNEL_RUN_TYPE\", \"\") != \"\"\nis_rerun = os.environ.get('KAGGLE_IS_COMPETITION_RERUN', \"\") != \"\"\nis_interactive = os.environ.get('KAGGLE_KERNEL_RUN_TYPE', '') == 'Interactive'  # Interactive Mode\nis_batch = os.environ.get('KAGGLE_KERNEL_RUN_TYPE','') == 'Batch'\nis_local = not is_kaggle\nis_background = is_rerun or is_batch\nis_dialog = not is_background\n\n! export PATH=\"${HOME}/.local/bin:${PATH}\" && uv pip uninstall --system jax\n! export PATH=\"${HOME}/.local/bin:${PATH}\" && uv pip install --system -U \\\n    keras tensorflow tensorflow-tpu tifffile imagecodecs scikit-image \\\n    albumentations pybind11[global] connected-components-3d surface-distance \\\n    --find-links https://storage.googleapis.com/libtpu-tf-releases/index.html\n    \n! mkdir whls\n! cp -r /kaggle/input/medic-ai /kaggle/working\n! cd /kaggle/working/medic-ai && pip install build -q\n! cd /kaggle/working/medic-ai && python -m build\n! ls /kaggle/working/medic-ai/dist/*.whl\n! cp medic-ai/dist/*.whl whls/\n! rm -r /kaggle/working/medic-ai\nvar = \"/kaggle/working/whls\"\ninstall_libraries = [\n        'pip', 'install', '-U', '--no-index',\n        '--find-links', f'{var}',\n        'medicai',\n    ]\nsubprocess.run(install_libraries, check=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T12:58:43.522957Z","iopub.execute_input":"2026-02-08T12:58:43.523110Z","iopub.status.idle":"2026-02-08T12:59:22.462087Z","shell.execute_reply.started":"2026-02-08T12:58:43.523093Z","shell.execute_reply":"2026-02-08T12:59:22.461170Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf; print('TensorFlow version' + tf.__version__)\n# Detect and initialize TPU\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver(tpu=\"local\")  \n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.TPUStrategy(tpu)\n    print(f\"Running on TPU: {tpu.master()}\")\nexcept Exception as e:\n    print(f\"Could not initialize TPU: {e}\")\n    strategy = tf.distribute.get_strategy()  # Fallback to default CPU/GPU\n    assert 1 == -1\n\n# Print number of available TPU cores\nprint(\"Number of devices:\", strategy.num_replicas_in_sync, \"🚀\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T12:59:22.462706Z","iopub.execute_input":"2026-02-08T12:59:22.462909Z","execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, warnings\nwarnings.filterwarnings('ignore')\n#os.environ[\"KERAS_BACKEND\"] = \"jax\"\nimport glob\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport zipfile\nimport tifffile\nimport matplotlib.pyplot as plt\n# mainly for training API\nimport keras\nfrom keras import ops\nfrom keras.optimizers import SGD, AdamW, Muon\nfrom keras.optimizers.schedules import CosineDecay, PolynomialDecay\n# 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 (UNet, SegFormer, TransUNet, SwinUNETR, UPerNet, ConvNeXtV2Tiny, UNETRPlusPlus)\nfrom medicai.losses import (SparseDiceCELoss, SparseTverskyLoss, SparseCenterlineDiceLoss)\nfrom medicai.metrics import SparseDiceMetric\nfrom medicai.callbacks import SlidingWindowInferenceCallback\nfrom medicai.utils import SlidingWindowInference\nfrom medicai.utils import soft_skeletonize\nprint(keras.version(), keras.config.backend(), medicai.version())\n\n# due to distributed training only\nkeras.config.disable_flash_attention()\n# reproducibility\nkeras.utils.set_random_seed(101)\n# distributed config\ndevices = keras.distribution.list_devices()\ntotal_device = len(devices)\nprint(f'detected devices: {devices} total device: {total_device}')\n#data_parallel = keras.distribution.DataParallel(devices=devices)\n#keras.distribution.set_distribution(data_parallel)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"input_shape=(160, 160, 160)\nbatch_size=1 * total_device\nnum_classes=3\n# Each tfrecord contains 6 samples, total 786 samples.\nnum_samples = 780\nepochs = 1\n\ndef 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\n\ndef 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    # Convert to float32\n    image = tf.cast(image, tf.float32)\n    label = tf.cast(label, tf.float32)\n    return image, label\n\ndef 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        RandRotate(\n            keys=[\"image\", \"label\"], \n            factor=0.2, \n            prob=0.7, \n            fill_mode=\"crop\",\n        ),\n\n        ## Intensiry transformation\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n        RandShiftIntensity(\n            keys=[\"image\"], offsets=0.10, prob=0.5\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.8,\n            num_cuts=5,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\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\"]\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"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\n\nall_tfrec = sorted(\n    glob.glob(\"/kaggle/input/vesuvius-tfrecords/*.tfrec\"),\n    key=lambda x: int(x.split(\"_\")[-1].replace(\".tfrec\", \"\"))\n)\nval_idx = -1\nval_patterns = [all_tfrec[val_idx]]\ntrain_patterns = [f for i, f in enumerate(all_tfrec) if i != len(all_tfrec) + val_idx]\ntrain_ds = tfrecord_loader(train_patterns, batch_size=batch_size, shuffle=True)\nval_ds = tfrecord_loader(val_patterns, batch_size=1, shuffle=False)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"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    # Decide which slices to plot\n    step = max(1, D // max_slices)\n    slices = range(0, D, step)\n    n_slices = len(slices)\n    fig, axes = plt.subplots(2, n_slices, figsize=(3*n_slices, 6))\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        axes[1, i].imshow(mask[s], cmap='gray')\n        axes[1, i].set_title(f\"Mask {s}\")\n        axes[1, i].axis('off')\n    plt.suptitle(f\"Sample {sample_idx}\")\n    plt.tight_layout()\n    plt.show()\nx, y = next(iter(val_ds))\nx.shape, y.shape\n#plot_sample(x, y, sample_idx=0, max_slices=4)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"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()\n\n#plot_planes(\n#    np.squeeze(x[0]), # picking one sample\n#    np.squeeze(y[0])  # picking one sample\n#)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"soft_skel = soft_skeletonize(ops.cast(y == 1, 'float32'), iters=10)\nprint(soft_skel.shape)\n#plot_sample(y, soft_skel, sample_idx=0, max_slices=4)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## check available models (classification + segmentation)\n# medicai.models.list_models()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Pre-build encoder\n# model = SegFormer(\n#     input_shape=input_shape + (1,),\n#     encoder_name='mit_b0',\n#     classifier_activation='softmax',\n#     num_classes=num_classes,\n# )\n\n# model = UPerNet(\n#     input_shape=input_shape + (1,),\n#     encoder_name=\"convnext_base\",\n#     classifier_activation='softmax',\n#     num_classes=num_classes,\n# )\n\n# model = UNETRPlusPlus(\n#     input_shape=input_shape + (1,),\n#     encoder_name='unetr_plusplus_encoder',\n#     classifier_activation='softmax',\n#     num_classes=num_classes,\n# )\n\nmodel = TransUNet(\n    encoder_name='seresnext50',\n    input_shape=input_shape + (1,),\n    num_classes=num_classes,\n    classifier_activation='softmax',\n\n    # encoder=None,\n    # encoder_depth=5,\n    # num_vit_layers=12,\n    # num_heads=8,\n    # num_queries=100,\n    # embed_dim=512,\n    # mlp_dim=1024,\n    # dropout_rate=0.1,\n    # decoder_activation=\"leaky_relu\",\n    # decoder_filters=(256, 128, 64, 32, 16),\n    # name=None,\n)\nprint(model.count_params() / 1e6)\n\n# ## Custom encoder\n# # backbone with 4 skip connection\n# backbone = ConvNeXtV2Tiny(\n#     input_shape=input_shape + (1,),\n#     include_top=False\n# )\n# # print(backbone.pyramid_outputs)\n\n# # segmentator\n# segmentor = TransUNet(\n#     encoder=backbone,\n#     encoder_depth=4,\n#     num_classes=num_classes,\n#     classifier_activation='softmax'\n# )\n# inputs = keras.Input(shape=input_shape + (1,))\n# x = segmentor(inputs)\n\n# # final supsampling, 2x tmes.\n# outputs = ResizingND(\n#     target_shape=input_shape,\n#     interpolation='trilinear',\n#     align_corners=False\n# )(x)\n# model = keras.Model(inputs=inputs, outputs=outputs)\n# model.count_params() / 1e6\n\n# ALERT: This attributes only available in medicai (not in core keras)\ntry:\n    print(model.instance_describe())\nexcept AttributeError:\n    pass\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null},{"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)\n# define optomizer, loss, metrics\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\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=50\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)\nmetrics = [\n    SparseDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        ignore_class_ids=2,\n        name='dice'\n    ),\n]\nmodel.compile(\n    optimizer=optim,\n    loss=combined_loss_fn,\n    metrics=metrics,\n)\nswi_callback_metric = SparseDiceMetric(\n    from_logits=False,\n    ignore_class_ids=2,\n    num_classes=num_classes,\n    name='val_dice',\n)\nswi_callback = SlidingWindowInferenceCallback(\n    model,\n    dataset=val_ds,\n    metrics=swi_callback_metric,\n    num_classes=num_classes,\n    interval=5,\n    overlap=0.5,\n    mode='gaussian',\n    roi_size=input_shape,\n    sw_batch_size=1 * total_device,\n    save_path=\"model.weights.h5\"\n)\n# ALERT: Starting may take time.\nmodel.fit(\n    train_ds,\n    epochs=epochs,\n    callbacks=[\n        swi_callback\n    ]\n)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_weights(\"model.weights.h5\")\nswi = SlidingWindowInference(\n    model,\n    num_classes=num_classes,\n    roi_size=input_shape,\n    mode='gaussian',\n    sw_batch_size=1 * total_device,\n    overlap=0.5,\n)\ndice = SparseDiceMetric(\n    from_logits=False,\n    num_classes=num_classes,\n    ignore_class_ids=2,\n    name='dice',\n)\nfor sample in val_ds:\n    x, y = sample\n    output = swi(x)\n    y = ops.convert_to_tensor(y)\n    output = ops.convert_to_tensor(output)\n    dice.update_state(y, output)\n\ndice_score = float(ops.convert_to_numpy(dice.result()))\nprint(f\"Dice Score: {dice_score:.4f}\")\ndice.reset_state()\n\nx, y = next(iter(val_ds))\nprint(x.shape, y.shape)\ny_pred = swi(x)\nprint(y_pred.shape)\nsegment = y_pred.argmax(-1).astype(np.uint8)\nprint(segment.shape, np.unique(segment))\nplot_sample(x, segment, sample_idx=0, max_slices=4)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-08T12:59:29.135Z"}},"outputs":[],"execution_count":null}]}