{"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":14068685,"sourceType":"datasetVersion","datasetId":8955007},{"sourceId":14137536,"sourceType":"datasetVersion","datasetId":9009083},{"sourceId":284794863,"sourceType":"kernelVersion"},{"sourceId":681152,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":516822,"modelId":510647},{"sourceId":681333,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":516996,"modelId":531656},{"sourceId":689100,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":522336,"modelId":536361}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !mkdir whls\n\n# !pip download keras-nightly --dest whls\n# !pip download tifffile imagecodecs --dest whls\n\n# # medic-ai (GitHub)\n# !git clone https://github.com/innat/medic-ai.git\n# !pip install build -q\n# !cd medic-ai && python -m build\n\n# # copy wheel to whls folder\n# !cp medic-ai/dist/*.whl whls/\n# !rm -r /kaggle/working/medic-ai","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install --no-index --find-links=/kaggle/input/xyz-installer/whls \\\n#     keras-nightly \\\n#     tifffile \\\n#     imagecodecs \\\n#     medicai\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T16:59:11.932413Z","iopub.execute_input":"2025-12-16T16:59:11.932603Z","iopub.status.idle":"2025-12-16T16:59:18.624942Z","shell.execute_reply.started":"2025-12-16T16:59:11.932586Z","shell.execute_reply":"2025-12-16T16:59:18.623470Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import medicai\n# import keras\n# import tifffile\n# import imagecodecs\n# import wrapt\n\n# print(\"medicai OK\")\n# print(\"keras:\", keras.__version__)\n# print(\"wrapt:\", wrapt.__version__)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T16:59:18.626565Z","iopub.execute_input":"2025-12-16T16:59:18.626929Z","iopub.status.idle":"2025-12-16T16:59:36.489386Z","shell.execute_reply.started":"2025-12-16T16:59:18.626896Z","shell.execute_reply":"2025-12-16T16:59:36.488733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index --find-links=\"/kaggle/input/surface-package-scraper\" -q pytorch_lightning monai albumentations imagecodecs --no-deps # \"numpy==1.26.4\" \"scipy==1.15.3\"\n!pip uninstall -q -y tensorflow \n\nimport os\nimport warnings\nfrom pathlib import Path\nimport random\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport pytorch_lightning as pl\nfrom tqdm.auto import tqdm\nimport tifffile\nimport zipfile\n\nwarnings.filterwarnings(\"ignore\")\n\n# ===========================\n# SIMPLE CONFIG\n# ===========================\nclass Config:\n    # Paths\n    DATA_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\n    TRAIN_IMAGES_DIR = DATA_DIR / \"train_images\"\n    TRAIN_LABELS_DIR = DATA_DIR / \"train_labels\"\n    TEST_IMAGES_DIR = DATA_DIR / \"test_images\"\n    OUTPUT_DIR = Path(\".\")\n    \n    # Model - SIMPLE\n    INPUT_SIZE = (160, 160, 160)\n    IN_CHANNELS = 1\n    OUT_CHANNELS = 2  # Background + surface\n    \n    # Training - SIMPLE\n    BATCH_SIZE = 2\n    NUM_WORKERS = 4\n    MAX_EPOCHS = 100\n    LR = 1e-4\n    VAL_SPLIT = 0.2\n    \n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nprint(\"=\"*70)\nprint(\"VESUVIUS - Simple DynUNet\")\nprint(\"=\"*70)\nprint(f\"Input: {Config.INPUT_SIZE}\")\nprint(f\"Batch: {Config.BATCH_SIZE}\")\nprint(\"=\"*70)\n\n# ===========================\n# DATASET - SIMPLE\n# ===========================\nclass SimpleDataset(Dataset):\n    def __init__(self, images_dir, labels_dir, files):\n        self.images_dir = images_dir\n        self.labels_dir = labels_dir\n        self.files = files\n        print(f\"Dataset: {len(files)} volumes\")\n    \n    def __len__(self):\n        return len(self.files)\n    \n    def __getitem__(self, idx):\n        fname = self.files[idx]\n        \n        # Load\n        img = tifffile.imread(str(self.images_dir / fname)).astype(np.float32)\n        \n        mask = None\n        if self.labels_dir:\n            mask_path = self.labels_dir / fname\n            if mask_path.exists():\n                mask = tifffile.imread(str(mask_path)).astype(np.uint8)\n                mask = (mask > 0).astype(np.uint8)  # Binary\n        \n        # To tensor\n        img = torch.from_numpy(img).half().div_(255.0).unsqueeze(0)\n        \n        if mask is not None:\n            mask = torch.from_numpy(mask).long().unsqueeze(0)\n        else:\n            mask = torch.zeros_like(img, dtype=torch.long)\n        \n        return img, mask, Path(fname).stem\n\n# ===========================\n# DATAMODULE - SIMPLE\n# ===========================\nfrom monai import transforms as MT\n\ndef collate(batch):\n    return batch\n\nclass SimpleDataModule(pl.LightningDataModule):\n    def __init__(self, train_dir, label_dir, size, val_split, batch_size, num_workers):\n        super().__init__()\n        self.train_dir = train_dir\n        self.label_dir = label_dir\n        self.size = size\n        self.val_split = val_split\n        self.batch_size = batch_size\n        self.num_workers = num_workers\n        \n        # Simple augmentations\n        self.train_transform = MT.Compose([\n            MT.Resized(keys=[\"image\", \"label\"], spatial_size=size, mode=[\"trilinear\", \"nearest\"]),\n            MT.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=0),\n            MT.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=1),\n            MT.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=2),\n        ])\n        \n        self.val_transform = MT.Compose([\n            MT.Resized(keys=[\"image\", \"label\"], spatial_size=size, mode=[\"trilinear\", \"nearest\"])\n        ])\n    \n    def setup(self, stage=None):\n        files = sorted([f.name for f in self.train_dir.glob(\"*.tif\")])\n        \n        random.seed(42)\n        random.shuffle(files)\n        split_idx = int(len(files) * (1 - self.val_split))\n        \n        train_files = files[:split_idx]\n        val_files = files[split_idx:]\n        \n        print(f\"Train: {len(train_files)}, Val: {len(val_files)}\")\n        \n        self.train_dataset = SimpleDataset(self.train_dir, self.label_dir, train_files)\n        self.val_dataset = SimpleDataset(self.train_dir, self.label_dir, val_files)\n    \n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, batch_size=self.batch_size, shuffle=True,\n                         num_workers=self.num_workers, pin_memory=True, collate_fn=collate)\n    \n    def val_dataloader(self):\n        return DataLoader(self.val_dataset, batch_size=self.batch_size, shuffle=False,\n                         num_workers=self.num_workers, pin_memory=True, collate_fn=collate)\n    \n    def on_after_batch_transfer(self, batch, dataloader_idx):\n        if not isinstance(batch, list):\n            return batch\n        \n        x_list, y_list, ids = [], [], []\n        device = self.trainer.strategy.root_device if self.trainer else torch.device(\"cuda\")\n        transform = self.train_transform if self.trainer.training else self.val_transform\n        \n        for x, y, fid in batch:\n            x = x.to(device, non_blocking=True)\n            y = y.to(device, non_blocking=True)\n            \n            data = transform({\"image\": x, \"label\": y})\n            x_list.append(data[\"image\"])\n            y_list.append(data[\"label\"])\n            ids.append(fid)\n        \n        return torch.stack(x_list), torch.stack(y_list), ids\n\n# ===========================\n# MODEL - SIMPLE DYNUNET\n# ===========================\nfrom monai.networks.nets import DynUNet\nfrom monai.losses import DiceCELoss\n\nclass SimpleDynUNet(pl.LightningModule):\n    def __init__(self, in_channels=1, out_channels=2, lr=1e-4):\n        super().__init__()\n        self.save_hyperparameters()\n        \n        # DynUNet - automatically calculates architecture\n        spatial_dims = 3\n        kernel_size = [[3, 3, 3]] * 5\n        strides = [[1, 1, 1]] + [[2, 2, 2]] * 4\n        upsample_kernel_size = strides[1:]\n        \n        self.model = DynUNet(\n            spatial_dims=spatial_dims,\n            in_channels=in_channels,\n            out_channels=out_channels,\n            kernel_size=kernel_size,\n            strides=strides,\n            upsample_kernel_size=upsample_kernel_size,\n            norm_name=\"instance\",\n            deep_supervision=False,\n            dropout=0.1\n        )\n        \n        self.loss_fn = DiceCELoss(softmax=True, to_onehot_y=False, include_background=True)\n        self.lr = lr\n    \n    def forward(self, x):\n        return self.model(x)\n    \n    def _compute_metrics(self, preds, targets):\n        preds_hard = torch.argmax(torch.softmax(preds, dim=1), dim=1)\n        \n        # Dice\n        pred_fg = (preds_hard == 1).float()\n        target_fg = (targets.squeeze(1) == 1).float()\n        \n        inter = (pred_fg * target_fg).sum()\n        union = pred_fg.sum() + target_fg.sum()\n        dice = (2 * inter + 1e-6) / (union + 1e-6)\n        \n        return {\"dice\": dice}\n    \n    def training_step(self, batch, batch_idx):\n        x, y, _ = batch\n        \n        # One-hot encode targets\n        y_onehot = F.one_hot(y.squeeze(1).long(), num_classes=self.hparams.out_channels).permute(0, 4, 1, 2, 3).float()\n        \n        logits = self(x)\n        loss = self.loss_fn(logits, y_onehot)\n        metrics = self._compute_metrics(logits, y)\n        \n        self.log(\"train_loss\", loss, prog_bar=True)\n        self.log(\"train_dice\", metrics[\"dice\"], prog_bar=True)\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        x, y, _ = batch\n        \n        y_onehot = F.one_hot(y.squeeze(1).long(), num_classes=self.hparams.out_channels).permute(0, 4, 1, 2, 3).float()\n        \n        logits = self(x)\n        loss = self.loss_fn(logits, y_onehot)\n        metrics = self._compute_metrics(logits, y)\n        \n        self.log(\"val_loss\", loss, prog_bar=True)\n        self.log(\"val_dice\", metrics[\"dice\"], prog_bar=True)\n        return loss\n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.parameters(), lr=self.lr, weight_decay=1e-4)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100, eta_min=1e-6)\n        return {\"optimizer\": optimizer, \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"epoch\"}}\n\n# ===========================\n# TRAIN\n# ===========================\ndef train():\n    print(\"\\nTraining...\")\n    \n    # Data\n    dm = SimpleDataModule(\n        Config.TRAIN_IMAGES_DIR,\n        Config.TRAIN_LABELS_DIR,\n        Config.INPUT_SIZE,\n        Config.VAL_SPLIT,\n        Config.BATCH_SIZE,\n        Config.NUM_WORKERS\n    )\n    dm.setup()\n    \n    # Model\n    model = SimpleDynUNet(\n        in_channels=Config.IN_CHANNELS,\n        out_channels=Config.OUT_CHANNELS,\n        lr=Config.LR\n    )\n    \n    print(f\"\\nModel: DynUNet\")\n    \n    # Callbacks\n    from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\n    \n    checkpoint = ModelCheckpoint(\n        dirpath=Config.OUTPUT_DIR,\n        filename=\"dynunet-{epoch:02d}-{val_dice:.4f}\",\n        monitor=\"val_dice\",\n        mode=\"max\",\n        save_top_k=2,\n        verbose=True\n    )\n    \n    early_stop = EarlyStopping(monitor=\"val_dice\", patience=8, mode=\"max\", verbose=True)\n    \n    # Trainer\n    trainer = pl.Trainer(\n        max_epochs=Config.MAX_EPOCHS,\n        accelerator=\"auto\",\n        devices=\"auto\",\n        callbacks=[checkpoint, early_stop],\n        precision=\"16-mixed\",\n        log_every_n_steps=10,\n        accumulate_grad_batches=4\n    )\n    \n    trainer.fit(model, dm)\n    \n    print(\"\\n✓ Training done!\")\n    return trainer, model, dm\n\n# ===========================\n# INFERENCE\n# ===========================\ndef inference(checkpoint_path, dm):\n    print(\"\\nInference...\")\n    \n    model = SimpleDynUNet.load_from_checkpoint(checkpoint_path)\n    model.eval()\n    model.to(Config.DEVICE)\n    \n    test_files = sorted([f.name for f in Config.TEST_IMAGES_DIR.glob(\"*.tif\")])\n    test_dataset = SimpleDataset(Config.TEST_IMAGES_DIR, None, test_files)\n    \n    print(f\"Processing {len(test_files)} volumes...\")\n    \n    predictions = []\n    \n    for img, _, fid in tqdm(test_dataset):\n        img_device = img.to(model.device)\n        processed = dm.val_transform({\"image\": img_device, \"label\": torch.zeros_like(img_device)})\n        inputs = processed[\"image\"].unsqueeze(0)\n        \n        with torch.no_grad():\n            logits = model(inputs)\n            pred = torch.argmax(torch.softmax(logits, dim=1), dim=1)[0]\n        \n        # Resize to original\n        orig_path = Config.TEST_IMAGES_DIR / f\"{fid}.tif\"\n        with tifffile.TiffFile(str(orig_path)) as tif:\n            orig_shape = tif.series[0].shape\n        \n        pred_resized = F.interpolate(\n            pred.float().unsqueeze(0).unsqueeze(0),\n            size=orig_shape,\n            mode='nearest'\n        ).squeeze().byte().cpu().numpy()\n        \n        save_path = Config.OUTPUT_DIR / f\"{fid}.tif\"\n        tifffile.imwrite(str(save_path), pred_resized, compression='lzw')\n        predictions.append(f\"{fid}.tif\")\n    \n    # Zip\n    with zipfile.ZipFile('submission.zip', 'w', zipfile.ZIP_DEFLATED) as z:\n        for fname in tqdm(predictions):\n            if os.path.exists(fname):\n                z.write(fname)\n                os.remove(fname)\n    \n    print(\"✓ Submission ready!\")\n\n# ===========================\n# MAIN\n# ===========================\nif __name__ == \"__main__\":\n    trainer, model, dm = train()\n    \n    # Find best checkpoint\n    import re\n    ckpts = list(Config.OUTPUT_DIR.glob(\"dynunet-*.ckpt\"))\n    if ckpts:\n        pattern = re.compile(r\"val_dice=?([0-9]+\\.[0-9]+)\")\n        best = max(ckpts, key=lambda p: float(pattern.search(p.name).group(1)))\n        print(f\"\\nBest: {best}\")\n        \n        inference(str(best), dm)\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"✓ DONE!\")\n    print(\"=\"*70)\n    print(\"\\n📊 CURRENT SETUP:\")\n    print(\"   Model: DynUNet (simple, auto-architecture)\")\n    print(\"   Size: 160³\")\n    print(\"   Batch: 2 × 4 = 8 (with gradient accumulation)\")\n    print(\"   Augment: Just flips\")\n    print(\"   Expected: 0.54-0.57\")\n    print(\"\\n💡 TO IMPROVE SCORE (Easy → Hard):\")\n    print(\"\\n1. EASY WINS (+0.02-0.03):\")\n    print(\"   - Change INPUT_SIZE to (192, 192, 192)\")\n    print(\"   - Add: MT.RandRotated(...) in augmentations\")\n    print(\"   - Change LR to 2e-4\")\n    print(\"\\n2. MEDIUM IMPROVEMENTS (+0.02-0.04):\")\n    print(\"   - Add TverskyLoss (alpha=0.3, beta=0.7) for better recall\")\n    print(\"   - Increase BATCH_SIZE to 4 if memory allows\")\n    print(\"   - Add dropout=0.2 in DynUNet\")\n    print(\"\\n3. BIGGER CHANGES (+0.03-0.06):\")\n    print(\"   - Switch to SegResNet (proven better)\")\n    print(\"   - Train 2 models and ensemble predictions\")\n    print(\"   - Add test-time augmentation (TTA)\")\n    print(\"\\n4. ADVANCED (+0.05-0.08):\")\n    print(\"   - Use SwinUNETR (best but complex)\")\n    print(\"   - Deep supervision = True\")\n    print(\"   - Post-processing (morphological operations)\")\n    print(\"=\"*70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T17:04:42.437739Z","iopub.execute_input":"2025-12-16T17:04:42.438458Z","iopub.status.idle":"2025-12-16T17:05:36.652749Z","shell.execute_reply.started":"2025-12-16T17:04:42.438433Z","shell.execute_reply":"2025-12-16T17:05:36.651945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}