{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13962683,"sourceType":"datasetVersion","datasetId":8751895}],"isInternetEnabled":false,"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 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":"2025-12-07T15:11:26.674277Z","iopub.execute_input":"2025-12-07T15:11:26.675173Z","iopub.status.idle":"2025-12-07T15:11:27.011108Z","shell.execute_reply.started":"2025-12-07T15:11:26.675138Z","shell.execute_reply":"2025-12-07T15:11:27.010512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras\nfrom keras import ops\n\nimport torch\nimport tensorflow as tf\n\nkeras.version(), torch.__version__, keras.config.backend()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:11:27.012055Z","iopub.execute_input":"2025-12-07T15:11:27.012384Z","iopub.status.idle":"2025-12-07T15:11:35.413328Z","shell.execute_reply.started":"2025-12-07T15:11:27.012349Z","shell.execute_reply":"2025-12-07T15:11:35.412692Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loader","metadata":{}},{"cell_type":"code","source":"from medicai.transforms import (\n    Compose,\n    ScaleIntensityRange,\n    Resize,\n    RandShiftIntensity,\n    RandRotate90,\n    RandFlip,\n    RandSpatialCrop\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:11:35.414263Z","iopub.execute_input":"2025-12-07T15:11:35.414793Z","iopub.status.idle":"2025-12-07T15:11:35.422976Z","shell.execute_reply.started":"2025-12-07T15:11:35.414769Z","shell.execute_reply":"2025-12-07T15:11:35.422486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nimport numpy as np\nfrom medicai.transforms import (\n    Compose,\n    ScaleIntensityRange,\n    RandShiftIntensity,\n    RandRotate90,\n    RandFlip,\n    RandSpatialCrop\n)\n\n# ==================== FIXED CUSTOM TRANSFORMS ====================\n\nclass RandAdjustContrast:\n    \"\"\"Random Contrast Adjustment (Gamma Correction) - TF Graph Compatible\"\"\"\n    def __init__(self, keys, prob=0.3, gamma_range=(0.7, 1.3)):\n        self.keys = keys\n        self.prob = prob\n        self.gamma_range = gamma_range\n    \n    def __call__(self, data):\n        for key in self.keys:\n            # Use tf.cond instead of Python if\n            data[key] = tf.cond(\n                tf.random.uniform([]) < self.prob,\n                lambda: self._apply_gamma(data[key]),\n                lambda: data[key]\n            )\n        return data\n    \n    def _apply_gamma(self, image):\n        gamma = tf.random.uniform(\n            [], \n            self.gamma_range[0], \n            self.gamma_range[1]\n        )\n        # Clip to [0, 1] before applying gamma\n        image_clipped = tf.clip_by_value(image, 0.0, 1.0)\n        return tf.pow(image_clipped, gamma)\n\n\nclass RandUnsharpMask:\n    \"\"\"Random Unsharp Masking - TF Graph Compatible\"\"\"\n    def __init__(self, keys, prob=0.3, amount_range=(0.5, 1.5), kernel_size=3):\n        self.keys = keys\n        self.prob = prob\n        self.amount_range = amount_range\n        self.kernel_size = kernel_size\n    \n    def __call__(self, data):\n        for key in self.keys:\n            data[key] = tf.cond(\n                tf.random.uniform([]) < self.prob,\n                lambda: self._apply_unsharp_mask(data[key]),\n                lambda: data[key]\n            )\n        return data\n    \n    def _apply_unsharp_mask(self, image):\n        amount = tf.random.uniform(\n            [], \n            self.amount_range[0], \n            self.amount_range[1]\n        )\n        \n        # Simple box blur untuk 3D\n        blurred = self._box_blur_3d(image, self.kernel_size)\n        \n        # Unsharp mask formula\n        sharpened = image + amount * (image - blurred)\n        return tf.clip_by_value(sharpened, 0.0, 1.0)\n    \n    def _box_blur_3d(self, volume, kernel_size):\n        \"\"\"Box blur menggunakan average pooling 3D\"\"\"\n        # Add batch dimension\n        volume_batched = volume[tf.newaxis, ...]\n        \n        # Apply 3D average pooling\n        ksize = [1, kernel_size, kernel_size, kernel_size, 1]\n        strides = [1, 1, 1, 1, 1]\n        blurred = tf.nn.avg_pool3d(\n            volume_batched,\n            ksize=ksize,\n            strides=strides,\n            padding='SAME'\n        )\n        \n        # Remove batch dimension\n        return blurred[0]\n\n\nclass RandAddGaussianNoise:\n    \"\"\"Random Gaussian Noise - TF Graph Compatible\"\"\"\n    def __init__(self, keys, prob=0.2, noise_std_range=(0.0, 0.05)):\n        self.keys = keys\n        self.prob = prob\n        self.noise_std_range = noise_std_range\n    \n    def __call__(self, data):\n        for key in self.keys:\n            data[key] = tf.cond(\n                tf.random.uniform([]) < self.prob,\n                lambda: self._add_noise(data[key]),\n                lambda: data[key]\n            )\n        return data\n    \n    def _add_noise(self, image):\n        noise_std = tf.random.uniform(\n            [], \n            self.noise_std_range[0], \n            self.noise_std_range[1]\n        )\n        \n        noise = tf.random.normal(\n            tf.shape(image),\n            mean=0.0,\n            stddev=noise_std\n        )\n        \n        noisy = image + noise\n        return tf.clip_by_value(noisy, 0.0, 1.0)\n\n\nclass RandBrightnessContrast:\n    \"\"\"Random Brightness and Contrast adjustment - TF Graph Compatible\"\"\"\n    def __init__(self, keys, prob=0.3, \n                 brightness_range=(-0.1, 0.1),\n                 contrast_range=(0.8, 1.2)):\n        self.keys = keys\n        self.prob = prob\n        self.brightness_range = brightness_range\n        self.contrast_range = contrast_range\n    \n    def __call__(self, data):\n        for key in self.keys:\n            data[key] = tf.cond(\n                tf.random.uniform([]) < self.prob,\n                lambda: self._adjust_brightness_contrast(data[key]),\n                lambda: data[key]\n            )\n        return data\n    \n    def _adjust_brightness_contrast(self, image):\n        # Random brightness adjustment\n        brightness_delta = tf.random.uniform(\n            [],\n            self.brightness_range[0],\n            self.brightness_range[1]\n        )\n        image = image + brightness_delta\n        \n        # Random contrast adjustment\n        contrast_factor = tf.random.uniform(\n            [],\n            self.contrast_range[0],\n            self.contrast_range[1]\n        )\n        mean = tf.reduce_mean(image)\n        image = (image - mean) * contrast_factor + mean\n        \n        return tf.clip_by_value(image, 0.0, 1.0)\n\n\nclass RandHistogramShift:\n    \"\"\"Random Histogram Shifting - Simplified TF Graph Compatible\"\"\"\n    def __init__(self, keys, prob=0.3, shift_range=(-0.1, 0.1)):\n        self.keys = keys\n        self.prob = prob\n        self.shift_range = shift_range\n    \n    def __call__(self, data):\n        for key in self.keys:\n            data[key] = tf.cond(\n                tf.random.uniform([]) < self.prob,\n                lambda: self._shift_histogram(data[key]),\n                lambda: data[key]\n            )\n        return data\n    \n    def _shift_histogram(self, image):\n        \"\"\"Simplified histogram shift using intensity scaling\"\"\"\n        shift = tf.random.uniform(\n            [],\n            self.shift_range[0],\n            self.shift_range[1]\n        )\n        \n        # Non-linear transformation\n        shifted = image + shift * (image - 0.5)\n        return tf.clip_by_value(shifted, 0.0, 1.0)\n\n\n# ==================== PIPELINE DENGAN FIXED TRANSFORMS ====================\n\ndef train_transformation(image, label):\n    \"\"\"Training transformation dengan custom enhancement - FIXED VERSION\"\"\"\n    data = {\"image\": image, \"label\": label}\n    \n    pipeline = Compose([\n        # 1. Normalisasi dasar\n        ScaleIntensityRange(\n            keys=[\"image\"],\n            a_min=0,\n            a_max=255,\n            clip=True,\n        ),\n        \n        # 2. Spatial augmentation (dari medicai)\n        RandSpatialCrop(\n            keys=[\"image\", \"label\"],\n            roi_size=(128, 128, 128),\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        # 3. CUSTOM ENHANCEMENT TRANSFORMS (FIXED)\n        RandAdjustContrast(\n            keys=[\"image\"],\n            prob=0.3,\n            gamma_range=(0.7, 1.3)\n        ),\n        \n        RandUnsharpMask(\n            keys=[\"image\"],\n            prob=0.25,\n            amount_range=(0.5, 1.2),\n            kernel_size=3\n        ),\n        \n        RandBrightnessContrast(\n            keys=[\"image\"],\n            prob=0.3,\n            brightness_range=(-0.1, 0.1),\n            contrast_range=(0.85, 1.15)\n        ),\n        \n        RandAddGaussianNoise(\n            keys=[\"image\"],\n            prob=0.2,\n            noise_std_range=(0.0, 0.03)\n        ),\n        \n        RandHistogramShift(\n            keys=[\"image\"],\n            prob=0.2,\n            shift_range=(-0.05, 0.05)\n        ),\n        \n        # 4. Intensity shift (dari medicai)\n        RandShiftIntensity(\n            keys=[\"image\"], \n            offsets=0.10, \n            prob=0.5\n        ),\n    ])\n    \n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n\n\ndef val_transformation(image, label):\n    \"\"\"Validation transform - hanya normalisasi\"\"\"\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        ScaleIntensityRange(\n            keys=[\"image\"],\n            a_min=0,\n            a_max=255,\n            clip=True,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:11:38.949365Z","iopub.execute_input":"2025-12-07T15:11:38.950037Z","iopub.status.idle":"2025-12-07T15:11:38.970973Z","shell.execute_reply.started":"2025-12-07T15:11:38.950014Z","shell.execute_reply":"2025-12-07T15:11:38.970222Z"}},"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-12-07T15:11:39.198559Z","iopub.execute_input":"2025-12-07T15:11:39.199143Z","iopub.status.idle":"2025-12-07T15:11:39.204249Z","shell.execute_reply.started":"2025-12-07T15:11:39.199113Z","shell.execute_reply":"2025-12-07T15:11:39.203436Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_inputs(image, label):\n    # Only take gt 1 \n    label = (label == 1)\n    \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    \n    return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:11:40.803136Z","iopub.execute_input":"2025-12-07T15:11:40.803739Z","iopub.status.idle":"2025-12-07T15:11:40.807804Z","shell.execute_reply.started":"2025-12-07T15:11:40.803714Z","shell.execute_reply":"2025-12-07T15:11:40.806975Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_tfrecord_dataset(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).prefetch(tf.data.AUTOTUNE)\n    return dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:11:41.003666Z","iopub.execute_input":"2025-12-07T15:11:41.004261Z","iopub.status.idle":"2025-12-07T15:11:41.009214Z","shell.execute_reply.started":"2025-12-07T15:11:41.004236Z","shell.execute_reply":"2025-12-07T15:11:41.008362Z"}},"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\nbatch_size = 6\ntrain_ds = load_tfrecord_dataset(\n    train_pattern, batch_size=batch_size, shuffle=True\n)\n\n# Keep the batch size 1 for validation to perform sliding-window-inference\nval_ds = load_tfrecord_dataset(\n    val_pattern, batch_size=1, shuffle=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:11:43.259924Z","iopub.execute_input":"2025-12-07T15:11:43.26087Z","iopub.status.idle":"2025-12-07T15:11:46.075616Z","shell.execute_reply.started":"2025-12-07T15:11:43.260837Z","shell.execute_reply":"2025-12-07T15:11:46.074987Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(train_ds))\nx.shape, y.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:11:46.076922Z","iopub.execute_input":"2025-12-07T15:11:46.077582Z","iopub.status.idle":"2025-12-07T15:12:41.849058Z","shell.execute_reply.started":"2025-12-07T15:11:46.077557Z","shell.execute_reply":"2025-12-07T15:12:41.848177Z"}},"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":"2025-12-07T15:12:43.821327Z","iopub.execute_input":"2025-12-07T15:12:43.821595Z","iopub.status.idle":"2025-12-07T15:12:46.459788Z","shell.execute_reply.started":"2025-12-07T15:12:43.821572Z","shell.execute_reply":"2025-12-07T15:12:46.459034Z"}},"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},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_sample(\n    x, y, sample_idx=0, max_slices=4\n)","metadata":{"trusted":true},"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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"import medicai\nfrom medicai.models import SegFormer\nfrom medicai.losses import BinaryDiceCELoss\nfrom medicai.metrics import BinaryDiceMetric\nfrom medicai.utils.inference import SlidingWindowInference","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"medicai.models.list_models()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = 1\nmodel = SegFormer(\n    input_shape=(128, 128, 128, 1),\n    encoder_name='mit_b0',\n    classifier_activation='sigmoid',\n    num_classes=num_classes,\n)\nmodel.count_params() / 1e6","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Pipeline","metadata":{}},{"cell_type":"code","source":"# define optomizer, loss\noptim = keras.optimizers.AdamW(\n    learning_rate=1e-4,\n    weight_decay=1e-5,\n)\nloss_fn = BinaryDiceCELoss(\n    from_logits=False, \n    num_classes=num_classes\n)\n\n# define sliding-window-inferencer for validation\nswi = SlidingWindowInference(\n    model,\n    num_classes=num_classes,\n    roi_size=(128, 128, 128),\n    sw_batch_size=4,\n    overlap=0.5,\n)","metadata":{"trusted":true},"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 = 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        y_true = ops.convert_to_tensor(labels)\n        y_pred = ops.convert_to_tensor(outputs)\n        \n        # Kalau 5D: (B, D, H, W, C) → flatten jadi (B, D*H*W, C)\n        if len(y_true.shape) == 5:\n            b, d, h, w, c = y_true.shape\n            y_true = ops.reshape(y_true, (b, d * h * w, c))\n            y_pred = ops.reshape(y_pred, (b, d * h * w, c))\n        \n        metrics.update_state(y_true, y_pred)\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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, dataloader, metrics):\n    # optional tapi bagus: set ke eval\n    model.eval()\n\n    for x, y in dataloader:\n        # kalau pakai torch backend, sebaiknya no_grad biar hemat memori\n        with torch.no_grad():\n            output = swi(x)  # atau model(x), sesuai punyamu\n\n        # konversi ke tensor keras.ops\n        y = ops.convert_to_tensor(y)\n        output = ops.convert_to_tensor(output)\n\n        # EXPECTED SHAPE awal:\n        # y      : (B, D, H, W, C)\n        # output : (B, D, H, W, C)\n        if len(y.shape) == 5:\n            b, d, h, w, c = y.shape  # (batch, depth, height, width, channel)\n\n            # flatten dim spasial → (B, N, C), N = D*H*W\n            y = ops.reshape(y, (b, d * h * w, c))\n            output = ops.reshape(output, (b, d * h * w, c))\n\n        # urutan: (y_true, y_pred)\n        metrics.update_state(y, output)\n\n    return metrics\n","metadata":{"trusted":true},"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 = BinaryDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        name='dice'\n    )\n\n    # metrics for validation\n    val_metrics = BinaryDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \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) % 5 == 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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_training(\n    train_ds, val_ds, model, epochs=1\n)","metadata":{"trusted":true},"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 > 0.35).astype(np.uint8)\nsegment.shape, np.unique(segment)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_sample(\n    x, y, sample_idx=0, max_slices=4\n)","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":"import tifffile as tiff\nimport numpy as np\n\ndef load_tiff_3d_np(path):\n    volume = tiff.imread(path)      # bisa return (D, H, W) atau (H, W, D) atau (H, W)\n    volume = volume.astype(\"float32\")\n\n    # Normalisasi 0–1, asumsi 16-bit TIFF (Vesuvius biasanya 0–65535)\n    max_val = volume.max() if volume.max() > 0 else 1.0\n    volume = volume / max_val\n\n    # Pastikan ada channel dimension\n    if volume.ndim == 3:\n        # asumsikan (D, H, W) → tambahkan channel=1 → (D, H, W, 1)\n        volume = volume[..., np.newaxis]\n    elif volume.ndim == 2:\n        # (H, W) → (1, H, W, 1) misal dianggap depth=1\n        volume = volume[np.newaxis, ..., np.newaxis]\n\n    return volume\n\npath = \"/kaggle/input/vesuvius-challenge-surface-detection/test_images/1407735.tif\"\nvol_np = load_tiff_3d_np(path)\nprint(\"numpy shape:\", vol_np.shape)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"res = swi(np.expand_dims(vol_np, axis=0))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"segres = (res > 0.35).astype(np.uint8)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tiff.imwrite(\"/kaggle/working/1407735.tif\", segres[0])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import zipfile\nimport os\n\noutput_zip = \"submission.zip\"\npred_dir = \"/kaggle/working\"\n\nwith zipfile.ZipFile(output_zip, \"w\") as zipf:\n    for filename in os.listdir(pred_dir):\n        if filename.endswith(\".tif\"):\n            zipf.write(\n                os.path.join(pred_dir, filename),\n                arcname=filename  # penting: hanya nama file, tanpa path\n            )\n\nprint(\"ZIP created:\", output_zip)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}