{
  "id": 679565,
  "title": "123rd place — not where I aimed, but here's what I learned",
  "url": "/competitions/vesuvius-challenge-surface-detection/writeups/135th-place-not-where-i-aimed-but-heres-what-i",
  "author_name": "",
  "post_date": "2026-03-02T10:27:35.093Z",
  "votes": 8,
  "comment_count": 2,
  "views": 0,
  "content": "<h1>123rd place — not where I aimed, but here's what I learned</h1>\n<ul>\n<li><strong>Competition:</strong> <a href=\"https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection\" target=\"_blank\">Vesuvius Challenge Surface Detection</a></li>\n<li><strong>Result:</strong> 123rd / 1427 · Score: 0.579 · 🥉 Bronze Medal</li>\n<li><strong>Training code:</strong> <a href=\"https://www.kaggle.com/code/osman0/vesuvius-surface-train-tpu\" target=\"_blank\">kaggle.com/code/osman0/vesuvius-surface-train-tpu</a></li>\n</ul>\n<hr>\n<h2>Personal Reflection</h2>\n<p>I'll be honest.  I didn't do as well as I had hoped. I was aiming much higher like top 10 But the result didn't reflect that effort, and I think I know whys.</p>\n<p>I'll be real about where I was coming from: I'm relatively new to competitive machine learning, TPUs, and JAX all at once. There was a steep learning curve just to get anything running, let alone running well.</p>\n<p>I won't pretend the motivation was purely academic either. Top 10 meant prize money, and I genuinely needed it. That drove me to put in a lot of hours — probably more than I should have at certain points. I tried hard. I just didn't have the experience yet to convert that effort into a top result.</p>\n<p>And that's okay. I learned an enormous amount. I went from knowing nothing about distributed TPU training to having a working end-to-end 3D segmentation pipeline. That's real progress, even if the leaderboard didn't fully show it. Next time, I'll come in with that foundation already in place.</p>\n<p>A big thank you to <a href=\"https://www.kaggle.com/ipythonx\" target=\"_blank\">Innat</a> — his <a href=\"https://www.kaggle.com/code/ipythonx/train-vesuvius-surface-3d-detection-on-tpu\" target=\"_blank\">notebook</a> was what got me started with the <code>medicai</code> framework and TPU-compatible 3D training.</p>\n<hr>\n<h2>Acknowledgements</h2>\n<p>This training was done as part of the <strong>TRC (TPU Research Cloud) program</strong>. I'm genuinely grateful to Google for providing access to TPU v4-32 hardware through TRC training 160³ 3D models at this scale would have been completely infeasible otherwise. Thank you.</p>\n<hr>\n<h2>TL;DR</h2>\n<p>I built a 3D segmentation pipeline to detect ancient papyrus surfaces from CT scans of the Herculaneum scrolls. My final solution uses <strong>TransUNet</strong> <a href=\"#references\">[1]</a> with a <strong>SEResNeXt</strong> <a href=\"#references\">[2]</a> encoder, trained on <strong>160³ voxel</strong> patches with a <strong>Tversky</strong> <a href=\"#references\">[3]</a> <strong>+ CenterlineDice</strong> <a href=\"#references\">[4]</a> loss on <strong>TPU v4-32</strong>. Inference combines <strong>sliding window</strong>, <strong>6-fold TTA</strong>, <strong>dual checkpoint ensemble</strong>, and <strong>morphological post-processing</strong>. My biggest personal takeaway was learning to train 3D medical segmentation models end-to-end on TPUs with JAX/Keras.</p>\n<hr>\n<h2>Solution Architecture</h2>\n<p><img src=\"https://www.googleapis.com/download/storage/v1/b/kaggle-user-content/o/inbox%2F20259944%2F27dc43a53e40e5e9ff69b8ab1eff4495%2FScreenshot%202026-03-02%20at%2012.37.40.png?generation=1772444684940610&amp;alt=media\" alt=\"/Users/osman/Desktop/Screenshot 2026-03-02 at 12.37.40.png\"></p>\n<p><em>End-to-end pipeline: 3D data prep → TransUNet training → dual checkpoint blending → post-processing → binary ink mask.</em></p>\n<h3>Model: TransUNet + SEResNeXt</h3>\n<pre><code>Input (160³×1) → SEResNeXt CNN Encoder → Transformer Bottleneck → UNet Decoder → Softmax (3 classes)\n</code></pre>\n<p><strong>Why this architecture?</strong></p>\n<ul>\n<li><strong>160³ input</strong> — 4.6× larger receptive field than 96³, critical for thin papyrus surfaces spanning hundreds of voxels</li>\n<li><strong>SEResNeXt encoder</strong> <a href=\"#references\">[2]</a> — squeeze-excitation blocks provide channel-wise attention, proven on 3D CT data</li>\n<li><strong>Hybrid CNN + Transformer</strong> <a href=\"#references\">[1]</a> — CNN encoder captures local features, transformer bottleneck captures global context</li>\n<li><strong>HBM-efficient</strong> — CNN encoder is lighter on TPU HBM than Swin attention maps at 160³ resolution</li>\n</ul>\n<p>I experimented with two encoder sizes:</p>\n<table>\n<thead>\n<tr>\n<th>Variant</th>\n<th>Training</th>\n<th>Inference</th>\n</tr>\n</thead>\n<tbody>\n<tr>\n<td><code>seresnext50</code></td>\n<td>✅ (train-tpu.py)</td>\n<td>—</td>\n</tr>\n<tr>\n<td><code>seresnext101</code></td>\n<td>✅ (final model)</td>\n<td>✅ (inference-kaggle6.py)</td>\n</tr>\n</tbody>\n</table>\n<hr>\n<h2>Training Details</h2>\n<table>\n<thead>\n<tr>\n<th>Parameter</th>\n<th>Value</th>\n</tr>\n</thead>\n<tbody>\n<tr>\n<td><strong>Hardware</strong></td>\n<td>Kaggle TPU v4-32 (32 cores)</td>\n</tr>\n<tr>\n<td><strong>Framework</strong></td>\n<td>Keras 3 + JAX backend</td>\n</tr>\n<tr>\n<td><strong>Batch size</strong></td>\n<td><code>total_devices // 2</code> (16 on TPU v4-32)</td>\n</tr>\n<tr>\n<td><strong>Epochs</strong></td>\n<td>800 (CosineDecay + EarlyStopping patience=60)</td>\n</tr>\n<tr>\n<td><strong>Base LR</strong></td>\n<td>2e-4 with sqrt batch scaling</td>\n</tr>\n<tr>\n<td><strong>LR Schedule</strong></td>\n<td>Single-cycle CosineDecay, 5-epoch warmup, alpha=0.005</td>\n</tr>\n<tr>\n<td><strong>Optimizer</strong></td>\n<td>AdamW (weight_decay=1e-5, clipnorm=1.0)</td>\n</tr>\n<tr>\n<td><strong>Loss</strong></td>\n<td>SparseTverskyLoss + 0.2 × SparseCenterlineDiceLoss</td>\n</tr>\n<tr>\n<td><strong>Val split</strong></td>\n<td>5 TFRecord shards (~30 samples), seed=42</td>\n</tr>\n</tbody>\n</table>\n<h3>Loss Function</h3>\n<pre><code># Tversky loss — recall-focused (alpha=0.3 FP, beta=0.7 FN)\ntversky = SparseTverskyLoss(alpha=0.3, beta=0.7, ignore_class_ids=2)\n\n# CenterlineDice — topology preservation via soft-skeletonisation\ncldice = SparseCenterlineDiceLoss(target_class_ids=1, ignore_class_ids=2, iters=20)\n\n# Combined\nloss = tversky + 0.2 * cldice\n</code></pre>\n<p><strong>Rationale:</strong> Tversky with β=0.7 penalizes missed surfaces (FN) more than false positives <a href=\"#references\">[3]</a>, aligning with the competition metric's sensitivity to gaps. CenterlineDice <a href=\"#references\">[4]</a> preserves topological continuity via soft-skeletonisation on predictions — directly targeting the TopoScore sub-metric. The <code>ignore_class_ids=2</code> automatically masks out unlabeled regions.</p>\n<h3>Data Pipeline &amp; Augmentation</h3>\n<p>Data is stored as TFRecords (780 samples). Training augmentation includes:</p>\n<ul>\n<li><strong>RandSpatialCrop</strong> (160³, <code>min_valid_ratio=0.3</code> — rejects patches &gt;70% ignore-class)</li>\n<li><strong>RandFlip</strong> (all 3 axes, prob=0.5)</li>\n<li><strong>RandRotate90</strong> (prob=0.4) + <strong>RandRotate</strong> (factor=0.2, prob=0.25)</li>\n<li><strong>NormalizeIntensity</strong> (<code>nonzero=True</code>) — adapts to each volume's actual intensity range rather than fixed [0,255]→[0,1] scaling</li>\n<li><strong>RandShiftIntensity</strong> (offsets=0.10, prob=0.5)</li>\n<li><strong>RandCutOut</strong> (volume mode, prob=0.2) — occlusion regularisation</li>\n</ul>\n<h3>LR Schedule Evolution</h3>\n<p>Earlier experiments used <code>CosineDecayRestarts(m_mul=0.9)</code>, but the learning rate collapsed too early (~epoch 300). Switching to a <strong>single-cycle CosineDecay</strong> over the full 800-epoch budget produced much more stable convergence.</p>\n<h3>TPU-Specific Implementation</h3>\n<ul>\n<li><strong>DataParallel</strong> distribution across 32 TPU v4 cores with auto-sharding</li>\n<li><strong>Flash attention disabled</strong> — required for distributed training compatibility</li>\n<li><strong>Custom HBM Memory Monitor callback</strong> — logs TPU HBM usage and host RAM every N epochs</li>\n<li><strong>Auto-resume from preemption</strong> — checkpoint + CSV log continuation for Kaggle TPU interrupts</li>\n</ul>\n<hr>\n<h2>Inference Pipeline</h2>\n<h3>Sliding Window Inference</h3>\n<pre><code>SlidingWindowInference(\n    roi_size=(160, 160, 160),  # must match ViT positional embeddings (5³=125 patches)\n    overlap=0.5,\n    blend_mode=\"gaussian\",     # smooth boundary artifacts\n)\n</code></pre>\n<h3>Test-Time Augmentation (6-fold)</h3>\n<pre><code>TTA_AXES = [(), (0,), (1,), (2,), (0,1), (1,2)]\n# Original + 3 single-axis flips + 2 pair flips → averaged\n</code></pre>\n<h3>Dual Checkpoint Blend</h3>\n<pre><code>pred = 0.80 * pred_best_val_dice + 0.20 * pred_best_train_loss\n</code></pre>\n<p>Two checkpoints from the same training run — <code>best_val_dice</code> generalizes better, <code>best_train_loss</code> is more expressive. Blending consistently improved over either single checkpoint.</p>\n<h3>Post-Processing</h3>\n<ol>\n<li><strong>Ignore-class suppression</strong>: <code>ink_prob -= 0.15 × ignore_prob</code> — leverages class-2 predictions to reduce false positives in uncertain regions</li>\n<li><strong>Thresholding</strong>: prob ≥ 0.5 → binary mask</li>\n<li><strong>Morphological closing</strong> (radius=1) — fills small gaps</li>\n<li><strong>Connected component filtering</strong> (min_size=50) — removes noise</li>\n</ol>\n<hr>\n<h2>Learnings</h2>\n<h3>TPU for 3D Medical Segmentation</h3>\n<p>This competition was my first serious experience training on TPUs. It required several specific codebase adjustments to run successfully:</p>\n<p><strong>1. JAX backend over TensorFlow:</strong>\nJAX is significantly more stable and faster on TPUs than TF eager mode.</p>\n<pre><code>import os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n</code></pre>\n<p><strong>2. Disabling Preallocation (Critical for Inference):</strong>\nBy default, JAX reserves ~75% of VRAM upfront. For a 160³ 3D volume on inference, this guarantees an OOM error before the model even spins up.</p>\n<pre><code>os.environ[\"XLA_PYTHON_CLIENT_PREALLOCATE\"] = \"false\"\nos.environ[\"XLA_PYTHON_CLIENT_ALLOCATOR\"] = \"platform\"\n</code></pre>\n<p><strong>3. Distributed Training Strategy:</strong>\nScaling across 32 cores requires the <code>DataParallel</code> distribution. Tuning the batch size is critical; too large wastes HBM, too small underutilizes the v4-32 hardware.</p>\n<pre><code>import jax\nimport keras\n\njax.distributed.initialize()\ndevices = keras.distribution.list_devices()\ndata_parallel = keras.distribution.DataParallel(devices=devices)\nkeras.distribution.set_distribution(data_parallel)\ndata_parallel.auto_shard_dataset = False # Essential for custom tf.data pipelines\n</code></pre>\n<p><strong>4. Flash Attention Incompatibility:</strong>\nDistributed training over multiple TPU workers currently requires disabling Flash Attention to prevent <code>XlaRuntimeError</code>.</p>\n<pre><code>keras.config.disable_flash_attention()\n</code></pre>\n<p><strong>5. Custom HBM Management:</strong>\n160³ volumes plus TransUNet push even the v4-32's generous memory boundaries. I wrote a custom Keras callback using <code>jax.local_devices()[i].memory_stats()</code> to monitor actual HBM usage dynamically during training.</p>\n<p><strong>6. Preemption Resilience:</strong>\nKaggle TPU sessions get preempted frequently. An auto-resume script that saves checkpoint weights and CSV logs per-epoch is mandatory, not optional.</p>\n<p>---### Competition Strategy</p>\n<ul>\n<li><strong>~60 submissions</strong> — aggressive but deliberate experimentation</li>\n<li><strong>Fast iteration</strong> — TPU speed enabled multiple experiments per day</li>\n<li><strong>Post-processing matters</strong> — the inference pipeline is as important as the model itself</li>\n<li><strong>Loss-metric alignment</strong> — designing loss functions that mirror competition metrics (Surface Dice + TopoScore → Tversky + clDice) is crucial</li>\n</ul>\n<h3>General Insights</h3>\n<ul>\n<li>In 3D segmentation, <strong>input resolution</strong> often matters more than model complexity</li>\n<li><strong>Ensemble techniques</strong> have low cost but high return</li>\n<li><strong>Weight loading issues</strong> across library versions are common — always save version metadata alongside checkpoints</li>\n</ul>\n<hr>\n<h2>Tech Stack</h2>\n<table>\n<thead>\n<tr>\n<th>Category</th>\n<th>Tool</th>\n</tr>\n</thead>\n<tbody>\n<tr>\n<td>Framework</td>\n<td>Keras 3 + JAX</td>\n</tr>\n<tr>\n<td>Hardware</td>\n<td>Kaggle TPU v4-32</td>\n</tr>\n<tr>\n<td>Model library</td>\n<td>medicai (TransUNet, losses, metrics, transforms)</td>\n</tr>\n<tr>\n<td>Data pipeline</td>\n<td>TFRecords + tf.data</td>\n</tr>\n<tr>\n<td>3D transforms</td>\n<td>medicai.transforms (MONAI-inspired)</td>\n</tr>\n<tr>\n<td>Inference</td>\n<td>SlidingWindowInference + TTA</td>\n</tr>\n<tr>\n<td>Post-processing</td>\n<td>scipy.ndimage</td>\n</tr>\n</tbody>\n</table>\n<hr>\n<h2>Conclusion</h2>\n<p>I will definitely focus more on reviewing open-source projects first, and I will continue to give my best, even in difficult situations.\nLastly, I would really appreciate it if you could leave your advice in the comments.</p>",
  "messages": [
    {
      "id": "3416203",
      "postDate": "03/02/2026 09:47:44",
      "content": "<h1>123rd place — not where I aimed, but here's what I learned</h1>\n<ul>\n<li><strong>Competition:</strong> <a href=\"https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection\" target=\"_blank\">Vesuvius Challenge Surface Detection</a></li>\n<li><strong>Result:</strong> 123rd / 1427 · Score: 0.579 · 🥉 Bronze Medal</li>\n<li><strong>Training code:</strong> <a href=\"https://www.kaggle.com/code/osman0/vesuvius-surface-train-tpu\" target=\"_blank\">kaggle.com/code/osman0/vesuvius-surface-train-tpu</a></li>\n</ul>\n<hr>\n<h2>Personal Reflection</h2>\n<p>I'll be honest.  I didn't do as well as I had hoped. I was aiming much higher like top 10 But the result didn't reflect that effort, and I think I know whys.</p>\n<p>I'll be real about where I was coming from: I'm relatively new to competitive machine learning, TPUs, and JAX all at once. There was a steep learning curve just to get anything running, let alone running well.</p>\n<p>I won't pretend the motivation was purely academic either. Top 10 meant prize money, and I genuinely needed it. That drove me to put in a lot of hours — probably more than I should have at certain points. I tried hard. I just didn't have the experience yet to convert that effort into a top result.</p>\n<p>And that's okay. I learned an enormous amount. I went from knowing nothing about distributed TPU training to having a working end-to-end 3D segmentation pipeline. That's real progress, even if the leaderboard didn't fully show it. Next time, I'll come in with that foundation already in place.</p>\n<p>A big thank you to <a href=\"https://www.kaggle.com/ipythonx\" target=\"_blank\">Innat</a> — his <a href=\"https://www.kaggle.com/code/ipythonx/train-vesuvius-surface-3d-detection-on-tpu\" target=\"_blank\">notebook</a> was what got me started with the <code>medicai</code> framework and TPU-compatible 3D training.</p>\n<hr>\n<h2>Acknowledgements</h2>\n<p>This training was done as part of the <strong>TRC (TPU Research Cloud) program</strong>. I'm genuinely grateful to Google for providing access to TPU v4-32 hardware through TRC training 160³ 3D models at this scale would have been completely infeasible otherwise. Thank you.</p>\n<hr>\n<h2>TL;DR</h2>\n<p>I built a 3D segmentation pipeline to detect ancient papyrus surfaces from CT scans of the Herculaneum scrolls. My final solution uses <strong>TransUNet</strong> <a href=\"#references\">[1]</a> with a <strong>SEResNeXt</strong> <a href=\"#references\">[2]</a> encoder, trained on <strong>160³ voxel</strong> patches with a <strong>Tversky</strong> <a href=\"#references\">[3]</a> <strong>+ CenterlineDice</strong> <a href=\"#references\">[4]</a> loss on <strong>TPU v4-32</strong>. Inference combines <strong>sliding window</strong>, <strong>6-fold TTA</strong>, <strong>dual checkpoint ensemble</strong>, and <strong>morphological post-processing</strong>. My biggest personal takeaway was learning to train 3D medical segmentation models end-to-end on TPUs with JAX/Keras.</p>\n<hr>\n<h2>Solution Architecture</h2>\n<p><img src=\"https://www.googleapis.com/download/storage/v1/b/kaggle-user-content/o/inbox%2F20259944%2F27dc43a53e40e5e9ff69b8ab1eff4495%2FScreenshot%202026-03-02%20at%2012.37.40.png?generation=1772444684940610&amp;alt=media\" alt=\"/Users/osman/Desktop/Screenshot 2026-03-02 at 12.37.40.png\"></p>\n<p><em>End-to-end pipeline: 3D data prep → TransUNet training → dual checkpoint blending → post-processing → binary ink mask.</em></p>\n<h3>Model: TransUNet + SEResNeXt</h3>\n<pre><code>Input (160³×1) → SEResNeXt CNN Encoder → Transformer Bottleneck → UNet Decoder → Softmax (3 classes)\n</code></pre>\n<p><strong>Why this architecture?</strong></p>\n<ul>\n<li><strong>160³ input</strong> — 4.6× larger receptive field than 96³, critical for thin papyrus surfaces spanning hundreds of voxels</li>\n<li><strong>SEResNeXt encoder</strong> <a href=\"#references\">[2]</a> — squeeze-excitation blocks provide channel-wise attention, proven on 3D CT data</li>\n<li><strong>Hybrid CNN + Transformer</strong> <a href=\"#references\">[1]</a> — CNN encoder captures local features, transformer bottleneck captures global context</li>\n<li><strong>HBM-efficient</strong> — CNN encoder is lighter on TPU HBM than Swin attention maps at 160³ resolution</li>\n</ul>\n<p>I experimented with two encoder sizes:</p>\n<table>\n<thead>\n<tr>\n<th>Variant</th>\n<th>Training</th>\n<th>Inference</th>\n</tr>\n</thead>\n<tbody>\n<tr>\n<td><code>seresnext50</code></td>\n<td>✅ (train-tpu.py)</td>\n<td>—</td>\n</tr>\n<tr>\n<td><code>seresnext101</code></td>\n<td>✅ (final model)</td>\n<td>✅ (inference-kaggle6.py)</td>\n</tr>\n</tbody>\n</table>\n<hr>\n<h2>Training Details</h2>\n<table>\n<thead>\n<tr>\n<th>Parameter</th>\n<th>Value</th>\n</tr>\n</thead>\n<tbody>\n<tr>\n<td><strong>Hardware</strong></td>\n<td>Kaggle TPU v4-32 (32 cores)</td>\n</tr>\n<tr>\n<td><strong>Framework</strong></td>\n<td>Keras 3 + JAX backend</td>\n</tr>\n<tr>\n<td><strong>Batch size</strong></td>\n<td><code>total_devices // 2</code> (16 on TPU v4-32)</td>\n</tr>\n<tr>\n<td><strong>Epochs</strong></td>\n<td>800 (CosineDecay + EarlyStopping patience=60)</td>\n</tr>\n<tr>\n<td><strong>Base LR</strong></td>\n<td>2e-4 with sqrt batch scaling</td>\n</tr>\n<tr>\n<td><strong>LR Schedule</strong></td>\n<td>Single-cycle CosineDecay, 5-epoch warmup, alpha=0.005</td>\n</tr>\n<tr>\n<td><strong>Optimizer</strong></td>\n<td>AdamW (weight_decay=1e-5, clipnorm=1.0)</td>\n</tr>\n<tr>\n<td><strong>Loss</strong></td>\n<td>SparseTverskyLoss + 0.2 × SparseCenterlineDiceLoss</td>\n</tr>\n<tr>\n<td><strong>Val split</strong></td>\n<td>5 TFRecord shards (~30 samples), seed=42</td>\n</tr>\n</tbody>\n</table>\n<h3>Loss Function</h3>\n<pre><code># Tversky loss — recall-focused (alpha=0.3 FP, beta=0.7 FN)\ntversky = SparseTverskyLoss(alpha=0.3, beta=0.7, ignore_class_ids=2)\n\n# CenterlineDice — topology preservation via soft-skeletonisation\ncldice = SparseCenterlineDiceLoss(target_class_ids=1, ignore_class_ids=2, iters=20)\n\n# Combined\nloss = tversky + 0.2 * cldice\n</code></pre>\n<p><strong>Rationale:</strong> Tversky with β=0.7 penalizes missed surfaces (FN) more than false positives <a href=\"#references\">[3]</a>, aligning with the competition metric's sensitivity to gaps. CenterlineDice <a href=\"#references\">[4]</a> preserves topological continuity via soft-skeletonisation on predictions — directly targeting the TopoScore sub-metric. The <code>ignore_class_ids=2</code> automatically masks out unlabeled regions.</p>\n<h3>Data Pipeline &amp; Augmentation</h3>\n<p>Data is stored as TFRecords (780 samples). Training augmentation includes:</p>\n<ul>\n<li><strong>RandSpatialCrop</strong> (160³, <code>min_valid_ratio=0.3</code> — rejects patches &gt;70% ignore-class)</li>\n<li><strong>RandFlip</strong> (all 3 axes, prob=0.5)</li>\n<li><strong>RandRotate90</strong> (prob=0.4) + <strong>RandRotate</strong> (factor=0.2, prob=0.25)</li>\n<li><strong>NormalizeIntensity</strong> (<code>nonzero=True</code>) — adapts to each volume's actual intensity range rather than fixed [0,255]→[0,1] scaling</li>\n<li><strong>RandShiftIntensity</strong> (offsets=0.10, prob=0.5)</li>\n<li><strong>RandCutOut</strong> (volume mode, prob=0.2) — occlusion regularisation</li>\n</ul>\n<h3>LR Schedule Evolution</h3>\n<p>Earlier experiments used <code>CosineDecayRestarts(m_mul=0.9)</code>, but the learning rate collapsed too early (~epoch 300). Switching to a <strong>single-cycle CosineDecay</strong> over the full 800-epoch budget produced much more stable convergence.</p>\n<h3>TPU-Specific Implementation</h3>\n<ul>\n<li><strong>DataParallel</strong> distribution across 32 TPU v4 cores with auto-sharding</li>\n<li><strong>Flash attention disabled</strong> — required for distributed training compatibility</li>\n<li><strong>Custom HBM Memory Monitor callback</strong> — logs TPU HBM usage and host RAM every N epochs</li>\n<li><strong>Auto-resume from preemption</strong> — checkpoint + CSV log continuation for Kaggle TPU interrupts</li>\n</ul>\n<hr>\n<h2>Inference Pipeline</h2>\n<h3>Sliding Window Inference</h3>\n<pre><code>SlidingWindowInference(\n    roi_size=(160, 160, 160),  # must match ViT positional embeddings (5³=125 patches)\n    overlap=0.5,\n    blend_mode=\"gaussian\",     # smooth boundary artifacts\n)\n</code></pre>\n<h3>Test-Time Augmentation (6-fold)</h3>\n<pre><code>TTA_AXES = [(), (0,), (1,), (2,), (0,1), (1,2)]\n# Original + 3 single-axis flips + 2 pair flips → averaged\n</code></pre>\n<h3>Dual Checkpoint Blend</h3>\n<pre><code>pred = 0.80 * pred_best_val_dice + 0.20 * pred_best_train_loss\n</code></pre>\n<p>Two checkpoints from the same training run — <code>best_val_dice</code> generalizes better, <code>best_train_loss</code> is more expressive. Blending consistently improved over either single checkpoint.</p>\n<h3>Post-Processing</h3>\n<ol>\n<li><strong>Ignore-class suppression</strong>: <code>ink_prob -= 0.15 × ignore_prob</code> — leverages class-2 predictions to reduce false positives in uncertain regions</li>\n<li><strong>Thresholding</strong>: prob ≥ 0.5 → binary mask</li>\n<li><strong>Morphological closing</strong> (radius=1) — fills small gaps</li>\n<li><strong>Connected component filtering</strong> (min_size=50) — removes noise</li>\n</ol>\n<hr>\n<h2>Learnings</h2>\n<h3>TPU for 3D Medical Segmentation</h3>\n<p>This competition was my first serious experience training on TPUs. It required several specific codebase adjustments to run successfully:</p>\n<p><strong>1. JAX backend over TensorFlow:</strong>\nJAX is significantly more stable and faster on TPUs than TF eager mode.</p>\n<pre><code>import os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n</code></pre>\n<p><strong>2. Disabling Preallocation (Critical for Inference):</strong>\nBy default, JAX reserves ~75% of VRAM upfront. For a 160³ 3D volume on inference, this guarantees an OOM error before the model even spins up.</p>\n<pre><code>os.environ[\"XLA_PYTHON_CLIENT_PREALLOCATE\"] = \"false\"\nos.environ[\"XLA_PYTHON_CLIENT_ALLOCATOR\"] = \"platform\"\n</code></pre>\n<p><strong>3. Distributed Training Strategy:</strong>\nScaling across 32 cores requires the <code>DataParallel</code> distribution. Tuning the batch size is critical; too large wastes HBM, too small underutilizes the v4-32 hardware.</p>\n<pre><code>import jax\nimport keras\n\njax.distributed.initialize()\ndevices = keras.distribution.list_devices()\ndata_parallel = keras.distribution.DataParallel(devices=devices)\nkeras.distribution.set_distribution(data_parallel)\ndata_parallel.auto_shard_dataset = False # Essential for custom tf.data pipelines\n</code></pre>\n<p><strong>4. Flash Attention Incompatibility:</strong>\nDistributed training over multiple TPU workers currently requires disabling Flash Attention to prevent <code>XlaRuntimeError</code>.</p>\n<pre><code>keras.config.disable_flash_attention()\n</code></pre>\n<p><strong>5. Custom HBM Management:</strong>\n160³ volumes plus TransUNet push even the v4-32's generous memory boundaries. I wrote a custom Keras callback using <code>jax.local_devices()[i].memory_stats()</code> to monitor actual HBM usage dynamically during training.</p>\n<p><strong>6. Preemption Resilience:</strong>\nKaggle TPU sessions get preempted frequently. An auto-resume script that saves checkpoint weights and CSV logs per-epoch is mandatory, not optional.</p>\n<p>---### Competition Strategy</p>\n<ul>\n<li><strong>~60 submissions</strong> — aggressive but deliberate experimentation</li>\n<li><strong>Fast iteration</strong> — TPU speed enabled multiple experiments per day</li>\n<li><strong>Post-processing matters</strong> — the inference pipeline is as important as the model itself</li>\n<li><strong>Loss-metric alignment</strong> — designing loss functions that mirror competition metrics (Surface Dice + TopoScore → Tversky + clDice) is crucial</li>\n</ul>\n<h3>General Insights</h3>\n<ul>\n<li>In 3D segmentation, <strong>input resolution</strong> often matters more than model complexity</li>\n<li><strong>Ensemble techniques</strong> have low cost but high return</li>\n<li><strong>Weight loading issues</strong> across library versions are common — always save version metadata alongside checkpoints</li>\n</ul>\n<hr>\n<h2>Tech Stack</h2>\n<table>\n<thead>\n<tr>\n<th>Category</th>\n<th>Tool</th>\n</tr>\n</thead>\n<tbody>\n<tr>\n<td>Framework</td>\n<td>Keras 3 + JAX</td>\n</tr>\n<tr>\n<td>Hardware</td>\n<td>Kaggle TPU v4-32</td>\n</tr>\n<tr>\n<td>Model library</td>\n<td>medicai (TransUNet, losses, metrics, transforms)</td>\n</tr>\n<tr>\n<td>Data pipeline</td>\n<td>TFRecords + tf.data</td>\n</tr>\n<tr>\n<td>3D transforms</td>\n<td>medicai.transforms (MONAI-inspired)</td>\n</tr>\n<tr>\n<td>Inference</td>\n<td>SlidingWindowInference + TTA</td>\n</tr>\n<tr>\n<td>Post-processing</td>\n<td>scipy.ndimage</td>\n</tr>\n</tbody>\n</table>\n<hr>\n<h2>Conclusion</h2>\n<p>I will definitely focus more on reviewing open-source projects first, and I will continue to give my best, even in difficult situations.\nLastly, I would really appreciate it if you could leave your advice in the comments.</p>",
      "rawMarkdown": "# 123rd place — not where I aimed, but here's what I learned\n\n- **Competition:** [Vesuvius Challenge Surface Detection](https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection)\n- **Result:** 123rd / 1427 · Score: 0.579 · 🥉 Bronze Medal\n- **Training code:** [kaggle.com/code/osman0/vesuvius-surface-train-tpu](https://www.kaggle.com/code/osman0/vesuvius-surface-train-tpu)\n\n---\n\n## Personal Reflection\n\nI'll be honest.  I didn't do as well as I had hoped. I was aiming much higher like top 10 But the result didn't reflect that effort, and I think I know whys.\n\nI'll be real about where I was coming from: I'm relatively new to competitive machine learning, TPUs, and JAX all at once. There was a steep learning curve just to get anything running, let alone running well.\n\nI won't pretend the motivation was purely academic either. Top 10 meant prize money, and I genuinely needed it. That drove me to put in a lot of hours — probably more than I should have at certain points. I tried hard. I just didn't have the experience yet to convert that effort into a top result.\n\nAnd that's okay. I learned an enormous amount. I went from knowing nothing about distributed TPU training to having a working end-to-end 3D segmentation pipeline. That's real progress, even if the leaderboard didn't fully show it. Next time, I'll come in with that foundation already in place.\n\nA big thank you to [Innat](https://www.kaggle.com/ipythonx) — his [notebook](https://www.kaggle.com/code/ipythonx/train-vesuvius-surface-3d-detection-on-tpu) was what got me started with the `medicai` framework and TPU-compatible 3D training.\n\n---\n\n## Acknowledgements\n\nThis training was done as part of the **TRC (TPU Research Cloud) program**. I'm genuinely grateful to Google for providing access to TPU v4-32 hardware through TRC training 160³ 3D models at this scale would have been completely infeasible otherwise. Thank you.\n\n---\n\n## TL;DR\n\nI built a 3D segmentation pipeline to detect ancient papyrus surfaces from CT scans of the Herculaneum scrolls. My final solution uses **TransUNet** [[1]](#references) with a **SEResNeXt** [[2]](#references) encoder, trained on **160³ voxel** patches with a **Tversky** [[3]](#references) **+ CenterlineDice** [[4]](#references) loss on **TPU v4-32**. Inference combines **sliding window**, **6-fold TTA**, **dual checkpoint ensemble**, and **morphological post-processing**. My biggest personal takeaway was learning to train 3D medical segmentation models end-to-end on TPUs with JAX/Keras.\n\n---\n\n## Solution Architecture\n\n![/Users/osman/Desktop/Screenshot 2026-03-02 at 12.37.40.png](https://www.googleapis.com/download/storage/v1/b/kaggle-user-content/o/inbox%2F20259944%2F27dc43a53e40e5e9ff69b8ab1eff4495%2FScreenshot%202026-03-02%20at%2012.37.40.png?generation=1772444684940610&alt=media)\n\n*End-to-end pipeline: 3D data prep → TransUNet training → dual checkpoint blending → post-processing → binary ink mask.*\n\n### Model: TransUNet + SEResNeXt\n\n```\nInput (160³×1) → SEResNeXt CNN Encoder → Transformer Bottleneck → UNet Decoder → Softmax (3 classes)\n```\n\n**Why this architecture?**\n- **160³ input** — 4.6× larger receptive field than 96³, critical for thin papyrus surfaces spanning hundreds of voxels\n- **SEResNeXt encoder** [[2]](#references) — squeeze-excitation blocks provide channel-wise attention, proven on 3D CT data\n- **Hybrid CNN + Transformer** [[1]](#references) — CNN encoder captures local features, transformer bottleneck captures global context\n- **HBM-efficient** — CNN encoder is lighter on TPU HBM than Swin attention maps at 160³ resolution\n\nI experimented with two encoder sizes:\n\n| Variant | Training | Inference |\n|---------|----------|-----------|\n| `seresnext50` | ✅ (train-tpu.py) | — |\n| `seresnext101` | ✅ (final model) | ✅ (inference-kaggle6.py) |\n\n---\n\n## Training Details\n\n| Parameter | Value |\n|-----------|-------|\n| **Hardware** | Kaggle TPU v4-32 (32 cores) |\n| **Framework** | Keras 3 + JAX backend |\n| **Batch size** | `total_devices // 2` (16 on TPU v4-32) |\n| **Epochs** | 800 (CosineDecay + EarlyStopping patience=60) |\n| **Base LR** | 2e-4 with sqrt batch scaling |\n| **LR Schedule** | Single-cycle CosineDecay, 5-epoch warmup, alpha=0.005 |\n| **Optimizer** | AdamW (weight_decay=1e-5, clipnorm=1.0) |\n| **Loss** | SparseTverskyLoss + 0.2 × SparseCenterlineDiceLoss |\n| **Val split** | 5 TFRecord shards (~30 samples), seed=42 |\n\n### Loss Function\n\n```python\n# Tversky loss — recall-focused (alpha=0.3 FP, beta=0.7 FN)\ntversky = SparseTverskyLoss(alpha=0.3, beta=0.7, ignore_class_ids=2)\n\n# CenterlineDice — topology preservation via soft-skeletonisation\ncldice = SparseCenterlineDiceLoss(target_class_ids=1, ignore_class_ids=2, iters=20)\n\n# Combined\nloss = tversky + 0.2 * cldice\n```\n\n**Rationale:** Tversky with β=0.7 penalizes missed surfaces (FN) more than false positives [[3]](#references), aligning with the competition metric's sensitivity to gaps. CenterlineDice [[4]](#references) preserves topological continuity via soft-skeletonisation on predictions — directly targeting the TopoScore sub-metric. The `ignore_class_ids=2` automatically masks out unlabeled regions.\n\n### Data Pipeline & Augmentation\n\nData is stored as TFRecords (780 samples). Training augmentation includes:\n\n- **RandSpatialCrop** (160³, `min_valid_ratio=0.3` — rejects patches >70% ignore-class)\n- **RandFlip** (all 3 axes, prob=0.5)\n- **RandRotate90** (prob=0.4) + **RandRotate** (factor=0.2, prob=0.25)\n- **NormalizeIntensity** (`nonzero=True`) — adapts to each volume's actual intensity range rather than fixed [0,255]→[0,1] scaling\n- **RandShiftIntensity** (offsets=0.10, prob=0.5)\n- **RandCutOut** (volume mode, prob=0.2) — occlusion regularisation\n\n### LR Schedule Evolution\n\nEarlier experiments used `CosineDecayRestarts(m_mul=0.9)`, but the learning rate collapsed too early (~epoch 300). Switching to a **single-cycle CosineDecay** over the full 800-epoch budget produced much more stable convergence.\n\n### TPU-Specific Implementation\n\n- **DataParallel** distribution across 32 TPU v4 cores with auto-sharding\n- **Flash attention disabled** — required for distributed training compatibility\n- **Custom HBM Memory Monitor callback** — logs TPU HBM usage and host RAM every N epochs\n- **Auto-resume from preemption** — checkpoint + CSV log continuation for Kaggle TPU interrupts\n\n---\n\n## Inference Pipeline\n\n### Sliding Window Inference\n```python\nSlidingWindowInference(\n    roi_size=(160, 160, 160),  # must match ViT positional embeddings (5³=125 patches)\n    overlap=0.5,\n    blend_mode=\"gaussian\",     # smooth boundary artifacts\n)\n```\n\n### Test-Time Augmentation (6-fold)\n```python\nTTA_AXES = [(), (0,), (1,), (2,), (0,1), (1,2)]\n# Original + 3 single-axis flips + 2 pair flips → averaged\n```\n\n### Dual Checkpoint Blend\n```python\npred = 0.80 * pred_best_val_dice + 0.20 * pred_best_train_loss\n```\n\nTwo checkpoints from the same training run — `best_val_dice` generalizes better, `best_train_loss` is more expressive. Blending consistently improved over either single checkpoint.\n\n### Post-Processing\n1. **Ignore-class suppression**: `ink_prob -= 0.15 × ignore_prob` — leverages class-2 predictions to reduce false positives in uncertain regions\n2. **Thresholding**: prob ≥ 0.5 → binary mask\n3. **Morphological closing** (radius=1) — fills small gaps\n4. **Connected component filtering** (min_size=50) — removes noise\n\n---\n## Learnings\n\n### TPU for 3D Medical Segmentation\n\nThis competition was my first serious experience training on TPUs. It required several specific codebase adjustments to run successfully:\n\n**1. JAX backend over TensorFlow:**\nJAX is significantly more stable and faster on TPUs than TF eager mode.\n```python\nimport os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n```\n\n**2. Disabling Preallocation (Critical for Inference):**\nBy default, JAX reserves ~75% of VRAM upfront. For a 160³ 3D volume on inference, this guarantees an OOM error before the model even spins up.\n```python\nos.environ[\"XLA_PYTHON_CLIENT_PREALLOCATE\"] = \"false\"\nos.environ[\"XLA_PYTHON_CLIENT_ALLOCATOR\"] = \"platform\"\n```\n\n**3. Distributed Training Strategy:**\nScaling across 32 cores requires the `DataParallel` distribution. Tuning the batch size is critical; too large wastes HBM, too small underutilizes the v4-32 hardware.\n```python\nimport jax\nimport keras\n\njax.distributed.initialize()\ndevices = keras.distribution.list_devices()\ndata_parallel = keras.distribution.DataParallel(devices=devices)\nkeras.distribution.set_distribution(data_parallel)\ndata_parallel.auto_shard_dataset = False # Essential for custom tf.data pipelines\n```\n\n**4. Flash Attention Incompatibility:**\nDistributed training over multiple TPU workers currently requires disabling Flash Attention to prevent `XlaRuntimeError`.\n```python\nkeras.config.disable_flash_attention()\n```\n\n**5. Custom HBM Management:**\n160³ volumes plus TransUNet push even the v4-32's generous memory boundaries. I wrote a custom Keras callback using `jax.local_devices()[i].memory_stats()` to monitor actual HBM usage dynamically during training.\n\n**6. Preemption Resilience:**\nKaggle TPU sessions get preempted frequently. An auto-resume script that saves checkpoint weights and CSV logs per-epoch is mandatory, not optional.\n\n---### Competition Strategy\n- **~60 submissions** — aggressive but deliberate experimentation\n- **Fast iteration** — TPU speed enabled multiple experiments per day\n- **Post-processing matters** — the inference pipeline is as important as the model itself\n- **Loss-metric alignment** — designing loss functions that mirror competition metrics (Surface Dice + TopoScore → Tversky + clDice) is crucial\n\n### General Insights\n- In 3D segmentation, **input resolution** often matters more than model complexity\n- **Ensemble techniques** have low cost but high return\n- **Weight loading issues** across library versions are common — always save version metadata alongside checkpoints\n\n---\n## Tech Stack\n\n| Category | Tool |\n|----------|------|\n| Framework | Keras 3 + JAX |\n| Hardware | Kaggle TPU v4-32 |\n| Model library | medicai (TransUNet, losses, metrics, transforms) |\n| Data pipeline | TFRecords + tf.data |\n| 3D transforms | medicai.transforms (MONAI-inspired) |\n| Inference | SlidingWindowInference + TTA |\n| Post-processing | scipy.ndimage |\n\n---\n## Conclusion\n\nI will definitely focus more on reviewing open-source projects first, and I will continue to give my best, even in difficult situations.\nLastly, I would really appreciate it if you could leave your advice in the comments.",
      "votes": null
    },
    {
      "id": "3416232",
      "postDate": "03/02/2026 11:03:18",
      "content": "<p><a href=\"https://www.kaggle.com/osman0\" target=\"_blank\">@osman0</a> Great work. </p>\n<p>Throughout this competition, we identified several missing components that could strengthen <code>medicai</code>, and we managed to incorporate some of them within the limited timeframe.</p>\n<p>We also gained a deeper appreciation for the generalization strength of the <code>nnUNet</code> framework. Whether we like it or not, <code>nnUNet</code> clearly outperformed others on the leaderboard. I’m planning to integrate this framework into <code>medicai</code> through a mid- to high-level API to make it easier to use (<a href=\"https://github.com/innat/medic-ai/issues/30\" target=\"_blank\">ticket</a>). However, this won’t be a straightforward task. We need to ensure a proper and validated translation of <code>nnUNet</code>’s design into <code>medicai</code>, while maintaining overall API consistency.</p>\n<p>Now, you've understanding about segmentation, I would suggest to explote <a href=\"https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/5th-place-solution\" target=\"_blank\">5th place</a> solution, quite interesting and solid kaggle thing. </p>\n<p>Another valuable lesson from this competition was learning how to effectively interact with LLMs. Many participants found them incredibly helpful, and understanding how to communicate with them properly made a significant difference; <a href=\"https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/bronze-medal-chatgpt-vibe-coding\" target=\"_blank\">discussion</a>.</p>",
      "rawMarkdown": "osman0 Great work. \n\nThroughout this competition, we identified several missing components that could strengthen `medicai`, and we managed to incorporate some of them within the limited timeframe.\n\nWe also gained a deeper appreciation for the generalization strength of the `nnUNet` framework. Whether we like it or not, `nnUNet` clearly outperformed others on the leaderboard. I’m planning to integrate this framework into `medicai` through a mid- to high-level API to make it easier to use ([ticket](https://github.com/innat/medic-ai/issues/30)). However, this won’t be a straightforward task. We need to ensure a proper and validated translation of `nnUNet`’s design into `medicai`, while maintaining overall API consistency.\n\nNow, you've understanding about segmentation, I would suggest to explote [5th place](https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/5th-place-solution) solution, quite interesting and solid kaggle thing. \n\nAnother valuable lesson from this competition was learning how to effectively interact with LLMs. Many participants found them incredibly helpful, and understanding how to communicate with them properly made a significant difference; [discussion](https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/bronze-medal-chatgpt-vibe-coding).",
      "votes": null
    },
    {
      "id": "3416238",
      "postDate": "03/02/2026 11:34:17",
      "content": "<p>Thank you so much for your kind words and for reaching out!</p>\n<p>I especially want to thank you because you and your work were incredibly helpful in my journey to learn about TPUs. I also had the opportunity to try out some of the models from your medic-ai repository, and it was a truly amazing experience for me.</p>\n<p>I really appreciate your suggestion about the 5th place solution—I will definitely take a close look at it. and olso def I look for your nnUNet framework. </p>\n<p>Thanks again.  </p>",
      "rawMarkdown": "Thank you so much for your kind words and for reaching out!\n\nI especially want to thank you because you and your work were incredibly helpful in my journey to learn about TPUs. I also had the opportunity to try out some of the models from your medic-ai repository, and it was a truly amazing experience for me.\n\nI really appreciate your suggestion about the 5th place solution—I will definitely take a close look at it. and olso def I look for your nnUNet framework. \n\nThanks again.",
      "votes": null
    }
  ],
  "comments": [
    {
      "id": 3416232,
      "author_name": "ipythonx",
      "author_url": "",
      "post_date": "03/02/2026 11:03:18",
      "content": "<p><a href=\"https://www.kaggle.com/osman0\" target=\"_blank\">@osman0</a> Great work. </p>\n<p>Throughout this competition, we identified several missing components that could strengthen <code>medicai</code>, and we managed to incorporate some of them within the limited timeframe.</p>\n<p>We also gained a deeper appreciation for the generalization strength of the <code>nnUNet</code> framework. Whether we like it or not, <code>nnUNet</code> clearly outperformed others on the leaderboard. I’m planning to integrate this framework into <code>medicai</code> through a mid- to high-level API to make it easier to use (<a href=\"https://github.com/innat/medic-ai/issues/30\" target=\"_blank\">ticket</a>). However, this won’t be a straightforward task. We need to ensure a proper and validated translation of <code>nnUNet</code>’s design into <code>medicai</code>, while maintaining overall API consistency.</p>\n<p>Now, you've understanding about segmentation, I would suggest to explote <a href=\"https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/5th-place-solution\" target=\"_blank\">5th place</a> solution, quite interesting and solid kaggle thing. </p>\n<p>Another valuable lesson from this competition was learning how to effectively interact with LLMs. Many participants found them incredibly helpful, and understanding how to communicate with them properly made a significant difference; <a href=\"https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/bronze-medal-chatgpt-vibe-coding\" target=\"_blank\">discussion</a>.</p>",
      "votes": null,
      "replies": [
        {
          "id": 3416238,
          "author_name": "osman0",
          "author_url": "",
          "post_date": "03/02/2026 11:34:17",
          "content": "<p>Thank you so much for your kind words and for reaching out!</p>\n<p>I especially want to thank you because you and your work were incredibly helpful in my journey to learn about TPUs. I also had the opportunity to try out some of the models from your medic-ai repository, and it was a truly amazing experience for me.</p>\n<p>I really appreciate your suggestion about the 5th place solution—I will definitely take a close look at it. and olso def I look for your nnUNet framework. </p>\n<p>Thanks again.  </p>",
          "votes": null,
          "replies": []
        }
      ]
    }
  ],
  "raw_markdown_by_id": {
    "3416203": "# 123rd place — not where I aimed, but here's what I learned\n\n- **Competition:** [Vesuvius Challenge Surface Detection](https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection)\n- **Result:** 123rd / 1427 · Score: 0.579 · 🥉 Bronze Medal\n- **Training code:** [kaggle.com/code/osman0/vesuvius-surface-train-tpu](https://www.kaggle.com/code/osman0/vesuvius-surface-train-tpu)\n\n---\n\n## Personal Reflection\n\nI'll be honest.  I didn't do as well as I had hoped. I was aiming much higher like top 10 But the result didn't reflect that effort, and I think I know whys.\n\nI'll be real about where I was coming from: I'm relatively new to competitive machine learning, TPUs, and JAX all at once. There was a steep learning curve just to get anything running, let alone running well.\n\nI won't pretend the motivation was purely academic either. Top 10 meant prize money, and I genuinely needed it. That drove me to put in a lot of hours — probably more than I should have at certain points. I tried hard. I just didn't have the experience yet to convert that effort into a top result.\n\nAnd that's okay. I learned an enormous amount. I went from knowing nothing about distributed TPU training to having a working end-to-end 3D segmentation pipeline. That's real progress, even if the leaderboard didn't fully show it. Next time, I'll come in with that foundation already in place.\n\nA big thank you to [Innat](https://www.kaggle.com/ipythonx) — his [notebook](https://www.kaggle.com/code/ipythonx/train-vesuvius-surface-3d-detection-on-tpu) was what got me started with the `medicai` framework and TPU-compatible 3D training.\n\n---\n\n## Acknowledgements\n\nThis training was done as part of the **TRC (TPU Research Cloud) program**. I'm genuinely grateful to Google for providing access to TPU v4-32 hardware through TRC training 160³ 3D models at this scale would have been completely infeasible otherwise. Thank you.\n\n---\n\n## TL;DR\n\nI built a 3D segmentation pipeline to detect ancient papyrus surfaces from CT scans of the Herculaneum scrolls. My final solution uses **TransUNet** [[1]](#references) with a **SEResNeXt** [[2]](#references) encoder, trained on **160³ voxel** patches with a **Tversky** [[3]](#references) **+ CenterlineDice** [[4]](#references) loss on **TPU v4-32**. Inference combines **sliding window**, **6-fold TTA**, **dual checkpoint ensemble**, and **morphological post-processing**. My biggest personal takeaway was learning to train 3D medical segmentation models end-to-end on TPUs with JAX/Keras.\n\n---\n\n## Solution Architecture\n\n![/Users/osman/Desktop/Screenshot 2026-03-02 at 12.37.40.png](https://www.googleapis.com/download/storage/v1/b/kaggle-user-content/o/inbox%2F20259944%2F27dc43a53e40e5e9ff69b8ab1eff4495%2FScreenshot%202026-03-02%20at%2012.37.40.png?generation=1772444684940610&alt=media)\n\n*End-to-end pipeline: 3D data prep → TransUNet training → dual checkpoint blending → post-processing → binary ink mask.*\n\n### Model: TransUNet + SEResNeXt\n\n```\nInput (160³×1) → SEResNeXt CNN Encoder → Transformer Bottleneck → UNet Decoder → Softmax (3 classes)\n```\n\n**Why this architecture?**\n- **160³ input** — 4.6× larger receptive field than 96³, critical for thin papyrus surfaces spanning hundreds of voxels\n- **SEResNeXt encoder** [[2]](#references) — squeeze-excitation blocks provide channel-wise attention, proven on 3D CT data\n- **Hybrid CNN + Transformer** [[1]](#references) — CNN encoder captures local features, transformer bottleneck captures global context\n- **HBM-efficient** — CNN encoder is lighter on TPU HBM than Swin attention maps at 160³ resolution\n\nI experimented with two encoder sizes:\n\n| Variant | Training | Inference |\n|---------|----------|-----------|\n| `seresnext50` | ✅ (train-tpu.py) | — |\n| `seresnext101` | ✅ (final model) | ✅ (inference-kaggle6.py) |\n\n---\n\n## Training Details\n\n| Parameter | Value |\n|-----------|-------|\n| **Hardware** | Kaggle TPU v4-32 (32 cores) |\n| **Framework** | Keras 3 + JAX backend |\n| **Batch size** | `total_devices // 2` (16 on TPU v4-32) |\n| **Epochs** | 800 (CosineDecay + EarlyStopping patience=60) |\n| **Base LR** | 2e-4 with sqrt batch scaling |\n| **LR Schedule** | Single-cycle CosineDecay, 5-epoch warmup, alpha=0.005 |\n| **Optimizer** | AdamW (weight_decay=1e-5, clipnorm=1.0) |\n| **Loss** | SparseTverskyLoss + 0.2 × SparseCenterlineDiceLoss |\n| **Val split** | 5 TFRecord shards (~30 samples), seed=42 |\n\n### Loss Function\n\n```python\n# Tversky loss — recall-focused (alpha=0.3 FP, beta=0.7 FN)\ntversky = SparseTverskyLoss(alpha=0.3, beta=0.7, ignore_class_ids=2)\n\n# CenterlineDice — topology preservation via soft-skeletonisation\ncldice = SparseCenterlineDiceLoss(target_class_ids=1, ignore_class_ids=2, iters=20)\n\n# Combined\nloss = tversky + 0.2 * cldice\n```\n\n**Rationale:** Tversky with β=0.7 penalizes missed surfaces (FN) more than false positives [[3]](#references), aligning with the competition metric's sensitivity to gaps. CenterlineDice [[4]](#references) preserves topological continuity via soft-skeletonisation on predictions — directly targeting the TopoScore sub-metric. The `ignore_class_ids=2` automatically masks out unlabeled regions.\n\n### Data Pipeline & Augmentation\n\nData is stored as TFRecords (780 samples). Training augmentation includes:\n\n- **RandSpatialCrop** (160³, `min_valid_ratio=0.3` — rejects patches >70% ignore-class)\n- **RandFlip** (all 3 axes, prob=0.5)\n- **RandRotate90** (prob=0.4) + **RandRotate** (factor=0.2, prob=0.25)\n- **NormalizeIntensity** (`nonzero=True`) — adapts to each volume's actual intensity range rather than fixed [0,255]→[0,1] scaling\n- **RandShiftIntensity** (offsets=0.10, prob=0.5)\n- **RandCutOut** (volume mode, prob=0.2) — occlusion regularisation\n\n### LR Schedule Evolution\n\nEarlier experiments used `CosineDecayRestarts(m_mul=0.9)`, but the learning rate collapsed too early (~epoch 300). Switching to a **single-cycle CosineDecay** over the full 800-epoch budget produced much more stable convergence.\n\n### TPU-Specific Implementation\n\n- **DataParallel** distribution across 32 TPU v4 cores with auto-sharding\n- **Flash attention disabled** — required for distributed training compatibility\n- **Custom HBM Memory Monitor callback** — logs TPU HBM usage and host RAM every N epochs\n- **Auto-resume from preemption** — checkpoint + CSV log continuation for Kaggle TPU interrupts\n\n---\n\n## Inference Pipeline\n\n### Sliding Window Inference\n```python\nSlidingWindowInference(\n    roi_size=(160, 160, 160),  # must match ViT positional embeddings (5³=125 patches)\n    overlap=0.5,\n    blend_mode=\"gaussian\",     # smooth boundary artifacts\n)\n```\n\n### Test-Time Augmentation (6-fold)\n```python\nTTA_AXES = [(), (0,), (1,), (2,), (0,1), (1,2)]\n# Original + 3 single-axis flips + 2 pair flips → averaged\n```\n\n### Dual Checkpoint Blend\n```python\npred = 0.80 * pred_best_val_dice + 0.20 * pred_best_train_loss\n```\n\nTwo checkpoints from the same training run — `best_val_dice` generalizes better, `best_train_loss` is more expressive. Blending consistently improved over either single checkpoint.\n\n### Post-Processing\n1. **Ignore-class suppression**: `ink_prob -= 0.15 × ignore_prob` — leverages class-2 predictions to reduce false positives in uncertain regions\n2. **Thresholding**: prob ≥ 0.5 → binary mask\n3. **Morphological closing** (radius=1) — fills small gaps\n4. **Connected component filtering** (min_size=50) — removes noise\n\n---\n## Learnings\n\n### TPU for 3D Medical Segmentation\n\nThis competition was my first serious experience training on TPUs. It required several specific codebase adjustments to run successfully:\n\n**1. JAX backend over TensorFlow:**\nJAX is significantly more stable and faster on TPUs than TF eager mode.\n```python\nimport os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n```\n\n**2. Disabling Preallocation (Critical for Inference):**\nBy default, JAX reserves ~75% of VRAM upfront. For a 160³ 3D volume on inference, this guarantees an OOM error before the model even spins up.\n```python\nos.environ[\"XLA_PYTHON_CLIENT_PREALLOCATE\"] = \"false\"\nos.environ[\"XLA_PYTHON_CLIENT_ALLOCATOR\"] = \"platform\"\n```\n\n**3. Distributed Training Strategy:**\nScaling across 32 cores requires the `DataParallel` distribution. Tuning the batch size is critical; too large wastes HBM, too small underutilizes the v4-32 hardware.\n```python\nimport jax\nimport keras\n\njax.distributed.initialize()\ndevices = keras.distribution.list_devices()\ndata_parallel = keras.distribution.DataParallel(devices=devices)\nkeras.distribution.set_distribution(data_parallel)\ndata_parallel.auto_shard_dataset = False # Essential for custom tf.data pipelines\n```\n\n**4. Flash Attention Incompatibility:**\nDistributed training over multiple TPU workers currently requires disabling Flash Attention to prevent `XlaRuntimeError`.\n```python\nkeras.config.disable_flash_attention()\n```\n\n**5. Custom HBM Management:**\n160³ volumes plus TransUNet push even the v4-32's generous memory boundaries. I wrote a custom Keras callback using `jax.local_devices()[i].memory_stats()` to monitor actual HBM usage dynamically during training.\n\n**6. Preemption Resilience:**\nKaggle TPU sessions get preempted frequently. An auto-resume script that saves checkpoint weights and CSV logs per-epoch is mandatory, not optional.\n\n---### Competition Strategy\n- **~60 submissions** — aggressive but deliberate experimentation\n- **Fast iteration** — TPU speed enabled multiple experiments per day\n- **Post-processing matters** — the inference pipeline is as important as the model itself\n- **Loss-metric alignment** — designing loss functions that mirror competition metrics (Surface Dice + TopoScore → Tversky + clDice) is crucial\n\n### General Insights\n- In 3D segmentation, **input resolution** often matters more than model complexity\n- **Ensemble techniques** have low cost but high return\n- **Weight loading issues** across library versions are common — always save version metadata alongside checkpoints\n\n---\n## Tech Stack\n\n| Category | Tool |\n|----------|------|\n| Framework | Keras 3 + JAX |\n| Hardware | Kaggle TPU v4-32 |\n| Model library | medicai (TransUNet, losses, metrics, transforms) |\n| Data pipeline | TFRecords + tf.data |\n| 3D transforms | medicai.transforms (MONAI-inspired) |\n| Inference | SlidingWindowInference + TTA |\n| Post-processing | scipy.ndimage |\n\n---\n## Conclusion\n\nI will definitely focus more on reviewing open-source projects first, and I will continue to give my best, even in difficult situations.\nLastly, I would really appreciate it if you could leave your advice in the comments.",
    "3416232": "osman0 Great work. \n\nThroughout this competition, we identified several missing components that could strengthen `medicai`, and we managed to incorporate some of them within the limited timeframe.\n\nWe also gained a deeper appreciation for the generalization strength of the `nnUNet` framework. Whether we like it or not, `nnUNet` clearly outperformed others on the leaderboard. I’m planning to integrate this framework into `medicai` through a mid- to high-level API to make it easier to use ([ticket](https://github.com/innat/medic-ai/issues/30)). However, this won’t be a straightforward task. We need to ensure a proper and validated translation of `nnUNet`’s design into `medicai`, while maintaining overall API consistency.\n\nNow, you've understanding about segmentation, I would suggest to explote [5th place](https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/5th-place-solution) solution, quite interesting and solid kaggle thing. \n\nAnother valuable lesson from this competition was learning how to effectively interact with LLMs. Many participants found them incredibly helpful, and understanding how to communicate with them properly made a significant difference; [discussion](https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/bronze-medal-chatgpt-vibe-coding).",
    "3416238": "Thank you so much for your kind words and for reaching out!\n\nI especially want to thank you because you and your work were incredibly helpful in my journey to learn about TPUs. I also had the opportunity to try out some of the models from your medic-ai repository, and it was a truly amazing experience for me.\n\nI really appreciate your suggestion about the 5th place solution—I will definitely take a close look at it. and olso def I look for your nnUNet framework. \n\nThanks again."
  },
  "source": "meta"
}