OSCR

Single-cell-level perturbation-induced and condition-related signal estimation with batch effect removal using NDreamer.

Code ↔ Paper

24 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [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. [17] § Results › Evaluation metrics ↔ ndreamer/metrics.py, lines 164–278 · score 0.57 · bLISI, asw_batch, kBET, metrics, cell
  18. [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. [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. [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. [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. [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. [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. [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

  1. import os
  2. import numpy as np
  3. import pandas as pd
  4. import torch
  5. import torch.nn as nn
  6. import torch.nn.functional as F
  7. from torch.cuda import device
  8. from ndreamer.model_DL import NDreamer_generator,Discriminator,set_seed
  9. from ndreamer.DL_loss_func import CrossEntropy, create_triplets_within_groups, IndependenceLoss, kl_divergence_loss, \
  10. IndependenceLoss_label, OrthogonalityRegularization, EntropyPenalty, create_triplets_within_groups_logits, \
  11. IndependenceLoss_between_matrix, compute_mmd, compute_local_neighborhood_loss, create_triplets_within_groups_tensor
  12. from ndreamer.data_preprocess import process_adata,generate_balanced_dataloader,generate_adata_to_dataloader
  13. from ndreamer.DL_loss_func import reconstruction_error
  14. from ndreamer.statistics import *
  15. from sklearn.decomposition import PCA
  16. # Dynamic import of tqdm based on the environment
  17. import sys
  18. if 'ipykernel' in sys.modules:
  19. from tqdm.notebook import tqdm
  20. else:
  21. from tqdm import tqdm
  22. from ndreamer.plot import plot_distribution
  23. class NDreamer_DL(nn.Module):
  24. def __init__(self, adata, condition_key, contorl_name, num_hvg, require_batch=False, batch_key=None,
  25. resolution_low=0.5, resolution_high=7, cluster_method="Leiden", embedding_dim=512, codebooks=None,
  26. codebook_dim=8, encoder_hidden=None, decoder_hidden=None, z_dim=256,
  27. cos_loss_scaler=20, random_seed=123, batch_size=2048, epoches=10, lr=1e-3, triplet_margin=5,
  28. independent_loss_scaler=1000, save_pth="./model/", developer_test_mode=False,
  29. library_size_normalize_adata=False, save_preprocessed_adata_path="./model/preprocessed.h5ad",
  30. KL_scaler=5e-3, reconstruct_scaler=1, triplet_scaler=5, num_triplets_per_label=15,
  31. tau=0.01, commitment_loss_scaler=1, cluster_correlation_scaler=50,
  32. try_identify_perturb_escaped_cell=False, reset_threshold=1 / 1024, reset_interval=30,
  33. try_identify_cb_specific_subtypes=False, local_neighborhood_loss_scaler=1,
  34. local_neighbor_sigma=1, n_neighbors=20, local_neighbor_across_cluster_scaler=20,
  35. have_negative_data=False):
  36. super(NDreamer_DL, self).__init__()
  37. input_dim=num_hvg
  38. if codebooks is None:
  39. codebooks = [1024 for i in range(32)]
  40. if encoder_hidden is None:
  41. encoder_hidden = [2048, 1024]
  42. if decoder_hidden is None:
  43. decoder_hidden = [512, 1024]
  44. self.developer_test_mode=developer_test_mode
  45. self.require_batch=require_batch
  46. self.batch_key = batch_key
  47. self.condition_key = condition_key
  48. self.triplet_margin=triplet_margin
  49. if not os.path.exists(save_pth):
  50. os.mkdir(save_pth)
  51. self.save_pth=save_pth
  52. self.random_seed = random_seed
  53. self.batch_size = batch_size
  54. self.epoches = epoches
  55. self.lr = lr
  56. self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
  57. print("Using device:",self.device)
  58. self.kl_scaler = KL_scaler
  59. self.reconstruct_scaler = reconstruct_scaler
  60. self.triplet_scaler = triplet_scaler
  61. self.num_triplets_per_label=num_triplets_per_label
  62. self.independent_loss_scaler = independent_loss_scaler
  63. self.codebooks=codebooks
  64. self.local_neighborhood_loss_scaler=local_neighborhood_loss_scaler
  65. print(local_neighborhood_loss_scaler)
  66. self.local_neighbor_sigma=local_neighbor_sigma
  67. self.try_identify_cb_specific_subtypes=try_identify_cb_specific_subtypes
  68. self.try_identify_perturb_escaped_cell=try_identify_perturb_escaped_cell
  69. self.n_neighbors=n_neighbors
  70. self.local_neighbor_across_cluster_scaler=local_neighbor_across_cluster_scaler
  71. #data preprocessing
  72. if not developer_test_mode:
  73. print("Start data preprocessing")
  74. adata, batch_dict, condition_dict = process_adata(adata=adata, condition_key=condition_key,
  75. input_dim=input_dim, control_name=contorl_name,
  76. require_batch=require_batch, batch_key=batch_key,
  77. resolution_low=resolution_low,
  78. resolution_high=resolution_high,
  79. cluster_method=cluster_method,
  80. library_size_normalize_adata=library_size_normalize_adata)
  81. self.adata = adata
  82. self.batch_dict = batch_dict
  83. self.condition_dict = condition_dict
  84. print("Data preprocessing done")
  85. if save_preprocessed_adata_path is not None:
  86. adata.write(save_preprocessed_adata_path)
  87. else:
  88. self.adata=adata
  89. print("Remaining number of cells:",self.adata.shape[0])
  90. calcluated_epoches=15*self.adata.shape[0]//(self.batch_size*len(np.unique(self.adata.obs["group"])))+1
  91. if calcluated_epoches>self.epoches:
  92. if len(np.unique(self.adata.obs["group"]))<=10:
  93. calcluated_epoches = calcluated_epoches + reset_interval - calcluated_epoches % reset_interval - 2
  94. self.epoches=calcluated_epoches
  95. print("Too few epoches (steps, if rigorously speaking). Changing epoch to", self.epoches, "to adjust for number of cells")
  96. num_batches = []#np.unique(adata.obs["batch"]).shape[0]
  97. if not self.developer_test_mode and self.require_batch:
  98. if isinstance(batch_key, str):
  99. num_batches = [max(self.batch_dict[batch_key].values()) + 1]
  100. elif isinstance(batch_key, list):
  101. num_batches=[]
  102. for batch_keyi in batch_key:
  103. num_batches.append(max(self.batch_dict[batch_keyi].values()) + 1)
  104. else:
  105. num_batches=[]
  106. self.num_batches=num_batches
  107. num_treatments = np.unique(adata.obs["condition"]).shape[0]
  108. if not developer_test_mode:
  109. num_treatments=max(num_treatments,max(self.condition_dict.values())+1)
  110. # define the model
  111. self.VQ_VAE = NDreamer_generator(input_dim=input_dim, num_treatments=num_treatments,z_dim=z_dim,
  112. num_batches=num_batches,embedding_dim=embedding_dim,tau=tau,
  113. codebooks=codebooks,codebook_dim=codebook_dim,encoder_hidden=encoder_hidden,
  114. decoder_hidden=decoder_hidden,commitment_loss_scaler=commitment_loss_scaler,
  115. reset_threshold=reset_threshold,reset_interval=reset_interval,
  116. try_identify_cb_specific_subtypes=try_identify_cb_specific_subtypes,
  117. try_identify_perturb_escaped_cell=try_identify_perturb_escaped_cell,
  118. have_negative_data=have_negative_data)
  119. self.require_batch = require_batch
  120. print("Require batch:",self.require_batch)
  121. self.independence_loss=IndependenceLoss(scaler=independent_loss_scaler)
  122. self.independence_loss_label=IndependenceLoss_label(scaler=cluster_correlation_scaler/num_treatments)
  123. self.independence_loss_between_codebook=IndependenceLoss_between_matrix(scaler=independent_loss_scaler/100)
  124. self.cross_entropy = CrossEntropy()
  125. self.cos_loss_scaler = cos_loss_scaler
  126. #self.MINE=MINE(x_dim=input_dim,z_dim=z_dim,hidden_dim=1024)
  127. # init the model
  128. self.VQ_VAE.to(self.device)
  129. self.independence_loss.to(self.device)
  130. self.independence_loss_label.to(self.device)
  131. #self.cross_entropy.to(self.device)
  132. #self.orthogonality_regularization.to(self.device)
  133. #self.entropy_penalty.to(self.device)
  134. self.independence_loss_between_codebook.to(self.device)
  135. #self.MINE.to(self.device)
  136. self.logits=None
  137. self.df_latent=None
  138. def train_model(self):
  139. optimizer_G = torch.optim.AdamW(self.VQ_VAE.parameters(), lr=self.lr)
  140. progress_bar = tqdm(total=self.epoches, desc="Overall Progress", leave=True, miniters=1, mininterval=0)
  141. for epoch in range(self.epoches):
  142. set_seed(self.random_seed+epoch)
  143. data_loader = generate_balanced_dataloader(self.adata, batch_size=self.batch_size)
  144. self.VQ_VAE.train()
  145. all_losses = 0
  146. T_loss = 0
  147. V_loss = 0
  148. I_loss = 0
  149. K_loss = 0
  150. N_loss = 0
  151. Commitment_loss=0
  152. Dependent_loss=0
  153. for i, (exp, condition, batch, labels_low, labels_high) in enumerate(data_loader):
  154. #print(exp.shape, condition.shape, batch.shape, labels_low.shape, labels_high.shape)
  155. # convert to cuda
  156. exp = exp.to(self.device)
  157. condition = condition.to(self.device)
  158. batch = batch.to(self.device)
  159. labels_low = labels_low.to(self.device)
  160. labels_high = labels_high.to(self.device)
  161. # run VQ-VAE
  162. z, variance, reconstructed, logits, commitment_loss,escape_judger_choice, cb_specifc_embedding=self.VQ_VAE(exp=exp, treatment=condition, batch=batch)
  163. commitment_loss=commitment_loss/len(self.codebooks)
  164. if self.VQ_VAE.encoder.codebooks[0].just_reset_codebook:
  165. print("Finish resetting codebook embeddings, current step (epoch):",epoch)
  166. continue
  167. # calculate the reconstruction loss
  168. reconst_loss_mse = reconstruction_error(exp, reconstructed)*self.reconstruct_scaler
  169. reconst_loss_cos = (1 - torch.sum(F.normalize(reconstructed, p=2) * F.normalize(exp, p=2), 1)).mean()
  170. reconst_loss_cos = self.cos_loss_scaler * reconst_loss_cos
  171. reconst_loss = reconst_loss_mse + reconst_loss_cos
  172. reconst_loss = torch.clamp(reconst_loss, max=1e5)
  173. KL_loss=kl_divergence_loss(z_mean=z,z_var=variance)*self.kl_scaler
  174. KL_loss = torch.clamp(KL_loss, max=1e5)
  175. # independence loss for VAE
  176. independent_loss=0
  177. independent_loss=independent_loss+self.independence_loss(condition, logits)
  178. for i in range(len(self.num_batches)):
  179. independent_loss = independent_loss + self.independence_loss(batch[:,i], logits)
  180. independent_loss = torch.clamp(independent_loss, max=1e5)
  181. # convert matrix batch to unique int number vector
  182. base = batch.max() + 1 # Base larger than the largest number in the matrix
  183. powers = base ** torch.arange(batch.size(1), device=exp.device).flip(0).float() # Positional encoding powers
  184. batch_vector = torch.matmul(batch.float(), powers).long() # Encoded as unique integers
  185. dependent_loss = 0
  186. dependent_loss=dependent_loss-self.independence_loss_label(assignments=labels_high,
  187. probabilities=logits,
  188. condition=condition,
  189. batch=batch_vector)
  190. dependent_loss=dependent_loss-self.independence_loss_label(assignments=labels_low,
  191. probabilities=logits,
  192. condition=condition,
  193. batch=batch_vector)
  194. # triplet loss
  195. triplet_loss = create_triplets_within_groups_tensor(z, labels_low, labels_high, batch_vector,
  196. condition, margin=self.triplet_margin,
  197. num_triplets=self.num_triplets_per_label)
  198. triplet_loss=triplet_loss*self.triplet_scaler
  199. neighbor_loss = compute_local_neighborhood_loss(latent=z, inputs=exp,
  200. condition=condition, batch=batch_vector,
  201. sigma=self.local_neighbor_sigma,
  202. n_neighbors=self.n_neighbors,
  203. label_low=labels_low,
  204. label_high=labels_high,
  205. local_neighbor_across_cluster_scaler=self.local_neighbor_across_cluster_scaler
  206. ) * self.local_neighborhood_loss_scaler
  207. all_loss = reconst_loss + independent_loss +KL_loss + triplet_loss + neighbor_loss + commitment_loss + dependent_loss# + embedding_loss +entropy_loss
  208. all_loss = torch.clamp(all_loss, max=1e5)
  209. all_loss.backward()
  210. optimizer_G.step()
  211. optimizer_G.zero_grad()
  212. all_losses += all_loss.item()/len(data_loader)
  213. N_loss += neighbor_loss.item()/len(data_loader)
  214. T_loss += triplet_loss.item()/len(data_loader)
  215. V_loss += reconst_loss.item()/len(data_loader)
  216. I_loss +=independent_loss.item()/len(data_loader)
  217. K_loss +=KL_loss.item()/len(data_loader)
  218. Dependent_loss+=dependent_loss.item()/len(data_loader)
  219. #Entropy_loss+=entropy_loss.item()
  220. #Embedding_loss += 0#embedding_loss.item()
  221. Commitment_loss+=commitment_loss.item()/len(data_loader)
  222. print(f"Epoch: {epoch + 1}/{self.epoches} | "
  223. f"All Loss: {all_losses:.4f} | "
  224. f"Neighborhood Loss: {N_loss:.4f} | "
  225. f"Triplet Loss: {T_loss:.4f} | "
  226. f"Reconstruction Loss: {V_loss:.4f} | "
  227. f"Independent Loss: {I_loss:.4f} | "
  228. f"KL Loss: {K_loss:.4f} | "
  229. f"Commitment Loss: {Commitment_loss:.4f} | "
  230. f"Dependent Loss: {Dependent_loss:.4f}")
  231. progress_bar.update(1) # Increment the progress bar by one for each batch processed
  232. 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)
  233. progress_bar.close()
  234. torch.save(self.state_dict(), os.path.join(self.save_pth, 'ndreamer' + '.pth'))
  235. def get_modifier_space(self,dim=50, save_path_adata=None, save_path_latent_df=None,
  236. print_codebook_statisitics=False, kBET_test=False, save_others=False):
  237. self.VQ_VAE.eval()
  238. data_loader = generate_adata_to_dataloader(self.adata)
  239. device = 'cuda' if torch.cuda.is_available() else 'cpu'
  240. all_z = []
  241. all_indices = []
  242. logits=[]
  243. cb_specific_subtypes=[]
  244. escape_probs=[]
  245. with torch.no_grad():
  246. for i, (x, indices, condition, batch) in enumerate(data_loader):
  247. x = x.to(device)
  248. condition = condition.to(device)
  249. batch = batch.to(device)
  250. z, variance, reconstructed, logit, commitment_loss, escape_judger_choice, cb_specifc_embedding = self.VQ_VAE(
  251. exp=x, treatment=condition, batch=batch)
  252. all_z.append(z.cpu().detach())
  253. all_indices.extend(indices.tolist())
  254. if self.try_identify_perturb_escaped_cell:
  255. escape_probs.append(escape_judger_choice.detach().cpu())
  256. if self.try_identify_cb_specific_subtypes:
  257. cb_specific_subtypes.append(cb_specifc_embedding.detach().cpu())
  258. if len(logits) == 0:
  259. logits = [i.detach().cpu() for i in logit]
  260. else:
  261. logits = [torch.concatenate([logits[i], logit[i].detach().cpu()], dim=0) for i in
  262. range(len(logits))]
  263. all_z_combined = torch.cat(all_z, dim=0)
  264. all_indices_tensor = torch.tensor(all_indices)
  265. all_z_reordered = all_z_combined[all_indices_tensor.argsort()]
  266. all_z_np = all_z_reordered.numpy()
  267. # Create anndata object with reordered embeddings
  268. self.adata.obsm['X_effect_modifier_space'] = all_z_np
  269. pca = PCA(n_components=dim)
  270. # Fit and transform the data
  271. X_pca = pca.fit_transform(self.adata.obsm['X_effect_modifier_space'])
  272. # Store the PCA-reduced data back into adata.obsm
  273. self.adata.obsm['X_effect_modifier_space_PCA'] = X_pca
  274. if save_path_adata is not None:
  275. if save_path_adata.find(".h5ad")<0:
  276. save_path_adata = save_path_adata+".h5ad"
  277. self.adata.write(save_path_adata)
  278. else:
  279. self.adata.write(os.path.join(self.save_pth, 'adata' + '.h5ad'))
  280. print("Effect modifier space saved.")
  281. if self.try_identify_cb_specific_subtypes:
  282. print(cb_specific_subtypes)
  283. cb_specific_subtypes = torch.concat(cb_specific_subtypes, dim=0)
  284. plot_distribution(cb_specific_subtypes.reshape(-1).numpy())
  285. self.adata.obs["original_order"]=np.array(range(self.adata.shape[0]))
  286. if not save_others:
  287. return
  288. torch.save(logits, os.path.join(self.save_pth, "logits.pth"))
  289. self.logits = logits
  290. if self.try_identify_perturb_escaped_cell:
  291. escape_probs = torch.concat(escape_probs, dim=0)
  292. torch.save(escape_probs, os.path.join(self.save_pth, "escape_probs.pth"))
  293. if self.try_identify_cb_specific_subtypes:
  294. cb_specific_subtypes = torch.concat(cb_specific_subtypes, dim=0)
  295. torch.save(cb_specific_subtypes, os.path.join(self.save_pth, "cb_specific_subtypes.pth"))
  296. 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])])
  297. df_latent["batch"]=np.array(self.adata.obs["batch"])
  298. df_latent["condition"]=np.array(self.adata.obs["condition"])
  299. df_latent["group"]=np.array(self.adata.obs["group"])
  300. if save_path_latent_df is not None:
  301. if save_path_latent_df.find(".csv")<0:
  302. save_path_latent_df = save_path_latent_df+".csv"
  303. df_latent.to_csv(save_path_latent_df)
  304. else:
  305. df_latent.to_csv(os.path.join(self.save_pth, 'latent.csv'))
  306. self.df_latent=df_latent
  307. if print_codebook_statisitics:
  308. print(torch.max(self.logits[0],dim=-1))
  309. print("Statistic values of the codebook selection")
  310. print("mean:\n",[torch.mean(logits[i], dim=0) for i in range(len(logits))])
  311. print("variance:\n",[torch.std(logits[i], dim=0) for i in range(len(logits))])
  312. if kBET_test:
  313. run_kbet(df_latent,do_PCA=False)
  314. def run_kBET_test(self):
  315. run_kbet(self.df_latent, do_PCA=False)
  316. if __name__ == '__main__':
  317. test_run_mode=False
  318. train_model=True
  319. if test_run_mode:
  320. import scanpy as sc
  321. import anndata as ad
  322. '''exp = torch.abs(torch.randn((2048, 2000)))
  323. condition = torch.randint(low=0, high=3, size=(2048,), dtype=torch.long)
  324. batch1 = torch.randint(low=0, high=4, size=(2048,), dtype=torch.long) # np.zeros(2048)
  325. batch2 = torch.randint(low=0, high=2, size=(2048,), dtype=torch.long)
  326. adata = ad.AnnData(X=exp.numpy())
  327. adata.obs["condition"] = condition.numpy()
  328. adata.obs["group"] = condition.numpy()
  329. adata.obs["batch1"] = batch1
  330. adata.obs["batch"] = batch1
  331. adata.obs["batch2"] = batch2
  332. adata.obs["leiden1"] = torch.randint(low=0, high=3, size=(2048,), dtype=torch.long).numpy()
  333. adata.obs["leiden1"] = adata.obs["leiden1"].astype("category")
  334. adata.obs["leiden2"] = torch.randint(low=0, high=9, size=(2048,), dtype=torch.long).numpy()
  335. adata.obs["leiden2"] = adata.obs["leiden2"].astype("category")
  336. model = NDreamer_DL(adata=adata, condition_key="condition", contorl_name=0, num_hvg=2000,
  337. developer_test_mode=False, require_batch=True, batch_key="batch1",#["batch1","batch2"],
  338. batch_size=512)
  339. model.train_model()'''
  340. adata = sc.read("../ECCITE_preprocessed.h5ad")
  341. adata.obsm["batch"]=np.expand_dims(adata.obs["batch"].copy(),axis=-1)
  342. print(adata)
  343. print(np.unique(adata.obs["condition"]))
  344. model = NDreamer_DL(adata, condition_key='perturbation', contorl_name='NT', num_hvg=2000, require_batch=True,
  345. batch_key='replicate',
  346. resolution_low=0.5, resolution_high=7, cluster_method="Leiden", embedding_dim=512,
  347. codebooks=[1024 for i in range(32)],
  348. codebook_dim=8, encoder_hidden=[1024, 512], decoder_hidden=[512, 1024], z_dim=256,
  349. cos_loss_scaler=20, random_seed=123, batch_size=1024, epoches=1, lr=1e-3,
  350. triplet_margin=5, independent_loss_scaler=1000, save_pth="./model/",
  351. developer_test_mode=True,
  352. library_size_normalize_adata=False,
  353. save_preprocessed_adata_path=None,
  354. KL_scaler=5e-3, reconstruct_scaler=1, triplet_scaler=5, num_triplets_per_label=15,
  355. tau=0.01, commitment_loss_scaler=1, cluster_correlation_scaler=50, reset_threshold=1 / 1024,
  356. reset_interval=30, try_identify_cb_specific_subtypes=False,
  357. local_neighborhood_loss_scaler=1, local_neighbor_sigma=1,
  358. try_identify_perturb_escaped_cell=False, n_neighbors=20,
  359. local_neighbor_across_cluster_scaler=20)
  360. model.train_model()
  361. print()
  362. model.get_modifier_space(save_path_adata="../ECCITE_results.h5ad")
  363. adata1 = model.adata.copy()
  364. sc.pp.neighbors(adata1, use_rep='X_effect_modifier_space_PCA', n_neighbors=15)
  365. sc.tl.umap(adata1)
  366. sc.pl.umap(adata1, color=["MULTI_ID","HTO_classification",
  367. "gene_target","perturbation",
  368. "replicate","Phase"], frameon=False, ncols=1)
  369. else:
  370. import os
  371. os.environ["CUDA_LAUNCH_BLOCKING"] = "1"
  372. import scanpy as sc
  373. #adata=sc.read_h5ad("../data/PBMC.h5ad")
  374. #adata = sc.read("../PBMC_preprocessed.h5ad")
  375. adata = sc.read("../virus_preprocessed.h5ad")
  376. #adata=sc.read_h5ad("../PBMC_imbalance_preprocessed.h5ad")
  377. #adata=sc.read_h5ad("../common1.h5ad")
  378. if "condition" not in adata.obsm.keys():
  379. adata.obsm["batch"] = np.ones((adata.shape[0],1))#np.expand_dims(adata.obs["condition"].copy(), axis=1)
  380. print(adata)
  381. print(np.unique(adata.obs["condition"]))
  382. model = NDreamer_DL(adata, condition_key="condition", contorl_name="control", num_hvg=min(adata.shape[1],3608),
  383. require_batch=False,
  384. batch_key=None,
  385. resolution_low=0.5, resolution_high=7, cluster_method="Leiden", embedding_dim=512,
  386. codebooks=[1024 for i in range(32)],
  387. codebook_dim=8, encoder_hidden=[2048, 1024], decoder_hidden=[512, 1024], z_dim=256,
  388. cos_loss_scaler=20, random_seed=123, batch_size=1024, epoches=100, lr=1e-3,
  389. triplet_margin=5,independent_loss_scaler=1000, save_pth="./model/",
  390. developer_test_mode=True,
  391. library_size_normalize_adata=False,
  392. save_preprocessed_adata_path=None,
  393. KL_scaler=5e-3, reconstruct_scaler=1, triplet_scaler=5, num_triplets_per_label=15,
  394. tau=0.01, commitment_loss_scaler=1, cluster_correlation_scaler=50,reset_threshold=1/1024,
  395. reset_interval=30,try_identify_cb_specific_subtypes=False,
  396. local_neighborhood_loss_scaler=1,local_neighbor_sigma=1,
  397. try_identify_perturb_escaped_cell=False,n_neighbors=20,
  398. local_neighbor_across_cluster_scaler=20, have_negative_data=True)
  399. '''model = NDreamer_DL(adata, condition_key="condition", contorl_name="control", num_hvg=min(adata.shape[1], 3608),
  400. require_batch=False,
  401. batch_key=None,have_negative_data=True,developer_test_mode=True)'''
  402. if train_model:
  403. model.train_model()
  404. print()
  405. model.load_state_dict(torch.load(os.path.join(model.save_pth, 'ndreamer' + '.pth')))
  406. model.get_modifier_space(save_path_adata="../PBMC_results.h5ad")
  407. adata1 = model.adata.copy()
  408. sc.pp.neighbors(adata1, use_rep='X_effect_modifier_space_PCA',n_neighbors=15)
  409. sc.tl.umap(adata1)
  410. try:
  411. sc.pl.umap(adata1, color=['condition', 'cell_type'], frameon=False, ncols=1)
  412. except:
  413. sc.pl.umap(adata1, color=['condition', 'cell_type1021'], frameon=False, ncols=1)
  414. sc.pp.neighbors(adata1, use_rep='X_effect_modifier_space',n_neighbors=15)
  415. sc.tl.umap(adata1)
  416. try:
  417. sc.pl.umap(adata1, color=['condition', 'cell_type'], frameon=False, ncols=1)
  418. except:
  419. sc.pl.umap(adata1, color=['condition', 'cell_type1021'], frameon=False, ncols=1)
  420. #model.run_kBET_test()

model.py at commit dfff8e9, under GPL-3.0 · at the source

Overview

Authors: Xiao Xiao1, Hongyu Zhao1,2,3, Zuoheng Wang1,4
  1. Department of Biostatistics, Yale University School of Public Health, 300 George Street, New Haven, CT 06511, United States
  2. Department of Genetics, Yale University School of Medicine, 333 Cedar Street, New Haven, CT 06520, United States
  3. Department of Statistics and Data Science, Yale University, 219 Prospect Street, New Haven, CT 06511, United States
  4. Department of Biomedical Informatics & Data Science, Yale University School of Medicine, 101 College Street, New Haven, CT 06510, United States
Institutions: Yale University (United States)
Journal: Briefings in bioinformatics, volume 27, issue 5, article bbag485
Dates: received 17 February 2026; accepted 6 August 2026; published online 14 September 2026; in print September 2026
Type: Research article · Language: English
License: CC BY-NC
Identifiers: DOI 10.1093/bib/bbag485 · PMID 42734923 · PMCID PMC13573602 · OpenAlex W7213096958
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: genetics / omics (modality), human (organism), Alzheimer's / dementia (population)
Methods: Connectivity, Preprocessing, Machine learning
Keywords: scRNA-seq, batch effect, neural discrete representation learning, counterfactual causal matching, perturbation-induced signal, disease-associated signal
MeSH: Alzheimer Disease*, Single-Cell Analysis*, Algorithms, Animals, Gene Expression Profiling, Humans, Single-Cell Gene Expression Analysis (* major topic)
Journal subjects: Problem Solving Protocol
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: NIH (R01LM014087)
Citations: not cited yet (Europe PMC); 50 references in the paper

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

License: GPL-3.0
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: dfff8e9fb8c0a9bb246468d06d42aee5e666b1bf, 20 March 2025
Languages: Python (13), Jupyter (2)
Size: 20 files, 15 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (setup.py), 2 notebooks
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (12 files), PyTorch (11 files), Scanpy (9 files), SciPy (6 files), pandas (5 files), scikit-learn (5 files), anndata (4 files), Matplotlib (4 files), rpy2 (3 files), seaborn (2 files), statsmodels (1 file), UMAP (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
17 files

lugia-xiao/NDreamer_reproducible

License: GPL-3.0
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 118684aa70fe23fcb60985c54121984a27ede9d2, 16 May 2026
Languages: Jupyter (136), Python (65), R (13)
Size: 471 files, 214 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (case_control_benchmark/scDisInFact_code/scDisInFact-main/setup.py), tests, 136 notebooks
Not found: CITATION.cff, continuous integration, documentation
Tools: NumPy (15 files), Scanpy (15 files), Matplotlib (12 files), rpy2 (11 files), ggplot2 (9 files), PyTorch (9 files), scikit-learn (9 files), pandas (7 files), anndata (6 files), seaborn (6 files), tidyverse (6 files), clusterProfiler (5 files), SciPy (3 files), reshape2 (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
29 files

Code availability

The NDreamer package is available at https://github.com/lugia-xiao/NDreamer with detailed tutorials. All code used to reproduce the results shown in the article are available at https://github.com/lugia-xiao/NDreamer_reproducible.

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

Data availability

The PBMC dataset is available from the GEO under accession number GSE96583 and can be downloaded from the scGen tutorial https://drive.google.com/uc?id=1r87vhoLLq6PXAYdmyyd89zG90eJOFYLk. The rhinovirus infection dataset and another PBMC dataset are available at https://datadryad.org/stash/dataset/doi:10.5061/dryad.4xgxd25g1 (https://doi.org/10.5061/dryad.4xgxd25g1). The ECCITE-seq dataset is available from the GEO under accession number GSE153056 and can be downloaded from pertpy [2] using the code ‘pertpy.data.papalexi_2021()’. The ASD dataset is available from the GEO under accession number GSE157977. The T1D dataset is available from the GEO under accession number GSE148073 and can be downloaded from https://cellxgene.cziscience.com/collections/51544e44-293b-4c2b-8c26-560678423380. The kidney dataset is available from the GEO under accession number GSE211785, and we use the version ‘GSE211785_Susztak_SC_SN_ATAC_merged_PreSCVI_final.h5ad.gz’. The mouse radiation therapy dataset is available from the GEO under accession number GSE280883. The SEA-AD MTG snRNA-seq dataset is available at https://cellxgene.cziscience.com/collections/1ca90a2d-2943-483d-b678-b809bf464c30. The SEA-AD dataset is available at https://cellxgene.cziscience.com/collections/1ca90a2d-2943-483d-b678-b809bf464c30.

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://doi.org/10.1093/bib/bbag485

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/bib/bbag485},
url = {https://doi.org/10.1093/bib/bbag485},
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/09/01
VL - 27
IS - 5
SP - bbag485
SN - 1467-5463
PB - Oxford University Press
DO - 10.1093/bib/bbag485
UR - https://doi.org/10.1093/bib/bbag485
LA - en
ER -

CSL-JSON

{
"id": "10.1093/bib/bbag485",
"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": "Brief Bioinform",
"volume": "27",
"issue": "5",
"page": "bbag485",
"DOI": "10.1093/bib/bbag485",
"PMID": "42734923",
"PMCID": "PMC13573602",
"ISSN": "1467-5463",
"publisher": "Oxford University Press",
"URL": "https://doi.org/10.1093/bib/bbag485",
"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 bioinformatics
In 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 blue
In 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: Nature
In 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: iScience
In 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. Medicine
In 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 aging
In 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: Nature
In 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 biology
In 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 reports
In 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 Technologies
Journal: 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.

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.