A deep representation learning model to predict response to vagus nerve stimulation.
The 3 matches
- [1] § Results › Latent representations support accurate prediction and spatial interpretation of VNS outcome ↔ evaluation/eval_svm.py, lines 333–386 · score 0.65 · dot product, Grad CAM, SHAP, activations, zeros, Gradients
- [2] § Results › Self-supervised representation learning captures anatomical patterns associated with VNS outcome ↔ demo/utils/loss_functions.py, lines 83–135 · score 0.58 · spectral loss, loss function, perceptual loss, VQ VAE, reconstructions, predict
- [3] § Results › Self-supervised representation learning captures anatomical patterns associated with VNS outcome ↔ utils/loss_functions.py, lines 83–135 · score 0.58 · spectral loss, loss function, perceptual loss, VQ VAE, reconstructions, predict
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 631 lines · 26 KB · no license · 1 match
- ##### This script will evaluate the perforamnce of the SVM and then grad cam and other stuff
- #### Hrishikesh Suresh
- #### Ibrahim Lab 2025
- # Import necessary libraries
- import sys
- import os
- if not os.path.isdir('/hpf/projects/'):
- sys.path.insert(0, '../')
- import torch
- import lightning.pytorch as pl
- from torch.utils.data import DataLoader
- from torchvision import transforms
- from sklearn.model_selection import train_test_split
- import pandas as pd
- import numpy as np
- import matplotlib.pyplot as plt
- import seaborn as sns
- import nibabel as nib
- from tqdm.auto import tqdm
- from datasets.t1_datamodule import ImagingDataModule
- from sklearn.model_selection import train_test_split
- from models.vqvae import BrainVQVAE
- from models.classifiers import TransformerEncoderClassifier
- import wandb
- from lightning.pytorch.callbacks import EarlyStopping, ModelCheckpoint
- from lightning.pytorch.loggers import WandbLogger, TensorBoardLogger
- import shutil
- from monai.networks.nets.patchgan_discriminator import PatchDiscriminator
- from utils.loss_functions import Recon_Loss, Recon_Loss_GradMod, gradnorm
- import yaml
- import argparse
- torch.set_float32_matmul_precision('high')
- import pickle
- import warnings
- import shap
- from sklearn.metrics import roc_curve, confusion_matrix, roc_auc_score
- from utils.misc_functions import save_model_space_image
- warnings.filterwarnings(
- "ignore",
- message=".*`torch.cuda.amp.autocast*",
- category=FutureWarning
- )
- if os.environ.get('DEBUGGING') == '1':
- DEBUGGING = True
- else:
- DEBUGGING = False
- #Plot it all nicely
- def plot_curve_and_matrix(preds, probas, labels, dataset_name, intermediates_dir):
- #Adjust samplew weight to account for the imbalance
- non_responder_count = np.sum(labels == 0)
- responder_count = np.sum(labels == 1)
- pos_weight = non_responder_count / responder_count
- sample_weight = np.ones_like(labels)
- sample_weight[labels == 1] = pos_weight
- fpr, tpr, thresholds = roc_curve(labels, probas, sample_weight=sample_weight)
- fig, ax = plt.subplots(1, 2, figsize=(12, 6))
- sns.heatmap(confusion_matrix(labels, preds), annot=True, fmt='d', ax=ax[0], cmap='Blues')
- ax[0].set_title(f'{dataset_name} Confusion Matrix')
- ax[0].set_xlabel('Predicted')
- ax[0].set_ylabel('True')
- auc = roc_auc_score(labels, probas, sample_weight=sample_weight)
- ax[1].plot(fpr, tpr, label=f'AUC: {auc:.4f}')
- ax[1].plot([0, 1], [0, 1], linestyle='--', label='Random Classifier')
- ax[1].set_title(f'{dataset_name} ROC Curve')
- ax[1].set_xlabel('False Positive Rate')
- ax[1].set_ylabel('True Positive Rate')
- #add legend for random classifier and the model
- ax[1].legend(loc='lower right')
- plt.savefig(os.path.join(intermediates_dir, f'{dataset_name}_roc_curve.png'))
- plt.show()
- #Print the optimal threshold
- optimal_idx = np.argmax(tpr - fpr)
- optimal_threshold = thresholds[optimal_idx]
- return auc, optimal_threshold
- def eval(config):
- images_dir = str(config['images_dir'])
- metadata_csv = str(config['metadata_csv'])
- intermediates_dir= str(config['intermediates_dir'])
- outcome_col = str(config['outcome_col'])
- if os.path.exists(intermediates_dir):
- shutil.rmtree(intermediates_dir)
- os.makedirs(intermediates_dir)
- # Load the metadata
- df = pd.read_csv(metadata_csv)
- if not outcome_col in df.columns:
- raise ValueError('The metadata file must contain an outcome column. Since this is a classification task')
- if len(np.unique(df[outcome_col])) != 2:
- if 'binarization_threshold' in config:
- df['raw_outcome'] = df[outcome_col].copy()
- df[outcome_col] = df[outcome_col].apply(lambda x: 1 if x > float(config['binarization_threshold']) else 0)
- else:
- raise ValueError('The outcome column must be binary. Please binarize the outcome columnor provide a binarization threshold in the config file')
- # Define universal params
- batch_size = int(config['batch_size'])
- num_workers = int(config['num_workers'])
- #Set up the data module
- use_clinical = bool(config['use_clinical'])
- use_train_transform = bool(config['use_train_transform'])
- preload = bool(config['preload'])
- image_dim = tuple(map(int, config['image_dim']))
- #Check if the debugger is on, if so make num_workers = 1
- if DEBUGGING:
- num_workers = 1
- #Load the checkpoint
- vqvae = BrainVQVAE.load_from_checkpoint(config['vqvae_checkpoint_path'])
- vqvae.intermediates_dir = intermediates_dir
- #Load the SVM classifier
- classifier = pickle.load(open(config['svm_path'], 'rb'))
- features_to_keep = np.load(config['svm_feature_indices'])
- data_module = ImagingDataModule(images_dir = images_dir,
- metadata_df = df,
- train_ids = [],
- val_ids = [],
- test_ids = df['study_id'].values,
- batch_size = batch_size,
- num_workers = num_workers,
- use_clinical = use_clinical,
- use_train_transform = use_train_transform,
- preload = preload,
- crop_or_pad_dim=image_dim,
- outcome_col=outcome_col)
- data_module.setup()
- test_loader = data_module.test_dataloader()
- vqvae.eval()
- predictions = []
- probabilities = []
- labels = []
- with torch.no_grad():
- for batch in tqdm(test_loader, desc='Generating predictions', total=len(test_loader)):
- x, label = batch
- x = x.to(vqvae.device)
- label = label.to(vqvae.device)
- latents = vqvae.encode(x)
- quantized_latents = vqvae.quantize(latents)[0]
- mean_latent = torch.mean(quantized_latents, dim=[2,3,4])
- detached_latent = mean_latent.detach().cpu().numpy()[:, features_to_keep].reshape(1, -1)
- outcome_pred = classifier.predict(detached_latent)
- outcome_prob = classifier.predict_proba(detached_latent)[:,1]
- predictions.append(outcome_pred)
- probabilities.append(outcome_prob)
- labels.append(label.cpu().detach().numpy())
- #Save the predictions and labels
- predictions = np.array(predictions)
- probabilities = np.array(probabilities)
- labels = np.array(labels)
- np.save(os.path.join(intermediates_dir, 'predictions.npy'), predictions)
- np.save(os.path.join(intermediates_dir, 'probabilities.npy'), probabilities)
- np.save(os.path.join(intermediates_dir, 'labels.npy'), labels)
- #Plot the confusion matrix and ROC curve
- auc, optimal_threshold = plot_curve_and_matrix(predictions, probabilities, labels, 'test', intermediates_dir)
- print(f'Optimal threshold: {optimal_threshold}')
- print(f'AUC: {auc}')
- def run_gradcam_analysis(config):
- images_dir = str(config['images_dir'])
- intermediates_dir= str(config['intermediates_dir'])
- if not os.path.exists(intermediates_dir):
- os.makedirs(intermediates_dir)
- train_df = pd.read_csv(os.path.join(config['dataframes_folder'], 'train_df.csv'))
- val_df = pd.read_csv(os.path.join(config['dataframes_folder'], 'val_df.csv'))
- test_df = pd.read_csv(os.path.join(config['dataframes_folder'], 'test_df.csv'))
- train_df['split'] = 'train'
- val_df['split'] = 'val'
- test_df['split'] = 'test'
- df = pd.concat([train_df, val_df, test_df], axis=0)
- if not 'outcome' in df.columns:
- raise ValueError('The metadata file must contain an outcome column. Since this is a classification task')
- # Define universal params
- batch_size = int(config['batch_size'])
- num_workers = int(config['num_workers'])
- #Set up the data module
- use_clinical = bool(config['use_clinical'])
- use_train_transform = bool(config['use_train_transform'])
- preload = bool(config['preload'])
- image_dim = tuple(map(int, config['image_dim']))
- outcome_col = str(config['outcome_col'])
- #Check if the debugger is on, if so make num_workers = 1
- if DEBUGGING:
- num_workers = 1
- data_module = ImagingDataModule(images_dir = images_dir,
- metadata_df = df,
- train_ids = [],
- val_ids = df[df['split'] == 'train']['study_id'].values,
- test_ids = df['study_id'].values,
- batch_size = 1,
- num_workers = num_workers,
- use_clinical = use_clinical,
- use_train_transform = use_train_transform,
- preload = preload,
- crop_or_pad_dim=image_dim,
- outcome_col=outcome_col)
- #Load the checkpoint
- vqvae = BrainVQVAE.load_from_checkpoint(config['vqvae_checkpoint_path'])
- vqvae.intermediates_dir = intermediates_dir
- #Load the SVM classifier
- classifer = pickle.load(open(config['svm_path'], 'rb'))
- features_to_keep = np.load(config['svm_feature_indices'])
- #Target layer for gradcam
- target_layer = vqvae.model.encoder.blocks[16].conv
- #Get the dataloaders
- data_module.setup()
- test_loader = data_module.test_dataloader()
- train_loader = data_module.val_dataloader()
- vqvae.eval()
- train_predictions = []
- full_quantized_latents = []
- grad_cam_dir = os.path.join(intermediates_dir, 'gradcam_v2')
- if not os.path.exists(grad_cam_dir):
- os.makedirs(grad_cam_dir)
- with torch.no_grad():
- #First run all the training data (in the validation data loader) to get the median of each feature
- for batch in tqdm(train_loader, desc='Generating predictions for training data', total=len(train_loader)):
- x, label = batch
- x = x.to(vqvae.device)
- label = label.to(vqvae.device)
- latents = vqvae.encode(x)
- quantized_latents = vqvae.quantize(latents)[0]
- mean_latent = torch.mean(quantized_latents, dim=[2,3,4])
- mean_latent = mean_latent.view(1, 32, -1)
- detached_latent = mean_latent.detach().cpu().numpy().reshape(-1)[features_to_keep]
- train_predictions.append(detached_latent)
- full_quantized_latents.append(quantized_latents.detach().cpu().numpy())
- #Make it a numpy array
- train_predictions = np.array(train_predictions)
- full_quantized_latents = np.array(full_quantized_latents).squeeze()
- #Get the median of each feature
- median_latent = np.median(train_predictions, axis=0).reshape(1, -1)
- #Make the explainer
- def get_pred(x):
- return classifer.predict_proba(x)[:,1]
- explainer = shap.Explainer(get_pred, median_latent)
- activations = []
- gradients = []
- def forward_hook(module, input, output):
- activations.append(output)
- def backward_hook(module, grad_input, grad_output):
- gradients.append(grad_output[0])
- target_layer.register_forward_hook(forward_hook)
- target_layer.register_full_backward_hook(backward_hook)
- latent_heatmaps = []
- heatmaps = []
- images = []
- labels = []
- preds = []
- for batch in tqdm(test_loader, desc='Generating gradient maps', total=len(test_loader)):
- activations.clear()
- gradients.clear()
- x, label = batch
- x = x.to(vqvae.device)
- label = label.to(vqvae.device)
- latents = vqvae.encode(x)
- quantized_latents = vqvae.quantize(latents)[0]
- mean_latent = torch.mean(quantized_latents, dim=[2,3,4])
- mean_latent = torch.abs(mean_latent) # Take the absolute value of the mean latent
- mean_latent = mean_latent.view(1, 32, -1)
- detached_latent = mean_latent.detach().cpu().numpy()[:, features_to_keep, :].reshape(1, -1)
- output = classifer.predict_proba(detached_latent)[:,1]
- shap_values = explainer(detached_latent)
- shap_values = np.array(shap_values.values).reshape(1, -1)
- # Clone the mean latent
- shap_weight_latent = mean_latent.clone()
- shap_weight_latent[:, :, :] = 0 # Zero out all values
- shap_weight_latent[:, features_to_keep, :] = torch.tensor(shap_values.reshape(1, -1, 1), device=shap_weight_latent.device, dtype=shap_weight_latent.dtype) # Set SHAP values
- dot_product = torch.dot(mean_latent.view(-1), shap_weight_latent.view(-1))
- vqvae.zero_grad()
- dot_product.backward()
- import numpy as np
- weights = torch.mean(gradients[0], dim=[2,3,4]).view(1,32,1,1,1)
- heatmap = torch.sum(weights * activations[0], dim=1).squeeze()
- heatmap = heatmap.cpu().detach().numpy()
- heatmap /= np.max(np.abs(heatmap))
- upsampled_heamap = torch.nn.Upsample(size=(x.shape[2], x.shape[3], x.shape[4]), mode='trilinear')(torch.tensor(heatmap).unsqueeze(0).unsqueeze(0)).squeeze()
- squeezed_img = x.squeeze().cpu().detach().numpy()
- latent_heatmaps.append(heatmap)
- heatmaps.append(upsampled_heamap)
- images.append(squeezed_img)
- preds.append(output)
- labels.append(label.cpu().detach().numpy())
- #Save the heatmaps and images
- latent_heatmaps = np.array(latent_heatmaps)
- heatmaps = np.array(heatmaps)
- images = np.array(images)
- preds = np.array(preds)
- labels = np.array(labels)
- latent_heatmaps_dir = os.path.join(grad_cam_dir, 'latent_heatmaps')
- heatmaps_dir = os.path.join(grad_cam_dir, 'heatmaps')
- images_dir = os.path.join(grad_cam_dir, 'images')
- if not os.path.exists(latent_heatmaps_dir):
- os.makedirs(latent_heatmaps_dir)
- if not os.path.exists(heatmaps_dir):
- os.makedirs(heatmaps_dir)
- if not os.path.exists(images_dir):
- os.makedirs(images_dir)
- for idx in tqdm(range(len(heatmaps)), desc='Saving Heatmaps and Images'):
- save_model_space_image(heatmaps[idx], os.path.join(heatmaps_dir, f'heatmap_{idx}.nii.gz'))
- save_model_space_image(images[idx], os.path.join(images_dir, f'image_{idx}.nii.gz'))
- latent_heatmap = latent_heatmaps[idx].squeeze()
- latent_heatmap_img = nib.Nifti1Image(latent_heatmap, np.eye(4))
- nib.save(latent_heatmap_img, os.path.join(latent_heatmaps_dir, f'latent_heatmap_{idx}.nii.gz'))
- #Save the predictions and labels
- np.save(os.path.join(grad_cam_dir, 'preds.npy'), preds)
- np.save(os.path.join(grad_cam_dir, 'labels.npy'), labels)
- registered_images_dir = os.path.join(grad_cam_dir, 'registered_images')
- registered_heatmaps_dir = os.path.join(grad_cam_dir, 'registered_heatmaps')
- registration_transform_dir = os.path.join(grad_cam_dir, 'registration_transforms')
- if not os.path.exists(registration_transform_dir):
- os.makedirs(registration_transform_dir)
- if not os.path.exists(registered_images_dir):
- os.makedirs(registered_images_dir)
- if not os.path.exists(registered_heatmaps_dir):
- os.makedirs(registered_heatmaps_dir)
- #Register the images and heatmaps
- import ants
- mni_1mm_brain = '/usr/local/fsl/data/standard/MNI152_T1_1mm_brain.nii.gz'
- mni_1mm_brain = ants.image_read(mni_1mm_brain)
- for idx in tqdm(range(len(os.listdir(heatmaps_dir))), desc='Registering Images and Heatmaps'):
- image = ants.image_read(os.path.join(images_dir, f'image_{idx}.nii.gz'))
- heatmap = ants.image_read(os.path.join(heatmaps_dir, f'heatmap_{idx}.nii.gz'))
- registered_image = ants.registration(fixed=mni_1mm_brain, moving=image, type_of_transform='SyN')
- registered_heatmap = ants.apply_transforms(fixed=mni_1mm_brain, moving=heatmap, transformlist=registered_image['fwdtransforms'])
- ants.image_write(registered_image['warpedmovout'], os.path.join(registered_images_dir, f'registered_image_{idx}.nii.gz'))
- ants.image_write(registered_heatmap, os.path.join(registered_heatmaps_dir, f'registered_heatmap_{idx}.nii.gz'))
- for transform in registered_image['fwdtransforms']:
- if '.nii.gz' in transform:
- shutil.move(transform, os.path.join(registration_transform_dir, f'warp_{idx}.nii.gz'))
- elif '.mat' in transform:
- shutil.move(transform, os.path.join(registration_transform_dir, f'affine_{idx}.mat'))
- else:
- raise ValueError('Unknown transform type')
- for transform in registered_image['invtransforms']:
- if '.nii.gz' in transform:
- shutil.move(transform, os.path.join(registration_transform_dir, f'inv_warp_{idx}.nii.gz'))
- responders = np.where(labels == 1)[0]
- non_responders = np.where(labels == 0)[0]
- responder_images = np.zeros(mni_1mm_brain.shape)
- for idx in tqdm(responders, desc='Averaging Responder Images'):
- image = ants.image_read(os.path.join(registered_heatmaps_dir, f'registered_heatmap_{idx}.nii.gz'))
- responder_images += image.numpy()
- responder_images /= len(responders)
- non_responder_images = np.zeros(mni_1mm_brain.shape)
- for idx in tqdm(non_responders, desc='Averaging Non-Responder Images'):
- image = ants.image_read(os.path.join(registered_heatmaps_dir, f'registered_heatmap_{idx}.nii.gz'))
- non_responder_images += image.numpy()
- non_responder_images /= len(non_responders)
- responder_average = ants.from_numpy(responder_images, origin=mni_1mm_brain.origin, spacing=mni_1mm_brain.spacing, direction=mni_1mm_brain.direction)
- non_responder_average = ants.from_numpy(non_responder_images, origin=mni_1mm_brain.origin, spacing=mni_1mm_brain.spacing, direction=mni_1mm_brain.direction)
- ants.image_write(responder_average, os.path.join(grad_cam_dir, 'responder_average.nii.gz'))
- ants.image_write(non_responder_average, os.path.join(grad_cam_dir, 'non_responder_average.nii.gz'))
- all_subjects_average = np.zeros(mni_1mm_brain.shape)
- for idx in tqdm(range(len(os.listdir(registered_heatmaps_dir))), desc='Averaging All Subjects Images'):
- image = ants.image_read(os.path.join(registered_heatmaps_dir, f'registered_heatmap_{idx}.nii.gz'))
- all_subjects_average += image.numpy()
- all_subjects_average /= len(os.listdir(registered_heatmaps_dir))
- all_subjects_average = ants.from_numpy(all_subjects_average, origin=mni_1mm_brain.origin, spacing=mni_1mm_brain.spacing, direction=mni_1mm_brain.direction)
- ants.image_write(all_subjects_average, os.path.join(grad_cam_dir, 'all_subjects_average.nii.gz'))
- all_subjects_abs_average = np.abs(all_subjects_average.numpy())
- all_subjects_abs_average_img = ants.from_numpy(all_subjects_abs_average, origin=mni_1mm_brain.origin, spacing=mni_1mm_brain.spacing, direction=mni_1mm_brain.direction)
- ants.image_write(all_subjects_abs_average_img, os.path.join(grad_cam_dir, 'all_subjects_abs_average.nii.gz'))
- all_subjects_abs_average_masked = all_subjects_abs_average * (mni_1mm_brain.numpy() > 0).astype(np.float32)
- all_subjects_abs_average_masked_img = ants.from_numpy(all_subjects_abs_average_masked, origin=mni_1mm_brain.origin, spacing=mni_1mm_brain.spacing, direction=mni_1mm_brain.direction)
- ants.image_write(all_subjects_abs_average_masked_img, os.path.join(grad_cam_dir, 'all_subjects_abs_average_masked.nii.gz'))
- def median_decoding(config):
- '''
- This function decodes and saves the median responder and non-responder images for contrastive analysis.
- '''
- images_dir = str(config['images_dir'])
- intermediates_dir = str(config['intermediates_dir'])
- if not os.path.exists(intermediates_dir):
- os.makedirs(intermediates_dir)
- train_df = pd.read_csv(os.path.join(config['dataframes_folder'], 'train_df.csv'))
- val_df = pd.read_csv(os.path.join(config['dataframes_folder'], 'val_df.csv'))
- test_df = pd.read_csv(os.path.join(config['dataframes_folder'], 'test_df.csv'))
- train_df['split'] = 'train'
- val_df['split'] = 'val'
- test_df['split'] = 'test'
- df = pd.concat([train_df, val_df, test_df], axis=0)
- if not 'outcome' in df.columns:
- raise ValueError('The metadata file must contain an outcome column. Since this is a classification task')
- # Define universal params
- batch_size = int(config['batch_size'])
- num_workers = int(config['num_workers'])
- # Set up the data module
- use_clinical = bool(config['use_clinical'])
- use_train_transform = bool(config['use_train_transform'])
- preload = bool(config['preload'])
- image_dim = tuple(map(int, config['image_dim']))
- outcome_col = str(config['outcome_col'])
- # Check if the debugger is on, if so make num_workers = 1
- if DEBUGGING:
- num_workers = 1
- # Load the checkpoint
- vqvae = BrainVQVAE.load_from_checkpoint(config['vqvae_checkpoint_path'])
- vqvae.intermediates_dir = intermediates_dir
- # Load the SVM classifier
- classifier = pickle.load(open(config['svm_path'], 'rb'))
- features_to_keep = np.load(config['svm_feature_indices'])
- data_module = ImagingDataModule(
- images_dir=images_dir,
- metadata_df=df,
- train_ids=[],
- val_ids=df[df['split'] == 'train']['study_id'].values,
- test_ids=df['study_id'].values,
- batch_size=1,
- num_workers=num_workers,
- use_clinical=use_clinical,
- use_train_transform=use_train_transform,
- preload=preload,
- crop_or_pad_dim=image_dim,
- outcome_col=outcome_col
- )
- # Get the dataloaders
- data_module.setup()
- train_loader = data_module.val_dataloader() #
- train_mean_latent_predictions = []
- encoded_volumes = []
- train_labels = []
- grad_cam_dir = os.path.join(intermediates_dir, 'gradcam')
- if not os.path.exists(grad_cam_dir):
- os.makedirs(grad_cam_dir)
- with torch.no_grad():
- # Run all the training data (in the validation data loader) to get the median of each feature
- for batch in tqdm(train_loader, desc='Generating predictions for training data', total=len(train_loader)):
- x, label = batch
- x = x.to(vqvae.device)
- label = label.to(vqvae.device)
- latents = vqvae.encode(x)
- quantized_latents = vqvae.quantize(latents)[0]
- mean_latent = torch.mean(quantized_latents, dim=[2,3,4])
- mean_latent = mean_latent.view(1, 32, -1)
- detached_latent = mean_latent.detach().cpu().numpy().reshape(-1)[features_to_keep]
- train_mean_latent_predictions.append(detached_latent)
- encoded_volumes.append(latents.detach())
- train_labels.append(label.detach())
- # Make it a numpy array
- train_mean_latent_predictions = torch.tensor(train_mean_latent_predictions)
- train_labels = torch.tensor(train_labels)
- # Generate mean responder and non-responder latents
- encoded_volumes = torch.stack(encoded_volumes).squeeze()
- responder_indices = torch.where(train_labels == 1)[0]
- non_responder_indices = torch.where(train_labels == 0)[0]
- responder_encoded = encoded_volumes[responder_indices]
- non_responder_encoded = encoded_volumes[non_responder_indices]
- median_responder_encoded = torch.median(responder_encoded, dim=0).values.unsqueeze(0)
- median_non_responder_encoded = torch.median(non_responder_encoded, dim=0).values.unsqueeze(0)
- median_responder_quantized = vqvae.quantize(median_responder_encoded)[0]
- median_non_responder_quantized = vqvae.quantize(median_non_responder_encoded)[0]
- responder_decoded = vqvae.decode(median_responder_quantized)
- non_responder_decoded = vqvae.decode(median_non_responder_quantized)
- save_model_space_image(responder_decoded.squeeze().cpu().detach().numpy(), os.path.join(grad_cam_dir, 'median_responder_image.nii.gz'))
- save_model_space_image(non_responder_decoded.squeeze().cpu().detach().numpy(), os.path.join(grad_cam_dir, 'median_non_responder_image.nii.gz'))
- def main():
- parser = argparse.ArgumentParser(description='Evaluate SVM, run Grad-CAM analysis, or median decoding for VQ-VAE model')
- parser.add_argument('--config', type=str, required=True, help='Path to the config file')
- parser.add_argument('--svm_path', type=str, help='Path to the svm checkpoint', required=False)
- parser.add_argument('--svm_feature_indices', type=str, help='Path to the svm feature indices', required=False)
- parser.add_argument('--dataframes_folder', type=str, help='Path to the dataframes folder with train.df, val.df, test.df', required=False)
- parser.add_argument('--mode', type=str, choices=['eval', 'gradcam', 'median_decoding'], default='eval', help="Which mode to run: 'eval' for evaluation, 'gradcam' for Grad-CAM analysis, 'median_decoding' to save median responder/non-responder images")
- # Accept any additional arguments
- parser.add_argument('args', nargs=argparse.REMAINDER)
- args, unknown = parser.parse_known_args()
- config_file = args.config
- # Load the config file
- with open(config_file) as file:
- config = yaml.load(file, Loader=yaml.FullLoader)
- # Override config with command-line arguments if provided
- if args.svm_path:
- config['svm_path'] = args.svm_path
- if args.svm_feature_indices:
- config['svm_feature_indices'] = args.svm_feature_indices
- if args.dataframes_folder:
- config['dataframes_folder'] = args.dataframes_folder
- # Copy the config file to the intermediates directory
- intermediates_dir = str(config['intermediates_dir'])
- if not os.path.exists(intermediates_dir):
- os.makedirs(intermediates_dir)
- shutil.copy(config_file, os.path.join(intermediates_dir, 'config.yaml'))
- if args.mode == 'eval':
- eval(config)
- elif args.mode == 'gradcam':
- run_gradcam_analysis(config)
- elif args.mode == 'median_decoding':
- median_decoding(config)
- else:
- raise ValueError(f"Unknown mode: {args.mode}")
- if __name__ == '__main__':
- main()
eval_svm.py at commit ab62c28, no license · at the source
Overview
and 27 other authors
Emefa Akwayena9, Dewi Schrader10, Robert J. Bollo11, Matthew D. Smyth12, Diana Aum13, Sean M. Lew14, Shelly Wang15, Toba N. Niazi15, Aria Fallah16, Jeffrey S. Raskin17, Howard L. Weiner18, Nisha Gadgil18, Gregory W. Albert19, Aristides Hadjinicolaou20, Philippe Major20, Farbod Niazi21, Guillaume Theaud21, Sami Obaid22, Elysa Widjaja23, Birgit Ertl-Wagner24, Logi Vidarsson24, Margot J. Taylor24, Alexandre Boutet25, James T. Rutka3,26, Melissa A. LoPresti27, Puneet Jain6, George M. Ibrahim1,2,3,26,2828 affiliations
- Institute of Biomedical Engineering, University of Toronto,Toronto, ON Canada
- Program in Neuroscience and Mental Health, The Hospital for Sick Children Research Institute,Toronto, ON Canada
- Division of Neurosurgery, Department of Surgery, University of Toronto,Toronto, ON Canada
- Department of Pediatrics, Cincinnati Children’s Hospital,Cincinnati, OH USA
- Krembil Brain Institute,Toronto, ON Canada
- Division of Neurology, The Hospital for Sick Children,Toronto, ON Canada
- Division of Neurosurgery, CHU Sainte-Justine and Centre hospitalier de l’Université de Montréal,Montréal, QC Canada
- Deparment of Neurosurgery, Riley Hospital for Children,Indianapolis, IN USA
- Department of Neurosurgery, UPMC Children’s Hospital of Pittsburgh,Pittsburgh, PA USA
- Division of Neurology, BC Children’s Hospital,Vancouver, BC Canada
- Department of Neurosurgery, University of Utah Health,Salt Lake City, UT USA
- Johns Hopkins University, Department of Neurosurgery, Johns Hopkins All Children’s Hospital,St. Petersburg, FL USA
- Department of Neurosurgery, St. Louis Children’s Hospital,St. Louis, MO USA
- Department of Neurosurgery, Medical College Wisconsin,Milwaukee, WI USA
- Division of Neurosurgery, Nicklaus Children’s Hospital,Miami, FL USA
- Department of Neurosurgery, UCLA Mattel Children’s Hospital,Los Angeles, CA USA
- Department of Neurosurgery, Children’s Hospital of Chicago,Chicago, IL USA
- Department of Neurosurgery, Texas Children’s Hospital,Houston, TX USA
- Department of Neurosurgery, Arkansas Children’s Hospital,Little Rock, AR USA
- Division of Neurology, CHU Sainte-Justine,Montréal, QC Canada
- Centre de recherche du CHUM,Montréal, QC Canada
- Department of Surgery, Université de Montréal,Montréal, QC Canada
- Department of Medical Imaging, Children’s Hospital of Chicago,Chicago, IL USA
- Department of Diagnostic & Interventional Radiology, The Hospital for Sick Children,Toronto, ON Canada
- Joint Department of Medical Imaging, University Health Network,Toronto, ON Canada
- Division of Neurosurgery, The Hospital for Sick Children,Toronto, ON Canada
- Department of Neurosurgery, University of Rochester Medical Center,Rochester, NY USA
- Institute of Medical Science, University of Toronto,Toronto, ON Canada
Abstract
The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.
Repositories
Its files are read in the Code ↔ Paper reader above, with 3 matches between paragraphs and lines of code.
gmilab/VQVNS
ab62c28f297629eb8b0d3dd716101cd2182636ff, 25 November 2025Availability: 1 check, the latest on 29 September 2026: the link answers
- 29 September 2026: the link answers
25 files
- datasets/
t1_datamodule.py , Python, 165 lines - demo/
generic_t1_preprocess.py , Python, 290 lines - demo/
infer_single_subject.py , Python, 83 lines - demo/
model_inference_demo.ipy , Jupyter, 186 linesnb - demo/
models/ , Python, 2 lines__init__.py - demo/
models/ , Python, 182 linesquantizers.py - demo/
models/ , Python, 120 linesvae.py - demo/
models/ , Python, 356 linesvqvae.py - demo/
utils/ , Python, 210 lines, 1 matchloss_functions.py - demo/
utils/ , Python, 123 linesmisc_functions.py - evaluation/
eval_recon_fidelity.py , Python, 153 lines - evaluation/
eval_svm.py , Python, 631 lines, 1 match - evaluation/
eval_vqvae.py , Python, 216 lines - infer_single_subject.py, Python, 74 lines
- models/
quantizers.py , Python, 182 lines - models/
vqvae.py , Python, 371 lines - training/
train_svm_classifier.py , Python, 528 lines - training/
train_vqvae.py , Python, 309 lines - utils/
identify_smallest_image_ , Python, 72 linescrop.py - utils/
loss_functions.py , Python, 210 lines, 1 match - utils/
metrics.py , Python, 50 lines - utils/
misc_functions.py , Python, 167 lines - utils/
mni_dbm_vns_transforms.p , Python, 96 linesy - utils/
preprocessing/ , Python, 279 linesgeneric_t1_preprocess.py - README.md, Text, 163 lines
huggingface.co/hsuresh/vqvns
75438ce6e41ad6eecb5c2cbfa601aa20f199e050, 5 November 2025Availability: 1 check, the latest on 29 September 2026: the link answers
- 29 September 2026: the link answers
1 file
- README.md, Text, 32 lines
Zenodo 18510266
Availability: 1 check, the latest on 29 September 2026: the link answers (HTTP 200)
- 29 September 2026: the link answers (HTTP 200)
25 files
- datasets/
t1_datamodule.py , Python, 165 lines - demo/
generic_t1_preprocess.py , Python, 290 lines - demo/
infer_single_subject.py , Python, 83 lines - demo/
model_inference_demo.ipy , Jupyter, 186 linesnb - demo/
models/ , Python, 2 lines__init__.py - demo/
models/ , Python, 182 linesquantizers.py - demo/
models/ , Python, 120 linesvae.py - demo/
models/ , Python, 356 linesvqvae.py - demo/
utils/ , Python, 210 linesloss_functions.py - demo/
utils/ , Python, 123 linesmisc_functions.py - evaluation/
eval_recon_fidelity.py , Python, 153 lines - evaluation/
eval_svm.py , Python, 631 lines - evaluation/
eval_vqvae.py , Python, 216 lines - infer_single_subject.py, Python, 74 lines
- models/
quantizers.py , Python, 182 lines - models/
vqvae.py , Python, 371 lines - training/
train_svm_classifier.py , Python, 528 lines - training/
train_vqvae.py , Python, 309 lines - utils/
identify_smallest_image_ , Python, 72 linescrop.py - utils/
loss_functions.py , Python, 210 lines - utils/
metrics.py , Python, 50 lines - utils/
misc_functions.py , Python, 167 lines - utils/
mni_dbm_vns_transforms.p , Python, 96 linesy - utils/
preprocessing/ , Python, 279 linesgeneric_t1_preprocess.py - README.md, Text, 163 lines
Code availability statement
The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:
- it points to the authors' code: gmilab/
VQVNS , huggingface.co/hsuresh/ vqvns
Read it in the paper: doi.org/10.1038/s41467-026-71555-0.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 3 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 48 scripts, each with its path and the digest of its content;
- 3 matches between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
Datasets cited
- doi:10.18112/
openneuro.ds005602.v1.0. , at OpenNeuro; found in the references0 - doi:10.7910/
dvn/ , at the source; found in the referencesilxiks
Data availability statement
The paper has a data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:
- no repository, dataset or request procedure was recognized in it
Read it in the paper: doi.org/10.1038/s41467-026-71555-0.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 29 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 47 authors, 5 keywords, 14 MeSH terms, 1 funder, 89 references.
Cite
This paper
Suresh, H., Mithani, K., Li, V., Latypov, T. H., Warsi, N. M., Wong, S. M., Erdman, L., Kang, J., Germann, J., Gouveia, F. V., Coleman, S. C., Berger, A., Chau, V., Weiss, S., Gorodetsky, C., Donner, E., Weil, A. G., Tailor, J., Abel, T. J., . . . Ibrahim, G. M. (2026). A deep representation learning model to predict response to vagus nerve stimulation. Nature communications, 17(1), 4932. https://
BibTeX
@article{suresh2026deep,
author = {Suresh, Hrishikesh and Mithani, Karim and Li, Vicki and Latypov, Timur H. and Warsi, Nebras M. and Wong, Simeon M. and Erdman, Lauren and Kang, Jaeyoung and Germann, Jurgen and Gouveia, Flavia Venetucci and Coleman, Sebastian C. and Berger, Alexandre and Chau, Vann and Weiss, Shelly and Gorodetsky, Carolina and Donner, Elizabeth and Weil, Alexander G. and Tailor, Jignesh and Abel, Taylor J. and Remick, Madison and Akwayena, Emefa and Schrader, Dewi and Bollo, Robert J. and Smyth, Matthew D. and Aum, Diana and Lew, Sean M. and Wang, Shelly and Niazi, Toba N. and Fallah, Aria and Raskin, Jeffrey S. and Weiner, Howard L. and Gadgil, Nisha and Albert, Gregory W. and Hadjinicolaou, Aristides and Major, Philippe and Niazi, Farbod and Theaud, Guillaume and Obaid, Sami and Widjaja, Elysa and Ertl-Wagner, Birgit and Vidarsson, Logi and Taylor, Margot J. and Boutet, Alexandre and Rutka, James T. and LoPresti, Melissa A. and Jain, Puneet and Ibrahim, George M.},
title = {{A deep representation learning model to predict response to vagus nerve stimulation}},
journal = {Nature communications},
year = {2026},
month = apr,
volume = {17},
number = {1},
pages = {4932},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/
url = {https://
pmid = {41946715},
pmcid = {PMC13234157}
}
RIS
TY - JOUR
AU - Suresh, Hrishikesh
AU - Mithani, Karim
AU - Li, Vicki
AU - Latypov, Timur H.
AU - Warsi, Nebras M.
AU - Wong, Simeon M.
AU - Erdman, Lauren
AU - Kang, Jaeyoung
AU - Germann, Jurgen
AU - Gouveia, Flavia Venetucci
AU - Coleman, Sebastian C.
AU - Berger, Alexandre
AU - Chau, Vann
AU - Weiss, Shelly
AU - Gorodetsky, Carolina
AU - Donner, Elizabeth
AU - Weil, Alexander G.
AU - Tailor, Jignesh
AU - Abel, Taylor J.
AU - Remick, Madison
AU - Akwayena, Emefa
AU - Schrader, Dewi
AU - Bollo, Robert J.
AU - Smyth, Matthew D.
AU - Aum, Diana
AU - Lew, Sean M.
AU - Wang, Shelly
AU - Niazi, Toba N.
AU - Fallah, Aria
AU - Raskin, Jeffrey S.
AU - Weiner, Howard L.
AU - Gadgil, Nisha
AU - Albert, Gregory W.
AU - Hadjinicolaou, Aristides
AU - Major, Philippe
AU - Niazi, Farbod
AU - Theaud, Guillaume
AU - Obaid, Sami
AU - Widjaja, Elysa
AU - Ertl-Wagner, Birgit
AU - Vidarsson, Logi
AU - Taylor, Margot J.
AU - Boutet, Alexandre
AU - Rutka, James T.
AU - LoPresti, Melissa A.
AU - Jain, Puneet
AU - Ibrahim, George M.
TI - A deep representation learning model to predict response to vagus nerve stimulation
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/
VL - 17
IS - 1
SP - 4932
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"type": "article-journal",
"title": "A deep representation learning model to predict response to vagus nerve stimulation",
"container-title": "Nature communications",
"author": [
{
"family": "Suresh",
"given": "Hrishikesh"
},
{
"family": "Mithani",
"given": "Karim"
},
{
"family": "Li",
"given": "Vicki"
},
{
"family": "Latypov",
"given": "Timur H."
},
{
"family": "Warsi",
"given": "Nebras M."
},
{
"family": "Wong",
"given": "Simeon M."
},
{
"family": "Erdman",
"given": "Lauren"
},
{
"family": "Kang",
"given": "Jaeyoung"
},
{
"family": "Germann",
"given": "Jurgen"
},
{
"family": "Gouveia",
"given": "Flavia Venetucci"
},
{
"family": "Coleman",
"given": "Sebastian C."
},
{
"family": "Berger",
"given": "Alexandre"
},
{
"family": "Chau",
"given": "Vann"
},
{
"family": "Weiss",
"given": "Shelly"
},
{
"family": "Gorodetsky",
"given": "Carolina"
},
{
"family": "Donner",
"given": "Elizabeth"
},
{
"family": "Weil",
"given": "Alexander G."
},
{
"family": "Tailor",
"given": "Jignesh"
},
{
"family": "Abel",
"given": "Taylor J."
},
{
"family": "Remick",
"given": "Madison"
},
{
"family": "Akwayena",
"given": "Emefa"
},
{
"family": "Schrader",
"given": "Dewi"
},
{
"family": "Bollo",
"given": "Robert J."
},
{
"family": "Smyth",
"given": "Matthew D."
},
{
"family": "Aum",
"given": "Diana"
},
{
"family": "Lew",
"given": "Sean M."
},
{
"family": "Wang",
"given": "Shelly"
},
{
"family": "Niazi",
"given": "Toba N."
},
{
"family": "Fallah",
"given": "Aria"
},
{
"family": "Raskin",
"given": "Jeffrey S."
},
{
"family": "Weiner",
"given": "Howard L."
},
{
"family": "Gadgil",
"given": "Nisha"
},
{
"family": "Albert",
"given": "Gregory W."
},
{
"family": "Hadjinicolaou",
"given": "Aristides"
},
{
"family": "Major",
"given": "Philippe"
},
{
"family": "Niazi",
"given": "Farbod"
},
{
"family": "Theaud",
"given": "Guillaume"
},
{
"family": "Obaid",
"given": "Sami"
},
{
"family": "Widjaja",
"given": "Elysa"
},
{
"family": "Ertl-Wagner",
"given": "Birgit"
},
{
"family": "Vidarsson",
"given": "Logi"
},
{
"family": "Taylor",
"given": "Margot J."
},
{
"family": "Boutet",
"given": "Alexandre"
},
{
"family": "Rutka",
"given": "James T."
},
{
"family": "LoPresti",
"given": "Melissa A."
},
{
"family": "Jain",
"given": "Puneet"
},
{
"family": "Ibrahim",
"given": "George M."
}
],
"container-title-short":
"volume": "17",
"issue": "1",
"page": "4932",
"DOI": "10.1038/
"PMID": "41946715",
"PMCID": "PMC13234157",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
7
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1038/s41586-026-10631-3 [code]
- A prognostic human brain network for diffuse midline glioma.Journal: NatureIn common: ANTs, FreeSurfer, FSL, 7 other tools, DOI 10.7910/dvn/ilxiks, clinical / translational, 8 references
- [2] doi:10.1126/sciadv.aeb5842 [code]
- Regulation of autism-related self-injurious behavior by electrical stimulation of corticostriatal circuits in mice and humans.Journal: Science advancesIn common: ANTs, SciPy, NumPy, 2 references, 4 authors
- [3] doi:10.1162/imag.a.1352 [code]
- Brain-age in ultra-low-field MRI: How well does it work?Journal: Imaging neuroscience (Cambridge, Mass.)In common: MONAI, ANTs, FreeSurfer, 6 other tools, structural MRI / diffusion, 3 references
- [4] doi:10.1371/journal.pbio.3003856 [code]
- Aging and metabolism contribute separately to brain-body health.Journal: PLoS biologyIn common: ANTs, FreeSurfer, FSL, 8 other tools, structural MRI / diffusion, clinical / translational, 2 references
- [5] doi:10.1162/imag.a.1164 [code]
- Bias and generalizability of brain age prediction models: A multi-cohort evaluation with anatomical and interpretability insights.Journal: Imaging neuroscience (Cambridge, Mass.)In common: MONAI, ANTs, FreeSurfer, 9 other tools, structural MRI / diffusion
- [6] doi:10.1002/alz.71649 [code]
- Postmortem brain MRI reveals differential associations of subcortical and limbic volumes with cortical thinning and neurodegenerative pathologies.Journal: Alzheimer's & dementia : the journal of the Alzheimer's AssociationIn common: ANTs, FreeSurfer, FSL, 8 other tools, structural MRI / diffusion, 2 references
- [7] doi:10.1126/sciadv.adu9309 [code]
- Variations of global brain asymmetry are associated with aging and related diseases.Journal: Science advancesIn common: ANTs, FreeSurfer, FSL, 6 other tools, 4 references
- [8] doi:10.1162/imag.a.1252 [code]
- Does the brain's E:I balance really shape long-range temporal correlations? Lessons learned from 3T MRI.Journal: Imaging neuroscience (Cambridge, Mass.)In common: ANTs, FreeSurfer, FSL, 7 other tools, structural MRI / diffusion, 3 references
- [9] doi:10.1038/s41467-026-73996-z [code]
- Genetic architecture of white matter microstructure captured by unsupervised deep representation learning of fractional anisotropy maps.Journal: Nature communicationsIn common: MONAI, PyTorch Lightning, NiBabel, 7 other tools, structural MRI / diffusion, 2 references
- [10] doi:10.1038/s41467-026-71918-7 [code]
- Developmental disinhibition gates language lateralization in childhood.Journal: Nature communicationsIn common: ANTs, FreeSurfer, FSL, 8 other tools, 2 references
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 3 repositories of the authors' code, each at its verified commit and with its license, 48 scripts, and 3 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:54ee74937381b5a7…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
