{"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":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# train_3d_unet.py\nimport os\nimport glob\nimport random\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom torch.optim import Adam\nimport nibabel as nib\nfrom scipy.ndimage import zoom, gaussian_filter\nfrom sklearn.model_selection import KFold\nfrom tqdm import tqdm\n\n# -------------------------\n# Simple 3D UNet (small)\n# -------------------------\nclass DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, 3, padding=1),\n            nn.InstanceNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch, out_ch, 3, padding=1),\n            nn.InstanceNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.net(x)\n\nclass Down(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.net = nn.Sequential(nn.MaxPool3d(2), DoubleConv(in_ch, out_ch))\n    def forward(self, x): return self.net(x)\n\nclass Up(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.up = nn.ConvTranspose3d(in_ch, in_ch//2, 2, stride=2)\n        self.conv = DoubleConv(in_ch, out_ch)\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        diff = [x2.size(i) - x1.size(i) for i in range(2,5)]\n        x1 = F.pad(x1, [d//2 for d in diff[::-1]])\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\nclass UNet3D(nn.Module):\n    def __init__(self, in_ch=1, out_ch=1, fmaps=32):\n        super().__init__()\n        self.inc = DoubleConv(in_ch, fmaps)\n        self.down1 = Down(fmaps, fmaps*2)\n        self.down2 = Down(fmaps*2, fmaps*4)\n        self.down3 = Down(fmaps*4, fmaps*8)\n        self.up1 = Up(fmaps*8, fmaps*4)\n        self.up2 = Up(fmaps*4, fmaps*2)\n        self.up3 = Up(fmaps*2, fmaps)\n        self.outc = nn.Conv3d(fmaps, out_ch, 1)\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x = self.up1(x4, x3)\n        x = self.up2(x, x2)\n        x = self.up3(x, x1)\n        x = self.outc(x)\n        return x\n\n# -------------------------\n# Dataset (NIfTI volumes)\n# -------------------------\nclass ScrollDataset(Dataset):\n    def __init__(self, images, masks, patch_size=(96,96,64), augment=True):\n        self.images = images\n        self.masks = masks\n        self.patch_size = np.array(patch_size)\n        self.augment = augment\n    def __len__(self):\n        return len(self.images)\n    def random_crop(self, img, msk):\n        shape = np.array(img.shape)\n        start = [np.random.randint(0, max(1, shape[i]-self.patch_size[i])) for i in range(3)]\n        slices = tuple(slice(start[i], start[i]+self.patch_size[i]) for i in range(3))\n        return img[slices], msk[slices]\n    def __getitem__(self, idx):\n        img = nib.load(self.images[idx]).get_fdata().astype(np.float32)\n        msk = nib.load(self.masks[idx]).get_fdata().astype(np.uint8)\n        # Normalize\n        img = np.clip(img, np.percentile(img,1), np.percentile(img,99))\n        img = (img - img.mean()) / (img.std() + 1e-8)\n        # Crop patch\n        if any(np.array(img.shape) < self.patch_size):\n            pad = np.maximum(self.patch_size - np.array(img.shape), 0)\n            pad_before = (pad//2).tolist(); pad_after = (pad - pad//2).tolist()\n            img = np.pad(img, tuple(zip(pad_before, pad_after)), mode='constant')\n            msk = np.pad(msk, tuple(zip(pad_before, pad_after)), mode='constant')\n        img_p, msk_p = self.random_crop(img, msk)\n        # Augment\n        if self.augment:\n            if random.random() < 0.5:\n                img_p = np.flip(img_p, axis=0); msk_p = np.flip(msk_p, axis=0)\n            if random.random() < 0.5:\n                img_p = np.flip(img_p, axis=1); msk_p = np.flip(msk_p, axis=1)\n            if random.random() < 0.2:\n                img_p = gaussian_filter(img_p, sigma=np.random.uniform(0,1.0))\n        # To tensor\n        img_p = torch.from_numpy(img_p[None]).float()\n        msk_p = torch.from_numpy(msk_p[None]).float()\n        return img_p, msk_p\n\n# -------------------------\n# Utilities\n# -------------------------\ndef load_paths(data_dir):\n    imgs = sorted(glob.glob(os.path.join(data_dir, 'images','*.nii*')))\n    msks = sorted(glob.glob(os.path.join(data_dir, 'masks','*.nii*')))\n    return imgs, msks\n\ndef dice_coef(pred, target, eps=1e-6):\n    p = (pred>0.5).float()\n    inter = (p*target).sum()\n    return (2*inter) / (p.sum()+target.sum()+eps)\n\n# -------------------------\n# Training loop\n# -------------------------\ndef train(data_dir, out_dir, epochs=200, batch_size=1, lr=1e-4):\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    imgs, msks = load_paths(data_dir)\n    kf = KFold(n_splits=5, shuffle=True, random_state=42)\n    train_idx, val_idx = list(kf.split(imgs))[0]\n    train_imgs = [imgs[i] for i in train_idx]; train_msks = [msks[i] for i in train_idx]\n    val_imgs = [imgs[i] for i in val_idx]; val_msks = [msks[i] for i in val_idx]\n    train_ds = ScrollDataset(train_imgs, train_msks, augment=True)\n    val_ds = ScrollDataset(val_imgs, val_msks, augment=False)\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=4)\n    val_loader = DataLoader(val_ds, batch_size=1, shuffle=False, num_workers=2)\n    model = UNet3D(in_ch=1, out_ch=1, fmaps=32).to(device)\n    opt = Adam(model.parameters(), lr=lr)\n    for epoch in range(epochs):\n        model.train()\n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\")\n        train_loss = 0.0\n        for x,y in pbar:\n            x = x.to(device); y = y.to(device)\n            logits = model(x)\n            probs = torch.sigmoid(logits)\n            dice_loss = 1 - (2*(probs*y).sum() / (probs.sum()+y.sum()+1e-6))\n            bce = F.binary_cross_entropy(probs, y)\n            loss = dice_loss + 0.5*bce\n            opt.zero_grad(); loss.backward(); opt.step()\n            train_loss += loss.item()\n            pbar.set_postfix(loss= train_loss/(pbar.n+1))\n        # validation\n        model.eval()\n        dices = []\n        with torch.no_grad():\n            for x,y in val_loader:\n                x = x.to(device); y = y.to(device)\n                pred = torch.sigmoid(model(x))\n                dices.append(dice_coef(pred, y).item())\n        mean_dice = np.mean(dices)\n        print(f\"[Epoch {epoch+1}] val_dice={mean_dice:.4f}\")\n        # Save\n        torch.save(model.state_dict(), os.path.join(out_dir, f\"model_epoch{epoch+1:03d}.pt\"))\n\nif __name__ == \"__main__\":\n    # Example usage:\n    # prepare folders: data/images/*.nii data/masks/*.nii\n    train(data_dir='./data', out_dir='./out', epochs=200, batch_size=1, lr=1e-4)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}