-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathVNET.py
More file actions
296 lines (241 loc) · 10.6 KB
/
Copy pathVNET.py
File metadata and controls
296 lines (241 loc) · 10.6 KB
1
import osimport randomfrom pathlib import Pathimport numpy as npimport torchimport torch.nn as nnimport torch.optim as optimfrom torch.utils.data import Dataset, DataLoaderimport pytorch_lightning as plfrom pytorch_lightning.callbacks import ModelCheckpointfrom pytorch_lightning.loggers import TensorBoardLoggerimport matplotlib.pyplot as pltfrom sklearn.metrics import precision_score, recall_score, f1_score, confusion_matrix, accuracy_scoreimport seaborn as snsclass LungTumorDataset(Dataset): def __init__(self, data_dir, transform=None): self.data_dir = Path(data_dir) self.ct_slices = sorted(list(self.data_dir.glob("ct_slice_*.npy"))) self.masks = sorted(list(self.data_dir.glob("mask_mask_*.npy"))) self.transform = transform def __len__(self): return len(self.ct_slices) def __getitem__(self, idx): ct_slice = np.load(self.ct_slices[idx]) mask = np.load(self.masks[idx]) ct_slice = torch.from_numpy(ct_slice).float().unsqueeze(0) mask = torch.from_numpy(mask).long() if self.transform: ct_slice = self.transform(ct_slice) return ct_slice, maskclass ConvBlock(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=5, padding=2) self.bn1 = nn.BatchNorm2d(out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=5, padding=2) self.bn2 = nn.BatchNorm2d(out_channels) self.relu = nn.ReLU(inplace=True) def forward(self, x): out = self.relu(self.bn1(self.conv1(x))) out = self.relu(self.bn2(self.conv2(out))) return outclass VNet(nn.Module): def __init__(self, in_channels=1, out_channels=2): super().__init__() self.encoder_1 = ConvBlock(in_channels, 16) self.encoder_2 = ConvBlock(16, 32) self.encoder_3 = ConvBlock(32, 64) self.encoder_4 = ConvBlock(64, 128) self.bottleneck = ConvBlock(128, 256) self.upconv_4 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2) self.decoder_4 = ConvBlock(256, 128) self.upconv_3 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2) self.decoder_3 = ConvBlock(128, 64) self.upconv_2 = nn.ConvTranspose2d(64, 32, kernel_size=2, stride=2) self.decoder_2 = ConvBlock(64, 32) self.upconv_1 = nn.ConvTranspose2d(32, 16, kernel_size=2, stride=2) self.decoder_1 = ConvBlock(32, 16) self.final_conv = nn.Conv2d(16, out_channels, kernel_size=1) self.pool = nn.MaxPool2d(2) def forward(self, x): # Encoder enc1 = self.encoder_1(x) enc2 = self.encoder_2(self.pool(enc1)) enc3 = self.encoder_3(self.pool(enc2)) enc4 = self.encoder_4(self.pool(enc3)) # Bottleneck bottleneck = self.bottleneck(self.pool(enc4)) # Decoder dec4 = self.upconv_4(bottleneck) dec4 = torch.cat((dec4, enc4), dim=1) dec4 = self.decoder_4(dec4) dec3 = self.upconv_3(dec4) dec3 = torch.cat((dec3, enc3), dim=1) dec3 = self.decoder_3(dec3) dec2 = self.upconv_2(dec3) dec2 = torch.cat((dec2, enc2), dim=1) dec2 = self.decoder_2(dec2) dec1 = self.upconv_1(dec2) dec1 = torch.cat((dec1, enc1), dim=1) dec1 = self.decoder_1(dec1) return self.final_conv(dec1)class VNetLightning(pl.LightningModule): def __init__(self, in_channels=1, out_channels=2, learning_rate=1e-3): super().__init__() self.model = VNet(in_channels, out_channels) self.learning_rate = learning_rate self.criterion = nn.CrossEntropyLoss() def forward(self, x): return self.model(x) def training_step(self, batch, batch_idx): x, y = batch y_hat = self(x) loss = self.criterion(y_hat, y) self.log('train_loss', loss) return loss def validation_step(self, batch, batch_idx): x, y = batch y_hat = self(x) loss = self.criterion(y_hat, y) self.log('val_loss', loss) return loss def configure_optimizers(self): optimizer = optim.Adam(self.parameters(), lr=self.learning_rate) return optimizerdef dice_coefficient(pred, target): smooth = 1e-5 num = pred.size(0) pred = pred.view(num, -1) target = target.view(num, -1) intersection = (pred * target).sum(1) union = pred.sum(1) + target.sum(1) dice = (2. * intersection + smooth) / (union + smooth) return dice.mean()def iou_score(pred, target, threshold=0.5): pred = (pred > threshold).float() intersection = (pred * target).sum() union = pred.sum() + target.sum() - intersection return (intersection + 1e-5) / (union + 1e-5)def compute_metrics(all_preds, all_masks, thresholds=[0.5, 0.6, 0.7, 0.8, 0.9]): metrics = {} metrics['accuracy'] = accuracy_score(all_masks, all_preds) metrics['sensitivity'] = recall_score(all_masks, all_preds, average='binary', pos_label=1) metrics['specificity'] = recall_score(all_masks, all_preds, average='binary', pos_label=0) metrics['precision'] = precision_score(all_masks, all_preds, average='binary') metrics['recall'] = recall_score(all_masks, all_preds, average='binary') metrics['f1_score'] = f1_score(all_masks, all_preds, average='binary') metrics['dcs'] = dice_coefficient(torch.tensor(all_preds), torch.tensor(all_masks)) iou_scores = [] for threshold in thresholds: iou = iou_score(torch.tensor(all_preds), torch.tensor(all_masks), threshold) iou_scores.append(iou.item()) metrics['iou_scores'] = iou_scores return metricsdef main(): # Set random seed for reproducibility pl.seed_everything(42) # Define data paths train_dir = Path("/Users/zubairsaeed/Downloads/PhD/Code/1.Segmentation/Data/Dataset/Preprocessed/segmentation_data/train") val_dir = Path("/Users/zubairsaeed/Downloads/PhD/Code/1.Segmentation/Data/Dataset/Preprocessed/segmentation_data/val") # Create datasets train_dataset = LungTumorDataset(train_dir) val_dataset = LungTumorDataset(val_dir) # Create data loaders train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=1, persistent_workers=True) val_loader = DataLoader(val_dataset, batch_size=16, shuffle=False, num_workers=1, persistent_workers=True) # Initialize model model = VNetLightning() # Define callbacks checkpoint_callback = ModelCheckpoint( dirpath='checkpoints', filename='vnet-{epoch:02d}-{val_loss:.2f}', save_top_k=3, monitor='val_loss', mode='min' ) # Define logger logger = TensorBoardLogger("lightning_logs", name="vnet") # Initialize trainer device = torch.device("mps" if torch.backends.mps.is_available() else "cuda" if torch.cuda.is_available() else "cpu") trainer = pl.Trainer( max_epochs=100, callbacks=[checkpoint_callback], logger=logger, accelerator='auto', devices=1 if device.type in ['cuda', 'mps'] else None ) # Train the model trainer.fit(model, train_loader, val_loader) # Evaluate on validation set model.eval() val_dice_score = 0 all_preds = [] all_masks = [] all_confidences = [] with torch.no_grad(): for ct_slices, masks in val_loader: ct_slices, masks = ct_slices.to(model.device), masks.to(model.device) outputs = model(ct_slices) preds = torch.argmax(outputs, dim=1) confidences = torch.softmax(outputs, dim=1)[:, 1] # Confidence for positive class val_dice_score += dice_coefficient(preds, masks) all_preds.extend(preds.cpu().numpy().flatten()) all_masks.extend(masks.cpu().numpy().flatten()) all_confidences.extend(confidences.cpu().numpy().flatten()) val_dice_score /= len(val_loader) print(f"Validation Dice Score: {val_dice_score:.4f}") # Compute metrics metrics = compute_metrics(all_preds, all_masks) for key, value in metrics.items(): if key != 'iou_scores': print(f"{key.capitalize()}: {value:.4f}") print("IOU Scores:") for threshold, iou in zip([0.5, 0.6, 0.7, 0.8, 0.9], metrics['iou_scores']): print(f" Threshold {threshold}: {iou:.4f}") # Plot confusion matrix cm = confusion_matrix(all_masks, all_preds) plt.figure(figsize=(8, 6)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues') plt.xlabel('Predicted') plt.ylabel('True') plt.title('Confusion Matrix') plt.show() # Box plot of metrics plt.figure(figsize=(10, 6)) sns.boxplot(data=[metrics['accuracy'], metrics['sensitivity'], metrics['specificity'], metrics['precision'], metrics['recall'], metrics['f1_score'], metrics['dcs']] + metrics['iou_scores']) plt.xticks(range(13), ['Accuracy', 'Sensitivity', 'Specificity', 'Precision', 'Recall', 'F1-score', 'DCS'] + [f'IOU-{t}' for t in [0.5, 0.6, 0.7, 0.8, 0.9]]) plt.title('Distribution of Metrics') plt.show() # Visualize some predictions with confidence scores def visualize_prediction(ct_slice, true_mask, pred_mask, confidence): fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(15, 5)) ax1.imshow(ct_slice.squeeze(), cmap='bone') ax1.set_title('CT Slice') ax1.axis('off') ax2.imshow(ct_slice.squeeze(), cmap='bone') ax2.imshow(true_mask, alpha=0.5, cmap='autumn') ax2.set_title('True Mask') ax2.axis('off') ax3.imshow(ct_slice.squeeze(), cmap='bone') ax3.imshow(pred_mask, alpha=0.5, cmap='autumn') ax3.set_title(f'Predicted Mask\nConfidence: {confidence:.4f}') ax3.axis('off') plt.tight_layout() plt.show() # Visualize 4 random validation predictions with confidence scores model.eval() with torch.no_grad(): indices = random.sample(range(len(val_dataset)), 4) for i in indices: ct_slice, true_mask = val_dataset[i] ct_slice = ct_slice.unsqueeze(0).to(model.device) output = model(ct_slice) pred_mask = torch.argmax(output, dim=1).squeeze().cpu().numpy() confidence = torch.softmax(output, dim=1)[0, 1].item() # Confidence for positive class visualize_prediction(ct_slice.cpu().squeeze(), true_mask.squeeze(), pred_mask, confidence)if __name__ == '__main__': main()