Single-cell-level perturbation-induced and condition-related signal estimation with batch effect removal using NDreamer.
The 24 matches · 5 of them tie a paragraph to a whole file, not to given lines: weak matches, whose lines are not tinted
- [1] § Results › Evaluation metrics ↔ R_plot_codes/experimental_benchmark_plot_codes/ablation.R, the whole file · a weak match · score 0.85 · ASW_label, bLISI, mixing metrics, ASW_batch, cLISI, kBET
- [2] § Results › Evaluation metrics ↔ R_plot_codes/experimental_benchmark_plot_codes/plot_experimental.R, lines 13–86 · score 0.81 · ASW_label, bLISI, ASW_batch, cLISI, kBET, biological conservation
- [3] § Methods › Estimation of the effect modifier embedding › Normalized mutual information loss ↔ ndreamer/DL_loss_func.py, lines 386–461 · score 0.75 · hot encoded, joint probability, assignment vector, matrix, mutual, loss
- [4] § Results › Evaluation metrics ↔ R_plot_codes/case_control_code/radar.R, lines 14–122 · score 0.75 · EM bio conservation, denoised batch mixing, EM batch, signal preservation, Radar, distortion
- [5] § Results › Evaluation metrics › Simulation dataset ↔ R_plot_codes/experimental_benchmark_plot_codes/overall_plot_experimental.R, lines 37–92 · score 0.74 · scCAPE, scGen, CINEMA OT, biological conservation, Mixscape, NDreamer
- [6] § Results › Exploration of batch effects, individual treatment effect, and sensitivity test › Ablation study ↔ ndreamer/model.py, lines 221–306 · score 0.72 · VQ VAE, local neighborhood loss, triplet loss, independent loss, model, NDreamer
- [7] § Results › NDreamer demonstrates robustly improved performance across experimental perturbation datasets ↔ R_plot_codes/experimental_benchmark_plot_codes/overall_plot_experimental.R, lines 37–92 · score 0.67 · scCAPE, scGen, CINEMA OT, Mixscape, NDreamer, benchmarking
- [8] § Results › NDreamer demonstrates robustly improved performance across experimental perturbation datasets ↔ R_plot_codes/experimental_benchmark_plot_codes/ablation.R, the whole file · a weak match · score 0.66 · bLISI, ASW_batch, mixing metrics, kBET, benchmarking, NDreamer
- [9] § Results › NDreamer uncovers differential gene expression patterns in a large Alzheimer's disease cohort ↔ R_plot_codes/AD_ADNC_code/AD_neurons_GO.R, lines 1–10 · score 0.65 · PCDH9 AS2, PPP1R9A AS1, GO, genes, NDreamer
- [10] § Methods › Overview of the NDreamer deep learning model ↔ tutorial_experimental_PBMC.ipynb, lines 25–180 · score 0.63 · local neighborhood loss, Gaussian kernel, triplet loss function, NDreamer, modifier space, reconstruct
- [11] § Results › Evaluation metrics ↔ R_plot_codes/case_control_code/radar_mouse.R, lines 14–120 · score 0.63 · denoised batch mixing, EM batch, signal preservation, Radar, distortion, mouse
- [12] § Results › Exploration of batch effects, individual treatment effect, and sensitivity test › Exploration of batch effects ↔ ndreamer/pipeline.py, lines 292–362 · score 0.63 · CD8T, local neighborhood loss, independent loss, UMAP, NDreamer, PBMC
- [13] § Results › NDreamer demonstrates robustly improved performance across experimental perturbation datasets ↔ ndreamer/metrics.py, lines 164–278 · score 0.61 · bLISI, asw_batch, kBET, silhouette, metrics, neighbor
- [14] § Methods › Overview of the NDreamer deep learning model ↔ ndreamer/model.py, lines 221–306 · score 0.61 · local neighborhood loss, dependent loss, triplet loss, modifier space, VAE, reconstruct
- [15] § Methods › Estimation of the effect modifier embedding › Local neighborhood loss ↔ ndreamer/DL_loss_func.py, lines 464–551 · score 0.59 · Gaussian kernel, local neighborhood loss, nearest neighbors, distance, cluster, space
- [16] § Methods › Estimation of the effect modifier embedding › Normalized mutual information loss ↔ ndreamer/DL_loss_func.py, lines 386–461 · score 0.59 · mutual information loss, encourage independence, entropy, assignments, probability, batch
- [17] § Results › Evaluation metrics ↔ ndreamer/metrics.py, lines 164–278 · score 0.57 · bLISI, asw_batch, kBET, metrics, cell
- [18] § Results › NDreamer uncovers differential gene expression patterns in a large Alzheimer's disease cohort ↔ R_plot_codes/case_control_code/simple_GO.R, the whole file · a weak match · score 0.57 · enriched GO terms, GO enrichment, KIT, TRPC3, genes
- [19] § Results › NDreamer demonstrates robustly improved performance across experimental perturbation datasets ↔ R_plot_codes/experimental_benchmark_plot_codes/plot_experimental_have_batch_ndreamer.R, lines 13–96 · score 0.55 · cLISI, conservation metrics, ARI, NMI, ASW, perturbation
- [20] § Results › NDreamer demonstrates robustly improved performance across experimental perturbation datasets ↔ R_plot_codes/experimental_benchmark_plot_codes/plot_experimental_have_batch.R, lines 13–96 · score 0.54 · cLISI, conservation metrics, ARI, NMI, ASW, perturbation
- [21] § Methods › Estimation of the effect modifier embedding › Encoder with neural discrete representation learning and variational inference ↔ ndreamer/model.py, lines 188–219 · score 0.53 · KL divergence, divergence loss, VAE
- [22] § Results › NDreamer uncovers differential gene expression patterns in a large Alzheimer's disease cohort ↔ R_plot_codes/AD_ADNC_code/AD_respiration_GO.R, the whole file · a weak match · score 0.53 · enriched GO terms, GO enrichment, PLCG2, genes
- [23] § Results › NDreamer uncovers differential gene expression patterns in a large Alzheimer's disease cohort ↔ R_plot_codes/AD_ADNC_code/AD_respiration_GO.R, the whole file · a weak match · score 0.52 · ATP6, COX2, CYTB, ND3, ND4, respiration
- [24] § Methods › Estimation of the effect modifier embedding › Encoder with neural discrete representation learning and variational inference ↔ ndreamer/model.py, lines 123–186 · score 0.51 · VQ VAE, commitment loss, codebook, Encoder, embeddings
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 · 482 lines · 26 KB · GPL-3.0 · 4 matches
- import os
- import numpy as np
- import pandas as pd
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- from torch.cuda import device
- from ndreamer.model_DL import NDreamer_generator,Discriminator,set_seed
- from ndreamer.DL_loss_func import CrossEntropy, create_triplets_within_groups, IndependenceLoss, kl_divergence_loss, \
- IndependenceLoss_label, OrthogonalityRegularization, EntropyPenalty, create_triplets_within_groups_logits, \
- IndependenceLoss_between_matrix, compute_mmd, compute_local_neighborhood_loss, create_triplets_within_groups_tensor
- from ndreamer.data_preprocess import process_adata,generate_balanced_dataloader,generate_adata_to_dataloader
- from ndreamer.DL_loss_func import reconstruction_error
- from ndreamer.statistics import *
- from sklearn.decomposition import PCA
- # Dynamic import of tqdm based on the environment
- import sys
- if 'ipykernel' in sys.modules:
- from tqdm.notebook import tqdm
- else:
- from tqdm import tqdm
- from ndreamer.plot import plot_distribution
- class NDreamer_DL(nn.Module):
- def __init__(self, adata, condition_key, contorl_name, num_hvg, require_batch=False, batch_key=None,
- resolution_low=0.5, resolution_high=7, cluster_method="Leiden", embedding_dim=512, codebooks=None,
- codebook_dim=8, encoder_hidden=None, decoder_hidden=None, z_dim=256,
- cos_loss_scaler=20, random_seed=123, batch_size=2048, epoches=10, lr=1e-3, triplet_margin=5,
- independent_loss_scaler=1000, save_pth="./model/", developer_test_mode=False,
- library_size_normalize_adata=False, save_preprocessed_adata_path="./model/preprocessed.h5ad",
- KL_scaler=5e-3, reconstruct_scaler=1, triplet_scaler=5, num_triplets_per_label=15,
- tau=0.01, commitment_loss_scaler=1, cluster_correlation_scaler=50,
- try_identify_perturb_escaped_cell=False, reset_threshold=1 / 1024, reset_interval=30,
- try_identify_cb_specific_subtypes=False, local_neighborhood_loss_scaler=1,
- local_neighbor_sigma=1, n_neighbors=20, local_neighbor_across_cluster_scaler=20,
- have_negative_data=False):
- super(NDreamer_DL, self).__init__()
- input_dim=num_hvg
- if codebooks is None:
- codebooks = [1024 for i in range(32)]
- if encoder_hidden is None:
- encoder_hidden = [2048, 1024]
- if decoder_hidden is None:
- decoder_hidden = [512, 1024]
- self.developer_test_mode=developer_test_mode
- self.require_batch=require_batch
- self.batch_key = batch_key
- self.condition_key = condition_key
- self.triplet_margin=triplet_margin
- if not os.path.exists(save_pth):
- os.mkdir(save_pth)
- self.save_pth=save_pth
- self.random_seed = random_seed
- self.batch_size = batch_size
- self.epoches = epoches
- self.lr = lr
- self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
- print("Using device:",self.device)
- self.kl_scaler = KL_scaler
- self.reconstruct_scaler = reconstruct_scaler
- self.triplet_scaler = triplet_scaler
- self.num_triplets_per_label=num_triplets_per_label
- self.independent_loss_scaler = independent_loss_scaler
- self.codebooks=codebooks
- self.local_neighborhood_loss_scaler=local_neighborhood_loss_scaler
- print(local_neighborhood_loss_scaler)
- self.local_neighbor_sigma=local_neighbor_sigma
- self.try_identify_cb_specific_subtypes=try_identify_cb_specific_subtypes
- self.try_identify_perturb_escaped_cell=try_identify_perturb_escaped_cell
- self.n_neighbors=n_neighbors
- self.local_neighbor_across_cluster_scaler=local_neighbor_across_cluster_scaler
- #data preprocessing
- if not developer_test_mode:
- print("Start data preprocessing")
- adata, batch_dict, condition_dict = process_adata(adata=adata, condition_key=condition_key,
- input_dim=input_dim, control_name=contorl_name,
- require_batch=require_batch, batch_key=batch_key,
- resolution_low=resolution_low,
- resolution_high=resolution_high,
- cluster_method=cluster_method,
- library_size_normalize_adata=library_size_normalize_adata)
- self.adata = adata
- self.batch_dict = batch_dict
- self.condition_dict = condition_dict
- print("Data preprocessing done")
- if save_preprocessed_adata_path is not None:
- adata.write(save_preprocessed_adata_path)
- else:
- self.adata=adata
- print("Remaining number of cells:",self.adata.shape[0])
- calcluated_epoches=15*self.adata.shape[0]//(self.batch_size*len(np.unique(self.adata.obs["group"])))+1
- if calcluated_epoches>self.epoches:
- if len(np.unique(self.adata.obs["group"]))<=10:
- calcluated_epoches = calcluated_epoches + reset_interval - calcluated_epoches % reset_interval - 2
- self.epoches=calcluated_epoches
- print("Too few epoches (steps, if rigorously speaking). Changing epoch to", self.epoches, "to adjust for number of cells")
- num_batches = []#np.unique(adata.obs["batch"]).shape[0]
- if not self.developer_test_mode and self.require_batch:
- if isinstance(batch_key, str):
- num_batches = [max(self.batch_dict[batch_key].values()) + 1]
- elif isinstance(batch_key, list):
- num_batches=[]
- for batch_keyi in batch_key:
- num_batches.append(max(self.batch_dict[batch_keyi].values()) + 1)
- else:
- num_batches=[]
- self.num_batches=num_batches
- num_treatments = np.unique(adata.obs["condition"]).shape[0]
- if not developer_test_mode:
- num_treatments=max(num_treatments,max(self.condition_dict.values())+1)
- # define the model
- self.VQ_VAE = NDreamer_generator(input_dim=input_dim, num_treatments=num_treatments,z_dim=z_dim,
- num_batches=num_batches,embedding_dim=embedding_dim,tau=tau,
- codebooks=codebooks,codebook_dim=codebook_dim,encoder_hidden=encoder_hidden,
- decoder_hidden=decoder_hidden,commitment_loss_scaler=commitment_loss_scaler,
- reset_threshold=reset_threshold,reset_interval=reset_interval,
- try_identify_cb_specific_subtypes=try_identify_cb_specific_subtypes,
- try_identify_perturb_escaped_cell=try_identify_perturb_escaped_cell,
- have_negative_data=have_negative_data)
- self.require_batch = require_batch
- print("Require batch:",self.require_batch)
- self.independence_loss=IndependenceLoss(scaler=independent_loss_scaler)
- self.independence_loss_label=IndependenceLoss_label(scaler=cluster_correlation_scaler/num_treatments)
- self.independence_loss_between_codebook=IndependenceLoss_between_matrix(scaler=independent_loss_scaler/100)
- self.cross_entropy = CrossEntropy()
- self.cos_loss_scaler = cos_loss_scaler
- #self.MINE=MINE(x_dim=input_dim,z_dim=z_dim,hidden_dim=1024)
- # init the model
- self.VQ_VAE.to(self.device)
- self.independence_loss.to(self.device)
- self.independence_loss_label.to(self.device)
- #self.cross_entropy.to(self.device)
- #self.orthogonality_regularization.to(self.device)
- #self.entropy_penalty.to(self.device)
- self.independence_loss_between_codebook.to(self.device)
- #self.MINE.to(self.device)
- self.logits=None
- self.df_latent=None
- def train_model(self):
- optimizer_G = torch.optim.AdamW(self.VQ_VAE.parameters(), lr=self.lr)
- progress_bar = tqdm(total=self.epoches, desc="Overall Progress", leave=True, miniters=1, mininterval=0)
- for epoch in range(self.epoches):
- set_seed(self.random_seed+epoch)
- data_loader = generate_balanced_dataloader(self.adata, batch_size=self.batch_size)
- self.VQ_VAE.train()
- all_losses = 0
- T_loss = 0
- V_loss = 0
- I_loss = 0
- K_loss = 0
- N_loss = 0
- Commitment_loss=0
- Dependent_loss=0
- for i, (exp, condition, batch, labels_low, labels_high) in enumerate(data_loader):
- #print(exp.shape, condition.shape, batch.shape, labels_low.shape, labels_high.shape)
- # convert to cuda
- exp = exp.to(self.device)
- condition = condition.to(self.device)
- batch = batch.to(self.device)
- labels_low = labels_low.to(self.device)
- labels_high = labels_high.to(self.device)
- # run VQ-VAE
- z, variance, reconstructed, logits, commitment_loss,escape_judger_choice, cb_specifc_embedding=self.VQ_VAE(exp=exp, treatment=condition, batch=batch)
- commitment_loss=commitment_loss/len(self.codebooks)
- if self.VQ_VAE.encoder.codebooks[0].just_reset_codebook:
- print("Finish resetting codebook embeddings, current step (epoch):",epoch)
- continue
- # calculate the reconstruction loss
- reconst_loss_mse = reconstruction_error(exp, reconstructed)*self.reconstruct_scaler
- reconst_loss_cos = (1 - torch.sum(F.normalize(reconstructed, p=2) * F.normalize(exp, p=2), 1)).mean()
- reconst_loss_cos = self.cos_loss_scaler * reconst_loss_cos
- reconst_loss = reconst_loss_mse + reconst_loss_cos
- reconst_loss = torch.clamp(reconst_loss, max=1e5)
- KL_loss=kl_divergence_loss(z_mean=z,z_var=variance)*self.kl_scaler
- KL_loss = torch.clamp(KL_loss, max=1e5)
- # independence loss for VAE
- independent_loss=0
- independent_loss=independent_loss+self.independence_loss(condition, logits)
- for i in range(len(self.num_batches)):
- independent_loss = independent_loss + self.independence_loss(batch[:,i], logits)
- independent_loss = torch.clamp(independent_loss, max=1e5)
- # convert matrix batch to unique int number vector
- base = batch.max() + 1 # Base larger than the largest number in the matrix
- powers = base ** torch.arange(batch.size(1), device=exp.device).flip(0).float() # Positional encoding powers
- batch_vector = torch.matmul(batch.float(), powers).long() # Encoded as unique integers
- dependent_loss = 0
- dependent_loss=dependent_loss-self.independence_loss_label(assignments=labels_high,
- probabilities=logits,
- condition=condition,
- batch=batch_vector)
- dependent_loss=dependent_loss-self.independence_loss_label(assignments=labels_low,
- probabilities=logits,
- condition=condition,
- batch=batch_vector)
- # triplet loss
- triplet_loss = create_triplets_within_groups_tensor(z, labels_low, labels_high, batch_vector,
- condition, margin=self.triplet_margin,
- num_triplets=self.num_triplets_per_label)
- triplet_loss=triplet_loss*self.triplet_scaler
- neighbor_loss = compute_local_neighborhood_loss(latent=z, inputs=exp,
- condition=condition, batch=batch_vector,
- sigma=self.local_neighbor_sigma,
- n_neighbors=self.n_neighbors,
- label_low=labels_low,
- label_high=labels_high,
- local_neighbor_across_cluster_scaler=self.local_neighbor_across_cluster_scaler
- ) * self.local_neighborhood_loss_scaler
- all_loss = reconst_loss + independent_loss +KL_loss + triplet_loss + neighbor_loss + commitment_loss + dependent_loss# + embedding_loss +entropy_loss
- all_loss = torch.clamp(all_loss, max=1e5)
- all_loss.backward()
- optimizer_G.step()
- optimizer_G.zero_grad()
- all_losses += all_loss.item()/len(data_loader)
- N_loss += neighbor_loss.item()/len(data_loader)
- T_loss += triplet_loss.item()/len(data_loader)
- V_loss += reconst_loss.item()/len(data_loader)
- I_loss +=independent_loss.item()/len(data_loader)
- K_loss +=KL_loss.item()/len(data_loader)
- Dependent_loss+=dependent_loss.item()/len(data_loader)
- #Entropy_loss+=entropy_loss.item()
- #Embedding_loss += 0#embedding_loss.item()
- Commitment_loss+=commitment_loss.item()/len(data_loader)
- print(f"Epoch: {epoch + 1}/{self.epoches} | "
- f"All Loss: {all_losses:.4f} | "
- f"Neighborhood Loss: {N_loss:.4f} | "
- f"Triplet Loss: {T_loss:.4f} | "
- f"Reconstruction Loss: {V_loss:.4f} | "
- f"Independent Loss: {I_loss:.4f} | "
- f"KL Loss: {K_loss:.4f} | "
- f"Commitment Loss: {Commitment_loss:.4f} | "
- f"Dependent Loss: {Dependent_loss:.4f}")
- progress_bar.update(1) # Increment the progress bar by one for each batch processed
- progress_bar.set_postfix(epoch=f"{epoch + 1}/{self.epoches}", all_loss=all_losses, neigh_loss=N_loss, triplet_loss=T_loss, reconst_loss=V_loss, independent_loss=I_loss, KL_loss=K_loss, Commitment_loss=Commitment_loss, Dependent_loss=Dependent_loss)
- progress_bar.close()
- torch.save(self.state_dict(), os.path.join(self.save_pth, 'ndreamer' + '.pth'))
- def get_modifier_space(self,dim=50, save_path_adata=None, save_path_latent_df=None,
- print_codebook_statisitics=False, kBET_test=False, save_others=False):
- self.VQ_VAE.eval()
- data_loader = generate_adata_to_dataloader(self.adata)
- device = 'cuda' if torch.cuda.is_available() else 'cpu'
- all_z = []
- all_indices = []
- logits=[]
- cb_specific_subtypes=[]
- escape_probs=[]
- with torch.no_grad():
- for i, (x, indices, condition, batch) in enumerate(data_loader):
- x = x.to(device)
- condition = condition.to(device)
- batch = batch.to(device)
- z, variance, reconstructed, logit, commitment_loss, escape_judger_choice, cb_specifc_embedding = self.VQ_VAE(
- exp=x, treatment=condition, batch=batch)
- all_z.append(z.cpu().detach())
- all_indices.extend(indices.tolist())
- if self.try_identify_perturb_escaped_cell:
- escape_probs.append(escape_judger_choice.detach().cpu())
- if self.try_identify_cb_specific_subtypes:
- cb_specific_subtypes.append(cb_specifc_embedding.detach().cpu())
- if len(logits) == 0:
- logits = [i.detach().cpu() for i in logit]
- else:
- logits = [torch.concatenate([logits[i], logit[i].detach().cpu()], dim=0) for i in
- range(len(logits))]
- all_z_combined = torch.cat(all_z, dim=0)
- all_indices_tensor = torch.tensor(all_indices)
- all_z_reordered = all_z_combined[all_indices_tensor.argsort()]
- all_z_np = all_z_reordered.numpy()
- # Create anndata object with reordered embeddings
- self.adata.obsm['X_effect_modifier_space'] = all_z_np
- pca = PCA(n_components=dim)
- # Fit and transform the data
- X_pca = pca.fit_transform(self.adata.obsm['X_effect_modifier_space'])
- # Store the PCA-reduced data back into adata.obsm
- self.adata.obsm['X_effect_modifier_space_PCA'] = X_pca
- if save_path_adata is not None:
- if save_path_adata.find(".h5ad")<0:
- save_path_adata = save_path_adata+".h5ad"
- self.adata.write(save_path_adata)
- else:
- self.adata.write(os.path.join(self.save_pth, 'adata' + '.h5ad'))
- print("Effect modifier space saved.")
- if self.try_identify_cb_specific_subtypes:
- print(cb_specific_subtypes)
- cb_specific_subtypes = torch.concat(cb_specific_subtypes, dim=0)
- plot_distribution(cb_specific_subtypes.reshape(-1).numpy())
- self.adata.obs["original_order"]=np.array(range(self.adata.shape[0]))
- if not save_others:
- return
- torch.save(logits, os.path.join(self.save_pth, "logits.pth"))
- self.logits = logits
- if self.try_identify_perturb_escaped_cell:
- escape_probs = torch.concat(escape_probs, dim=0)
- torch.save(escape_probs, os.path.join(self.save_pth, "escape_probs.pth"))
- if self.try_identify_cb_specific_subtypes:
- cb_specific_subtypes = torch.concat(cb_specific_subtypes, dim=0)
- torch.save(cb_specific_subtypes, os.path.join(self.save_pth, "cb_specific_subtypes.pth"))
- df_latent=pd.DataFrame(data=self.adata.obsm['X_effect_modifier_space_PCA'],columns=['latent_'+str(i) for i in range(X_pca.shape[1])])
- df_latent["batch"]=np.array(self.adata.obs["batch"])
- df_latent["condition"]=np.array(self.adata.obs["condition"])
- df_latent["group"]=np.array(self.adata.obs["group"])
- if save_path_latent_df is not None:
- if save_path_latent_df.find(".csv")<0:
- save_path_latent_df = save_path_latent_df+".csv"
- df_latent.to_csv(save_path_latent_df)
- else:
- df_latent.to_csv(os.path.join(self.save_pth, 'latent.csv'))
- self.df_latent=df_latent
- if print_codebook_statisitics:
- print(torch.max(self.logits[0],dim=-1))
- print("Statistic values of the codebook selection")
- print("mean:\n",[torch.mean(logits[i], dim=0) for i in range(len(logits))])
- print("variance:\n",[torch.std(logits[i], dim=0) for i in range(len(logits))])
- if kBET_test:
- run_kbet(df_latent,do_PCA=False)
- def run_kBET_test(self):
- run_kbet(self.df_latent, do_PCA=False)
- if __name__ == '__main__':
- test_run_mode=False
- train_model=True
- if test_run_mode:
- import scanpy as sc
- import anndata as ad
- '''exp = torch.abs(torch.randn((2048, 2000)))
- condition = torch.randint(low=0, high=3, size=(2048,), dtype=torch.long)
- batch1 = torch.randint(low=0, high=4, size=(2048,), dtype=torch.long) # np.zeros(2048)
- batch2 = torch.randint(low=0, high=2, size=(2048,), dtype=torch.long)
- adata = ad.AnnData(X=exp.numpy())
- adata.obs["condition"] = condition.numpy()
- adata.obs["group"] = condition.numpy()
- adata.obs["batch1"] = batch1
- adata.obs["batch"] = batch1
- adata.obs["batch2"] = batch2
- adata.obs["leiden1"] = torch.randint(low=0, high=3, size=(2048,), dtype=torch.long).numpy()
- adata.obs["leiden1"] = adata.obs["leiden1"].astype("category")
- adata.obs["leiden2"] = torch.randint(low=0, high=9, size=(2048,), dtype=torch.long).numpy()
- adata.obs["leiden2"] = adata.obs["leiden2"].astype("category")
- model = NDreamer_DL(adata=adata, condition_key="condition", contorl_name=0, num_hvg=2000,
- developer_test_mode=False, require_batch=True, batch_key="batch1",#["batch1","batch2"],
- batch_size=512)
- model.train_model()'''
- adata = sc.read("../ECCITE_preprocessed.h5ad")
- adata.obsm["batch"]=np.expand_dims(adata.obs["batch"].copy(),axis=-1)
- print(adata)
- print(np.unique(adata.obs["condition"]))
- model = NDreamer_DL(adata, condition_key='perturbation', contorl_name='NT', num_hvg=2000, require_batch=True,
- batch_key='replicate',
- resolution_low=0.5, resolution_high=7, cluster_method="Leiden", embedding_dim=512,
- codebooks=[1024 for i in range(32)],
- codebook_dim=8, encoder_hidden=[1024, 512], decoder_hidden=[512, 1024], z_dim=256,
- cos_loss_scaler=20, random_seed=123, batch_size=1024, epoches=1, lr=1e-3,
- triplet_margin=5, independent_loss_scaler=1000, save_pth="./model/",
- developer_test_mode=True,
- library_size_normalize_adata=False,
- save_preprocessed_adata_path=None,
- KL_scaler=5e-3, reconstruct_scaler=1, triplet_scaler=5, num_triplets_per_label=15,
- tau=0.01, commitment_loss_scaler=1, cluster_correlation_scaler=50, reset_threshold=1 / 1024,
- reset_interval=30, try_identify_cb_specific_subtypes=False,
- local_neighborhood_loss_scaler=1, local_neighbor_sigma=1,
- try_identify_perturb_escaped_cell=False, n_neighbors=20,
- local_neighbor_across_cluster_scaler=20)
- model.train_model()
- print()
- model.get_modifier_space(save_path_adata="../ECCITE_results.h5ad")
- adata1 = model.adata.copy()
- sc.pp.neighbors(adata1, use_rep='X_effect_modifier_space_PCA', n_neighbors=15)
- sc.tl.umap(adata1)
- sc.pl.umap(adata1, color=["MULTI_ID","HTO_classification",
- "gene_target","perturbation",
- "replicate","Phase"], frameon=False, ncols=1)
- else:
- import os
- os.environ["CUDA_LAUNCH_BLOCKING"] = "1"
- import scanpy as sc
- #adata=sc.read_h5ad("../data/PBMC.h5ad")
- #adata = sc.read("../PBMC_preprocessed.h5ad")
- adata = sc.read("../virus_preprocessed.h5ad")
- #adata=sc.read_h5ad("../PBMC_imbalance_preprocessed.h5ad")
- #adata=sc.read_h5ad("../common1.h5ad")
- if "condition" not in adata.obsm.keys():
- adata.obsm["batch"] = np.ones((adata.shape[0],1))#np.expand_dims(adata.obs["condition"].copy(), axis=1)
- print(adata)
- print(np.unique(adata.obs["condition"]))
- model = NDreamer_DL(adata, condition_key="condition", contorl_name="control", num_hvg=min(adata.shape[1],3608),
- require_batch=False,
- batch_key=None,
- resolution_low=0.5, resolution_high=7, cluster_method="Leiden", embedding_dim=512,
- codebooks=[1024 for i in range(32)],
- codebook_dim=8, encoder_hidden=[2048, 1024], decoder_hidden=[512, 1024], z_dim=256,
- cos_loss_scaler=20, random_seed=123, batch_size=1024, epoches=100, lr=1e-3,
- triplet_margin=5,independent_loss_scaler=1000, save_pth="./model/",
- developer_test_mode=True,
- library_size_normalize_adata=False,
- save_preprocessed_adata_path=None,
- KL_scaler=5e-3, reconstruct_scaler=1, triplet_scaler=5, num_triplets_per_label=15,
- tau=0.01, commitment_loss_scaler=1, cluster_correlation_scaler=50,reset_threshold=1/1024,
- reset_interval=30,try_identify_cb_specific_subtypes=False,
- local_neighborhood_loss_scaler=1,local_neighbor_sigma=1,
- try_identify_perturb_escaped_cell=False,n_neighbors=20,
- local_neighbor_across_cluster_scaler=20, have_negative_data=True)
- '''model = NDreamer_DL(adata, condition_key="condition", contorl_name="control", num_hvg=min(adata.shape[1], 3608),
- require_batch=False,
- batch_key=None,have_negative_data=True,developer_test_mode=True)'''
- if train_model:
- model.train_model()
- print()
- model.load_state_dict(torch.load(os.path.join(model.save_pth, 'ndreamer' + '.pth')))
- model.get_modifier_space(save_path_adata="../PBMC_results.h5ad")
- adata1 = model.adata.copy()
- sc.pp.neighbors(adata1, use_rep='X_effect_modifier_space_PCA',n_neighbors=15)
- sc.tl.umap(adata1)
- try:
- sc.pl.umap(adata1, color=['condition', 'cell_type'], frameon=False, ncols=1)
- except:
- sc.pl.umap(adata1, color=['condition', 'cell_type1021'], frameon=False, ncols=1)
- sc.pp.neighbors(adata1, use_rep='X_effect_modifier_space',n_neighbors=15)
- sc.tl.umap(adata1)
- try:
- sc.pl.umap(adata1, color=['condition', 'cell_type'], frameon=False, ncols=1)
- except:
- sc.pl.umap(adata1, color=['condition', 'cell_type1021'], frameon=False, ncols=1)
- #model.run_kBET_test()
model.py at commit dfff8e9, under GPL-3.0 · at the source
Overview
- Department of Biostatistics, Yale University School of Public Health, 300 George Street, New Haven, CT 06511, United States
- Department of Genetics, Yale University School of Medicine, 333 Cedar Street, New Haven, CT 06520, United States
- Department of Statistics and Data Science, Yale University, 219 Prospect Street, New Haven, CT 06511, United States
- Department of Biomedical Informatics & Data Science, Yale University School of Medicine, 101 College Street, New Haven, CT 06510, United States
Abstract
Advances in sequencing technologies and the growing volume of single-cell data have created unprecedented opportunities for uncovering gene expression patterns causally induced by experimental perturbations or statistically associated, but not necessarily causal, with disease conditions. However, current analytical methods inadequately account for batch effects and data sparsity or fail to capture the inherent non-linearity in single-cell data, leading to biased estimation. To address these limitations, we developed NDreamer that combines neural discrete representation learning and matching to remove batch effects and estimate perturbation-induced or condition-associated signals at single-cell resolution. NDreamer outperformed existing methods by using mutual information loss on discrete latent variables to disentangle cells’ intrinsic features from conditions or batch effects, while preserving both global and local variance within batches and conditions via triplet and local neighborhood loss. We applied NDreamer to multiple datasets across platforms, organs, and species and validated and benchmarked its performance in removing batch effects and estimating perturbation-induced or condition-associated signals. In particular, we applied NDreamer to an Alzheimer’s disease cohort, revealing biologically relevant gene expression patterns that distinguish dementia patients from controls.
Reproduced under the paper's license (CC BY-NC), from the paper cited above.
Repositories
Its files are read in the Code ↔ Paper reader above, with 24 matches between paragraphs and lines of code.
lugia-xiao/NDreamer
dfff8e9fb8c0a9bb246468d06d42aee5e666b1bf, 20 March 2025Availability: 1 check, the latest on 26 September 2026: the link answers
- 26 September 2026: the link answers
17 files
- ndreamer/
DL_loss_func.py , Python, 701 lines, 3 matches - ndreamer/
__init__.py , Python, 3 lines - ndreamer/
data_preprocess.py , Python, 208 lines - ndreamer/
matching.py , Python, 347 lines - ndreamer/
metrics.py , Python, 278 lines, 2 matches - ndreamer/
model.py , Python, 482 lines, 4 matches - ndreamer/
model_DL.py , Python, 396 lines - ndreamer/
pipeline.py , Python, 362 lines, 1 match - ndreamer/
plot.py , Python, 72 lines - ndreamer/
single_cell_utils.py , Python, 124 lines - ndreamer/
statistics.py , Python, 142 lines - ndreamer/
test.py , Python, 337 lines - setup.py, Python, 27 lines
- tutorial_case_control_T1
D.ipynb , Jupyter, 107 lines - tutorial_experimental_PB
MC.ipynb , Jupyter, 267 lines, 1 match - LICENSE, License, 674 lines
- README.md, Text, 152 lines
lugia-xiao/NDreamer_reproducible
118684aa70fe23fcb60985c54121984a27ede9d2, 16 May 2026Availability: 1 check, the latest on 26 September 2026: the link answers
- 26 September 2026: the link answers
29 files
- R_plot_codes/
AD_ADNC_code/ , R, 51 lines, 1 matchAD_neurons_GO.R - R_plot_codes/
AD_ADNC_code/ , R, 36 lines, 2 matchesAD_respiration_GO.R - R_plot_codes/
AD_ADNC_code/ , R, 48 linesGO_AD.R - R_plot_codes/
AD_ADNC_code/ , R, 20 linescommon_genes.R - R_plot_codes/
case_control_code/ , R, 139 lines, 1 matchradar.R - R_plot_codes/
case_control_code/ , R, 127 lines, 1 matchradar_mouse.R - R_plot_codes/
case_control_code/ , R, 52 lines, 1 matchsimple_GO.R - R_plot_codes/
case_control_code/ , R, 46 linessimple_GO_mouse.R - R_plot_codes/
experimental_benchmark_p , R, 43 lines, 2 matcheslot_codes/ ablation.R - R_plot_codes/
experimental_benchmark_p , R, 92 lines, 2 matcheslot_codes/ overall_plot_experimenta l.R - R_plot_codes/
experimental_benchmark_p , R, 93 lines, 1 matchlot_codes/ plot_experimental.R - R_plot_codes/
experimental_benchmark_p , R, 104 lines, 1 matchlot_codes/ plot_experimental_have_b atch.R - R_plot_codes/
experimental_benchmark_p , R, 104 lines, 1 matchlot_codes/ plot_experimental_have_b atch_ndreamer.R - benchmark/
cellanova.ipynb , Jupyter, 160 lines - benchmark/
cellanova.nbconvert.ipyn , Jupyter, 160 linesb - benchmark/
cinema_ot.ipynb , Jupyter, 140 lines - benchmark/
cinema_ot.nbconvert.ipyn , Jupyter, 140 linesb - benchmark/
cinema_ot_ITE.ipynb , Jupyter, 113 lines - benchmark/
cinema_ot_ITE.nbconvert. , Jupyter, 113 linesipynb - benchmark/
evaluate_batch.ipynb , Jupyter, 147 lines - benchmark/
evaluate_batch.nbconvert , Jupyter, 147 lines.ipynb - benchmark/
evaluate_cellanova.ipynb , Jupyter, 268 lines - benchmark/
evaluate_scd.ipynb , Jupyter, 291 lines - benchmark/
get_cellanova_mse.ipynb , Jupyter, 144 lines - benchmark/
metrics.py , Python, 275 lines - benchmark/
mixscape.ipynb , Jupyter, 172 lines - benchmark/
mixscape.nbconvert.ipynb , Jupyter, 172 lines - benchmark/
plot_cellanova.ipynb , Jupyter, 104 lines - repository limit reached (2,000 files or 30 MB): the rest is at the source (187 files)
- LICENSE, License, 674 lines
Code availability
The NDreamer package is available at https://
Reproduced under the paper's license (CC BY-NC), from the paper cited above.
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:
- 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 43 scripts, each with its path and the digest of its content;
- 24 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.5061/
dryad.4xgxd25g1 , at Dryad; found in “Data availability”
Data availability
The PBMC dataset is available from the GEO under accession number GSE96583 and can be downloaded from the scGen tutorial https://
Reproduced under the paper's license (CC BY-NC), from the paper cited above.
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, 27 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 3 authors, 6 keywords, 7 MeSH terms, 1 funder, 50 references.
Cite
This paper
Xiao, X., Zhao, H., & Wang, Z. (2026). Single-cell-level perturbation-induced and condition-related signal estimation with batch effect removal using NDreamer. Briefings in bioinformatics, 27(5), bbag485. https://
BibTeX
@article{xiao2026single,
author = {Xiao, Xiao and Zhao, Hongyu and Wang, Zuoheng},
title = {{Single-cell-level perturbation-induced and condition-related signal estimation with batch effect removal using NDreamer}},
journal = {Briefings in bioinformatics},
year = {2026},
month = sep,
volume = {27},
number = {5},
pages = {bbag485},
publisher = {Oxford University Press},
issn = {1467-5463},
doi = {10.1093/
url = {https://
pmid = {42734923},
pmcid = {PMC13573602}
}
RIS
TY - JOUR
AU - Xiao, Xiao
AU - Zhao, Hongyu
AU - Wang, Zuoheng
TI - Single-cell-level perturbation-induced and condition-related signal estimation with batch effect removal using NDreamer
T2 - Briefings in bioinformatics
J2 - Brief Bioinform
PY - 2026
DA - 2026/
VL - 27
IS - 5
SP - bbag485
SN - 1467-5463
PB - Oxford University Press
DO - 10.1093/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1093/
"type": "article-journal",
"title": "Single-cell-level perturbation-induced and condition-related signal estimation with batch effect removal using NDreamer",
"container-title": "Briefings in bioinformatics",
"author": [
{
"family": "Xiao",
"given": "Xiao"
},
{
"family": "Zhao",
"given": "Hongyu"
},
{
"family": "Wang",
"given": "Zuoheng"
}
],
"container-title-short":
"volume": "27",
"issue": "5",
"page": "bbag485",
"DOI": "10.1093/
"PMID": "42734923",
"PMCID": "PMC13573602",
"ISSN": "1467-5463",
"publisher": "Oxford University Press",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
9,
1
]
]
}
}
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.1093/bib/bbag175 [code]
- Uncovering causal relationships in single-cell omic studies with causarray.Journal: Briefings in bioinformaticsIn common: anndata, clusterProfiler, Scanpy, 9 other tools, Alzheimer's / dementia, 4 references
- [2] doi:10.1016/j.cpblue.2026.100007 [code]
- An integrated single-cell and spatial proteotranscriptomics atlas of fibroblast-driven immunoregulation within the human adult oral cavity.Journal: Cell press blueIn common: rpy2, UMAP, anndata, 12 other tools, 1 reference
- [3] doi:10.1038/s41586-026-10490-y [code]
- Lineage and organ signals sequentially build organ intrinsic nervous systems.Journal: NatureIn common: UMAP, anndata, Scanpy, 10 other tools, author Hongyu Zhao
- [4] doi:10.1016/j.isci.2026.116055 [code]
- Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.Journal: iScienceIn common: rpy2, UMAP, anndata, 11 other tools, genetics / omics, 1 reference
- [5] doi:10.1016/j.xcrm.2026.102766 [code]
- A longitudinal single-cell and spatial multiomic atlas of pediatric high-grade glioma.Journal: Cell reports. MedicineIn common: rpy2, UMAP, anndata, 11 other tools, genetics / omics
- [6] doi:10.1038/s41514-026-00391-9 [code]
- Region-specific transcriptional signatures of brain aging in the absence of neuropathology at the single-cell level.Journal: npj agingIn common: anndata, Scanpy, reshape2, 10 other tools, genetics / omics, 2 references
- [7] doi:10.1038/s41586-026-10629-x [code]
- Whole-genome duplication shaped cell-type evolution in the vertebrate brain.Journal: NatureIn common: UMAP, anndata, clusterProfiler, 11 other tools, genetics / omics
- [8] doi:10.1371/journal.pcbi.1014327 [code]
- Supervised deep learning with gene functional annotation for cell classification.Journal: PLoS computational biologyIn common: anndata, Scanpy, reshape2, 10 other tools, Alzheimer's / dementia, genetics / omics, 1 reference
- [9] doi:10.1016/j.celrep.2026.117073 [code]
- Single-cell epigenomics uncovers heterochromatin instability and transcription factor dysfunction during mouse brain aging.Journal: Cell reportsIn common: rpy2, anndata, clusterProfiler, 9 other tools, genetics / omics, 1 reference
- [10] doi:10.21203/rs.3.rs-9676637/v1 [code]
- A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic TechnologiesJournal: Research Square (preprint)In common: rpy2, UMAP, anndata, 10 other tools, genetics / omics
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: 2 repositories of the authors' code, each at its verified commit and with its license, 43 scripts, and 24 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:84e334bca5c0203a…
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.
