OSCR

Supervised deep learning with gene functional annotation for cell classification.

Code ↔ Paper

16 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 16 matches · 1 of them tie a paragraph to a whole file, not to given lines: a weak match, whose lines are not tinted
  1. [1] § Methods › Implementation of SDAN ↔ tutorial.ipynb, lines 270–315 · score 0.81 · BIOGRID ORGANISM, Homo_sapiens, tab3.txt.gz, gene interaction graph, official, symbols
  2. [2] § Methods › Implementation of SDAN ↔ check/check_graph.ipynb, lines 32–166 · score 0.78 · BIOGRID ORGANISM, Homo_sapiens, tab3.txt.gz, official, node, symbols
  3. [3] § Results › Gene expression in astrocyte or microglia can distinguish dementia status ↔ SDAN_Comparison/SEA_AD_spectra_training.py, lines 1–5 · score 0.71 · Seattle Alzheimer, SEA AD, Micro PVM, RNA, Astro, Disease
  4. [4] § Methods › Implementation of sciRED, Spectra, and scNET ↔ SDAN_Comparison/SEA_AD_comparison.py, lines 227–338 · score 0.69 · sciRED, protocol, DE genes, Pearson, Poisson, residuals
  5. [5] § Methods › Implementation of sciRED, Spectra, and scNET ↔ SDAN_Comparison/Su_2020_comparison.py, lines 229–335 · score 0.69 · sciRED, protocol, DE genes, Pearson, Poisson, residuals
  6. [6] § Methods › Implementation of sciRED, Spectra, and scNET ↔ SDAN_Comparison/Su_2020_comparison.py, lines 439–520 · score 0.64 · cell embeddings, expression matrix, scNET, neighbor, zero, model
  7. [7] § Methods › Implementation of SDAN ↔ SDAN/train.py, lines 8–116 · score 0.62 · weight decay, Adam, PyTorch, epochs, optimized, hidden
  8. [8] § Results › Gene expression in astrocyte or microglia can distinguish dementia status ↔ SEA_AD/step7_check_patient_subset.Rmd, lines 205–253 · score 0.61 · Amyloid beta, pTau, AT8, dementia, donors
  9. [9] § Methods › Implementation of SDAN ↔ SDAN_Comparison/Su_2020_comparison.py, lines 439–520 · score 0.61 · weight decay, Adam, PyTorch, epochs, optimized, hidden
  10. [10] § Methods › Implementation of SDAN ↔ check/check_de.py, lines 1–37 · score 0.57 · FDR threshold, DE genes, sensitivity, selection, class, SDAN
  11. [11] § Methods › Implementation of sciRED, Spectra, and scNET ↔ SDAN_Comparison/SEA_AD_comparison.py, lines 390–434 · score 0.57 · balanced subsampling, scNET, concatenated, neighbor, training, graph
  12. [12] § Methods › Extensions ↔ SDAN/train.py, lines 8–116 · score 0.57 · cross entropy loss, graph loss, softmax, scores, classifier, SDAN
  13. [13] § Methods › Extensions ↔ check/check_graph.ipynb, lines 32–166 · score 0.55 · cross entropy loss, graph loss, error, classifier, SDAN, cell
  14. [14] § Methods › Implementation of sciRED, Spectra, and scNET ↔ SDAN_Comparison/Su_2020_spectra_evaluation.py, lines 1–9 · score 0.54 · latent space, union gene, Spectra, CD4, CD8, trained
  15. [15] § Results › Gene expression in astrocyte or microglia can distinguish dementia status ↔ SEA_AD_plot.py, the whole file · a weak match · score 0.52 · SEA AD, Micro PVM, Astro, status, dementia, weights
  16. [16] § Methods › Implementation of SDAN ↔ tutorial.ipynb, lines 229–234 · score 0.51 · unsupervised clustering, dimensionality reduction, resolution, neighbors, Leiden

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 · 23 KB · MIT · 3 matches

  1. """
  2. This file unified version for SDAN, sciRED, and scNET backends.
  3. Compatible with args.py and preprocess.py.
  4. ALL outputs saved into: ./Su_2020/output_comparison/
  5. The backend output folders are:
  6. SDAN: output_comparison/SDAN
  7. sciRED: output_comparison/sciRED
  8. scNET: output_comparison/scNET
  9. """
  10. # ------------------------- Imports & setup -------------------------
  11. import os as _os
  12. _os.environ.setdefault("OMP_NUM_THREADS", "1")
  13. _os.environ.setdefault("OPENBLAS_NUM_THREADS", "1")
  14. _os.environ.setdefault("MKL_NUM_THREADS", "1")
  15. _os.environ.setdefault("NUMEXPR_NUM_THREADS", "1")
  16. _os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
  17. _os.environ.setdefault("NUMBA_THREADING_LAYER", "workqueue")
  18. import os, math, warnings
  19. import numpy as np
  20. import pandas as pd
  21. import scanpy as sc
  22. import scipy.sparse as sp
  23. import torch
  24. from tqdm import tqdm
  25. from anndata import AnnData
  26. from SDAN.model import pipeline
  27. from SDAN.preprocess import qc, construct_gene_graph, construct_gene_list
  28. from SDAN.args import parse_args
  29. from SDAN.evaluation import set_eval_labels, eval_and_save as _eval_and_save
  30. try:
  31. import Spectra
  32. except Exception:
  33. Spectra = None
  34. warnings.simplefilter(action="ignore", category=FutureWarning)
  35. np.random.seed(888)
  36. torch.manual_seed(888)
  37. # sc.settings.verbosity = 3
  38. sc.settings.verbosity = 1
  39. try:
  40. sc.logging.print_header()
  41. except Exception as e:
  42. print(f"[WARN] Skipped Scanpy header due to: {type(e).__name__}: {e}")
  43. sc.settings.set_figure_params(figsize=(8, 6), dpi=80, facecolor="white")
  44. # ------------------------- Paths & args -------------------------
  45. args = parse_args()
  46. d = "./Su_2020/"
  47. cell_type_str = args.cell_type
  48. data_dir = f"{d}gex_{cell_type_str}.mtx.gz"
  49. genes_dir = f"{d}gex_{cell_type_str}_genes.txt"
  50. meta_ind_dir = f"{d}Table_S1.xlsx"
  51. meta_cell_dir = f"{d}cell_info_{cell_type_str}.csv"
  52. os.makedirs(f"{d}output_comparison/", exist_ok=True)
  53. # ------------------------- Load and label -------------------------
  54. def load_and_label_data():
  55. """Load Su_2020 data and assign mild/severe labels."""
  56. print("[INFO] Loading raw data...")
  57. data = sc.read(data_dir, cache=True)
  58. gene_names = pd.read_csv(genes_dir, header=None).iloc[:, 0].astype(str).to_numpy()
  59. if len(gene_names) != data.n_vars:
  60. raise ValueError(f"[ERROR] genes length {len(gene_names)} != n_vars {data.n_vars}")
  61. meta_cell = pd.read_csv(meta_cell_dir)
  62. meta_ind = pd.read_excel(meta_ind_dir, sheet_name="S1.1 Patient Clinical Data")
  63. data.var["gene_symbols"] = gene_names
  64. data.var_names = pd.Index(gene_names)
  65. data.obs["barcode"] = meta_cell["V1"].astype(str).to_numpy()
  66. data.obs_names = pd.Index(data.obs["barcode"])
  67. gene_mito = pd.read_csv("./Annotation/mito_genes.tsv", sep="\t")
  68. mito_col = "hgnc_symbol" if "hgnc_symbol" in gene_mito.columns else gene_mito.columns[0]
  69. data = data[:, ~data.var_names.isin(gene_mito[mito_col])]
  70. data_nonzero_prop = (data.X != 0).sum(axis=0) / data.shape[0]
  71. data = data[:, data_nonzero_prop > 0.02]
  72. wos = meta_ind['Who Ordinal Scale'].astype(str).str.replace('1 or 2', '2', regex=False)
  73. wos = pd.to_numeric(wos, errors='coerce')
  74. meta_ind = meta_ind.assign(WOS=wos)
  75. meta_ind_WOS = meta_ind.groupby('Study Subject ID')['WOS'].max().dropna()
  76. mild_ind = meta_ind_WOS[meta_ind_WOS <= 2].index.to_series()
  77. severe_ind = meta_ind_WOS[meta_ind_WOS >= 5].index.to_series()
  78. data.obs["cell_type"] = np.select(
  79. [(meta_cell["individual"].isin(mild_ind)),
  80. (meta_cell["individual"].isin(severe_ind))],
  81. ["mild", "severe"],
  82. default="moderate",
  83. )
  84. data.obs["individual"] = meta_cell["individual"].values
  85. print("[INFO] Label counts:")
  86. for k, v in data.obs["cell_type"].value_counts().items():
  87. print(f" {k:<9}: {v:>6}")
  88. return data, meta_cell, mild_ind, severe_ind
  89. # ------------------------- Split -------------------------
  90. def split_by_individual(data, meta_cell, mild_ind, severe_ind):
  91. print("[INFO] Splitting by individual...")
  92. test_ind = pd.concat([
  93. mild_ind.sample(n=math.floor(0.5 * len(mild_ind))),
  94. severe_ind.sample(n=math.floor(0.5 * len(severe_ind))),
  95. ])
  96. train_ind = pd.concat([mild_ind, severe_ind]).drop(test_ind.index)
  97. train_cell_id = meta_cell[meta_cell["individual"].isin(train_ind)]["V1"]
  98. test_cell_id = meta_cell[meta_cell["individual"].isin(test_ind)]["V1"]
  99. val_cell_id = train_cell_id.sample(n=math.floor(0.1 * len(train_cell_id)))
  100. train_cell_id = train_cell_id.drop(val_cell_id.index)
  101. train_data = data[train_cell_id].copy()
  102. val_data = data[val_cell_id].copy()
  103. test_data = data[test_cell_id].copy()
  104. return train_data, val_data, test_data
  105. # ------------------------- DE-gene selection -------------------------
  106. def get_de_gene_list(train_data, cell_type_list, tag, out_dir=None):
  107. if out_dir is None:
  108. out_dir = f"{d}output_comparison/"
  109. os.makedirs(out_dir, exist_ok=True)
  110. out_path = os.path.join(out_dir, f"gene_list_{tag}.npy")
  111. print("[INFO] Constructing DE gene list...")
  112. X = train_data.X.toarray() if sp.issparse(train_data.X) else np.asarray(train_data.X)
  113. X = np.asarray(X, dtype=np.float64)
  114. X[~np.isfinite(X)] = 0.0
  115. train_data = train_data.copy()
  116. train_data.X = X
  117. de_input = train_data.copy()
  118. if float(np.nanmax(X)) > 50:
  119. sc.pp.normalize_total(de_input, target_sum=1e4)
  120. sc.pp.log1p(de_input)
  121. gene_list = construct_gene_list(
  122. de_input,
  123. cell_type_list,
  124. n_top_genes=args.n_top_genes,
  125. method="fdr_bh",
  126. alpha=0.05,
  127. )
  128. gene_list = pd.Index(gene_list.astype(str))
  129. np.save(out_path, np.array(gene_list))
  130. print(f"[SAVED] {out_path}")
  131. return gene_list
  132. # ------------------------- SDAN -------------------------
  133. def run_sdan():
  134. print("[INFO] Running SDAN backend ...")
  135. data, meta_cell, mild_ind, severe_ind = load_and_label_data()
  136. qc(data)
  137. train_data, val_data, test_data = split_by_individual(data, meta_cell, mild_ind, severe_ind)
  138. out_dir = os.path.join(d, "output_comparison", "SDAN")
  139. os.makedirs(out_dir, exist_ok=True)
  140. # Keep SDAN feature selection inside `pipeline()` so the comparison run
  141. # uses the same DE/HVG logic as the standalone SDAN path.
  142. # gene_list = get_de_gene_list(train_data, ["mild", "severe"], args.cell_type, out_dir=out_dir)
  143. # print(f"[INFO] The number of DE genes: {len(gene_list)}")
  144. #
  145. # for ad in [train_data, val_data, test_data]:
  146. # ad._inplace_subset_var([g for g in gene_list if g in ad.var_names])
  147. # ad.obs["cell_type"] = ad.obs["cell_type"].astype("category")
  148. for ad in [train_data, val_data, test_data]:
  149. ad.obs["cell_type"] = ad.obs["cell_type"].astype("category")
  150. args.mc_weight = args.graph_weight
  151. args.o_weight = args.graph_weight
  152. # tag = f"{args.cell_type}_{args.graph_weight:.1f}"
  153. tag = f"{args.cell_type}_SDAN_{args.graph_weight}"
  154. print("[INFO] Launching SDAN pipeline...")
  155. for ad in [train_data, val_data, test_data]:
  156. if sp.issparse(ad.X):
  157. ad.X = ad.X.toarray().astype(np.float32)
  158. else:
  159. ad.X = ad.X.astype(np.float32)
  160. torch.set_default_dtype(torch.float32)
  161. (train_SDAN, val_SDAN, test_SDAN), _, cell_type_list, gene_list = pipeline(
  162. [train_data, val_data, test_data], args, d, tag)
  163. train_s = torch.tensor(np.load(f"{d}output/train_s_{tag}.npy"))
  164. np.save(os.path.join(out_dir, f"train_s_{tag}.npy"), train_s.detach().cpu().numpy())
  165. print("[INFO] Projecting cell embeddings using train_s...")
  166. Ztr = train_SDAN.x.t().detach().cpu().numpy() @ train_s.detach().cpu().numpy()
  167. Zte = test_SDAN.x.t().detach().cpu().numpy() @ train_s.detach().cpu().numpy()
  168. train_reduced_adata = AnnData(X=Ztr, obs=train_data.obs.copy())
  169. test_reduced_adata = AnnData(X=Zte, obs=test_data.obs.copy())
  170. train_reduced_adata.write(os.path.join(out_dir, f"train_reduced_{tag}.h5ad"))
  171. test_reduced_adata.write(os.path.join(out_dir, f"test_reduced_{tag}.h5ad"))
  172. # Define consistent label order
  173. cell_type_list = sorted(pd.unique(train_data.obs["cell_type"]))
  174. print(f"[INFO] Cell type list: {cell_type_list}")
  175. # Set which label is positive/negative
  176. set_eval_labels(pos_label="severe", neg_label="mild")
  177. _eval_and_save(
  178. AnnData(X=Ztr, obs=train_data.obs.copy()),
  179. AnnData(X=Zte, obs=test_data.obs.copy()),
  180. "SDAN",
  181. os.path.join(out_dir, f"results_classifiers_{tag}.csv"),
  182. os.path.join(out_dir, f"logreg_feature_importance_{tag}.csv"),
  183. sorted(pd.unique(train_data.obs["cell_type"])),
  184. mild_ind,
  185. severe_ind,
  186. )
  187. print(f"[INFO] train_s shape : {train_s.shape}")
  188. print(f"[INFO] train_reduced : {train_reduced_adata.X.shape}")
  189. print(f"[INFO] test_reduced : {test_reduced_adata.X.shape}")
  190. print(f"[SAVED] {os.path.join(out_dir, f'train_s_{tag}.npy')}")
  191. print(f"[SAVED] {os.path.join(out_dir, f'train_reduced_{tag}.h5ad')}")
  192. print(f"[SAVED] {os.path.join(out_dir, f'test_reduced_{tag}.h5ad')}")
  193. print("[INFO][SDAN] done.")
  194. # ------------------------- sciRED -------------------------
  195. def run_scired():
  196. print("[INFO] Running sciRED backend (raw counts, no normalization)...")
  197. from sklearn.decomposition import PCA
  198. from sklearn.preprocessing import StandardScaler
  199. from sklearn.pipeline import Pipeline as SkPipe
  200. from sciRED import glm as sc_glm
  201. from sciRED import rotations as rot
  202. data, meta_cell, mild_ind, severe_ind = load_and_label_data()
  203. train_data, _, test_data = split_by_individual(data, meta_cell, mild_ind, severe_ind)
  204. out_dir = os.path.join(d, "output_comparison", "sciRED")
  205. os.makedirs(out_dir, exist_ok=True)
  206. gene_list = get_de_gene_list(train_data, ["mild", "severe"], args.cell_type, out_dir=out_dir)
  207. train_data = train_data[:, [g for g in gene_list if g in train_data.var_names]].copy()
  208. test_data = test_data[:, [g for g in gene_list if g in test_data.var_names]].copy()
  209. def _dense64(adata):
  210. X = adata.X.toarray() if sp.issparse(adata.X) else np.asarray(adata.X)
  211. return X.astype(np.float64, copy=False)
  212. Xtr_counts = _dense64(train_data)
  213. Xte_counts = _dense64(test_data)
  214. genes = np.asarray(train_data.var_names)
  215. G = Xtr_counts.shape[1]
  216. if "protocol" in train_data.obs.columns:
  217. prot_tr = pd.get_dummies(train_data.obs["protocol"], drop_first=False).to_numpy(dtype=np.float64)
  218. prot_te = pd.get_dummies(test_data.obs["protocol"], drop_first=False).to_numpy(dtype=np.float64)
  219. else:
  220. prot_tr = np.empty((Xtr_counts.shape[0], 0), dtype=np.float64)
  221. prot_te = np.empty((Xte_counts.shape[0], 0), dtype=np.float64)
  222. lib_tr = Xtr_counts.sum(axis=1).reshape(-1, 1)
  223. lib_te = Xte_counts.sum(axis=1).reshape(-1, 1)
  224. Dtr = np.column_stack([np.ones((Xtr_counts.shape[0], 1)), lib_tr, prot_tr])
  225. Dte = np.column_stack([np.ones((Xte_counts.shape[0], 1)), lib_te, prot_te])
  226. print("[sciRED] GLM residuals (train)...")
  227. rtr = sc_glm.poissonGLM(y=Xtr_counts, x=Dtr)
  228. Ytr = rtr["resid_pearson"]
  229. print("[sciRED] GLM residuals (test)...")
  230. rte = sc_glm.poissonGLM(y=Xte_counts, x=Dte)
  231. Yte = rte["resid_pearson"]
  232. if Ytr.shape[1] != G:
  233. if Ytr.shape[0] == G:
  234. Ytr = Ytr.T
  235. else:
  236. raise ValueError(f"[sciRED] Unexpected Ytr shape {Ytr.shape}")
  237. if Yte.shape[1] != G:
  238. if Yte.shape[0] == G:
  239. Yte = Yte.T
  240. else:
  241. raise ValueError(f"[sciRED] Unexpected Yte shape {Yte.shape}")
  242. k = int(getattr(args, "n_comp", 40))
  243. pipe = SkPipe([
  244. ("scaler", StandardScaler(with_mean=True, with_std=True)),
  245. ("pca", PCA(n_components=min(k, G), random_state=888))
  246. ])
  247. Ztr_pca = pipe.fit_transform(Ytr)
  248. Zte_pca = pipe.transform(Yte)
  249. L_pca = pipe.named_steps["pca"].components_.T
  250. vr = rot.varimax(L_pca)
  251. L_varimax = vr["rotloading"]
  252. Ztr_full = rot.get_rotated_scores(Ztr_pca, vr["rotmat"])
  253. Zte_full = rot.get_rotated_scores(Zte_pca, vr["rotmat"])
  254. tag = f"{args.cell_type}_sciRED"
  255. np.save(os.path.join(out_dir, f"train_s_{tag}.npy"), L_varimax.astype(np.float64))
  256. pd.DataFrame(L_varimax, index=genes,
  257. columns=[f"sciRED_{i}" for i in range(L_varimax.shape[1])]) \
  258. .to_csv(os.path.join(out_dir, f"varimax_loading_{tag}.csv"))
  259. prog_names = [f"sciRED_{i}" for i in range(L_varimax.shape[1])]
  260. train_reduced_adata = AnnData(X=Ztr_full, obs=train_data.obs.copy(),
  261. var=pd.DataFrame(index=prog_names))
  262. test_reduced_adata = AnnData(X=Zte_full, obs=test_data.obs.copy(),
  263. var=pd.DataFrame(index=prog_names))
  264. train_reduced_adata.write(os.path.join(out_dir, f"train_reduced_{tag}.h5ad"))
  265. test_reduced_adata.write(os.path.join(out_dir, f"test_reduced_{tag}.h5ad"))
  266. # Define consistent label order
  267. cell_type_list = sorted(pd.unique(train_data.obs["cell_type"]))
  268. print(f"[INFO] Cell type list: {cell_type_list}")
  269. # Set which label is positive/negative
  270. set_eval_labels(pos_label="severe", neg_label="mild")
  271. _eval_and_save(
  272. train_reduced_adata,
  273. test_reduced_adata,
  274. "sciRED",
  275. os.path.join(out_dir, f"results_classifiers_{tag}.csv"),
  276. os.path.join(out_dir, f"logreg_feature_importance_{tag}.csv"),
  277. sorted(pd.unique(train_data.obs["cell_type"])),
  278. mild_ind,
  279. severe_ind,
  280. )
  281. print(f"[INFO] train_s shape : {L_varimax.shape}")
  282. print(f"[INFO] train_reduced : {train_reduced_adata.X.shape}")
  283. print(f"[INFO] test_reduced : {test_reduced_adata.X.shape}")
  284. print(f"[SAVED] {os.path.join(out_dir, f'train_s_{tag}.npy')}")
  285. print(f"[SAVED] {os.path.join(out_dir, f'train_reduced_{tag}.h5ad')}")
  286. print(f"[SAVED] {os.path.join(out_dir, f'test_reduced_{tag}.h5ad')}")
  287. print(f"[INFO][sciRED] done ({tag}).")
  288. # ------------------------- scNET -------------------------
  289. def run_scnet():
  290. import scNET
  291. from scNET.MultyGraphModel import scNET as scNET_model
  292. from scNET.Utils import save_obj
  293. from torch_geometric.data import Data
  294. from torch_geometric.utils import train_test_split_edges
  295. print("[INFO] scNET backend (custom graph)")
  296. # get args
  297. args = parse_args()
  298. # load data + split
  299. data, meta_cell, mild_ind, severe_ind = load_and_label_data()
  300. qc(data)
  301. train_data, val_data, test_data = split_by_individual(
  302. data, meta_cell, mild_ind, severe_ind
  303. )
  304. tag = f"{args.cell_type}_scNET75"
  305. out_dir = os.path.join(d, "output_comparison", "scNET")
  306. os.makedirs(out_dir, exist_ok=True)
  307. # DE genes (train only)
  308. gene_list = get_de_gene_list(
  309. train_data,
  310. ["mild", "severe"],
  311. args.cell_type,
  312. out_dir=out_dir,
  313. )
  314. print(f"[INFO] The number of DE genes: {len(gene_list)}")
  315. gene_list = pd.Index(gene_list.astype(str))
  316. for ad in [train_data, val_data, test_data]:
  317. ad._inplace_subset_var([g for g in gene_list if g in ad.var_names])
  318. ad.obs["cell_type"] = ad.obs["cell_type"].astype("category")
  319. print(f"[INFO] DE genes used: {len(gene_list)}")
  320. edge_index, gene_names = construct_gene_graph(
  321. gene_list.tolist()
  322. )
  323. # edge_index format
  324. if isinstance(edge_index, list):
  325. edge_index = np.array(edge_index)
  326. # Case 1: list of edges → shape (E, 2)
  327. if edge_index.ndim == 2 and edge_index.shape[1] == 2:
  328. edge_index = edge_index.T
  329. # Case 2: flattened → reshape
  330. elif edge_index.ndim == 1:
  331. if len(edge_index) % 2 != 0:
  332. raise ValueError(f"[ERROR] edge_index length not even: {len(edge_index)}")
  333. edge_index = edge_index.reshape(-1, 2).T
  334. # Final check
  335. assert edge_index.shape[0] == 2, f"[ERROR] edge_index wrong shape: {edge_index.shape}"
  336. print(f"[INFO] Fixed edge_index shape: {edge_index.shape}")
  337. # save original cell names BEFORE concatenate
  338. train_cells = train_data.obs_names.tolist()
  339. test_cells = test_data.obs_names.tolist()
  340. # Concatenate (transductive)
  341. train_data.obs["split"] = "train"
  342. test_data.obs["split"] = "test"
  343. obj = train_data.concatenate(test_data, index_unique=None)
  344. # Balanced subsampling
  345. train_sub = obj[obj.obs["split"] == "train"].copy()
  346. test_sub = obj[obj.obs["split"] == "test"].copy()
  347. sc.pp.subsample(train_sub, n_obs=7500, random_state=888)
  348. sc.pp.subsample(test_sub, n_obs=7500, random_state=888)
  349. obj = train_sub.concatenate(test_sub, index_unique=None)
  350. print(f"[INFO] After subsample: {obj.shape}")
  351. print(
  352. f"[INFO] scNET retained cells: train={train_sub.n_obs}, "
  353. f"test={test_sub.n_obs}, total={obj.n_obs}"
  354. )
  355. if sp.issparse(obj.X):
  356. obj.X = obj.X.toarray()
  357. obj.X = np.asarray(obj.X, dtype=np.float32)
  358. if obj.raw is None:
  359. obj.raw = obj.copy()
  360. print("Check data's scale:",
  361. "min =", obj.X.min(),
  362. "max =", obj.X.max(),
  363. "mean =", obj.X.mean())
  364. # scNET default graph
  365. sc.pp.neighbors(obj, n_neighbors=10, n_pcs=15)
  366. # align SDAN gene graph
  367. gene_to_idx = {g: i for i, g in enumerate(obj.var_names)}
  368. genes = list(gene_list)
  369. edges = []
  370. for i in range(edge_index.shape[1]):
  371. g1 = genes[edge_index[0, i]]
  372. g2 = genes[edge_index[1, i]]
  373. if g1 in gene_to_idx and g2 in gene_to_idx:
  374. edges.append([gene_to_idx[g1], gene_to_idx[g2]])
  375. if len(edges) == 0:
  376. raise ValueError("No valid edges after alignment")
  377. device = scNET.main.device
  378. ppi_edge_index = torch.tensor(edges, dtype=torch.long).T.to(device)
  379. print(f"[INFO] gene graph: {ppi_edge_index.shape}")
  380. # expression matrix
  381. node_feature = obj.X.T # genes x cells
  382. print(f"[INFO] expression: {node_feature.shape}")
  383. # KNN graph
  384. knn_edge_index, highly_variable_index = scNET.main.build_knn_graph(obj)
  385. print(f"[INFO] knn: {knn_edge_index.shape}")
  386. x = torch.tensor(node_feature, dtype=torch.float32)
  387. x = ((x.T - x.mean(dim=1)) / (x.std(dim=1) + 1e-5)).T
  388. data = Data(x=x, edge_index=ppi_edge_index.cpu())
  389. data = train_test_split_edges(data)
  390. data = data.to(device)
  391. x = x.to(device)
  392. ppi_edge_index = ppi_edge_index.to(device)
  393. batch_size = max(1, knn_edge_index.shape[1] // max(1, args.scnet_batches))
  394. loader = scNET.main.mini_batch_knn(knn_edge_index, batch_size)
  395. model_name = tag
  396. embedding_dim = args.n_comp
  397. model = scNET_model(
  398. x.shape[0], # number of genes
  399. x.shape[1], # number of cells
  400. 250, # hidden width for gene-side encoder
  401. embedding_dim, # final gene embedding dim
  402. 250, # hidden width for cell-side encoder
  403. embedding_dim, # final cell embedding dim
  404. lambda_rows=1,
  405. lambda_cols=1,
  406. num_layers=3,
  407. ).to(device)
  408. optimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-5)
  409. print("[INFO] Training scNET...")
  410. epoch_bar = tqdm(range(args.scnet_epochs), desc="scNET Training", total=args.scnet_epochs)
  411. for epoch in epoch_bar:
  412. model.train()
  413. for batch in loader:
  414. knn_edge_index_batch = batch.T.to(device)
  415. loss, _, _ = model.calculate_loss(
  416. x,
  417. knn_edge_index_batch,
  418. data.train_pos_edge_index,
  419. highly_variable_index,
  420. )
  421. optimizer.zero_grad()
  422. loss.backward()
  423. optimizer.step()
  424. epoch_bar.set_postfix(loss=f"{loss.item():.4f}")
  425. model_path = os.path.join(out_dir, f"models_{tag}.pt")
  426. torch.save({
  427. "model_state_dict": model.state_dict(),
  428. "gene_list": list(gene_list),
  429. "args": vars(args),
  430. }, model_path)
  431. print(f"[SAVED MODEL] {model_path}")
  432. # SAVE embeddings
  433. print("[INFO] Saving embeddings...")
  434. full_knn_edge_index = torch.cat([batch.T.to(device) for batch in loader], dim=1)
  435. model.eval()
  436. with torch.no_grad():
  437. row_embed, col_embed, out_features = model(
  438. x,
  439. full_knn_edge_index,
  440. data.train_pos_edge_index,
  441. )
  442. row_embed_np = row_embed.detach().cpu().numpy()
  443. col_embed_np = col_embed.detach().cpu().numpy()
  444. out_features_np = out_features.detach().cpu().numpy()
  445. print("[DEBUG] row_embed shape:", row_embed_np.shape)
  446. print("[DEBUG] embedded_cells shape:", col_embed_np.shape)
  447. node_features = pd.DataFrame(
  448. row_embed_np,
  449. index=obj.var_names,
  450. columns=[f"dim_{i}" for i in range(row_embed_np.shape[1])]
  451. )
  452. embedded_cells_df = pd.DataFrame(
  453. col_embed_np,
  454. index=obj.obs_names,
  455. columns=[f"dim_{i}" for i in range(col_embed_np.shape[1])]
  456. )
  457. import pkg_resources
  458. embed_dir = os.path.join(out_dir)
  459. os.makedirs(embed_dir, exist_ok=True)
  460. node_path = os.path.join(embed_dir, f"node_features_{model_name}.pkl")
  461. cell_path = os.path.join(embed_dir, f"embedded_cells_{model_name}.pkl")
  462. out_features_path = os.path.join(embed_dir, f"out_features_{model_name}.pkl")
  463. node_features.to_pickle(node_path)
  464. embedded_cells_df.to_pickle(cell_path)
  465. save_obj(out_features_np, os.path.join(embed_dir, f"out_features_{model_name}"))
  466. print(f"[SAVED] {node_path} shape={node_features.shape}")
  467. print(f"[SAVED] {cell_path} shape={embedded_cells_df.shape}")
  468. print(f"[SAVED] {out_features_path} shape={out_features_np.shape}")
  469. pkg_embed_dir = pkg_resources.resource_filename(scNET.__name__, "./Embedding/")
  470. os.makedirs(pkg_embed_dir, exist_ok=True)
  471. node_features.to_pickle(os.path.join(pkg_embed_dir, f"node_features_{model_name}"))
  472. save_obj(row_embed_np, os.path.join(pkg_embed_dir, f"row_embedding_{model_name}"))
  473. save_obj(col_embed_np, os.path.join(pkg_embed_dir, f"col_embedding_{model_name}"))
  474. save_obj(out_features_np, os.path.join(pkg_embed_dir, f"out_features_{model_name}"))
  475. _, embedded_cells, _, _ = scNET.load_embeddings(model_name)
  476. embedded_cells = embedded_cells.values if hasattr(embedded_cells, "values") else np.asarray(embedded_cells)
  477. # Split embedding
  478. cell_to_idx = {cell: i for i, cell in enumerate(obj.obs_names)}
  479. train_cells_sub = [c for c in train_cells if c in cell_to_idx]
  480. test_cells_sub = [c for c in test_cells if c in cell_to_idx]
  481. Ztr = embedded_cells[[cell_to_idx[c] for c in train_cells_sub]]
  482. Zte = embedded_cells[[cell_to_idx[c] for c in test_cells_sub]]
  483. print("Train cells after subsample:", len(Ztr))
  484. print("Test cells after subsample:", len(Zte))
  485. print(f"[INFO] train: {Ztr.shape}, test: {Zte.shape}")
  486. train_ad = AnnData(
  487. X=Ztr,
  488. obs=train_data.obs.loc[train_cells_sub].copy()
  489. )
  490. test_ad = AnnData(
  491. X=Zte,
  492. obs=test_data.obs.loc[test_cells_sub].copy()
  493. )
  494. train_ad.write(os.path.join(out_dir, f"train_reduced_{tag}.h5ad"))
  495. test_ad.write(os.path.join(out_dir, f"test_reduced_{tag}.h5ad"))
  496. # evaluation
  497. set_eval_labels(pos_label="severe", neg_label="mild")
  498. _eval_and_save(
  499. train_ad,
  500. test_ad,
  501. "scNET",
  502. os.path.join(out_dir, f"results_classifiers_{tag}.csv"),
  503. os.path.join(out_dir, f"logreg_feature_importance_{tag}.csv"),
  504. sorted(pd.unique(train_data.obs["cell_type"])),
  505. mild_ind,
  506. severe_ind,
  507. )
  508. print("[INFO][scNET] done.")
  509. # ------------------------- Dispatch -------------------------
  510. BACKENDS = {
  511. "SDAN": run_sdan,
  512. "sciRED": run_scired,
  513. "scNET": run_scnet,
  514. }
  515. if __name__ == "__main__":
  516. backend = args.backend
  517. if backend not in BACKENDS:
  518. raise ValueError(f"Unsupported backend: {backend}. Choose from {list(BACKENDS.keys())}")
  519. BACKENDS[backend]()

