{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, sys, warnings\n\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\nwarnings.filterwarnings('ignore')\n\n# ============================================================\n# TPU HBM + Host RAM Monitor Callback\n# ============================================================\nimport resource as _resource\nimport keras as _keras_for_cb\nimport numpy as _np_cb\n\nclass TPUMemoryMonitorCallback(_keras_for_cb.callbacks.Callback):\n    \"\"\"\n    Her `log_freq` epoch'ta (ve training başında/sonunda):\n      - TPU HBM kullanımı (jax runtime backend memory stats)\n      - Host RAM kullanımı (RSS via resource module / psutil)\n    loglar. Sadece process 0 yazdırır.\n    \"\"\"\n    def __init__(self, log_freq=5, log_every_n_steps=None):\n        super().__init__()\n        self.log_freq          = log_freq\n        self.log_every_n_steps = log_every_n_steps\n        self._step_counter     = 0\n        self.peak_hbm_bytes    = 0\n        self.peak_ram_bytes    = 0\n        try:\n            import psutil\n            self._psutil = psutil\n        except ImportError:\n            self._psutil = None\n\n    def _get_hbm_info(self):\n        import jax as _jax\n        results = []\n        try:\n            for dev in _jax.local_devices():\n                stats = dev.memory_stats()\n                if stats is None: continue\n                used  = stats.get(\"bytes_in_use\", 0)\n                limit = stats.get(\"bytes_limit\", 0)\n                peak  = stats.get(\"peak_bytes_in_use\", 0)\n                pct   = (used / limit * 100) if limit > 0 else 0\n                results.append({\n                    \"device\": str(dev),\n                    \"hbm_used_gb\":  used  / (1024**3),\n                    \"hbm_limit_gb\": limit / (1024**3),\n                    \"hbm_peak_gb\":  peak  / (1024**3),\n                    \"hbm_pct\":      pct,\n                    \"hbm_free_gb\":  (limit - used) / (1024**3),\n                    \"_bytes_used\":  used,\n                })\n        except Exception: pass\n        return results\n\n    def _get_ram_info(self):\n        info = {}\n        if self._psutil:\n            vm = self._psutil.virtual_memory()\n            info = {\n                \"ram_used_gb\":  vm.used  / (1024**3),\n                \"ram_total_gb\": vm.total / (1024**3),\n                \"ram_pct\":      vm.percent,\n                \"ram_free_gb\":  vm.available / (1024**3),\n                \"_bytes_used\":  vm.used,\n            }\n        else:\n            ru = _resource.getrusage(_resource.RUSAGE_SELF)\n            import platform\n            rss_bytes = ru.ru_maxrss if platform.system() == \"Darwin\" else ru.ru_maxrss * 1024\n            info = {\n                \"ram_rss_gb\": rss_bytes / (1024**3),\n                \"ram_total_gb\": float('nan'),\n                \"ram_pct\": float('nan'),\n                \"ram_free_gb\": float('nan'),\n                \"_bytes_used\": rss_bytes,\n                \"_note\": \"psutil not installed — showing process RSS only\",\n            }\n        return info\n\n    def _log_memory(self, tag=\"\"):\n        import jax as _jax\n        if _jax.process_index() != 0: return\n\n        hbm_infos = self._get_hbm_info()\n        ram_info  = self._get_ram_info()\n\n        if hbm_infos:\n            max_hbm = max(h[\"_bytes_used\"] for h in hbm_infos)\n            self.peak_hbm_bytes = max(self.peak_hbm_bytes, max_hbm)\n        if ram_info.get(\"_bytes_used\", 0) > 0:\n            self.peak_ram_bytes = max(self.peak_ram_bytes, ram_info[\"_bytes_used\"])\n\n        print(f\"\\n  {'='*60}\")\n        print(f\"  [MemMonitor] {tag}\")\n        print(f\"  {'─'*60}\")\n\n        if hbm_infos:\n            for h in hbm_infos:\n                print(f\"    {h['device']}: \"\n                      f\"HBM {h['hbm_used_gb']:.2f}/{h['hbm_limit_gb']:.2f} GB \"\n                      f\"({h['hbm_pct']:.1f}%) | \"\n                      f\"Free: {h['hbm_free_gb']:.2f} GB | \"\n                      f\"Peak: {h['hbm_peak_gb']:.2f} GB\")\n            total_used  = sum(h[\"hbm_used_gb\"]  for h in hbm_infos)\n            total_limit = sum(h[\"hbm_limit_gb\"] for h in hbm_infos)\n            total_free  = sum(h[\"hbm_free_gb\"]  for h in hbm_infos)\n            print(f\"    ── HBM Total: {total_used:.2f}/{total_limit:.2f} GB | \"\n                  f\"Free: {total_free:.2f} GB | \"\n                  f\"Peak(max): {self.peak_hbm_bytes/(1024**3):.2f} GB\")\n        else:\n            print(\"    HBM: stats not available\")\n\n        if \"ram_rss_gb\" in ram_info:\n            print(f\"    Host RAM (process RSS): {ram_info['ram_rss_gb']:.2f} GB \"\n                  f\"({ram_info.get('_note', '')})\")\n        else:\n            print(f\"    Host RAM: {ram_info['ram_used_gb']:.2f}/{ram_info['ram_total_gb']:.2f} GB \"\n                  f\"({ram_info['ram_pct']:.1f}%) | \"\n                  f\"Free: {ram_info['ram_free_gb']:.2f} GB\")\n\n        print(f\"  {'='*60}\\n\")\n\n    def on_train_begin(self, logs=None):\n        self._log_memory(\"Training START\")\n\n    def on_epoch_end(self, epoch, logs=None):\n        if (epoch + 1) % self.log_freq == 0:\n            self._log_memory(f\"Epoch {epoch+1}\")\n\n    def on_train_batch_end(self, batch, logs=None):\n        if self.log_every_n_steps and self._step_counter % self.log_every_n_steps == 0:\n            self._log_memory(f\"Step {self._step_counter}\")\n        self._step_counter += 1\n\n    def on_train_end(self, logs=None):\n        self._log_memory(\"Training END (final)\")\n\n# ============================================================\n\nimport jax\nimport glob\nimport numpy as np\nimport pandas as pd\n\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# mainly for training API\nimport keras\nfrom keras import ops\nfrom keras.optimizers import AdamW\nfrom keras.optimizers.schedules import CosineDecay\n\n# only for tf.data API\nimport tensorflow as tf\n\n# mainly for 3D or 2D models, transformation, loss, metrics etc\nimport medicai\nfrom medicai.transforms import (\n    Compose,\n    NormalizeIntensity,\n    Resize,\n    RandShiftIntensity,\n    RandRotate90,\n    RandRotate,\n    RandFlip,\n    RandCutOut,\n    RandSpatialCrop\n)\nfrom medicai.models import TransUNet\nfrom medicai.losses import SparseTverskyLoss, SparseCenterlineDiceLoss\nfrom medicai.metrics import SparseDiceMetric\n\n# Initialize JAX distributed for multi-worker setup\njax.distributed.initialize()\n\n# 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)\ndata_parallel.auto_shard_dataset = False\n\ntotal_device = len(devices)\n\nif jax.process_index() == 0:\n    print(f'detected devices: {devices}')\n    print(f'total device: {total_device}')\n\n# ── Config ────────────────────────────────────────────────────────────────────\n# 160³ gives a 4.6× larger receptive field than 96³ — critical for thin papyrus\n# surfaces that span hundreds of voxels.\n# TransUNet/seresnext50 CNN encoder is memory-lighter than SwinUNETR at same\n# resolution, so batch_size = total_device // 2 is safe (same as kaggle4 swin_small).\ninput_shape  = (160, 160, 160)\nbatch_size   = max(1, total_device // 2)\nnum_classes  = 3\nnum_samples  = 780\nepochs       = 800  # CosineDecay single cycle over 800 epochs; early_stop protects\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 = tf.io.parse_single_example(example, feature_description)\n    image  = tf.io.decode_raw(parsed[\"image\"], tf.uint8)\n    label  = tf.io.decode_raw(parsed[\"label\"], tf.uint8)\n    image  = tf.reshape(image, tf.cast(parsed[\"image_shape\"], tf.int64))\n    label  = tf.reshape(label, tf.cast(parsed[\"label_shape\"], tf.int64))\n    return image, label\n\ndef prepare_inputs(image, label):\n    image = tf.cast(image[..., None], tf.float32)  # (D, H, W, 1)\n    label = tf.cast(label[..., None], tf.float32)  # (D, H, W, 1)\n    return image, label\n\ndef train_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        # ── Geometric ──────────────────────────────────────────────────────\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.3,   # reject patches that are >70% ignore-class\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.25,             # kept light — avoid suppressing train ceiling\n            fill_mode=\"crop\",\n        ),\n        # ── Intensity ──────────────────────────────────────────────────────\n        # NormalizeIntensity adapts to each volume's actual range (nonzero=True\n        # ignores black background). Much better than fixed [0,255]→[0,1] scale\n        # for CT-like data where HU ranges vary across scans.\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        # ── Occlusion ──────────────────────────────────────────────────────\n        RandCutOut(\n            keys=[\"image\", \"label\"],\n            invalid_label=2,\n            mask_size=[\n                input_shape[1] // 5,\n                input_shape[2] // 5,\n            ],\n            fill_mode=\"constant\",\n            cutout_mode='volume',\n            prob=0.2,              # light — enough regularisation without killing dice\n            num_cuts=1,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n\n\ndef val_transformation(image, label):\n    \"\"\"Crop + normalise — volumes MUST be input_shape to match model.\"\"\"\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\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.3,\n            max_attempts=10,\n        ),\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\n\ndef tfrecord_loader(tfrecord_pattern, batch_size=1, shuffle=True,\n                    drop_remainder=True, transform_fn=None, is_training=None):\n    if is_training is None:\n        is_training = shuffle\n    dataset = tf.data.TFRecordDataset(tf.io.gfile.glob(tfrecord_pattern))\n    dataset = dataset.shuffle(buffer_size=100) if shuffle else dataset\n    dataset = dataset.map(parse_tfrecord_fn,  num_parallel_calls=tf.data.AUTOTUNE)\n    dataset = dataset.map(prepare_inputs,     num_parallel_calls=tf.data.AUTOTUNE)\n    if transform_fn is None:\n        transform_fn = train_transformation if is_training else val_transformation\n    dataset = dataset.map(transform_fn,       num_parallel_calls=tf.data.AUTOTUNE)\n    dataset = dataset.batch(batch_size, drop_remainder=drop_remainder)\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    return dataset\n\n\n# ── Train / val split ─────────────────────────────────────────────────────────\nall_tfrec = sorted(\n    glob.glob(\"tfrecords/*.tfrec\"),\n    key=lambda x: int(x.split(\"_\")[-1].replace(\".tfrec\", \"\"))\n)\n\n# 5 val shards → ~30 samples → stable val_dice metric (vs 3 shards = ~18 in kaggle4)\nNUM_VAL_SHARDS = 5\n_rng = np.random.default_rng(seed=42)\n_shuffled      = _rng.permutation(len(all_tfrec))\nval_indices    = sorted(_shuffled[:NUM_VAL_SHARDS])\ntrain_indices  = sorted(_shuffled[NUM_VAL_SHARDS:])\nval_patterns   = [all_tfrec[i] for i in val_indices]\ntrain_patterns = [all_tfrec[i] for i in train_indices]\n\nif jax.process_index() == 0:\n    print(f\"[Split] val shards  ({len(val_patterns)}): {[os.path.basename(p) for p in val_patterns]}\")\n    print(f\"[Split] train shards({len(train_patterns)}): {len(train_patterns)} shards\")\n\ntrain_ds = tfrecord_loader(train_patterns, batch_size=batch_size, shuffle=True,  is_training=True)\nval_ds   = tfrecord_loader(val_patterns,   batch_size=batch_size, shuffle=False, is_training=False)\n\n\n# ── Model — TransUNet + seresnext50 ──────────────────────────────────────────\n# Why TransUNet + seresnext50 over SwinUNETR here:\n#   • 160³ input (4.6× more voxels than 96³) → captures long-range surface structure\n#   • seresnext50 squeeze-excitation blocks → channel-wise attention proven on 3D CT\n#   • CNN encoder lighter on HBM than Swin attention maps at 160³\n#   • Hybrid: CNN encoder (local) + transformer bottleneck (global) — best of both worlds\nmodel = TransUNet(\n    input_shape=input_shape + (1,),\n    encoder_name=\"seresnext50\",\n    classifier_activation=\"softmax\",\n    num_classes=num_classes,\n)\nmodel.count_params() / 1e6\nif jax.process_index() == 0:\n    try:\n        print(model.instance_describe())\n    except AttributeError:\n        pass\n\n\n# ── LR & Optimizer ───────────────────────────────────────────────────────────\n# Single CosineDecay over full training budget (no restarts).\n# Lesson from kaggle4: CosineDecayRestarts with m_mul=0.9 killed LR by epoch 300.\n# 5-epoch warmup stabilises early training when weights are random.\n# CNN encoder (seresnext50) benefits from slightly higher BASE_LR than transformers.\nBASE_LR          = 2e-4\nscaled_lr        = BASE_LR * (batch_size / 2) ** 0.5  # sqrt batch scaling\nsteps_per_epoch  = num_samples // batch_size\ntotal_steps      = steps_per_epoch * epochs\n\nlr_schedule = CosineDecay(\n    initial_learning_rate=scaled_lr,\n    decay_steps=total_steps,\n    alpha=0.005,                          # floor at 0.5% of peak LR\n    warmup_steps=steps_per_epoch * 5,     # 5-epoch linear warmup\n)\n\noptim = AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n    clipnorm=1.0,   # gradient clipping — stabilises clDice gradients\n)\n\n\n# ── Loss ─────────────────────────────────────────────────────────────────────\n# SparseTverskyLoss (recall-focused) + SparseCenterlineDiceLoss (topology).\n# Replaces the SkeletonRecallPlusDiceLoss + DiceCE combo from the old baseline:\n#   • No need for a separate skeleton target channel in y_true\n#   • clDice covers topology implicitly via soft-skeletonisation on predictions\n#   • Tversky beta=0.7 ensures missed surfaces are penalised more than FPs\ntversky_loss_fn = SparseTverskyLoss(\n    from_logits=False,\n    num_classes=num_classes,\n    ignore_class_ids=2,\n    alpha=0.3,   # FP weight\n    beta=0.7,    # FN weight — higher = recall-focused\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=20,\n)\ncombined_loss_fn = lambda y_true, y_pred: (\n    tversky_loss_fn(y_true, y_pred) + 0.2 * cldice_loss_fn(y_true, y_pred)\n)\n\nmetrics = [\n    SparseDiceMetric(\n        from_logits=False,\n        num_classes=num_classes,\n        ignore_class_ids=2,\n        name='dice',\n    ),\n]\n\nmodel.compile(\n    optimizer=optim,\n    loss=combined_loss_fn,\n    metrics=metrics,\n)\n\n\n# ── Auto-resume after TPU preemption ─────────────────────────────────────────\nCKPT_PATH = \"model_transunet.weights.h5\"\nLOG_PATH  = \"training_log_transunet.csv\"\ninitial_epoch = 0\n\nif os.path.exists(CKPT_PATH):\n    model.load_weights(CKPT_PATH)\n    if jax.process_index() == 0:\n        print(f\"[Resume] Loaded weights from {CKPT_PATH}\")\n\nif os.path.exists(LOG_PATH):\n    try:\n        _log_df = pd.read_csv(LOG_PATH)\n        if len(_log_df) > 0 and \"epoch\" in _log_df.columns:\n            initial_epoch = int(_log_df[\"epoch\"].iloc[-1]) + 1\n            if jax.process_index() == 0:\n                print(f\"[Resume] Resuming from epoch {initial_epoch}/{epochs}\")\n        else:\n            if jax.process_index() == 0:\n                print(\"[Resume] Log empty or missing 'epoch' — starting from epoch 0\")\n    except pd.errors.EmptyDataError:\n        if jax.process_index() == 0:\n            print(\"[Resume] Log is empty (preempted before first epoch) — starting from epoch 0\")\n# ─────────────────────────────────────────────────────────────────────────────\n\n\n# ── Callbacks ────────────────────────────────────────────────────────────────\ncheckpoint_best = keras.callbacks.ModelCheckpoint(\n    filepath=\"model_transunet.weights.h5\",\n    monitor='val_dice',\n    save_best_only=True,\n    save_weights_only=True,\n    mode='max',\n    verbose=1 if jax.process_index() == 0 else 0,\n)\n\ncheckpoint_trainloss = keras.callbacks.ModelCheckpoint(\n    filepath=\"model_transunet_trainloss.weights.h5\",\n    monitor='loss',\n    save_best_only=True,\n    save_weights_only=True,\n    mode='min',\n    verbose=1 if jax.process_index() == 0 else 0,\n)\n\n# patience=60 — TransUNet needs more runway than SwinUNETR (CNN converges differently)\nearly_stop = keras.callbacks.EarlyStopping(\n    monitor='val_dice',\n    mode='max',\n    patience=60,\n    restore_best_weights=True,\n    verbose=1 if jax.process_index() == 0 else 0,\n)\n\nmem_monitor = TPUMemoryMonitorCallback(log_freq=5)\ncsv_logger  = keras.callbacks.CSVLogger(LOG_PATH, separator=\",\", append=True)\n\n\n# ALERT: Starting may take time.\nmodel.fit(\n    train_ds,\n    validation_data=val_ds,\n    epochs=epochs,\n    initial_epoch=initial_epoch,\n    verbose=1 if jax.process_index() == 0 else 0,\n    callbacks=[\n        mem_monitor,\n        csv_logger,\n        checkpoint_best,\n        checkpoint_trainloss,\n        early_stop,\n    ],\n)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}