{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":29653,"databundleVersionId":2420395,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\nimport pandas as pd\nimport numpy as np\nimport nibabel as nib\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models\nfrom tqdm import tqdm  # Import tqdm for progress bar\n\n# Custom Dataset Class for 3D MRI Data\nclass BrainTumorDataset(Dataset):\n    def __init__(self, image_dir, csv_file, transform=None):\n        self.image_dir = image_dir\n        self.data = pd.read_csv(csv_file)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        brats_id = self.data.iloc[idx, 0]  # BraTS21ID\n        label = torch.tensor(self.data.iloc[idx, 1], dtype=torch.long)  # MGMT_value (0 or 1)\n\n        img_path = os.path.join(self.image_dir, f\"{brats_id:05d}_FLAIR.nii.gz\")\n\n        try:\n            image = nib.load(img_path).get_fdata()\n            image = np.nan_to_num(image, nan=0.0)\n            image = torch.tensor(image, dtype=torch.float32).unsqueeze(0)\n\n            if self.transform:\n                image = self.transform(image)\n\n        except Exception as e:\n            print(f\"Error loading {img_path}: {e}\")\n            image = torch.zeros((1, 64, 128, 128), dtype=torch.float32)\n\n        return image, label\n\n# Define fixed size for 3D CNN input\nfixed_size = (64, 128, 128)\n\n# Transformations with NaN handling\ntransform = transforms.Compose([\n    transforms.Lambda(lambda img: img.unsqueeze(0) if img.ndimension() == 3 else img),\n    transforms.Lambda(lambda img: F.interpolate(img.unsqueeze(0), size=fixed_size, mode=\"trilinear\", align_corners=False).squeeze(0)),\n    transforms.Lambda(lambda img: torch.nan_to_num(img, nan=0.0)),\n    transforms.Lambda(lambda img: (img - img.min()) / (img.max() - img.min() + 1e-8)),  # Min-Max Normalize\n    transforms.Lambda(lambda img: (img - 0.5) / 0.5),  # Scale to [-1,1]\n    transforms.Lambda(lambda img: torch.clamp(img, min=-1.0, max=1.0))  # Prevent extreme values\n])\n\n# Dataset paths\nroot_dir = '/kaggle/'\ncsv_file = '/kaggle'\n\n# Load dataset\ndataset = BrainTumorDataset(root_dir, csv_file, transform=transform)\n\n# Reduce batch size to avoid OOM errors\nbatch_size = 4  \ndata_loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, num_workers=2, pin_memory=True)\n\n# Define 3D ResNet Model\nclass BrainTumorResNet3D(nn.Module):\n    def __init__(self, num_classes=2):\n        super(BrainTumorResNet3D, self).__init__()\n        self.model = models.video.r3d_18(pretrained=True)\n        self.model.stem[0] = nn.Conv3d(\n            in_channels=1,  \n            out_channels=64, \n            kernel_size=(3, 7, 7), \n            stride=(1, 2, 2), \n            padding=(1, 3, 3), \n            bias=True\n        )\n        self.bn = nn.BatchNorm3d(64)\n        self.dropout = nn.Dropout3d(0.3)\n        self.model.fc = nn.Linear(self.model.fc.in_features, num_classes)\n\n    def forward(self, x):\n        x = self.model.stem[0](x)\n        x = self.bn(x)\n        x = self.model.stem[1](x)  \n        x = self.model.stem[2](x)  \n        x = self.model.layer1(x)\n        x = self.dropout(x)\n        x = self.model.layer2(x)\n        x = self.model.layer3(x)\n        x = self.model.layer4(x)\n        x = self.model.avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.model.fc(x)\n        return x\n\n# Set up device\ndevice = torch.device(\"cuda:1\" if torch.cuda.is_available() else \"cpu\")\nmodel = BrainTumorResNet3D(num_classes=2).to(device)\n\n# Loss, Optimizer, and Scheduler\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.1)\noptimizer = optim.Adam(model.parameters(), lr=0.0001, weight_decay=1e-5)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=2, verbose=True)\n\n# Training function with tqdm progress bar\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss, correct, total = 0.0, 0, 0\n    \n    progress_bar = tqdm(loader, desc=\"Training\", leave=False)\n    \n    for images, labels in progress_bar:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        if torch.isnan(loss):\n            print(\"NaN loss detected. Skipping batch.\")\n            continue\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = torch.max(outputs, 1)\n        correct += (predicted == labels).sum().item()\n        total += labels.size(0)\n        \n        progress_bar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{(100.0 * correct / total):.2f}%\")\n    \n    return running_loss / len(loader), 100.0 * correct / total\n\n# Training loop with tqdm epoch progress\nnum_epochs = 30  \nbest_loss = float('inf')\n\nfor epoch in range(num_epochs):\n    print(f\"\\nEpoch {epoch + 1}/{num_epochs}\")\n    train_loss, train_acc = train_epoch(model, data_loader, criterion, optimizer, device)\n    scheduler.step(train_loss)\n    print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n    \n    if train_loss < best_loss:\n        best_loss = train_loss\n        torch.save(model.state_dict(), 'best_brain_tumor_model.pth')\n        print(\"✅ Model Saved!\")\n\nprint(\"🎉 Training Complete!\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}