Supervised deep learning with gene functional annotation for cell classification.
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] § 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] § 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] § 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] § 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] § 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] § 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] § Methods › Implementation of SDAN ↔ SDAN/train.py, lines 8–116 · score 0.62 · weight decay, Adam, PyTorch, epochs, optimized, hidden
- [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] § Methods › Implementation of SDAN ↔ SDAN_Comparison/Su_2020_comparison.py, lines 439–520 · score 0.61 · weight decay, Adam, PyTorch, epochs, optimized, hidden
- [10] § Methods › Implementation of SDAN ↔ check/check_de.py, lines 1–37 · score 0.57 · FDR threshold, DE genes, sensitivity, selection, class, SDAN
- [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] § Methods › Extensions ↔ SDAN/train.py, lines 8–116 · score 0.57 · cross entropy loss, graph loss, softmax, scores, classifier, SDAN
- [13] § Methods › Extensions ↔ check/check_graph.ipynb, lines 32–166 · score 0.55 · cross entropy loss, graph loss, error, classifier, SDAN, cell
- [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] § 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] § 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
- """
- This file unified version for SDAN, sciRED, and scNET backends.
- Compatible with args.py and preprocess.py.
- ALL outputs saved into: ./Su_2020/output_comparison/
- The backend output folders are:
- SDAN: output_comparison/SDAN
- sciRED: output_comparison/sciRED
- scNET: output_comparison/scNET
- """
- # ------------------------- Imports & setup -------------------------
- import os as _os
- _os.environ.setdefault("OMP_NUM_THREADS", "1")
- _os.environ.setdefault("OPENBLAS_NUM_THREADS", "1")
- _os.environ.setdefault("MKL_NUM_THREADS", "1")
- _os.environ.setdefault("NUMEXPR_NUM_THREADS", "1")
- _os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
- _os.environ.setdefault("NUMBA_THREADING_LAYER", "workqueue")
- import os, math, warnings
- import numpy as np
- import pandas as pd
- import scanpy as sc
- import scipy.sparse as sp
- import torch
- from tqdm import tqdm
- from anndata import AnnData
- from SDAN.model import pipeline
- from SDAN.preprocess import qc, construct_gene_graph, construct_gene_list
- from SDAN.args import parse_args
- from SDAN.evaluation import set_eval_labels, eval_and_save as _eval_and_save
- try:
- import Spectra
- except Exception:
- Spectra = None
- warnings.simplefilter(action="ignore", category=FutureWarning)
- np.random.seed(888)
- torch.manual_seed(888)
- # sc.settings.verbosity = 3
- sc.settings.verbosity = 1
- try:
- sc.logging.print_header()
- except Exception as e:
- print(f"[WARN] Skipped Scanpy header due to: {type(e).__name__}: {e}")
- sc.settings.set_figure_params(figsize=(8, 6), dpi=80, facecolor="white")
- # ------------------------- Paths & args -------------------------
- args = parse_args()
- d = "./Su_2020/"
- cell_type_str = args.cell_type
- data_dir = f"{d}gex_{cell_type_str}.mtx.gz"
- genes_dir = f"{d}gex_{cell_type_str}_genes.txt"
- meta_ind_dir = f"{d}Table_S1.xlsx"
- meta_cell_dir = f"{d}cell_info_{cell_type_str}.csv"
- os.makedirs(f"{d}output_comparison/", exist_ok=True)
- # ------------------------- Load and label -------------------------
- def load_and_label_data():
- """Load Su_2020 data and assign mild/severe labels."""
- print("[INFO] Loading raw data...")
- data = sc.read(data_dir, cache=True)
- gene_names = pd.read_csv(genes_dir, header=None).iloc[:, 0].astype(str).to_numpy()
- if len(gene_names) != data.n_vars:
- raise ValueError(f"[ERROR] genes length {len(gene_names)} != n_vars {data.n_vars}")
- meta_cell = pd.read_csv(meta_cell_dir)
- meta_ind = pd.read_excel(meta_ind_dir, sheet_name="S1.1 Patient Clinical Data")
- data.var["gene_symbols"] = gene_names
- data.var_names = pd.Index(gene_names)
- data.obs["barcode"] = meta_cell["V1"].astype(str).to_numpy()
- data.obs_names = pd.Index(data.obs["barcode"])
- gene_mito = pd.read_csv("./Annotation/mito_genes.tsv", sep="\t")
- mito_col = "hgnc_symbol" if "hgnc_symbol" in gene_mito.columns else gene_mito.columns[0]
- data = data[:, ~data.var_names.isin(gene_mito[mito_col])]
- data_nonzero_prop = (data.X != 0).sum(axis=0) / data.shape[0]
- data = data[:, data_nonzero_prop > 0.02]
- wos = meta_ind['Who Ordinal Scale'].astype(str).str.replace('1 or 2', '2', regex=False)
- wos = pd.to_numeric(wos, errors='coerce')
- meta_ind = meta_ind.assign(WOS=wos)
- meta_ind_WOS = meta_ind.groupby('Study Subject ID')['WOS'].max().dropna()
- mild_ind = meta_ind_WOS[meta_ind_WOS <= 2].index.to_series()
- severe_ind = meta_ind_WOS[meta_ind_WOS >= 5].index.to_series()
- data.obs["cell_type"] = np.select(
- [(meta_cell["individual"].isin(mild_ind)),
- (meta_cell["individual"].isin(severe_ind))],
- ["mild", "severe"],
- default="moderate",
- )
- data.obs["individual"] = meta_cell["individual"].values
- print("[INFO] Label counts:")
- for k, v in data.obs["cell_type"].value_counts().items():
- print(f" {k:<9}: {v:>6}")
- return data, meta_cell, mild_ind, severe_ind
- # ------------------------- Split -------------------------
- def split_by_individual(data, meta_cell, mild_ind, severe_ind):
- print("[INFO] Splitting by individual...")
- test_ind = pd.concat([
- mild_ind.sample(n=math.floor(0.5 * len(mild_ind))),
- severe_ind.sample(n=math.floor(0.5 * len(severe_ind))),
- ])
- train_ind = pd.concat([mild_ind, severe_ind]).drop(test_ind.index)
- train_cell_id = meta_cell[meta_cell["individual"].isin(train_ind)]["V1"]
- test_cell_id = meta_cell[meta_cell["individual"].isin(test_ind)]["V1"]
- val_cell_id = train_cell_id.sample(n=math.floor(0.1 * len(train_cell_id)))
- train_cell_id = train_cell_id.drop(val_cell_id.index)
- train_data = data[train_cell_id].copy()
- val_data = data[val_cell_id].copy()
- test_data = data[test_cell_id].copy()
- return train_data, val_data, test_data
- # ------------------------- DE-gene selection -------------------------
- def get_de_gene_list(train_data, cell_type_list, tag, out_dir=None):
- if out_dir is None:
- out_dir = f"{d}output_comparison/"
- os.makedirs(out_dir, exist_ok=True)
- out_path = os.path.join(out_dir, f"gene_list_{tag}.npy")
- print("[INFO] Constructing DE gene list...")
- X = train_data.X.toarray() if sp.issparse(train_data.X) else np.asarray(train_data.X)
- X = np.asarray(X, dtype=np.float64)
- X[~np.isfinite(X)] = 0.0
- train_data = train_data.copy()
- train_data.X = X
- de_input = train_data.copy()
- if float(np.nanmax(X)) > 50:
- sc.pp.normalize_total(de_input, target_sum=1e4)
- sc.pp.log1p(de_input)
- gene_list = construct_gene_list(
- de_input,
- cell_type_list,
- n_top_genes=args.n_top_genes,
- method="fdr_bh",
- alpha=0.05,
- )
- gene_list = pd.Index(gene_list.astype(str))
- np.save(out_path, np.array(gene_list))
- print(f"[SAVED] {out_path}")
- return gene_list
- # ------------------------- SDAN -------------------------
- def run_sdan():
- print("[INFO] Running SDAN backend ...")
- data, meta_cell, mild_ind, severe_ind = load_and_label_data()
- qc(data)
- train_data, val_data, test_data = split_by_individual(data, meta_cell, mild_ind, severe_ind)
- out_dir = os.path.join(d, "output_comparison", "SDAN")
- os.makedirs(out_dir, exist_ok=True)
- # Keep SDAN feature selection inside `pipeline()` so the comparison run
- # uses the same DE/HVG logic as the standalone SDAN path.
- # gene_list = get_de_gene_list(train_data, ["mild", "severe"], args.cell_type, out_dir=out_dir)
- # print(f"[INFO] The number of DE genes: {len(gene_list)}")
- #
- # for ad in [train_data, val_data, test_data]:
- # ad._inplace_subset_var([g for g in gene_list if g in ad.var_names])
- # ad.obs["cell_type"] = ad.obs["cell_type"].astype("category")
- for ad in [train_data, val_data, test_data]:
- ad.obs["cell_type"] = ad.obs["cell_type"].astype("category")
- args.mc_weight = args.graph_weight
- args.o_weight = args.graph_weight
- # tag = f"{args.cell_type}_{args.graph_weight:.1f}"
- tag = f"{args.cell_type}_SDAN_{args.graph_weight}"
- print("[INFO] Launching SDAN pipeline...")
- for ad in [train_data, val_data, test_data]:
- if sp.issparse(ad.X):
- ad.X = ad.X.toarray().astype(np.float32)
- else:
- ad.X = ad.X.astype(np.float32)
- torch.set_default_dtype(torch.float32)
- (train_SDAN, val_SDAN, test_SDAN), _, cell_type_list, gene_list = pipeline(
- [train_data, val_data, test_data], args, d, tag)
- train_s = torch.tensor(np.load(f"{d}output/train_s_{tag}.npy"))
- np.save(os.path.join(out_dir, f"train_s_{tag}.npy"), train_s.detach().cpu().numpy())
- print("[INFO] Projecting cell embeddings using train_s...")
- Ztr = train_SDAN.x.t().detach().cpu().numpy() @ train_s.detach().cpu().numpy()
- Zte = test_SDAN.x.t().detach().cpu().numpy() @ train_s.detach().cpu().numpy()
- train_reduced_adata = AnnData(X=Ztr, obs=train_data.obs.copy())
- test_reduced_adata = AnnData(X=Zte, obs=test_data.obs.copy())
- train_reduced_adata.write(os.path.join(out_dir, f"train_reduced_{tag}.h5ad"))
- test_reduced_adata.write(os.path.join(out_dir, f"test_reduced_{tag}.h5ad"))
- # Define consistent label order
- cell_type_list = sorted(pd.unique(train_data.obs["cell_type"]))
- print(f"[INFO] Cell type list: {cell_type_list}")
- # Set which label is positive/negative
- set_eval_labels(pos_label="severe", neg_label="mild")
- _eval_and_save(
- AnnData(X=Ztr, obs=train_data.obs.copy()),
- AnnData(X=Zte, obs=test_data.obs.copy()),
- "SDAN",
- os.path.join(out_dir, f"results_classifiers_{tag}.csv"),
- os.path.join(out_dir, f"logreg_feature_importance_{tag}.csv"),
- sorted(pd.unique(train_data.obs["cell_type"])),
- mild_ind,
- severe_ind,
- )
- print(f"[INFO] train_s shape : {train_s.shape}")
- print(f"[INFO] train_reduced : {train_reduced_adata.X.shape}")
- print(f"[INFO] test_reduced : {test_reduced_adata.X.shape}")
- print(f"[SAVED] {os.path.join(out_dir, f'train_s_{tag}.npy')}")
- print(f"[SAVED] {os.path.join(out_dir, f'train_reduced_{tag}.h5ad')}")
- print(f"[SAVED] {os.path.join(out_dir, f'test_reduced_{tag}.h5ad')}")
- print("[INFO][SDAN] done.")
- # ------------------------- sciRED -------------------------
- def run_scired():
- print("[INFO] Running sciRED backend (raw counts, no normalization)...")
- from sklearn.decomposition import PCA
- from sklearn.preprocessing import StandardScaler
- from sklearn.pipeline import Pipeline as SkPipe
- from sciRED import glm as sc_glm
- from sciRED import rotations as rot
- data, meta_cell, mild_ind, severe_ind = load_and_label_data()
- train_data, _, test_data = split_by_individual(data, meta_cell, mild_ind, severe_ind)
- out_dir = os.path.join(d, "output_comparison", "sciRED")
- os.makedirs(out_dir, exist_ok=True)
- gene_list = get_de_gene_list(train_data, ["mild", "severe"], args.cell_type, out_dir=out_dir)
- train_data = train_data[:, [g for g in gene_list if g in train_data.var_names]].copy()
- test_data = test_data[:, [g for g in gene_list if g in test_data.var_names]].copy()
- def _dense64(adata):
- X = adata.X.toarray() if sp.issparse(adata.X) else np.asarray(adata.X)
- return X.astype(np.float64, copy=False)
- Xtr_counts = _dense64(train_data)
- Xte_counts = _dense64(test_data)
- genes = np.asarray(train_data.var_names)
- G = Xtr_counts.shape[1]
- if "protocol" in train_data.obs.columns:
- prot_tr = pd.get_dummies(train_data.obs["protocol"], drop_first=False).to_numpy(dtype=np.float64)
- prot_te = pd.get_dummies(test_data.obs["protocol"], drop_first=False).to_numpy(dtype=np.float64)
- else:
- prot_tr = np.empty((Xtr_counts.shape[0], 0), dtype=np.float64)
- prot_te = np.empty((Xte_counts.shape[0], 0), dtype=np.float64)
- lib_tr = Xtr_counts.sum(axis=1).reshape(-1, 1)
- lib_te = Xte_counts.sum(axis=1).reshape(-1, 1)
- Dtr = np.column_stack([np.ones((Xtr_counts.shape[0], 1)), lib_tr, prot_tr])
- Dte = np.column_stack([np.ones((Xte_counts.shape[0], 1)), lib_te, prot_te])
- print("[sciRED] GLM residuals (train)...")
- rtr = sc_glm.poissonGLM(y=Xtr_counts, x=Dtr)
- Ytr = rtr["resid_pearson"]
- print("[sciRED] GLM residuals (test)...")
- rte = sc_glm.poissonGLM(y=Xte_counts, x=Dte)
- Yte = rte["resid_pearson"]
- if Ytr.shape[1] != G:
- if Ytr.shape[0] == G:
- Ytr = Ytr.T
- else:
- raise ValueError(f"[sciRED] Unexpected Ytr shape {Ytr.shape}")
- if Yte.shape[1] != G:
- if Yte.shape[0] == G:
- Yte = Yte.T
- else:
- raise ValueError(f"[sciRED] Unexpected Yte shape {Yte.shape}")
- k = int(getattr(args, "n_comp", 40))
- pipe = SkPipe([
- ("scaler", StandardScaler(with_mean=True, with_std=True)),
- ("pca", PCA(n_components=min(k, G), random_state=888))
- ])
- Ztr_pca = pipe.fit_transform(Ytr)
- Zte_pca = pipe.transform(Yte)
- L_pca = pipe.named_steps["pca"].components_.T
- vr = rot.varimax(L_pca)
- L_varimax = vr["rotloading"]
- Ztr_full = rot.get_rotated_scores(Ztr_pca, vr["rotmat"])
- Zte_full = rot.get_rotated_scores(Zte_pca, vr["rotmat"])
- tag = f"{args.cell_type}_sciRED"
- np.save(os.path.join(out_dir, f"train_s_{tag}.npy"), L_varimax.astype(np.float64))
- pd.DataFrame(L_varimax, index=genes,
- columns=[f"sciRED_{i}" for i in range(L_varimax.shape[1])]) \
- .to_csv(os.path.join(out_dir, f"varimax_loading_{tag}.csv"))
- prog_names = [f"sciRED_{i}" for i in range(L_varimax.shape[1])]
- train_reduced_adata = AnnData(X=Ztr_full, obs=train_data.obs.copy(),
- var=pd.DataFrame(index=prog_names))
- test_reduced_adata = AnnData(X=Zte_full, obs=test_data.obs.copy(),
- var=pd.DataFrame(index=prog_names))
- train_reduced_adata.write(os.path.join(out_dir, f"train_reduced_{tag}.h5ad"))
- test_reduced_adata.write(os.path.join(out_dir, f"test_reduced_{tag}.h5ad"))
- # Define consistent label order
- cell_type_list = sorted(pd.unique(train_data.obs["cell_type"]))
- print(f"[INFO] Cell type list: {cell_type_list}")
- # Set which label is positive/negative
- set_eval_labels(pos_label="severe", neg_label="mild")
- _eval_and_save(
- train_reduced_adata,
- test_reduced_adata,
- "sciRED",
- os.path.join(out_dir, f"results_classifiers_{tag}.csv"),
- os.path.join(out_dir, f"logreg_feature_importance_{tag}.csv"),
- sorted(pd.unique(train_data.obs["cell_type"])),
- mild_ind,
- severe_ind,
- )
- print(f"[INFO] train_s shape : {L_varimax.shape}")
- print(f"[INFO] train_reduced : {train_reduced_adata.X.shape}")
- print(f"[INFO] test_reduced : {test_reduced_adata.X.shape}")
- print(f"[SAVED] {os.path.join(out_dir, f'train_s_{tag}.npy')}")
- print(f"[SAVED] {os.path.join(out_dir, f'train_reduced_{tag}.h5ad')}")
- print(f"[SAVED] {os.path.join(out_dir, f'test_reduced_{tag}.h5ad')}")
- print(f"[INFO][sciRED] done ({tag}).")
- # ------------------------- scNET -------------------------
- def run_scnet():
- import scNET
- from scNET.MultyGraphModel import scNET as scNET_model
- from scNET.Utils import save_obj
- from torch_geometric.data import Data
- from torch_geometric.utils import train_test_split_edges
- print("[INFO] scNET backend (custom graph)")
- # get args
- args = parse_args()
- # load data + split
- data, meta_cell, mild_ind, severe_ind = load_and_label_data()
- qc(data)
- train_data, val_data, test_data = split_by_individual(
- data, meta_cell, mild_ind, severe_ind
- )
- tag = f"{args.cell_type}_scNET75"
- out_dir = os.path.join(d, "output_comparison", "scNET")
- os.makedirs(out_dir, exist_ok=True)
- # DE genes (train only)
- gene_list = get_de_gene_list(
- train_data,
- ["mild", "severe"],
- args.cell_type,
- out_dir=out_dir,
- )
- print(f"[INFO] The number of DE genes: {len(gene_list)}")
- gene_list = pd.Index(gene_list.astype(str))
- for ad in [train_data, val_data, test_data]:
- ad._inplace_subset_var([g for g in gene_list if g in ad.var_names])
- ad.obs["cell_type"] = ad.obs["cell_type"].astype("category")
- print(f"[INFO] DE genes used: {len(gene_list)}")
- edge_index, gene_names = construct_gene_graph(
- gene_list.tolist()
- )
- # edge_index format
- if isinstance(edge_index, list):
- edge_index = np.array(edge_index)
- # Case 1: list of edges → shape (E, 2)
- if edge_index.ndim == 2 and edge_index.shape[1] == 2:
- edge_index = edge_index.T
- # Case 2: flattened → reshape
- elif edge_index.ndim == 1:
- if len(edge_index) % 2 != 0:
- raise ValueError(f"[ERROR] edge_index length not even: {len(edge_index)}")
- edge_index = edge_index.reshape(-1, 2).T
- # Final check
- assert edge_index.shape[0] == 2, f"[ERROR] edge_index wrong shape: {edge_index.shape}"
- print(f"[INFO] Fixed edge_index shape: {edge_index.shape}")
- # save original cell names BEFORE concatenate
- train_cells = train_data.obs_names.tolist()
- test_cells = test_data.obs_names.tolist()
- # Concatenate (transductive)
- train_data.obs["split"] = "train"
- test_data.obs["split"] = "test"
- obj = train_data.concatenate(test_data, index_unique=None)
- # Balanced subsampling
- train_sub = obj[obj.obs["split"] == "train"].copy()
- test_sub = obj[obj.obs["split"] == "test"].copy()
- sc.pp.subsample(train_sub, n_obs=7500, random_state=888)
- sc.pp.subsample(test_sub, n_obs=7500, random_state=888)
- obj = train_sub.concatenate(test_sub, index_unique=None)
- print(f"[INFO] After subsample: {obj.shape}")
- print(
- f"[INFO] scNET retained cells: train={train_sub.n_obs}, "
- f"test={test_sub.n_obs}, total={obj.n_obs}"
- )
- if sp.issparse(obj.X):
- obj.X = obj.X.toarray()
- obj.X = np.asarray(obj.X, dtype=np.float32)
- if obj.raw is None:
- obj.raw = obj.copy()
- print("Check data's scale:",
- "min =", obj.X.min(),
- "max =", obj.X.max(),
- "mean =", obj.X.mean())
- # scNET default graph
- sc.pp.neighbors(obj, n_neighbors=10, n_pcs=15)
- # align SDAN gene graph
- gene_to_idx = {g: i for i, g in enumerate(obj.var_names)}
- genes = list(gene_list)
- edges = []
- for i in range(edge_index.shape[1]):
- g1 = genes[edge_index[0, i]]
- g2 = genes[edge_index[1, i]]
- if g1 in gene_to_idx and g2 in gene_to_idx:
- edges.append([gene_to_idx[g1], gene_to_idx[g2]])
- if len(edges) == 0:
- raise ValueError("No valid edges after alignment")
- device = scNET.main.device
- ppi_edge_index = torch.tensor(edges, dtype=torch.long).T.to(device)
- print(f"[INFO] gene graph: {ppi_edge_index.shape}")
- # expression matrix
- node_feature = obj.X.T # genes x cells
- print(f"[INFO] expression: {node_feature.shape}")
- # KNN graph
- knn_edge_index, highly_variable_index = scNET.main.build_knn_graph(obj)
- print(f"[INFO] knn: {knn_edge_index.shape}")
- x = torch.tensor(node_feature, dtype=torch.float32)
- x = ((x.T - x.mean(dim=1)) / (x.std(dim=1) + 1e-5)).T
- data = Data(x=x, edge_index=ppi_edge_index.cpu())
- data = train_test_split_edges(data)
- data = data.to(device)
- x = x.to(device)
- ppi_edge_index = ppi_edge_index.to(device)
- batch_size = max(1, knn_edge_index.shape[1] // max(1, args.scnet_batches))
- loader = scNET.main.mini_batch_knn(knn_edge_index, batch_size)
- model_name = tag
- embedding_dim = args.n_comp
- model = scNET_model(
- x.shape[0], # number of genes
- x.shape[1], # number of cells
- 250, # hidden width for gene-side encoder
- embedding_dim, # final gene embedding dim
- 250, # hidden width for cell-side encoder
- embedding_dim, # final cell embedding dim
- lambda_rows=1,
- lambda_cols=1,
- num_layers=3,
- ).to(device)
- optimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-5)
- print("[INFO] Training scNET...")
- epoch_bar = tqdm(range(args.scnet_epochs), desc="scNET Training", total=args.scnet_epochs)
- for epoch in epoch_bar:
- model.train()
- for batch in loader:
- knn_edge_index_batch = batch.T.to(device)
- loss, _, _ = model.calculate_loss(
- x,
- knn_edge_index_batch,
- data.train_pos_edge_index,
- highly_variable_index,
- )
- optimizer.zero_grad()
- loss.backward()
- optimizer.step()
- epoch_bar.set_postfix(loss=f"{loss.item():.4f}")
- model_path = os.path.join(out_dir, f"models_{tag}.pt")
- torch.save({
- "model_state_dict": model.state_dict(),
- "gene_list": list(gene_list),
- "args": vars(args),
- }, model_path)
- print(f"[SAVED MODEL] {model_path}")
- # SAVE embeddings
- print("[INFO] Saving embeddings...")
- full_knn_edge_index = torch.cat([batch.T.to(device) for batch in loader], dim=1)
- model.eval()
- with torch.no_grad():
- row_embed, col_embed, out_features = model(
- x,
- full_knn_edge_index,
- data.train_pos_edge_index,
- )
- row_embed_np = row_embed.detach().cpu().numpy()
- col_embed_np = col_embed.detach().cpu().numpy()
- out_features_np = out_features.detach().cpu().numpy()
- print("[DEBUG] row_embed shape:", row_embed_np.shape)
- print("[DEBUG] embedded_cells shape:", col_embed_np.shape)
- node_features = pd.DataFrame(
- row_embed_np,
- index=obj.var_names,
- columns=[f"dim_{i}" for i in range(row_embed_np.shape[1])]
- )
- embedded_cells_df = pd.DataFrame(
- col_embed_np,
- index=obj.obs_names,
- columns=[f"dim_{i}" for i in range(col_embed_np.shape[1])]
- )
- import pkg_resources
- embed_dir = os.path.join(out_dir)
- os.makedirs(embed_dir, exist_ok=True)
- node_path = os.path.join(embed_dir, f"node_features_{model_name}.pkl")
- cell_path = os.path.join(embed_dir, f"embedded_cells_{model_name}.pkl")
- out_features_path = os.path.join(embed_dir, f"out_features_{model_name}.pkl")
- node_features.to_pickle(node_path)
- embedded_cells_df.to_pickle(cell_path)
- save_obj(out_features_np, os.path.join(embed_dir, f"out_features_{model_name}"))
- print(f"[SAVED] {node_path} shape={node_features.shape}")
- print(f"[SAVED] {cell_path} shape={embedded_cells_df.shape}")
- print(f"[SAVED] {out_features_path} shape={out_features_np.shape}")
- pkg_embed_dir = pkg_resources.resource_filename(scNET.__name__, "./Embedding/")
- os.makedirs(pkg_embed_dir, exist_ok=True)
- node_features.to_pickle(os.path.join(pkg_embed_dir, f"node_features_{model_name}"))
- save_obj(row_embed_np, os.path.join(pkg_embed_dir, f"row_embedding_{model_name}"))
- save_obj(col_embed_np, os.path.join(pkg_embed_dir, f"col_embedding_{model_name}"))
- save_obj(out_features_np, os.path.join(pkg_embed_dir, f"out_features_{model_name}"))
- _, embedded_cells, _, _ = scNET.load_embeddings(model_name)
- embedded_cells = embedded_cells.values if hasattr(embedded_cells, "values") else np.asarray(embedded_cells)
- # Split embedding
- cell_to_idx = {cell: i for i, cell in enumerate(obj.obs_names)}
- train_cells_sub = [c for c in train_cells if c in cell_to_idx]
- test_cells_sub = [c for c in test_cells if c in cell_to_idx]
- Ztr = embedded_cells[[cell_to_idx[c] for c in train_cells_sub]]
- Zte = embedded_cells[[cell_to_idx[c] for c in test_cells_sub]]
- print("Train cells after subsample:", len(Ztr))
- print("Test cells after subsample:", len(Zte))
- print(f"[INFO] train: {Ztr.shape}, test: {Zte.shape}")
- train_ad = AnnData(
- X=Ztr,
- obs=train_data.obs.loc[train_cells_sub].copy()
- )
- test_ad = AnnData(
- X=Zte,
- obs=test_data.obs.loc[test_cells_sub].copy()
- )
- train_ad.write(os.path.join(out_dir, f"train_reduced_{tag}.h5ad"))
- test_ad.write(os.path.join(out_dir, f"test_reduced_{tag}.h5ad"))
- # evaluation
- set_eval_labels(pos_label="severe", neg_label="mild")
- _eval_and_save(
- train_ad,
- test_ad,
- "scNET",
- os.path.join(out_dir, f"results_classifiers_{tag}.csv"),
- os.path.join(out_dir, f"logreg_feature_importance_{tag}.csv"),
- sorted(pd.unique(train_data.obs["cell_type"])),
- mild_ind,
- severe_ind,
- )
- print("[INFO][scNET] done.")
- # ------------------------- Dispatch -------------------------
- BACKENDS = {
- "SDAN": run_sdan,
- "sciRED": run_scired,
- "scNET": run_scnet,
- }
- if __name__ == "__main__":
- backend = args.backend
- if backend not in BACKENDS:
- raise ValueError(f"Unsupported backend: {backend}. Choose from {list(BACKENDS.keys())}")
- BACKENDS[backend]()
Su_2020_comparison.py at commit 958cb97, under MIT · at the source
Overview
- Department of Statistics, University of California, Berkeley, California, United States of America
- Public Health Sciences Division, Fred Hutchinson Cancer Center, Seattle, Washington, United States of America
- Department of Biostatistics, University of Washington, Seattle, Washington, United States of America
- Department of Biostatistics, University of North Carolina, Chapel Hill, North Carolina, United States of America
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
958cb977d67b84d45a38801da4a1ea7fbf1111ea, 27 April 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
79 files
- Annotation/
step1_check_PPI_data.Rmd , R, 64 lines - SDAN/
args.py , Python, 33 lines - SDAN/
layers.py , Python, 41 lines - SDAN/
model.py , Python, 65 lines - SDAN/
preprocess.py , Python, 107 lines - SDAN/
train.py , Python, 129 lines, 2 matches - SDAN/
utils.py , Python, 150 lines - SDAN_Comparison/
SEA_AD_comparison.py , Python, 606 lines, 2 matches - SDAN_Comparison/
SEA_AD_sdan_evaluation.p , Python, 170 linesy - SDAN_Comparison/
SEA_AD_spectra_evaluatio , Python, 213 linesn.py - SDAN_Comparison/
SEA_AD_spectra_preproces , Python, 121 liness.py - SDAN_Comparison/
SEA_AD_spectra_training. , Python, 167 lines, 1 matchpy - SDAN_Comparison/
Su_2020_comparison.py , Python, 631 lines, 3 matches - SDAN_Comparison/
Su_2020_sdan_evaluation. , Python, 176 linespy - SDAN_Comparison/
Su_2020_spectra_evaluati , Python, 216 lines, 1 matchon.py - SDAN_Comparison/
Su_2020_spectra_preproce , Python, 130 linesss.py - SDAN_Comparison/
Su_2020_spectra_training , Python, 158 lines.py - SDAN_Comparison/
Yost_2019_comparison.py , Python, 1,054 lines - SDAN_Comparison/
Yost_2019_sdan_evaluatio , Python, 175 linesn.py - SDAN_Comparison/
args.py , Python, 119 lines - SDAN_Comparison/
combine_Astro_Micro-PVM. , Python, 102 linespy - SDAN_Comparison/
combine_cd4_cd8.py , Python, 125 lines - SDAN_Comparison/
evaluation.py , Python, 172 lines - SDAN_Comparison/
plots.ipynb , Jupyter, 1,194 lines - SEA_AD.py, Python, 125 lines
- SEA_AD/
step1_pseudo_bulk.py , Python, 56 lines - SEA_AD/
step2_check_donors.Rmd , R, 122 lines - SEA_AD/
step3_pseudo_bulk_DE_ast , R, 546 linesro.Rmd - SEA_AD/
step3_pseudo_bulk_DE_mic , R, 572 linesroglia.Rmd - SEA_AD/
step4_evaluate_gene_sets , R, 247 lines.Rmd - SEA_AD/
step4_evaluate_gene_sets , R, 27 lines_goseq_knit.R - SEA_AD/
step4b_summerize_goseq.R , R, 120 linesmd - SEA_AD/
step5_evaluate_gene_sets , R, 174 lines_cell_types.Rmd - SEA_AD/
step5_evaluate_gene_sets , R, 27 lines_cell_types_knit.R - SEA_AD/
step6_check_results.Rmd , R, 233 lines - SEA_AD/
step6_check_results.py , Python, 9 lines - SEA_AD/
step6_check_results_knit , R, 27 lines.R - SEA_AD/
step7_check_patient_subs , R, 330 lines, 1 matchet.Rmd - SEA_AD_plot.py, Python, 89 lines, 1 match
- SF_2018/
step1_prepare_data.Rmd , R, 277 lines - SF_2018/
step2_pseudo_bulk_DE.Rmd , R, 308 lines - SF_2018/
step3_evaluate_gene_sets , R, 206 lines.Rmd - SF_2018/
step4_evaluate_gene_sets , R, 174 lines_cell_types.Rmd - SF_2018/
step5_check_results.Rmd , R, 219 lines - Su_2020.py, Python, 132 lines
- Su_2020/
ArrayExpress/ , R, 80 lines_check_sdrf.R - Su_2020/
_collect_data.R , R, 219 lines - Su_2020/
_collect_data.sh , Shell, 5 lines - Su_2020/
_get_mito_genes.Rmd , R, 54 lines - Su_2020/
s1_evaluate_gene_sets_go , R, 233 linesseq.Rmd - Su_2020/
s1_evaluate_gene_sets_go , R, 28 linesseq_knit.R - Su_2020/
s1_summerize_goseq.Rmd , R, 124 lines - Su_2020/
s1b_evaluate_gene_sets_c , R, 187 linesell_types.Rmd - Su_2020/
s1b_evaluate_gene_sets_c , R, 28 linesell_types_knit.R - Su_2020/
s2b_check_results.Rmd , R, 250 lines - Su_2020/
s2b_check_results_knit.R , R, 27 lines - Su_2020_plot.py, Python, 91 lines
- Yost_2019.py, Python, 132 lines
- Yost_2019/
step1_prepare_data.Rmd , R, 230 lines - Yost_2019/
step2_evaluate_gene_sets , R, 229 lines_goseq.Rmd - Yost_2019/
step2_evaluate_gene_sets , R, 26 lines_goseq_knit.R - Yost_2019/
step2b_summerize_goseq.R , R, 122 linesmd - Yost_2019/
step2c_zoom_in_goseq_res , R, 125 linesults.Rmd - Yost_2019/
step3_check_results.Rmd , R, 235 lines - Yost_2019/
step3_check_results_knit , R, 21 lines.R - Yost_2019_plot.py, Python, 61 lines
- check/
check_de.ipynb , Jupyter, 585 lines - check/
check_de.py , Python, 265 lines, 1 match - check/
check_genes.ipynb , Jupyter, 65 lines - check/
check_graph.ipynb , Jupyter, 384 lines, 2 matches - check/
check_graph.py , Python, 334 lines - check/
check_weight.py , Python, 215 lines - check/
compute_loading_comparis , Python, 545 lineson.py - check/
compute_sdan_program_auc , Python, 107 lines.py - check/
de_enrichment.py , Python, 474 lines - check/
loading_comparison.ipynb , Jupyter, 236 lines - tutorial.ipynb, Jupyter, 351 lines, 2 matches
- LICENSE, License, 21 lines
- README.md, Text, 191 lines
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
- arrayexpress:E-MTAB-9357
, at ArrayExpress; found in “Data Availability” - data.mendeley.com/
datasets/ , at Mendeley Data; found in “Data Availability”tzydswhhb5 - geo:GSE120575, at NCBI GEO; found in “Data Availability”
- portal.brain-map.org/
explore/ , at Allen Brain Map; found in “Data Availability”seattle-alzheimers-disea se
Data Availability
“Su et al. COVID-19 dataset: Gene expression data were downloaded from ArrayExpress: https://
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://
BibTeX
@article{lin2026supervis
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/
url = {https://
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/
VL - 22
IS - 6
SP - e1014327
SN - 1553-734X
PB - PLOS
DO - 10.1371/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1371/
"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":
"volume": "22",
"issue": "6",
"page": "e1014327",
"DOI": "10.1371/
"PMID": "42224351",
"PMCID": "PMC13235940",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://
"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 biologyIn 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. MedicineIn 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 biologyIn 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 journalIn 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: CellIn 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 methodsIn 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 blueIn 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. MedicineIn 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 communicationsIn 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: NatureIn 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.
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: 1 repository of the authors' code, each at its verified commit and with its license, 77 scripts, and 16 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:90fa98dfea85a065…
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.
