{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":4098976,"sourceType":"datasetVersion","datasetId":2424221}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# USEFULL LINKS\n* [DATASET DISCRIPTION](https://www.med.upenn.edu/sbia/brats2017/data.html)\n* [REFERENCE PAPER](https://ieeexplore.ieee.org/stamp/stamp.jsp?tp=&arnumber=6975210)","metadata":{}},{"cell_type":"code","source":"! pip install torchio --upgrade --quiet\n! pip install torchio[plot] --upgrade --quiet\n! pip install torcheval --upgrade --quiet","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:43:42.826214Z","iopub.execute_input":"2025-08-24T09:43:42.826373Z","iopub.status.idle":"2025-08-24T09:45:29.796497Z","shell.execute_reply.started":"2025-08-24T09:43:42.826357Z","shell.execute_reply":"2025-08-24T09:45:29.795744Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms.functional as TF\nfrom  torchvision.ops import sigmoid_focal_loss\nfrom torcheval.metrics import MulticlassF1Score \nimport sklearn \nfrom sklearn.model_selection import train_test_split\nimport torchio as tio\nimport os\nfrom tqdm.auto import tqdm, trange","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:45:29.798407Z","iopub.execute_input":"2025-08-24T09:45:29.798753Z","iopub.status.idle":"2025-08-24T09:45:39.170264Z","shell.execute_reply.started":"2025-08-24T09:45:29.798724Z","shell.execute_reply":"2025-08-24T09:45:39.169485Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Conv3Dx3(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1):\n        super(Conv3Dx3, self).__init__()\n        \n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.kernel_size = kernel_size\n        self.stride = stride\n        self.padding = padding\n        \n        self.conv = nn.Sequential(\n            nn.Conv3d(self.in_channels, self.out_channels, self.kernel_size, self.stride, self.padding, bias=False),\n            nn.BatchNorm3d(self.out_channels),\n            nn.ReLU(inplace=True),\n            \n            nn.Conv3d(self.out_channels, self.out_channels, self.kernel_size, self.stride, self.padding, bias=False),\n            nn.BatchNorm3d(self.out_channels),\n            nn.ReLU(inplace=True),\n            \n            # nn.Conv3d(self.out_channels, self.out_channels, self.kernel_size, self.stride, self.padding, bias=False),\n            # nn.BatchNorm3d(self.out_channels),\n            # nn.ReLU(inplace=True),\n            \n             nn.Dropout(p=0.1)\n        )\n        \n       \n        \n    def forward(self, x):\n        return self.conv(x)","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:45:39.171078Z","iopub.execute_input":"2025-08-24T09:45:39.171449Z","iopub.status.idle":"2025-08-24T09:45:39.177288Z","shell.execute_reply.started":"2025-08-24T09:45:39.171419Z","shell.execute_reply":"2025-08-24T09:45:39.176497Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, in_channels=1, out_channels=2, encoder_depth=4):\n        super(UNet, self).__init__()\n        \n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.depth = encoder_depth\n        \n        self.left = nn.ModuleList()\n        self.right = nn.ModuleList()\n        self.upscaling = nn.ModuleList()\n        \n        START_N_CHANNELS = 16\n        for n_channels in [START_N_CHANNELS*(2**i) for i in range(encoder_depth)]:\n            self.left.append(Conv3Dx3(in_channels, n_channels))\n            self.right.insert(0,Conv3Dx3(2*n_channels, n_channels))\n            self.upscaling.insert(0,nn.ConvTranspose3d(2*n_channels, n_channels, kernel_size=2, stride=2))\n            in_channels = n_channels\n            \n        self.bottleneck = Conv3Dx3(in_channels, 2*in_channels)    \n        self.final_conv = nn.Conv3d(START_N_CHANNELS, out_channels, kernel_size=1)\n\n    def forward(self,x):\n        skip_connections = []\n        \n        for conv in self.left:\n            x = conv(x)\n            skip_connections.append(x)\n            x = F.max_pool3d(x, kernel_size=2, stride=2, padding=0)\n        \n        x = self.bottleneck(x)\n        \n        for up, conv in zip(self.upscaling,self.right):\n            x = up(x)\n            skip = skip_connections.pop()\n            x = torch.cat([x, skip], dim=tio.CHANNELS_DIMENSION)\n            x = conv(x)\n        \n        x = self.final_conv(x)\n#         x = F.softmax(x, dim=tio.CHANNELS_DIMENSION)\n        \n        return x ","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:45:39.178228Z","iopub.execute_input":"2025-08-24T09:45:39.178470Z","iopub.status.idle":"2025-08-24T09:45:39.200681Z","shell.execute_reply.started":"2025-08-24T09:45:39.178444Z","shell.execute_reply":"2025-08-24T09:45:39.200051Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EarlyStopper:\n    def __init__(self, patience=1, min_delta=0):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.counter = 0\n        self.min_val_loss = float('inf')\n\n    def __call__(self, validation_loss):\n        if validation_loss < self.min_val_loss:\n            self.min_val_loss = validation_loss\n            self.counter = 0\n        elif validation_loss > (self.min_val_loss + self.min_delta):\n            self.counter += 1\n            if self.counter >= self.patience:\n                return True\n        return False","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:45:39.202369Z","iopub.execute_input":"2025-08-24T09:45:39.202611Z","iopub.status.idle":"2025-08-24T09:45:39.218042Z","shell.execute_reply.started":"2025-08-24T09:45:39.202593Z","shell.execute_reply":"2025-08-24T09:45:39.217442Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n   \n    def __init__(self, num_classes, softmax_dim=None):\n        super().__init__()\n        self.num_classes = num_classes\n        self.softmax_dim = softmax_dim\n        \n    def forward(self, logits, targets):\n        \n        probabilities = logits\n        if self.softmax_dim is not None:\n            probabilities = nn.Softmax(dim=self.softmax_dim)(logits)\n            \n        if targets.dim() == probabilities.dim() - 1:\n            targets = F.one_hot(targets, num_classes=self.num_classes)\n            targets = targets.permute((0,4,1,2,3))\n            \n#         print(targets.shape)\n#         print(probabilities.shape)\n#         probabilities = probabilities.reshape((-1,self.num_classes))\n#         targets = targets.reshape((-1,self.num_classes))                                                         \n\n        intersection = (targets * probabilities).sum()\n        mod_a = probabilities.sum()\n        mod_b = targets.sum()\n        \n        dice_coefficient = 2. * intersection / (mod_a + mod_b)\n        dice_loss = -dice_coefficient.log()\n        return dice_loss","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:45:39.218727Z","iopub.execute_input":"2025-08-24T09:45:39.218891Z","iopub.status.idle":"2025-08-24T09:45:39.235553Z","shell.execute_reply.started":"2025-08-24T09:45:39.218877Z","shell.execute_reply":"2025-08-24T09:45:39.234951Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_DIR = '/kaggle/input/brats17/BRATS2017/Brats17TrainingData/HGG'\nTEST_DIR = '/kaggle/input/brats17/BRATS2017/Brats17TestingData'","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:45:39.236298Z","iopub.execute_input":"2025-08-24T09:45:39.236490Z","iopub.status.idle":"2025-08-24T09:45:39.249961Z","shell.execute_reply.started":"2025-08-24T09:45:39.236467Z","shell.execute_reply":"2025-08-24T09:45:39.249405Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subjects = []\nfor sample in (pbar:=tqdm(os.listdir(TRAIN_DIR))):\n    pbar.set_postfix({'Loading file': sample})\n    subject = tio.Subject(\n        id = sample,\n        flair=tio.ScalarImage(f'{TRAIN_DIR}/{sample}/{sample}_flair.nii'),\n        t1=tio.ScalarImage(f'{TRAIN_DIR}/{sample}/{sample}_t1.nii'),\n        t1ce=tio.ScalarImage(f'{TRAIN_DIR}/{sample}/{sample}_t1ce.nii'),\n        t2=tio.ScalarImage(f'{TRAIN_DIR}/{sample}/{sample}_t2.nii'),\n        label=tio.LabelMap(f'{TRAIN_DIR}/{sample}/{sample}_seg.nii'),\n        )\n    \n    try:\n        subject.flair.load()\n        \n    except FileNotFoundError:\n        subject.add_image(tio.ScalarImage(f'{TRAIN_DIR}/{sample}/{sample}_flair.nii/{sample}_flair.nii'), 'flair')\n            \n    subjects.append(subject)","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:45:39.250661Z","iopub.execute_input":"2025-08-24T09:45:39.250927Z","iopub.status.idle":"2025-08-24T09:46:05.300442Z","shell.execute_reply.started":"2025-08-24T09:45:39.250910Z","shell.execute_reply":"2025-08-24T09:46:05.299719Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LEARNING_RATE = 0.01\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n#DEVICE = \"cpu\"\nBATCH_SIZE = 4\nNUM_EPOCHS = 64\nNUM_WORKERS = 4\nIMAGE_HEIGHT = 240 \nIMAGE_WIDTH = 240\nIMAGE_DEPTH = 155\nOUT_CLASSES = 5","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:46:05.301426Z","iopub.execute_input":"2025-08-24T09:46:05.301980Z","iopub.status.idle":"2025-08-24T09:46:05.370662Z","shell.execute_reply.started":"2025-08-24T09:46:05.301959Z","shell.execute_reply":"2025-08-24T09:46:05.369936Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:46:05.371536Z","iopub.execute_input":"2025-08-24T09:46:05.372095Z","iopub.status.idle":"2025-08-24T09:46:05.384546Z","shell.execute_reply.started":"2025-08-24T09:46:05.372066Z","shell.execute_reply":"2025-08-24T09:46:05.383833Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transforms = [\n    tio.Resample('flair'),\n    tio.ToCanonical(),\n#     tio.transforms.EnsureShapeMultiple(16),\n#     tio.CropOrPad((IMAGE_WIDTH,IMAGE_HEIGHT,IMAGE_DEPTH)),\n#     tio.OneHot(),\n    tio.RescaleIntensity(out_min_max=(0, 1),\n                         percentiles=(0.5, 99.5)),\n#     tio.ZNormalization(),\n]\n\ntraining_subjects, validation_subjects = train_test_split(subjects,train_size=0.8)\n\ntraining_dataset = tio.SubjectsDataset(\n    training_subjects,\n    transform=tio.Compose(transforms),\n    )\n\nvalidation_dataset = tio.SubjectsDataset(\n    validation_subjects,\n    transform=tio.Compose(transforms),\n    )\n","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:46:05.385310Z","iopub.execute_input":"2025-08-24T09:46:05.385865Z","iopub.status.idle":"2025-08-24T09:46:05.401647Z","shell.execute_reply.started":"2025-08-24T09:46:05.385841Z","shell.execute_reply":"2025-08-24T09:46:05.400991Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patch_size = 96\nsamples_per_volume = 3\nqueue_length = BATCH_SIZE * samples_per_volume\nsampler = tio.UniformSampler(patch_size=patch_size)\n\npatches_queue_training = tio.Queue(\n    training_dataset,\n    queue_length,\n    samples_per_volume,\n    sampler,\n    num_workers=NUM_WORKERS,\n)\n\ntraining_loader = tio.SubjectsLoader(patches_queue_training, batch_size=BATCH_SIZE, num_workers=0)","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:46:05.402297Z","iopub.execute_input":"2025-08-24T09:46:05.402499Z","iopub.status.idle":"2025-08-24T09:46:05.651982Z","shell.execute_reply.started":"2025-08-24T09:46:05.402484Z","shell.execute_reply":"2025-08-24T09:46:05.649853Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patches_queue_validation = tio.Queue(\n    validation_dataset,\n    queue_length,\n    samples_per_volume,\n    sampler,\n    num_workers=NUM_WORKERS,\n    shuffle_subjects=False,\n    shuffle_patches=False,\n)\n\nvalidation_loader = tio.SubjectsLoader(patches_queue_validation, batch_size=BATCH_SIZE, num_workers=0)","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:46:05.654006Z","iopub.execute_input":"2025-08-24T09:46:05.654367Z","iopub.status.idle":"2025-08-24T09:46:06.034294Z","shell.execute_reply.started":"2025-08-24T09:46:05.654319Z","shell.execute_reply":"2025-08-24T09:46:06.030174Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNet(in_channels=4, out_channels=OUT_CLASSES, encoder_depth=3).to(DEVICE)\n\n#optimizer = torch.optim.SGD(model.parameters(), lr=LEARNING_RATE, momentum=0.9, weight_decay=LEARNING_RATE/100)\noptimizer = torch.optim.SGD(model.parameters(), lr=LEARNING_RATE, momentum=0.6, weight_decay=0.001)\n# scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.8)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.1, patience=2)\n# optimizer = torch.optim.Adagrad(model.parameters(), lr=LEARNING_RATE, lr_decay=0.9, weight_decay=LEARNING_RATE/100)\n# criterion = nn.CrossEntropyLoss(weight=torch.Tensor([0.19,0.22,0.197,0.199,0.199]))\n#criterion = nn.BCELoss()\n#criterion = sigmoid_focal_loss\n#criterion = nn.BCEWithLogitsLoss()\ncriterion = nn.CrossEntropyLoss()\n# criterion = FocalLoss()\n# criterion = DiceLoss(OUT_CLASSES, softmax_dim=tio.CHANNELS_DIMENSION)\ncriterion.to(DEVICE)\nmetric = MulticlassF1Score(num_classes=OUT_CLASSES)","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:46:06.055440Z","iopub.execute_input":"2025-08-24T09:46:06.055968Z","iopub.status.idle":"2025-08-24T09:46:06.567609Z","shell.execute_reply.started":"2025-08-24T09:46:06.055912Z","shell.execute_reply":"2025-08-24T09:46:06.560350Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"training_loss_per_epoch = []\nvalidation_loss_per_epoch = []\ntraining_metric_per_epoch = []\nvalidation_metric_per_epoch = []","metadata":{"execution":{"iopub.status.busy":"2025-08-24T09:46:06.568826Z","iopub.execute_input":"2025-08-24T09:46:06.569265Z","iopub.status.idle":"2025-08-24T09:46:06.587961Z","shell.execute_reply.started":"2025-08-24T09:46:06.569230Z","shell.execute_reply":"2025-08-24T09:46:06.585179Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# torch.load('best_model.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T09:46:06.589860Z","iopub.execute_input":"2025-08-24T09:46:06.593559Z","iopub.status.idle":"2025-08-24T09:46:06.631976Z","shell.execute_reply.started":"2025-08-24T09:46:06.593494Z","shell.execute_reply":"2025-08-24T09:46:06.621590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"early_stopper = EarlyStopper(patience=3, min_delta=0.001)\n\nfor epoch in (out_pbar:=trange(NUM_EPOCHS, unit='epoch')):\n    \n    out_pbar.set_description(f'Epoch:{epoch+1}/{NUM_EPOCHS} Best Validation Metric: {early_stopper.min_val_loss:.4f}\\r', refresh=True)\n    \n    running_loss = 0.0    \n    model.train()\n    with  tqdm(total=len(training_loader), unit='batch') as pbar:\n        pbar.reset()\n        for subjects_batch in training_loader:\n            pbar.update()\n            \n            flair = subjects_batch['flair'][tio.DATA].to(DEVICE)\n            t1 = subjects_batch['t1'][tio.DATA].to(DEVICE)\n            t1ce = subjects_batch['t1ce'][tio.DATA].to(DEVICE)\n            t2 = subjects_batch['t2'][tio.DATA].to(DEVICE)\n\n            inputs = torch.cat([flair,t1,t1ce,t2], dim=tio.CHANNELS_DIMENSION)\n            targets = subjects_batch['label'][tio.DATA].to(DEVICE)\n            targets = targets.squeeze(dim=tio.CHANNELS_DIMENSION).long()\n            \n            outputs = model(inputs)\n            \n            loss = criterion(outputs, targets)\n            loss.backward()\n                        \n            optimizer.step()\n            optimizer.zero_grad()\n            \n            metric.update(outputs.reshape((-1,OUT_CLASSES)), targets.reshape(-1))\n            \n            running_loss += loss.item()\n            pbar.set_description(f'Current Training Loss: {loss}\\r')\n        \n        current_metric = metric.compute().item()\n        pbar.set_postfix({type(metric).__name__: current_metric})\n        training_metric_per_epoch.append(current_metric)\n        metric.reset()\n        \n        running_loss /= len(training_loader)\n        training_loss_per_epoch.append(running_loss)\n        pbar.set_description(f'Epoch:{epoch+1}/{NUM_EPOCHS} Total Training Loss: {running_loss:.4f}\\r', refresh=True)\n        \n             \n    running_loss = 0.0\n    model.eval()\n    with tqdm(total=len(validation_loader), unit='batch') as pbar, torch.no_grad():\n        pbar.reset()\n        for subjects_batch in validation_loader:\n            pbar.update()\n            \n            flair = subjects_batch['flair'][tio.DATA].to(DEVICE)\n            t1 = subjects_batch['t1'][tio.DATA].to(DEVICE)\n            t1ce = subjects_batch['t1ce'][tio.DATA].to(DEVICE)\n            t2 = subjects_batch['t2'][tio.DATA].to(DEVICE)\n\n            inputs = torch.cat([flair,t1,t1ce,t2], dim=tio.CHANNELS_DIMENSION)      \n            targets = subjects_batch['label'][tio.DATA].to(DEVICE)\n            targets = targets.squeeze(dim=tio.CHANNELS_DIMENSION).long()\n            \n            outputs = model(inputs)\n            \n            loss = criterion(outputs, targets).item()               \n            metric.update(outputs.reshape((-1,OUT_CLASSES)), targets.reshape(-1))\n            \n            running_loss += loss\n            pbar.set_description(f'Current Validation Loss: {loss}\\r')\n    \n        current_metric = metric.compute().item()\n        pbar.set_postfix({type(metric).__name__: current_metric})\n        validation_metric_per_epoch.append(current_metric)\n        metric.reset()\n        \n        running_loss /= len(validation_loader)\n        validation_loss_per_epoch.append(running_loss)\n        pbar.set_description(f'Epoch:{epoch+1}/{NUM_EPOCHS} Total Validation Loss: {running_loss:.4f}\\r', refresh=True)\n        \n            \n    scheduler.step(current_metric)\n    if 1-current_metric < early_stopper.min_val_loss:\n        torch.save(model.state_dict(),'best_model.pt')\n        \n    \n    if early_stopper(1-current_metric):\n        print(f'Early stop, best val metric:{early_stopper.min_val_loss:.5f}')\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T09:46:06.633132Z","iopub.execute_input":"2025-08-24T09:46:06.634985Z","iopub.status.idle":"2025-08-24T10:14:15.985565Z","shell.execute_reply.started":"2025-08-24T09:46:06.633494Z","shell.execute_reply":"2025-08-24T10:14:15.984349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load('best_model.pt'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:14:15.990014Z","iopub.execute_input":"2025-08-24T10:14:15.990322Z","iopub.status.idle":"2025-08-24T10:14:16.302493Z","shell.execute_reply.started":"2025-08-24T10:14:15.990278Z","shell.execute_reply":"2025-08-24T10:14:16.301939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from matplotlib import pyplot as plt\n\nplt.plot(training_loss_per_epoch, 'r', label='training')\nplt.plot(validation_loss_per_epoch, 'b', label='validation')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:14:16.303209Z","iopub.execute_input":"2025-08-24T10:14:16.303468Z","iopub.status.idle":"2025-08-24T10:14:16.551284Z","shell.execute_reply.started":"2025-08-24T10:14:16.303449Z","shell.execute_reply":"2025-08-24T10:14:16.550634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.plot(training_metric_per_epoch, 'r', label='training')\nplt.plot(validation_metric_per_epoch, 'b', label='validation')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:14:16.552108Z","iopub.execute_input":"2025-08-24T10:14:16.552359Z","iopub.status.idle":"2025-08-24T10:14:16.685613Z","shell.execute_reply.started":"2025-08-24T10:14:16.552342Z","shell.execute_reply":"2025-08-24T10:14:16.684975Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from math import floor\n\npatch_overlap=[\n     floor((patch_size-IMAGE_WIDTH%patch_size)/(IMAGE_WIDTH//patch_size)/2)*2,\n     floor((patch_size-IMAGE_HEIGHT%patch_size)/(IMAGE_HEIGHT//patch_size)/2)*2,\n     floor((patch_size-IMAGE_DEPTH%patch_size)/(IMAGE_DEPTH//patch_size)/2)*2,\n]\n# patch_overlap = 2\nsubject = validation_dataset[3]\n# subject = training_dataset[42]\n# sample = 'Brats17_CBICA_AAA_1'\n# subject = tio.Subject(\n#         id = sample,\n#         flair=tio.ScalarImage(f'{TEST_DIR}/{sample}/{sample}_flair.nii'),\n#         )\n# subject=tio.Compose(transforms)(subject)\n\ngrid_sampler = tio.inference.GridSampler(\n    subject,\n    patch_size,\n    patch_overlap,\n)\n\npatch_loader = tio.SubjectsLoader(grid_sampler, batch_size=BATCH_SIZE)\naggregator = tio.inference.GridAggregator(grid_sampler, overlap_mode='crop')\nmodel.eval()\n\nwith torch.no_grad():\n    for patches_batch in patch_loader:\n        flair = patches_batch['flair'][tio.DATA].to(DEVICE)\n        t1 = patches_batch['t1'][tio.DATA].to(DEVICE)\n        t1ce = patches_batch['t1ce'][tio.DATA].to(DEVICE)\n        t2 = patches_batch['t2'][tio.DATA].to(DEVICE)\n        inputs = torch.cat([flair, t1, t1ce, t2],dim=tio.CHANNELS_DIMENSION)\n        targets = patches_batch['label'][tio.DATA].to(DEVICE).squeeze(1).long()\n        \n        locations = patches_batch[tio.LOCATION]\n        logits = model(inputs)\n        print(criterion(logits,targets))\n#         logits = F.softmax(model(inputs), dim=tio.CHANNELS_DIMENSION)             \n        labels = logits.argmax(dim=tio.CHANNELS_DIMENSION, keepdim=True)\n        outputs = logits\n        aggregator.add_batch(outputs, locations)\n\noutput_tensor = aggregator.get_output_tensor()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:07.072240Z","iopub.execute_input":"2025-08-24T10:19:07.073034Z","iopub.status.idle":"2025-08-24T10:19:09.950589Z","shell.execute_reply.started":"2025-08-24T10:19:07.073000Z","shell.execute_reply":"2025-08-24T10:19:09.949766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# criterion(output_tensor.unsqueeze(0), subject.label.tensor.long())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:09.951900Z","iopub.execute_input":"2025-08-24T10:19:09.952161Z","iopub.status.idle":"2025-08-24T10:19:09.955484Z","shell.execute_reply.started":"2025-08-24T10:19:09.952142Z","shell.execute_reply":"2025-08-24T10:19:09.954910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"output_tensor.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:13.568577Z","iopub.execute_input":"2025-08-24T10:19:13.568875Z","iopub.status.idle":"2025-08-24T10:19:13.573889Z","shell.execute_reply.started":"2025-08-24T10:19:13.568854Z","shell.execute_reply":"2025-08-24T10:19:13.573158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subject.label.tensor.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:14.606269Z","iopub.execute_input":"2025-08-24T10:19:14.606562Z","iopub.status.idle":"2025-08-24T10:19:14.611431Z","shell.execute_reply.started":"2025-08-24T10:19:14.606533Z","shell.execute_reply":"2025-08-24T10:19:14.610849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subject.id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:15.344302Z","iopub.execute_input":"2025-08-24T10:19:15.344610Z","iopub.status.idle":"2025-08-24T10:19:15.349675Z","shell.execute_reply.started":"2025-08-24T10:19:15.344586Z","shell.execute_reply":"2025-08-24T10:19:15.348746Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subject.flair.plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:15.766079Z","iopub.execute_input":"2025-08-24T10:19:15.766378Z","iopub.status.idle":"2025-08-24T10:19:16.107641Z","shell.execute_reply.started":"2025-08-24T10:19:15.766354Z","shell.execute_reply":"2025-08-24T10:19:16.106917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    subject.t1ce.plot()\nexcept:\n    pass","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:16.225837Z","iopub.execute_input":"2025-08-24T10:19:16.226561Z","iopub.status.idle":"2025-08-24T10:19:16.574998Z","shell.execute_reply.started":"2025-08-24T10:19:16.226538Z","shell.execute_reply":"2025-08-24T10:19:16.574233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subject.t2.plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:16.914666Z","iopub.execute_input":"2025-08-24T10:19:16.914950Z","iopub.status.idle":"2025-08-24T10:19:17.259623Z","shell.execute_reply.started":"2025-08-24T10:19:16.914930Z","shell.execute_reply":"2025-08-24T10:19:17.258886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if subject.label.tensor.shape[0] == 1:\n    tio.LabelMap(tensor=subject.label.tensor).plot()    \nelse:\n    tio.LabelMap(tensor=subject.label.tensor.argmax(dim=0, keepdim=True)).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:17.260813Z","iopub.execute_input":"2025-08-24T10:19:17.261083Z","iopub.status.idle":"2025-08-24T10:19:17.580802Z","shell.execute_reply.started":"2025-08-24T10:19:17.261058Z","shell.execute_reply":"2025-08-24T10:19:17.580147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if output_tensor.shape[0] == 1:    \n    tio.LabelMap(tensor=output_tensor).plot()\nelse:\n    tio.LabelMap(tensor=output_tensor.argmax(dim=0, keepdim=True)).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:17.803016Z","iopub.execute_input":"2025-08-24T10:19:17.803699Z","iopub.status.idle":"2025-08-24T10:19:19.706493Z","shell.execute_reply.started":"2025-08-24T10:19:17.803669Z","shell.execute_reply":"2025-08-24T10:19:19.705957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if subject.label.tensor.shape[0] == 1:\n    print(subject.label.tensor.unique())\nelse:    \n    print(subject.label.tensor.argmax(dim=0, keepdim=True).unique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:21.676296Z","iopub.execute_input":"2025-08-24T10:19:21.676589Z","iopub.status.idle":"2025-08-24T10:19:21.801094Z","shell.execute_reply.started":"2025-08-24T10:19:21.676565Z","shell.execute_reply":"2025-08-24T10:19:21.800392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if output_tensor.shape[0] == 1:    \n    print(output_tensor.unique())\nelse:\n    print(output_tensor.argmax(dim=0, keepdim=True).unique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:25.520203Z","iopub.execute_input":"2025-08-24T10:19:25.520689Z","iopub.status.idle":"2025-08-24T10:19:27.431472Z","shell.execute_reply.started":"2025-08-24T10:19:25.520665Z","shell.execute_reply":"2025-08-24T10:19:27.430695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"view_label=2\nif output_tensor.shape[0] == 1:    \n    tio.LabelMap(tensor=(output_tensor.eq(view_label))).plot()\nelse:\n    tio.LabelMap(tensor=(output_tensor.argmax(dim=0, keepdim=True).eq(view_label))).plot()\n    tio.LabelMap(tensor=output_tensor[view_label:view_label+1]>0.1).plot()\n    tio.ScalarImage(tensor=output_tensor[view_label:view_label+1]).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:27.432590Z","iopub.execute_input":"2025-08-24T10:19:27.432818Z","iopub.status.idle":"2025-08-24T10:19:29.967153Z","shell.execute_reply.started":"2025-08-24T10:19:27.432801Z","shell.execute_reply":"2025-08-24T10:19:29.966608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"view_label=2\nif subject.label.tensor.shape[0] == 1:\n    tio.LabelMap(tensor=subject.label.tensor.eq(view_label)).plot()\nelse:\n    tio.LabelMap(tensor=subject.label.tensor[view_label:view_label+1]).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:29.968375Z","iopub.execute_input":"2025-08-24T10:19:29.968656Z","iopub.status.idle":"2025-08-24T10:19:30.271184Z","shell.execute_reply.started":"2025-08-24T10:19:29.968636Z","shell.execute_reply":"2025-08-24T10:19:30.270568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tio.LabelMap(tensor=(subject.flair.tensor>0.42)&~(subject.t1ce.tensor>0.25)).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:30.272132Z","iopub.execute_input":"2025-08-24T10:19:30.272394Z","iopub.status.idle":"2025-08-24T10:19:30.599454Z","shell.execute_reply.started":"2025-08-24T10:19:30.272376Z","shell.execute_reply":"2025-08-24T10:19:30.598575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"view_label=4\nif subject.label.tensor.shape[0] == 1:\n    tio.LabelMap(tensor=subject.label.tensor.eq(view_label)).plot()\nelse:\n    tio.LabelMap(tensor=subject.label.tensor[view_label:view_label+1]).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:30.600779Z","iopub.execute_input":"2025-08-24T10:19:30.601078Z","iopub.status.idle":"2025-08-24T10:19:30.917189Z","shell.execute_reply.started":"2025-08-24T10:19:30.601052Z","shell.execute_reply":"2025-08-24T10:19:30.916526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tio.LabelMap(tensor=subject.t1ce.tensor>0.3).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:36.621292Z","iopub.execute_input":"2025-08-24T10:19:36.621896Z","iopub.status.idle":"2025-08-24T10:19:36.922953Z","shell.execute_reply.started":"2025-08-24T10:19:36.621872Z","shell.execute_reply":"2025-08-24T10:19:36.922163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"view_label=1\nif subject.label.tensor.shape[0] == 1:\n    tio.LabelMap(tensor=subject.label.tensor.eq(view_label)).plot()\nelse:\n    tio.LabelMap(tensor=subject.label.tensor[view_label:view_label+1]).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:38.940105Z","iopub.execute_input":"2025-08-24T10:19:38.940547Z","iopub.status.idle":"2025-08-24T10:19:39.253035Z","shell.execute_reply.started":"2025-08-24T10:19:38.940483Z","shell.execute_reply":"2025-08-24T10:19:39.252291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tio.LabelMap(tensor=(subject.t1ce.tensor<0.2)&(subject.t1ce.tensor>0.15)).plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:19:39.419718Z","iopub.execute_input":"2025-08-24T10:19:39.419995Z","iopub.status.idle":"2025-08-24T10:19:39.734416Z","shell.execute_reply.started":"2025-08-24T10:19:39.419975Z","shell.execute_reply":"2025-08-24T10:19:39.733680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#output_tensor = output_tensor.argmax(dim=0, keepdim=True)\naffine = validation_dataset[0].label.affine\noutput_tensor = tio.ScalarImage(tensor=output_tensor,affine=affine)\noutput_tensor.save(f'{subject.id}_predict.nii')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:17:27.867798Z","iopub.execute_input":"2025-08-24T10:17:27.868046Z","iopub.status.idle":"2025-08-24T10:17:31.656598Z","shell.execute_reply.started":"2025-08-24T10:17:27.868030Z","shell.execute_reply":"2025-08-24T10:17:31.655782Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(),f'checkpoint-cross_entropy{validation_loss_per_epoch[-1]:.4f}.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T10:14:22.983284Z","iopub.status.idle":"2025-08-24T10:14:22.983581Z","shell.execute_reply.started":"2025-08-24T10:14:22.983430Z","shell.execute_reply":"2025-08-24T10:14:22.983442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}