Su_2020_comparison.py at commit 958cb97, under MIT · at the source

Overview

Authors: Zhexiao Lin1, Yuanyuan Gao1, Wei Sun2,3,4
ORCID iDs: Wei Sun
  1. Department of Statistics, University of California, Berkeley, California, United States of America
  2. Public Health Sciences Division, Fred Hutchinson Cancer Center, Seattle, Washington, United States of America
  3. Department of Biostatistics, University of Washington, Seattle, Washington, United States of America
  4. Department of Biostatistics, University of North Carolina, Chapel Hill, North Carolina, United States of America
Institutions: University of California, Berkeley (United States); University of North Carolina at Chapel Hill (United States); University of Washington (United States); Fred Hutch Cancer Center (United States)
Journal: PLoS computational biology, volume 22, issue 6, article e1014327
Dates: received 27 January 2026; accepted 12 May 2026; published online 1 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pcbi.1014327 · PMID 42224351 · PMCID PMC13235940 · OpenAlex W7163028389
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: genetics / omics (modality), human (organism), other condition (population), Alzheimer's / dementia (population), cellular / molecular (subfield)
Methods: Statistics, Smoothing, state filtering, decompositions, Connectivity, Machine learning
MeSH: Deep Learning*, Molecular Sequence Annotation*, Supervised Machine Learning*, Classification Algorithms, Computational Biology, COVID-19, Gene Expression Profiling, Graph Neural Networks, Humans, Single-Cell Gene Expression Analysis (* major topic)
Journal subjects: Biology and Life Sciences, Genetics, Gene Expression, Medicine and Health Sciences, Medical Conditions, Infectious Diseases, Viral Diseases, Covid 19, Mental Health and Psychiatry, Dementia, Neurology, Cell biology, Cellular types, Animal cells, Blood cells, White blood cells, T cells, Cytotoxic T cells, Immune cells, Immunology, Glial Cells, Macroglial Cells, Astrocytes, Biochemistry, Bioenergetics, Energy-Producing Organelles, Mitochondria, Cellular Structures and Organelles, Alzheimer's Disease, Neurodegenerative Diseases, Immune Response
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: National Human Genome Research Institute (HG013177); NIGMS (GM105785)
Citations: not cited yet (Europe PMC); 52 references in the paper

Abstract

Gene-by-gene differential expression analysis is a widely used supervised approach for interpreting single-cell RNA-sequencing (scRNA-seq) data. However, modern scRNA-seq datasets often contain large numbers of cells, leading to the identification of many differentially expressed genes with extremely small p-values but negligible effect sizes, thus making biological interpretation difficult. To overcome this challenge, we developed Supervised Deep learning with gene functional ANnotation (SDAN), a method that integrates gene functional annotation information (e.g., protein-protein interaction) with gene-expression profiles through a graph neural network. SDAN identifies functionally coherent gene sets that optimally classify cells, and the resulting cell-level classification scores can be aggregated to make individual-level predictions. We evaluated SDAN alongside three representative existing methods in three real-data applications aimed at identifying gene sets associated with severe COVID-19, dementia, and cancer immunotherapy response. Across all applications, SDAN consistently outperformed the alternative approaches by achieving two objectives simultaneously: accurate outcome classification and clear assignment of genes to functionally related gene sets.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 16 matches between paragraphs and lines of code.

Sun-lab/SDAN

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 958cb977d67b84d45a38801da4a1ea7fbf1111ea, 27 April 2026
Languages: Python (36), R (34), Jupyter (6), Shell (1)
Size: 1,394 files, 77 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, license file, environment (requirements.txt), 30 notebooks
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: pandas (36 files), NumPy (34 files), Scanpy (29 files), PyTorch (28 files), data.table (26 files), tidyverse (25 files), ggplot2 (22 files), ggpubr (22 files), scikit-learn (19 files), SciPy (16 files), anndata (13 files), limma (12 files), Matplotlib (12 files), PyTorch Geometric (12 files), seaborn (9 files), reshape2 (5 files), reticulate (5 files), pROC (4 files), DESeq2 (3 files), statsmodels (3 files), cowplot (2 files), h5py (2 files), SingleCellExperiment (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
79 files

The paper's code and data availability statement is in the Data section.

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:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 77 scripts, each with its path and the digest of its content;
  • 16 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

“Su et al. COVID-19 dataset: Gene expression data were downloaded from ArrayExpress: https://www.ebi.ac.uk/biostudies/arrayexpress/studies/E-MTAB-9357 on 2/22/2022. Patient information was extracted from Table S1 of Su et al.: https://data.mendeley.com/datasets/tzydswhhb5/5 on 2/22/2022.” SEA-AD dataset: The snRNA-seq data were downloaded from cellxgene website https://cellxgene.cziscience.com/collections/1ca90a2d-2943-483d-b678-b809bf464c30 on 2021/11/21. Donor meta data and clinical data were downloaded from the SEA-AD website https://portal.brain-map.org/explore/seattle-alzheimers-disease/seattle-alzheimers-disease-brain-cell-atlas-download on 2022/1/7. Cancer immunotherapy dataset: The gene expression data of Sade-Feldman et al. 2018 were downloaded from https://www.ncbi.nlm.nih.gov/geo/query/acc.cgi?acc=GSE120575. on 2023/11/6. Cell and sample information was extracted from Supplementary Table 1 of Sade-Feldman et al. 2018. The gene expression and meta-data of Yost et al. 2019 were downloaded from https://www.ncbi.nlm.nih.gov/geo/query/acc.cgi?acc=GSE123813. Codes for data processing and running SDAN for all the datasets are available at https://github.com/Sun-lab/SDAN.

Reproduced under the paper's license (CC BY), 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, 10 MeSH terms, 2 funders, 47 references.

Cite

This paper

Lin, Z., Gao, Y., & Sun, W. (2026). Supervised deep learning with gene functional annotation for cell classification. PLoS computational biology, 22(6), e1014327. https://doi.org/10.1371/journal.pcbi.1014327

BibTeX

@article{lin2026supervised,
author = {Lin, Zhexiao and Gao, Yuanyuan and Sun, Wei},
title = {{Supervised deep learning with gene functional annotation for cell classification}},
journal = {PLoS computational biology},
year = {2026},
month = jun,
volume = {22},
number = {6},
pages = {e1014327},
publisher = {PLOS},
issn = {1553-734X},
doi = {10.1371/journal.pcbi.1014327},
url = {https://doi.org/10.1371/journal.pcbi.1014327},
pmid = {42224351},
pmcid = {PMC13235940}
}

RIS

TY - JOUR
AU - Lin, Zhexiao
AU - Gao, Yuanyuan
AU - Sun, Wei
TI - Supervised deep learning with gene functional annotation for cell classification
T2 - PLoS computational biology
J2 - PLoS Comput Biol
PY - 2026
DA - 2026/06/01
VL - 22
IS - 6
SP - e1014327
SN - 1553-734X
PB - PLOS
DO - 10.1371/journal.pcbi.1014327
UR - https://doi.org/10.1371/journal.pcbi.1014327
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pcbi.1014327",
"type": "article-journal",
"title": "Supervised deep learning with gene functional annotation for cell classification",
"container-title": "PLoS computational biology",
"author": [
{
"family": "Lin",
"given": "Zhexiao"
},
{
"family": "Gao",
"given": "Yuanyuan"
},
{
"family": "Sun",
"given": "Wei"
}
],
"container-title-short": "PLoS Comput Biol",
"volume": "22",
"issue": "6",
"page": "e1014327",
"DOI": "10.1371/journal.pcbi.1014327",
"PMID": "42224351",
"PMCID": "PMC13235940",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pcbi.1014327",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
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.1038/s42003-026-10034-0 [code]
Region- and cell type-specific changes in gene expression in the cerebellum after classical fear conditioning.
Journal: Communications biology
In common: SingleCellExperiment, reticulate, limma, 16 other tools, genetics / omics, cellular / molecular, 1 reference
[2] 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: pROC, SingleCellExperiment, limma, 16 other tools, genetics / omics, other condition, cellular / molecular
[3] doi:10.1038/s44320-026-00208-7 [code]
Interpretable deep generative ensemble learning for single-cell omics with Hydra.
Journal: Molecular systems biology
In common: SingleCellExperiment, reticulate, limma, 14 other tools, cellular / molecular, 2 references
[4] doi:10.1038/s44318-026-00818-9 [code]
FAM134B-mediated ER-phagy degrades APP and suppresses Alzheimer's disease pathology.
Journal: The EMBO journal
In common: SingleCellExperiment, reticulate, limma, 15 other tools, Alzheimer's / dementia, cellular / molecular
[5] doi:10.1016/j.cell.2026.05.026 [code]
The critical role of the endogenous immune compartment after CAR T cell therapy in recurrent GBM.
Journal: Cell
In common: SingleCellExperiment, limma, anndata, 13 other tools, genetics / omics, other condition, 2 references
[6] doi:10.1038/s41592-026-03194-8 [code]
Beyond benchmarking: an expert-guided consensus approach to spatially aware clustering.
Journal: Nature methods
In common: PyTorch Geometric, SingleCellExperiment, limma, 15 other tools, genetics / omics
[7] 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: PyTorch Geometric, SingleCellExperiment, reticulate, 15 other tools
[8] doi:10.1016/j.xcrm.2026.102651 [code]
Integrative CSF profiling identifies disease-specific immune responses in leptomeningeal disease.
Journal: Cell reports. Medicine
In common: SingleCellExperiment, reticulate, anndata, 14 other tools, genetics / omics, other condition, cellular / molecular
[9] doi:10.1038/s41467-026-71803-3 [code]
Charting the transition from in vitro gliogenesis to the in vivo maturation of human glial progenitor cells transplanted into the hypomyelinated mouse brain.
Journal: Nature communications
In common: reticulate, anndata, DESeq2, 14 other tools, genetics / omics, cellular / molecular, 1 reference
[10] doi:10.1038/s41586-026-10214-2 [code]
Multidimensional profiling of heterogeneity in supratentorial ependymomas.
Journal: Nature
In common: SingleCellExperiment, reticulate, anndata, 14 other tools, genetics / omics, other condition

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.