OSCR

A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic Technologies

Code ↔ Paper

20 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 20 matches · 2 of them tie a paragraph to a whole file, not to given lines: weak matches, whose lines are not tinted
  1. [1] § Results › Overview of spatial cell-type deconvolution benchmarking ↔ Experiments/_Composite_score_generation/Save_Composite_Score_Brain.ipynb, lines 18–34 · score 0.95 · CelloScope, SpatialDecon, STDeconvolve, Cell DART, SpiceMix, DestVI
  2. [2] § Results › Overview of spatial cell-type deconvolution benchmarking ↔ Experiments/_Deconvolution_Metrics_Calculation/Evaluation_Brain_Category.ipynb, lines 29–43 · score 0.95 · CelloScope, SpatialDecon, STDeconvolve, Cell DART, SpiceMix, DestVI
  3. [3] § Results › Benchmarking of spatial domain detection methods ↔ Experiments/_Domain_Detection_Metrics_Calculation/all_datasets_metric_computation.ipynb, lines 185–194 · score 0.93 · DR SC, DeepST, SpaSRL, SpatialPCA, BayesCafe, BayesSpace
  4. [4] § Results › Benchmarking of spatial domain detection methods ↔ Experiments/_Domain_Detection_Metrics_Calculation/Metrics.py, lines 351–421 · score 0.93 · DR SC, DeepST, SpaSRL, SpatialPCA, BayesCafe, BayesSpace
  5. [5] § Results › Benchmarking of spatial domain detection methods ↔ Experiments/_Domain_Detection_Metrics_Calculation/all_datasets_metric_computation.ipynb, lines 289–313 · score 0.91 · mouse breast cancer, osmFISH, chicken heart, prostate cancer, liver cancer, kidney cancer
  6. [6] § Results › Benchmarking of spatial domain detection methods ↔ Experiments/_Domain_Detection_Metrics_Calculation/Metrics.py, lines 944–1089 · score 0.86 · SpaSRL, SpatialPCA, BayesCafe, BayesSpace, SpaceFlow, GraphST
  7. [7] § Methods › Overview of SynthST ↔ SynthST/Simulator_CTP/SynthST/model.py, lines 56–99 · score 0.86 · cross entropy, weight decay, adjacency matrix, proportion matrix, L2, encoder
  8. [8] § Results › Benchmarking of spatial domain detection methods ↔ Experiments/_Domain_Detection_Metrics_Calculation/Metrics.py, lines 351–421 · score 0.84 · DeepST, BayesSpace, SpaceFlow, prostate cancer, GraphST, PRECAST
  9. [9] § Results › Overview of spatial cell-type deconvolution benchmarking ↔ Experiments/_Domain_Detection_Metrics_Calculation/Metrics.py, lines 107–248 · score 0.78 · chicken heart, prostate cancer, liver cancer, kidney cancer, breast cancer, ground truth
  10. [10] § Methods › Evaluation metrics for domain detection ↔ Experiments/_Domain_Detection_Metrics_Calculation/Metrics.py, lines 423–497 · score 0.78 · Normalized Mutual, Adjusted Rand, Completeness Score, Homogeneity, ground truth, ARI
  11. [11] § Results › Overview of spatial cell-type deconvolution benchmarking ↔ Experiments/_Composite_score_generation/plotting_files.py, lines 1–39 · score 0.72 · spatial Pearson, Spatial Metrics, composite score, Rare Cell, Cosine, Lee
  12. [12] § Methods › Overview of SynthST ↔ SynthST/Simulator_CTP/SynthST/SynthST.py, the whole file · a weak match · score 0.68 · weight decay, SynthST, GAT, optimization, inferred, loss
  13. [13] § Results › Cell2location, RCTD and SONAR achieve the highest accuracy in cell-type deconvolution ↔ Experiments/_Composite_score_generation/Save_Composite_Score_Brain.ipynb, lines 18–34 · score 0.67 · SpatialDWLS, DestVI, Cell2location, Redeconve, Composite Score, SSIM
  14. [14] § Results › Cell2location, RCTD and SONAR achieve the highest accuracy in cell-type deconvolution ↔ Experiments/_Deconvolution_Metrics_Calculation/Evaluation_Brain_Category.ipynb, lines 29–43 · score 0.66 · SpatialDWLS, DestVI, Cell2location, Redeconve, SSIM, Geary
  15. [15] § Extended Data ↔ SynthST/Simulator_CTP/SynthST/model.py, lines 56–99 · score 0.64 · cross entropy, adjacency matrix, proportion matrix, reconstructs, decoder, loss
  16. [16] § Methods › Overview of SynthST ↔ SynthST/Simulator_gene_expression/Simulate_gene_expression_notebook.ipynb, lines 58–86 · score 0.62 · SynthST, gene expression, Cell2location, signature, abundance, simulated
  17. [17] § Methods › Composite score metrics for cell-type deconvolution ↔ Experiments/_Composite_score_generation/plotting_files.py, lines 58–89 · score 0.58 · Spatial Score, Composite Rare, Composite Score, SSIM, Correlation, Moran
  18. [18] § Methods › Composite score metrics for cell-type deconvolution ↔ Experiments/_Deconvolution_Metrics_Calculation/metrics.py, lines 261–279 · score 0.55 · Recall curve, Precision, threshold, ground truth, AUPR, metric
  19. [19] § Results › Overview of spatial cell-type deconvolution benchmarking ↔ SynthST/Simulator_gene_expression/Simulate_gene_expression_notebook.ipynb, lines 58–86 · score 0.52 · simulated gene expression, SynthST, signature, single cell, matrix, spots
  20. [20] § Results › Overview of spatial cell-type deconvolution benchmarking ↔ benchmarking_domain_detection/CCST/run_CCST.py, the whole file · a weak match · score 0.52 · single cell resolution, deep graph, gene expression, MERFISH, training, ground truth

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 · 1,193 lines · 47 KB · MIT · 5 matches

  1. # ----------------------- Imports -----------------------
  2. import os
  3. import re
  4. import numpy as np
  5. import pandas as pd
  6. import scanpy as sc
  7. import squidpy as sq
  8. import anndata as an
  9. import seaborn as sns
  10. import matplotlib.pyplot as plt
  11. from PIL import Image
  12. from matplotlib import font_manager as fm
  13. from scipy.spatial import distance_matrix
  14. from scipy.spatial.distance import squareform, pdist
  15. from sklearn.metrics import (
  16. adjusted_rand_score,
  17. normalized_mutual_info_score,
  18. silhouette_score,
  19. homogeneity_score,
  20. completeness_score
  21. )
  22. from sklearn.preprocessing import StandardScaler
  23. # ----------------------- Utility -----------------------
  24. def fx_1NN(i,location_in):
  25. location_in = np.array(location_in)
  26. dist_array = distance_matrix(location_in[i,:][None,:],location_in)[0,:]
  27. dist_array[i] = np.inf
  28. return np.min(dist_array)
  29. def fx_kNN(i,location_in,k,cluster_in):
  30. location_in = np.array(location_in)
  31. cluster_in = np.array(cluster_in)
  32. dist_array = distance_matrix(location_in[i,:][None,:],location_in)[0,:]
  33. dist_array[i] = np.inf
  34. ind = np.argsort(dist_array)[:k]
  35. cluster_use = np.array(cluster_in)
  36. if np.sum(cluster_use[ind]!=cluster_in[i])>(k/2):
  37. return 0
  38. else:
  39. return 1
  40. def _compute_CHAOS(clusterlabel, location):
  41. clusterlabel = np.array(clusterlabel)
  42. location = np.array(location)
  43. matched_location = StandardScaler().fit_transform(location)
  44. clusterlabel_unique = np.unique(clusterlabel)
  45. dist_val = np.zeros(len(clusterlabel_unique))
  46. count = 0
  47. for k in clusterlabel_unique:
  48. location_cluster = matched_location[clusterlabel==k,:]
  49. if len(location_cluster)<=2:
  50. continue
  51. n_location_cluster = len(location_cluster)
  52. results = [fx_1NN(i,location_cluster) for i in range(n_location_cluster)]
  53. dist_val[count] = np.sum(results)
  54. count = count + 1
  55. chaos = np.sum(dist_val)/len(clusterlabel)
  56. return np.exp(-0.5*chaos)
  57. def _compute_PAS(clusterlabel,location):
  58. clusterlabel = np.array(clusterlabel)
  59. location = np.array(location)
  60. matched_location = location
  61. results = [fx_kNN(i,matched_location,k=10,cluster_in=clusterlabel) for i in range(matched_location.shape[0])]
  62. return np.sum(results)/len(clusterlabel)
  63. def compute_ASW(pred,spatial_coords):
  64. distance_matrix = squareform(pdist(spatial_coords))
  65. sil = silhouette_score(X=distance_matrix, labels=pred, metric='precomputed')
  66. return (sil + 1)/2
  67. def LISI(coords, meta, label, perplexity=30, nn_eps=0):
  68. import rpy2.robjects as robjects
  69. from rpy2.robjects import pandas2ri
  70. pandas2ri.activate()
  71. from rpy2.robjects.packages import importr
  72. importr("lisi")
  73. if not isinstance(coords, pd.DataFrame):
  74. coords = pd.DataFrame(coords)
  75. if not isinstance(meta, pd.DataFrame):
  76. meta = pd.DataFrame(meta)
  77. meta = meta.loc[:, [label]]
  78. meta[label] = meta[label].astype(str)
  79. coords = robjects.conversion.py2rpy(coords)
  80. meta = robjects.conversion.py2rpy(meta)
  81. as_matrix = robjects.r["as.matrix"]
  82. lisi = robjects.r["compute_lisi"](as_matrix(coords), meta, label, perplexity, nn_eps)
  83. if isinstance(lisi, pd.DataFrame):
  84. lisi = lisi.values
  85. elif isinstance(lisi, np.recarray):
  86. lisi = [item[0] for item in lisi]
  87. return lisi
  88. # ----------------------- Import functions -----------------------
  89. def import_dataset(dataset_name, paths, mode = "evaluation"):
  90. """
  91. Load dataset given its name and a dictionary of file paths.
  92. Parameters
  93. ----------
  94. dataset_name : str
  95. Name of the dataset (used for branching logic).
  96. paths : dict
  97. Dictionary of file paths needed for that dataset.
  98. Must contain at least 'st_path'.
  99. Can contain 'gnd_path' if ground truth is from a CSV/TSV file.
  100. Returns
  101. -------
  102. gnd : pd.Series or pd.DataFrame
  103. Ground truth clusters.
  104. locs : pd.DataFrame
  105. Spatial coordinates.
  106. """
  107. # Load main AnnData object
  108. ST = an.read_h5ad(paths["st_path"])
  109. if dataset_name.startswith("DLPFC"):
  110. gnd = pd.read_csv(paths["gnd_path"], sep="\t")['layer_guess_reordered']
  111. gnd.index.name = "Index"
  112. locs = pd.DataFrame({
  113. "array_row": ST.obs.array_row,
  114. "array_col": ST.obs.array_col
  115. })
  116. elif dataset_name.startswith("embryo"):
  117. gnd = pd.DataFrame(ST.obs['annotation']).rename(columns={"annotation": "Cluster"})
  118. gnd.index.name = "Index"
  119. locs = pd.DataFrame({
  120. "array_row": ST.obsm["spatial"][:, 0],
  121. "array_col": ST.obsm["spatial"][:, 1]
  122. }, index=ST.obs_names)
  123. elif dataset_name == "mouse_breast_cancer":
  124. gnd = pd.DataFrame(ST.obs['ground_truth']).rename(columns={"annotation": "Cluster"})
  125. gnd.index.name = "Index"
  126. locs = pd.DataFrame({
  127. "array_row": ST.obs.array_row,
  128. "array_col": ST.obs.array_col
  129. })
  130. elif dataset_name.startswith("MERFISH_brain"):
  131. gnd = pd.DataFrame(ST.obs['ground_truth']).rename(columns={"annotation": "Cluster"})
  132. gnd.index.name = "Index"
  133. locs = pd.DataFrame({
  134. "array_row": ST.obsm["spatial"][:, 0],
  135. "array_col": ST.obsm["spatial"][:, 1]
  136. }, index=ST.obs_names)
  137. elif dataset_name == "simulated_breast_atlas":
  138. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names)['Ground Truth']
  139. gnd.index.name = "Index"
  140. locs = pd.DataFrame({
  141. "array_row": ST.obsm["spatial"][:, 0],
  142. "array_col": ST.obsm["spatial"][:, 1]
  143. }, index=ST.obs_names)
  144. elif dataset_name == "osmFISH":
  145. gnd = pd.DataFrame(ST.obs['ground_truth']).rename(columns={"annotation": "Cluster"})
  146. gnd.index.name = "Index"
  147. locs = pd.DataFrame({
  148. "array_row": ST.obsm["spatial"][:, 0],
  149. "array_col": ST.obsm["spatial"][:, 1]
  150. }, index=ST.obs_names)
  151. elif dataset_name.startswith("simulated_kidney_cancer"):
  152. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names)['Ground Truth']
  153. gnd.index.name = "Index"
  154. locs = pd.DataFrame({
  155. "array_row": ST.obs.new_x,
  156. "array_col": ST.obs.new_y
  157. })
  158. elif dataset_name.startswith("simulated_breast_cancer"):
  159. ST.obs_names_make_unique()
  160. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names)['Ground Truth']
  161. gnd.index.name = "Index"
  162. locs = pd.DataFrame({
  163. "array_row": ST.obs.new_x,
  164. "array_col": ST.obs.new_y
  165. })
  166. elif dataset_name.startswith("simulated_liver_cancer"):
  167. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names)['Ground Truth']
  168. gnd.index.name = "Index"
  169. locs = pd.DataFrame({
  170. "array_row": ST.obs.new_x,
  171. "array_col": ST.obs.new_y
  172. })
  173. elif dataset_name.startswith("simulated_intestine"):
  174. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names)['Ground Truth']
  175. gnd.index.name = "Index"
  176. locs = pd.DataFrame({
  177. "array_row": ST.obs.new_x,
  178. "array_col": ST.obs.new_y
  179. })
  180. elif dataset_name == "simulated_chicken_heart":
  181. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names)['Ground Truth']
  182. gnd.index.name = "Index"
  183. locs = pd.DataFrame({
  184. "array_row": ST.obs.array_row,
  185. "array_col": ST.obs.array_col
  186. })
  187. elif dataset_name == "simulated_prostate_cancer":
  188. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names)['Ground Truth']
  189. gnd.index.name = "Index"
  190. locs = pd.DataFrame({
  191. "array_row": ST.obs.x_new,
  192. "array_col": ST.obs.y_new
  193. })
  194. elif dataset_name == "simulated_cerebellum":
  195. gnd = pd.read_csv(paths["gnd_path"]).set_index(ST.obs_names.astype(int))['Ground Truth']
  196. gnd.index.name = "Index"
  197. locs = pd.DataFrame({
  198. "array_row": ST.obs.xcoord,
  199. "array_col": ST.obs.ycoord
  200. }).set_index(ST.obs_names.astype(int))
  201. else:
  202. raise ValueError(f"Incorrect input: {dataset_name}")
  203. if mode == "plot":
  204. locs = pd.DataFrame({
  205. "array_row": ST.obsm["spatial"][:, 0],
  206. "array_col": ST.obsm["spatial"][:, 1]
  207. }).set_index(ST.obs_names)
  208. if dataset_name == "simulated_cerebellum":
  209. locs.index = locs.index.astype(int)
  210. locs.index.name = "Index"
  211. return gnd, locs
  212. def load_prediction(dataset_name, method_name, path):
  213. """
  214. Load output file given its name and file path.
  215. Parameters
  216. ----------
  217. dataset_name : str
  218. Name of the dataset (used for branching logic).
  219. method_name : str
  220. Name of the method (used for branching logic).
  221. path : str
  222. Path to output file
  223. Returns
  224. -------
  225. pred : predicted clusters
  226. """
  227. if os.path.exists(path):
  228. if dataset_name.startswith("simulated_breast_cancer") and method_name in ["BASS","PRECAST", "BayesSpace" , "DR_SC"]:
  229. pred = pd.read_csv(path)
  230. pred = pred.set_index(pred.columns[0])
  231. if method_name == "PRECAST":
  232. pred = pred['cluster']
  233. pred.index = pred.index.str.replace(r'_\d+$', '', regex=True)
  234. elif(method_name == "banksy"):
  235. pred = pd.read_csv(path)
  236. pred = pred.set_index(pred.columns[0])
  237. columns_to_grab = [col for col in pred.columns if col.startswith("clust_M1_lam0.")]
  238. pred = pred[columns_to_grab]
  239. elif method_name == "giotto":
  240. pred = pd.read_csv(path)
  241. if "leiden_clus" in pred.columns:
  242. columns_to_grab = ["leiden_clus"]
  243. pred = pred.set_index(pred.columns[1])
  244. else:
  245. columns_to_grab = ["cluster"]
  246. pred = pred.set_index(pred.columns[0])
  247. pred = pred[columns_to_grab]
  248. elif method_name == "PRECAST" or method_name == "BayesCafe":
  249. pred = pd.read_csv(path)
  250. pred = pred.set_index(pred.columns[0])
  251. columns_to_grab = ["cluster"]
  252. pred = pred[columns_to_grab]
  253. elif method_name == "SpaceFlow":
  254. pred = pd.read_csv(path)
  255. pred = pred.set_index(pred.columns[1])
  256. columns_to_grab = ["Predicted_cell_label"]
  257. pred = pred[columns_to_grab]
  258. else:
  259. pred = pd.read_csv(path)
  260. pred = pred.set_index(pred.columns[0])
  261. else:
  262. print(f"path doesn't exist : {path}")
  263. print(f"{method_name} output missing for {dataset_name}")
  264. pred = None
  265. return pred
  266. # ----------------------- Metric Computationn -----------------------
  267. def compute_metrics(dataset_names, dataset_paths, method_names, pred_paths, error_log_file, output_dir, metric_names="all"):
  268. """
  269. Compute metric values given dataset and method names and paths
  270. Parameters
  271. ----------
  272. dataset_names : list
  273. List of datasets to compute on
  274. Options for dataset names present at bottom of this file
  275. dataset_paths: dictionary of dictionaries
  276. Dictionaries mapping dataset names to dictionaries with st_path and gnd_path needed for importing dataset
  277. method_names: list
  278. List of methods to compute on
  279. Options for method names present at bottom of this file
  280. pred_paths : dictionary of dictionaries
  281. Dictionary mapping dataset names and method names to path of output file
  282. metric_names : list or str
  283. List of metrics to compute or "all"
  284. error_log_file : str
  285. Path to error log file
  286. output_dir : str
  287. Path to save final metric compute files to
  288. Returns
  289. -------
  290. pred : predicted clusters
  291. """
  292. # Available metrics
  293. available_metrics = ["ARI", "NMI", "CHAOS", "PAS", "ASW", "HOM", "COM"]
  294. if metric_names == "all":
  295. metric_names = available_metrics
  296. # Initialize pivoted DataFrame for each metric
  297. method_names_dict = {
  298. "SCANIT": "SCANIT",
  299. "CCST": "CCST",
  300. "DeepST": "DeepST",
  301. "GraphST": "GraphST",
  302. "PROST": "PROST",
  303. "SpaSRL": "SpaSRL",
  304. "STAGATE": "STAGATE",
  305. "SpatialPCA": "SpatialPCA",
  306. "banksy": "Banksy",
  307. "giotto": "Giotto",
  308. "DR_SC": "DR.SC",
  309. "ISC_MEB": "ISC.MEB",
  310. "BayesSpace": "BayesSpace",
  311. "PRECAST": "PRECAST",
  312. "BayesCafe": "BayesCafe",
  313. "BASS": "BASS",
  314. "SpaceFlow" : "SpaceFlow",
  315. "IRIS" : "IRIS"
  316. }
  317. dataset_names_dict = {
  318. "DLPFC151507": "DLPFC 151507",
  319. "DLPFC151508": "DLPFC 151508",
  320. "DLPFC151509": "DLPFC 151509",
  321. "DLPFC151510": "DLPFC 151510",
  322. "DLPFC151669": "DLPFC 151669",
  323. "DLPFC151670": "DLPFC 151670",
  324. "DLPFC151671": "DLPFC 151671",
  325. "DLPFC151672": "DLPFC 151672",
  326. "DLPFC151673": "DLPFC 151673",
  327. "DLPFC151674": "DLPFC 151674",
  328. "DLPFC151675": "DLPFC 151675",
  329. "DLPFC151676": "DLPFC 151676",
  330. "embryo9.5": "Embryo 9.5",
  331. "embryo14.5": "Embryo 14.5",
  332. "mouse_breast_cancer": "Mouse Breast Cancer",
  333. "MERFISH_brain0.04": "MERFISH Brain 0.04",
  334. "MERFISH_brain0.09": "MERFISH Brain 0.09",
  335. "MERFISH_brain0.14": "MERFISH Brain 0.14",
  336. "MERFISH_brain0.19": "MERFISH Brain 0.19",
  337. "MERFISH_brain0.24": "MERFISH Brain 0.24",
  338. "osmFISH": "osmFISH",
  339. "simulated_kidney_cancer410": "Kidney Cancer 410",
  340. "simulated_kidney_cancer411": "Kidney Cancer 411",
  341. "simulated_kidney_cancer506": "Kidney Cancer 506",
  342. "simulated_breast_cancerER+_CID4290": "Breast Cancer ER+ CID 4290",
  343. "simulated_breast_cancerTNBC_CID44971": "Breast Cancer TNBC CID 44971",
  344. "simulated_liver_cancerHCC-1L": "Liver Cancer HCC-1L",
  345. "simulated_liver_cancerHCC-2L": "Liver Cancer HCC-2L",
  346. "simulated_liver_cancerHCC-3L": "Liver Cancer HCC-3L",
  347. "simulated_liver_cancerHCC-4L": "Liver Cancer HCC-4L",
  348. "simulated_breast_atlas": "Breast Atlas",
  349. "simulated_intestineA1": "Intestine A1",
  350. "simulated_intestineA2": "Intestine A2",
  351. "simulated_chicken_heart": "Chicken Heart",
  352. "simulated_prostate_cancer" : "Prostate Cancer",
  353. "simulated_cerebellum" : "Cerebellum"
  354. }
  355. metric_results = {
  356. metric: pd.DataFrame(
  357. index=list(method_names_dict.values()), columns=list(dataset_names_dict.values()), dtype=float
  358. ) for metric in metric_names
  359. }
  360. os.makedirs(output_dir, exist_ok=True)
  361. print(f"saving to {output_dir}")
  362. # Process each dataset
  363. for dataset_name in dataset_names:
  364. try:
  365. # Load dataset (ground truth and spatial coordinates)
  366. gnd, locs = import_dataset(dataset_name , dataset_paths[dataset_name])
  367. gnd = gnd.dropna()
  368. except Exception as e:
  369. with open(error_log_file, "a") as log:
  370. log.write(f"Dataset loading error: {dataset_name} - {e}\n")
  371. continue
  372. for method_name in method_names:
  373. try:
  374. # Load predictions
  375. print(f"Running on {dataset_name}_{method_name}")
  376. pred = load_prediction(dataset_name, method_name, pred_paths[dataset_name][method_name])
  377. # If predictions are missing, skip to the next method
  378. if pred is None:
  379. continue
  380. pred = pred[~pred.index.duplicated(keep='first')]
  381. # Ensure first columns are indices
  382. intersect_idx = gnd.index.intersection(pred.index)
  383. # Filter gnd and pred based on the intersection
  384. gnd_filtered = gnd.loc[intersect_idx]
  385. pred_filtered = pred.loc[intersect_idx]
  386. # Replace pred and gnd for further computation
  387. gnd_values = gnd_filtered.values.flatten()
  388. pred_values = pred_filtered.iloc[:, 0].values
  389. # Placeholder for spatial coordinates
  390. spatial_coords = locs.loc[intersect_idx].values
  391. # Compute metrics and populate the pivoted DataFrame
  392. for metric in metric_names:
  393. try:
  394. if metric == "ARI":
  395. value = adjusted_rand_score(gnd_values, pred_values)
  396. elif metric == "NMI":
  397. value = normalized_mutual_info_score(gnd_values, pred_values)
  398. elif metric == "CHAOS":
  399. value = _compute_CHAOS(pred_values, spatial_coords)
  400. elif metric == "PAS":
  401. value = _compute_PAS(pred_values, spatial_coords)
  402. elif metric == "ASW":
  403. value = compute_ASW(pred_values, spatial_coords)
  404. elif metric == "HOM":
  405. value = homogeneity_score(gnd_values, pred_values)
  406. elif metric == "COM":
  407. value = completeness_score(gnd_values, pred_values)
  408. else:
  409. continue
  410. # Populate the DataFrame
  411. metric_results[metric].loc[method_names_dict[method_name], dataset_names_dict[dataset_name]] = value
  412. except Exception as e:
  413. with open(error_log_file, "a") as log:
  414. log.write(f"Metric error: Dataset={dataset_name}, Method={method_name}, Metric={metric} - {e}\n")
  415. except Exception as e:
  416. with open(error_log_file, "a") as log:
  417. log.write(f"Prediction loading error: Dataset={dataset_name}, Method={method_name} - {e}\n")
  418. # Save each metric result as a CSV
  419. for metric, df in metric_results.items():
  420. try:
  421. output_file = os.path.join(output_dir, f"{metric}_results.csv")
  422. df.to_csv(output_file)
  423. print(f"Saved {metric} results to {output_file}")
  424. except Exception as e:
  425. with open(error_log_file, "a") as log:
  426. log.write(f"Metric saving error: Metric={metric} - {e}\n")
  427. def update_metrics(dataset_names, dataset_paths, method_names, pred_paths, error_log_file , output_dir,metric_names="all"):
  428. """
  429. Modify existing output files with new metric values given dataset and method names and paths
  430. Parameters
  431. ----------
  432. dataset_names : list
  433. List of datasets to compute on
  434. Options for dataset names present at bottom of this file
  435. dataset_paths: dictionary of dictionaries
  436. Dictionaries mapping dataset names to dictionaries with st_path and gnd_path needed for importing dataset
  437. method_names: list
  438. List of methods to compute on
  439. Options for method names present at bottom of this file
  440. pred_paths : dictionary of dictionaries
  441. Dictionary mapping dataset names and method names to path of output file
  442. metric_names : list or str
  443. List of metrics to compute or "all"
  444. error_log_file : str
  445. Path to error log file
  446. output_dir : str
  447. Path to save final metric compute files to
  448. Returns
  449. -------
  450. pred : predicted clusters
  451. """
  452. # Available metrics
  453. available_metrics = ["ARI", "NMI", "CHAOS", "PAS", "ASW", "HOM", "COM"]
  454. if metric_names == "all":
  455. metric_names = available_metrics
  456. method_names_dict = {
  457. "SCANIT": "SCANIT",
  458. "CCST": "CCST",
  459. "DeepST": "DeepST",
  460. "GraphST": "GraphST",
  461. "PROST": "PROST",
  462. "SpaSRL": "SpaSRL",
  463. "STAGATE": "STAGATE",
  464. "SpatialPCA": "SpatialPCA",
  465. "banksy": "Banksy",
  466. "giotto": "Giotto",
  467. "DR_SC": "DR.SC",
  468. "ISC_MEB": "ISC.MEB",
  469. "BayesSpace": "BayesSpace",
  470. "PRECAST": "PRECAST",
  471. "BayesCafe": "BayesCafe",
  472. "BASS": "BASS",
  473. "SpaceFlow":"SpaceFlow",
  474. "IRIS":"IRIS"
  475. }
  476. dataset_names_dict = {
  477. "DLPFC151507": "DLPFC 151507",
  478. "DLPFC151508": "DLPFC 151508",
  479. "DLPFC151509": "DLPFC 151509",
  480. "DLPFC151510": "DLPFC 151510",
  481. "DLPFC151669": "DLPFC 151669",
  482. "DLPFC151670": "DLPFC 151670",
  483. "DLPFC151671": "DLPFC 151671",
  484. "DLPFC151672": "DLPFC 151672",
  485. "DLPFC151673": "DLPFC 151673",
  486. "DLPFC151674": "DLPFC 151674",
  487. "DLPFC151675": "DLPFC 151675",
  488. "DLPFC151676": "DLPFC 151676",
  489. "embryo9.5": "Embryo 9.5",
  490. "embryo14.5": "Embryo 14.5",
  491. "mouse_breast_cancer": "Mouse Breast Cancer",
  492. "MERFISH_brain0.04": "MERFISH Brain 0.04",
  493. "MERFISH_brain0.09": "MERFISH Brain 0.09",
  494. "MERFISH_brain0.14": "MERFISH Brain 0.14",
  495. "MERFISH_brain0.19": "MERFISH Brain 0.19",
  496. "MERFISH_brain0.24": "MERFISH Brain 0.24",
  497. "osmFISH": "osmFISH",
  498. "simulated_kidney_cancer410": "Kidney Cancer 410",
  499. "simulated_kidney_cancer411": "Kidney Cancer 411",
  500. "simulated_kidney_cancer506": "Kidney Cancer 506",
  501. "simulated_breast_cancerER+_CID4290": "Breast Cancer ER+ CID 4290",
  502. "simulated_breast_cancerTNBC_CID44971": "Breast Cancer TNBC CID 44971",
  503. "simulated_liver_cancerHCC-1L": "Liver Cancer HCC-1L",
  504. "simulated_liver_cancerHCC-2L": "Liver Cancer HCC-2L",
  505. "simulated_liver_cancerHCC-3L": "Liver Cancer HCC-3L",
  506. "simulated_liver_cancerHCC-4L": "Liver Cancer HCC-4L",
  507. "simulated_breast_atlas": "Breast Atlas",
  508. "simulated_intestineA1": "Intestine A1",
  509. "simulated_intestineA2": "Intestine A2",
  510. "simulated_chicken_heart": "Chicken Heart",
  511. "simulated_prostate_cancer" : "Prostate Cancer",
  512. "simulated_cerebellum" : "Cerebellum"
  513. }
  514. # Initialize pivoted DataFrame for each metric
  515. metric_results = {}
  516. for metric in metric_names:
  517. path = f"{output_dir}/{metric}_results.csv"
  518. try:
  519. df = pd.read_csv(path,skip_blank_lines=True)
  520. df.set_index(df.columns[0], inplace=True, drop=True) # Set the first column as index
  521. df.index.name = None # Optional: remove index name
  522. metric_results[metric] = df
  523. except FileNotFoundError:
  524. # Create a new DataFrame if not already present
  525. metric_results[metric] = pd.DataFrame(index=method_names_dict.values(), columns=dataset_names_dict.values())
  526. print(f"Created new DataFrame for metric {metric}")
  527. os.makedirs(output_dir, exist_ok=True)
  528. # Process each dataset
  529. for dataset_name in dataset_names:
  530. try:
  531. # Load dataset (ground truth and spatial coordinates)
  532. gnd, locs = import_dataset(dataset_name , dataset_paths[dataset_name])
  533. gnd = gnd.dropna()
  534. except Exception as e:
  535. with open(error_log_file, "a") as log:
  536. log.write(f"Dataset loading error: {dataset_name} - {e}\n")
  537. continue
  538. for method_name in method_names:
  539. try:
  540. method_display_name = method_names_dict[method_name]
  541. # Load predictions
  542. print(f"Running on {dataset_name}_{method_name}")
  543. pred = load_prediction(dataset_name, method_name , pred_paths[dataset_name][method_name])
  544. # If predictions are missing, skip to the next method
  545. if pred is None:
  546. continue
  547. pred = pred[~pred.index.duplicated(keep='first')]
  548. # Ensure first columns are indices
  549. intersect_idx = gnd.index.intersection(pred.index)
  550. # Filter gnd and pred based on the intersection
  551. gnd_filtered = gnd.loc[intersect_idx]
  552. pred_filtered = pred.loc[intersect_idx]
  553. # Replace pred and gnd for further computation
  554. gnd_values = gnd_filtered.values.flatten()
  555. pred_values = pred_filtered.iloc[:,0].values
  556. # Placeholder for spatial coordinates
  557. spatial_coords = locs.loc[intersect_idx].values
  558. # Compute metrics and populate the pivoted DataFrame
  559. for metric in metric_names:
  560. try:
  561. if metric == "ARI":
  562. value = adjusted_rand_score(gnd_values, pred_values)
  563. elif metric == "NMI":
  564. value = normalized_mutual_info_score(gnd_values, pred_values)
  565. elif metric == "CHAOS":
  566. value = _compute_CHAOS(pred_values, spatial_coords)
  567. elif metric == "PAS":
  568. value = _compute_PAS(pred_values, spatial_coords)
  569. elif metric == "ASW":
  570. value = compute_ASW(pred_values, spatial_coords)
  571. elif metric == "HOM":
  572. value = homogeneity_score(gnd_values, pred_values)
  573. elif metric == "COM":
  574. value = completeness_score(gnd_values, pred_values)
  575. else:
  576. continue
  577. # Populate the DataFrame
  578. if method_display_name not in metric_results[metric].index:
  579. metric_results[metric].loc[method_display_name] = np.nan
  580. metric_results[metric].loc[method_names_dict[method_name], dataset_names_dict[dataset_name]] = value
  581. print(f"set value of {metric} to {value}")
  582. except Exception as e:
  583. with open(error_log_file, "a") as log:
  584. log.write(f"Metric error: Dataset={dataset_name}, Method={method_name}, Metric={metric} - {e}\n")
  585. except Exception as e:
  586. with open(error_log_file, "a") as log:
  587. log.write(f"Prediction loading error: Dataset={dataset_name}, Method={method_name} - {e}\n")
  588. # Save each metric result as a CSV
  589. for metric, df in metric_results.items():
  590. try:
  591. output_file = os.path.join(output_dir, f"{metric}_results.csv")
  592. df = df.iloc[:len(method_names_dict)] # Dynamic cutoff (good)
  593. df.to_csv(output_file)
  594. print(f"Saved {metric} results to {output_file}")
  595. except Exception as e:
  596. with open(error_log_file, "a") as log:
  597. log.write(f"Metric saving error: Metric={metric} - {e}\n")
  598. # ----------------------- Visualization Functions -----------------------
  599. def plot_stitched_heatmaps(metric_files, dataset_names, dataset_type, output_file):
  600. """
  601. Produce stitched heatmaps for all metrics
  602. Parameters
  603. ----------
  604. metric_files : dictionary
  605. Mapping of metric names to output files
  606. dataset_names: list
  607. Names of datasets to include in heatmap
  608. dataset_type: str
  609. Represents name of selected group of datasets (for example : simulated/real/DLPFC)
  610. Used only for naming
  611. output_file : str
  612. Path to save output image to
  613. """
  614. num_metrics = len(metric_files) + 1
  615. individual_plot_height = 8 # Height of each individual heatmap
  616. total_height = individual_plot_height * num_metrics # Total height for all metrics
  617. fig, axes = plt.subplots(4, 2, figsize=(10, total_height)) # Adjust figure size dynamically
  618. axes = axes.flatten()
  619. for ax, (metric, file_path) in zip(axes, metric_files.items()):
  620. try:
  621. data = pd.read_csv(file_path, index_col=0)
  622. data_subset = data[dataset_names]
  623. sns.heatmap(
  624. data_subset,
  625. cmap="coolwarm",
  626. cbar_kws={'label': metric},
  627. annot=True, # Add annotations
  628. fmt=".2f", # Format for annotations
  629. annot_kws={"size": 6 ,"weight": "bold" }, # Small font size for annotations
  630. ax=ax
  631. )
  632. ax.set_title(f"{dataset_type} Dataset Heatmap - {metric}" , fontweight = "bold")
  633. ax.set_xlabel("Datasets" , fontweight = "bold")
  634. ax.set_ylabel("Methods" , fontweight = "bold")
  635. labels = ax.get_xticklabels()
  636. bold_font = fm.FontProperties(weight='bold')
  637. ax.set_xticklabels(labels, rotation=45, ha="right", fontproperties=bold_font)
  638. y_labels = ax.get_yticklabels()
  639. ax.set_yticklabels(y_labels, fontproperties=bold_font)
  640. except Exception as e:
  641. ax.set_visible(False) # Hide axes if there's an error
  642. print(f"Error processing metric {metric}: {e}")
  643. try:
  644. metrics_to_average = ["NMI", "ARI", "ASW", "HOM", "COM"]
  645. metric_dfs = []
  646. for metric in metrics_to_average:
  647. if metric in metric_files:
  648. data = pd.read_csv(metric_files[metric], index_col=0)
  649. data_subset = data[dataset_names]
  650. metric_dfs.append(data_subset)
  651. # Compute composite score
  652. composite_data = sum(metric_dfs) / len(metrics_to_average)
  653. sns.heatmap(
  654. composite_data,
  655. cmap="coolwarm",
  656. cbar_kws={'label': "composite"},
  657. annot=True, # Add annotations
  658. fmt=".2f", # Format for annotations
  659. annot_kws={"size": 6 ,"weight": "bold" }, # Small font size for annotations
  660. ax=axes[-1]
  661. )
  662. ax.set_title(f"{dataset_type} Dataset Heatmap - Composite" , fontweight = "bold")
  663. ax.set_xlabel("Datasets" , fontweight = "bold")
  664. ax.set_ylabel("Methods" , fontweight = "bold")
  665. labels = ax.get_xticklabels()
  666. bold_font = fm.FontProperties(weight='bold')
  667. ax.set_xticklabels(labels, rotation=45, ha="right", fontproperties=bold_font)
  668. y_labels = ax.get_yticklabels()
  669. ax.set_yticklabels(y_labels, fontproperties=bold_font)
  670. except Exception as e:
  671. axes[-1].set_visible(False) # Hide composite score plot if there's an error
  672. print(f"Error computing composite score: {e}")
  673. plt.tight_layout() # Ensure spacing between plots
  674. plt.savefig(output_file, bbox_inches="tight", dpi = 200)
  675. plt.close()
  676. def plot_individual_heatmaps(metric_files, dataset_names, dataset_type, output_dir):
  677. """
  678. Produce individual heatmaps for metrics of choice
  679. Parameters
  680. ----------
  681. metric_files : dictionary
  682. Mapping of metric names to output files
  683. dataset_names: list
  684. Names of datasets to include in heatmap
  685. dataset_type: str
  686. Represents name of selected group of datasets (for example : simulated/real/DLPFC)
  687. Used only for naming
  688. output_file : str
  689. Path to save output image to
  690. """
  691. os.makedirs(output_dir, exist_ok=True)
  692. metrics = ["ARI", "NMI", "CHAOS", "PAS", "ASW", "HOM", "COM"]
  693. metric_dfs = []
  694. for metric, file_path in metric_files.items():
  695. try:
  696. data = pd.read_csv(file_path, index_col=0)
  697. data_subset = data[dataset_names]
  698. plt.figure(figsize=(8, 6))
  699. ax = sns.heatmap(
  700. data_subset,
  701. cmap="coolwarm",
  702. cbar_kws={'label': metric},
  703. annot=True,
  704. fmt=".2f",
  705. annot_kws={"size": 6, "weight": "bold"}
  706. )
  707. ax.set_title(f"{dataset_type} Dataset Heatmap - {metric}", fontweight="bold")
  708. ax.set_xlabel("Datasets", fontweight="bold")
  709. ax.set_ylabel("Methods", fontweight="bold")
  710. labels = ax.get_xticklabels()
  711. bold_font = fm.FontProperties(weight='bold')
  712. ax.set_xticklabels(labels, rotation=45, ha="right", fontproperties=bold_font)
  713. y_labels = ax.get_yticklabels()
  714. ax.set_yticklabels(y_labels, fontproperties=bold_font)
  715. output_path = os.path.join(output_dir, f"{dataset_type}_Heatmap_{metric}.png")
  716. plt.tight_layout()
  717. plt.savefig(output_path, bbox_inches="tight")
  718. plt.close()
  719. if metric in metrics:
  720. metric_dfs.append(data_subset)
  721. except Exception as e:
  722. print(f"Error processing metric {metric}: {e}")
  723. try:
  724. if metric_dfs:
  725. composite_data_accuracy = sum(metric_dfs[0:4]) / 4
  726. composite_data_accuracy.to_csv(os.path.join(output_dir, f"{dataset_type}_Heatmap_Composite_accuracy.csv"))
  727. plt.figure(figsize=(8, 6))
  728. ax = sns.heatmap(
  729. composite_data_accuracy,
  730. cmap="coolwarm",
  731. cbar_kws={'label': "composite - accuracy"},
  732. annot=True,
  733. fmt=".2f",
  734. annot_kws={"size": 6, "weight": "bold"}
  735. )
  736. ax.set_title(f"{dataset_type} Dataset Heatmap - Composite (accuracy)", fontweight="bold")
  737. ax.set_xlabel("Datasets", fontweight="bold")
  738. ax.set_ylabel("Methods", fontweight="bold")
  739. labels = ax.get_xticklabels()
  740. bold_font = fm.FontProperties(weight='bold')
  741. ax.set_xticklabels(labels, rotation=45, ha="right", fontproperties=bold_font)
  742. y_labels = ax.get_yticklabels()
  743. ax.set_yticklabels(y_labels, fontproperties=bold_font)
  744. output_path = os.path.join(output_dir, f"{dataset_type}_Heatmap_Composite_accuracy.png")
  745. plt.tight_layout()
  746. plt.savefig(output_path, bbox_inches="tight", dpi=200)
  747. plt.close()
  748. except Exception as e:
  749. print(f"Error computing composite accuracy score: {e}")
  750. try:
  751. if metric_dfs:
  752. composite_data_consistency = sum(metric_dfs[4:7]) / 3
  753. composite_data_consistency.to_csv(os.path.join(output_dir, f"{dataset_type}_Heatmap_Composite_consistency.csv"))
  754. plt.figure(figsize=(8, 6))
  755. ax = sns.heatmap(
  756. composite_data_consistency,
  757. cmap="coolwarm",
  758. cbar_kws={'label': "composite - consistency"},
  759. annot=True,
  760. fmt=".2f",
  761. annot_kws={"size": 6, "weight": "bold"}
  762. )
  763. ax.set_title(f"{dataset_type} Dataset Heatmap - Composite (consistency)", fontweight="bold")
  764. ax.set_xlabel("Datasets", fontweight="bold")
  765. ax.set_ylabel("Methods", fontweight="bold")
  766. labels = ax.get_xticklabels()
  767. bold_font = fm.FontProperties(weight='bold')
  768. ax.set_xticklabels(labels, rotation=45, ha="right", fontproperties=bold_font)
  769. y_labels = ax.get_yticklabels()
  770. ax.set_yticklabels(y_labels, fontproperties=bold_font)
  771. output_path = os.path.join(output_dir, f"{dataset_type}_Heatmap_Composite_consistency.png")
  772. plt.tight_layout()
  773. plt.savefig(output_path, bbox_inches="tight", dpi=200)
  774. plt.close()
  775. except Exception as e:
  776. print(f"Error computing composite consistency score: {e}")
  777. try:
  778. if metric_dfs:
  779. composite_data = (composite_data_accuracy + composite_data_consistency)/2
  780. composite_data.to_csv(os.path.join(output_dir, f"{dataset_type}_Heatmap_Composite.csv"))
  781. plt.figure(figsize=(8, 6))
  782. ax = sns.heatmap(
  783. composite_data,
  784. cmap="coolwarm",
  785. cbar_kws={'label': "composite"},
  786. annot=True,
  787. fmt=".2f",
  788. annot_kws={"size": 6, "weight": "bold"}
  789. )
  790. ax.set_title(f"{dataset_type} Dataset Heatmap - Composite", fontweight="bold")
  791. ax.set_xlabel("Datasets", fontweight="bold")
  792. ax.set_ylabel("Methods", fontweight="bold")
  793. labels = ax.get_xticklabels()
  794. bold_font = fm.FontProperties(weight='bold')
  795. ax.set_xticklabels(labels, rotation=45, ha="right", fontproperties=bold_font)
  796. y_labels = ax.get_yticklabels()
  797. ax.set_yticklabels(y_labels, fontproperties=bold_font)
  798. output_path = os.path.join(output_dir, f"{dataset_type}_Heatmap_Composite.png")
  799. plt.tight_layout()
  800. plt.savefig(output_path, bbox_inches="tight", dpi=200)
  801. plt.close()
  802. except Exception as e:
  803. print(f"Error computing composite score: {e}")
  804. def make_embedding_plots(dataset_name,dataset_paths,pred_paths,save_dir, method_names = ["ground_truth","SCANIT", "CCST" , "DeepST" , "GraphST" , "PROST" , "SpaSRL" , "STAGATE","SpatialPCA" , "banksy" , "giotto" , "DR_SC" , "ISC_MEB" , "BayesSpace" , "PRECAST" , "BayesCafe" , "BASS"],stitch=True):
  805. """
  806. Produce embedding plots for select datasets and method outputs
  807. Parameters
  808. ----------
  809. dataset_name : str
  810. Name of dataset to create plot for
  811. dataset_paths : dict
  812. Dictionary of file paths needed for that dataset.
  813. Must contain at least 'st_path'.
  814. Can contain 'gnd_path' if ground truth is from a CSV/TSV file.
  815. pred_paths: dict of dicts
  816. Dictionary of output file paths for those datasets and names
  817. saved as a dataset x method dictionary
  818. save_dir : str
  819. Path to save output images to
  820. """
  821. method_names_dict = {
  822. "SCANIT": "SCANIT",
  823. "CCST": "CCST",
  824. "DeepST": "DeepST",
  825. "GraphST": "GraphST",
  826. "PROST": "PROST",
  827. "SpaSRL": "SpaSRL",
  828. "STAGATE": "STAGATE",
  829. "SpatialPCA": "SpatialPCA",
  830. "banksy": "Banksy",
  831. "giotto": "Giotto",
  832. "DR_SC": "DR.SC",
  833. "ISC_MEB": "ISC.MEB",
  834. "BayesSpace": "BayesSpace",
  835. "PRECAST": "PRECAST",
  836. "BayesCafe": "BayesCafe",
  837. "BASS": "BASS",
  838. "SpaceFlow":"SpaceFlow",
  839. "IRIS" : "IRIS"
  840. }
  841. dataset_names_dict = {
  842. "DLPFC151507": "DLPFC 151507",
  843. "DLPFC151508": "DLPFC 151508",
  844. "DLPFC151509": "DLPFC 151509",
  845. "DLPFC151510": "DLPFC 151510",
  846. "DLPFC151669": "DLPFC 151669",
  847. "DLPFC151670": "DLPFC 151670",
  848. "DLPFC151671": "DLPFC 151671",
  849. "DLPFC151672": "DLPFC 151672",
  850. "DLPFC151673": "DLPFC 151673",
  851. "DLPFC151674": "DLPFC 151674",
  852. "DLPFC151675": "DLPFC 151675",
  853. "DLPFC151676": "DLPFC 151676",
  854. "embryo9.5": "Embryo 9.5",
  855. "embryo14.5": "Embryo 14.5",
  856. "mouse_breast_cancer": "Mouse Breast Cancer",
  857. "MERFISH_brain0.04": "MERFISH Brain 0.04",
  858. "MERFISH_brain0.09": "MERFISH Brain 0.09",
  859. "MERFISH_brain0.14": "MERFISH Brain 0.14",
  860. "MERFISH_brain0.19": "MERFISH Brain 0.19",
  861. "MERFISH_brain0.24": "MERFISH Brain 0.24",
  862. "osmFISH": "osmFISH",
  863. "simulated_kidney_cancer410": "Kidney Cancer 410",
  864. "simulated_kidney_cancer411": "Kidney Cancer 411",
  865. "simulated_kidney_cancer506": "Kidney Cancer 506",
  866. "simulated_breast_cancerER+_CID4290": "Breast Cancer ER+ CID 4290",
  867. "simulated_breast_cancerTNBC_CID44971": "Breast Cancer TNBC CID 44971",
  868. "simulated_liver_cancerHCC-1L": "Liver Cancer HCC-1L",
  869. "simulated_liver_cancerHCC-2L": "Liver Cancer HCC-2L",
  870. "simulated_liver_cancerHCC-3L": "Liver Cancer HCC-3L",
  871. "simulated_liver_cancerHCC-4L": "Liver Cancer HCC-4L",
  872. "simulated_breast_atlas": "Breast Atlas",
  873. "simulated_intestineA1": "Intestine A1",
  874. "simulated_intestineA2": "Intestine A2",
  875. "simulated_chicken_heart": "Chicken Heart",
  876. "simulated_prostate_cancer" : "Prostate Cancer",
  877. "simulated_cerebellum" : "Cerebellum"
  878. }
  879. gnd, locs = import_dataset(dataset_name, dataset_paths[dataset_name], mode = "plot")
  880. for method_name in method_names:
  881. if method_name == "ground_truth":
  882. gnd = gnd.squeeze()
  883. adata = an.AnnData(obs=pd.DataFrame({"cluster": gnd.astype(str)}))
  884. adata.obsm["spatial"] = np.array(locs)
  885. else:
  886. pred = load_prediction(dataset_name,method_name, pred_paths[dataset_name][method_name])
  887. if pred is None:
  888. print(f"output for {dataset_name} {method_name} not present")
  889. continue
  890. pred = pred[~pred.index.duplicated(keep='first')]
  891. intersect_idx = gnd.index.intersection(pred.index)
  892. pred_filtered = pred.loc[intersect_idx]
  893. # Handle both DataFrame and Series cases safely
  894. if isinstance(pred_filtered, pd.DataFrame):
  895. if pred_filtered.shape[1] > 1:
  896. # If multiple columns (e.g. banksy), take the first or warn
  897. print(f"Warning: {method_name} prediction for {dataset_name} has multiple columns; using the first one.")
  898. pred_final = pred_filtered.iloc[:, 0]
  899. else:
  900. pred_final = pred_filtered
  901. pred_final = pred_final.squeeze() # Ensure Series, not DataFrame
  902. pred_final.index = intersect_idx # Ensure correct indexing
  903. # pred_values = pred_filtered.iloc[:, 0].values
  904. spatial_coords = locs.loc[intersect_idx].values
  905. adata = an.AnnData(obs=pd.DataFrame({"cluster": pred_final.astype(str)}))
  906. adata.obsm['spatial'] = np.array(spatial_coords)
  907. # Ensure the directory exists
  908. if not os.path.exists(save_dir):
  909. os.makedirs(save_dir)
  910. print(f"Created directory: {save_dir}")
  911. # Generate and save the plot
  912. plt.rcParams['font.weight'] = 'bold' # Bold for all text
  913. plt.rcParams['axes.titleweight'] = 'bold' # Bold for titles
  914. plt.rcParams['axes.labelweight'] = 'bold'
  915. fig = sc.pl.embedding(
  916. adata,
  917. basis="spatial", # Use the 'spatial' embedding
  918. color="cluster", # Color points by the 'cluster' column
  919. title=f"{dataset_names_dict[dataset_name]} - {"Ground Truth" if method_name == "ground_truth" else method_names_dict[method_name]}",
  920. size=100,
  921. alpha=1,
  922. cmap="coolwarm",
  923. show = True,
  924. return_fig=True
  925. # save=f"{dataset_name}_{method_name}.png" # File will be saved in the default directory by Scanpy
  926. )
  927. # Move the file to the desired directory
  928. current_dir = os.getcwd()
  929. plot_path = os.path.join(current_dir, f"figures/spatial{dataset_name}.png")
  930. if not os.path.exists(os.path.join(save_dir, f"{dataset_name}")):
  931. os.makedirs(os.path.join(save_dir, f"{dataset_name}"))
  932. new_path = os.path.join(save_dir, f"{dataset_name}/{method_name}.png")
  933. fig.savefig(new_path, dpi=200, bbox_inches="tight", facecolor="white")
  934. print(f"figure saved to {new_path}")
  935. plt.close(fig)
  936. # if os.path.exists(plot_path):
  937. # os.rename(plot_path, new_path)
  938. # print(f"Moved plot to: {new_path}")
  939. # else:
  940. # print(f"Plot file not found at: {plot_path}")
  941. # Helper functions to stitch embedding plots together
  942. def get_all_files(folder_path):
  943. file_names = [f for f in os.listdir(folder_path) if os.path.isfile(os.path.join(folder_path, f))]
  944. png_files = [f for f in file_names if f.lower().endswith('.png')]
  945. png_files.sort()
  946. ground_truth = 'ground_truth.png' # Replace with the actual name or pattern for the ground truth file
  947. if ground_truth in png_files:
  948. png_files.remove(ground_truth)
  949. png_files.insert(0, ground_truth)
  950. return png_files
  951. def save_grid_image(grid_rows, grid_cols, folder_path, dataset_name, save_path, width, height):
  952. png_files = get_all_files(folder_path)
  953. image_paths = [os.path.join(folder_path, p) for p in png_files]
  954. images = [Image.open(img) for img in image_paths]
  955. fig, axes = plt.subplots(grid_rows, grid_cols, figsize=(width, height), facecolor="white")
  956. fig.subplots_adjust(wspace=0.0, hspace=0.0) # Reduced spacing between rows and columns
  957. axes = axes.flatten()
  958. for ax, img in zip(axes, images + [None] * (grid_rows * grid_cols - len(images))):
  959. if img is not None:
  960. ax.imshow(img)
  961. ax.axis("off") # Remove axes
  962. output_filename = dataset_name + "_grid_image.png"
  963. plt.savefig(os.path.join(save_path, output_filename), dpi=200, bbox_inches="tight", facecolor="white")
  964. plt.close()
  965. print(f"Grid saved as {output_filename} with 200 DPI.")
  966. # ----------------------- Composite Score Computation (over data subsets) -----------------------
  967. def return_custom_score(metric_files, dataset_names):
  968. metric_dfs = []
  969. for metric, file_path in metric_files.items():
  970. try:
  971. data = pd.read_csv(file_path, index_col=0)
  972. data_subset = data[dataset_names]
  973. metric_dfs.append(data_subset)
  974. except Exception as e:
  975. print(f"Error processing metric {metric}: {e}")
  976. try:
  977. if metric_dfs:
  978. composite_data_accuracy = sum(metric_dfs[0:4]) / 4
  979. except Exception as e:
  980. print(f"Error computing composite accuracy score: {e}")
  981. try:
  982. if metric_dfs:
  983. composite_data_consistency = sum(metric_dfs[4:7]) / 3
  984. except Exception as e:
  985. print(f"Error computing composite consistency score: {e}")
  986. try:
  987. if metric_dfs:
  988. composite_data = (composite_data_accuracy + composite_data_consistency)/2
  989. except Exception as e:
  990. print(f"Error computing composite score: {e}")
  991. composite_data = np.sum(composite_data, axis = 1)/len(dataset_names)
  992. top_5_rows = (
  993. composite_data # Convert to Series with MultiIndex (row, column)
  994. .sort_values(ascending=False) # Sort in descending order
  995. .head(5) # Get top 5
  996. .index.get_level_values(0) # Extract row names
  997. .tolist() # Convert to list
  998. )
  999. return(composite_data, top_5_rows)
  1000. # ----------------------- Useful Variables (comment out to use) -----------------------
  1001. # dataset_names = ["DLPFC151507" , "DLPFC151508" , "DLPFC151509","DLPFC151510","DLPFC151669","DLPFC151670","DLPFC151671","DLPFC151672","DLPFC151673","DLPFC151674","DLPFC151675","DLPFC151676","embryo9.5" , "embryo14.5" , "mouse_breast_cancer", "MERFISH_brain0.04", "MERFISH_brain0.09", "MERFISH_brain0.14", "MERFISH_brain0.19", "MERFISH_brain0.24","osmFISH" , "simulated_kidney_cancer410" , "simulated_kidney_cancer411" , "simulated_kidney_cancer506" , "simulated_breast_cancerER+_CID4290" , "simulated_breast_cancerTNBC_CID44971" , "simulated_liver_cancerHCC-1L", "simulated_liver_cancerHCC-2L", "simulated_liver_cancerHCC-3L", "simulated_liver_cancerHCC-4L" , "simulated_breast_atlas" , "simulated_intestineA1", "simulated_intestineA2","simulated_chicken_heart", "simulated_prostate_cancer", "simulated_cerebellum"]
  1002. # DLPFC_dataset_names = ["DLPFC151507" , "DLPFC151508" , "DLPFC151509","DLPFC151510","DLPFC151669","DLPFC151670","DLPFC151671","DLPFC151672","DLPFC151673","DLPFC151674","DLPFC151675","DLPFC151676"]
  1003. # real_dataset_names = ["embryo9.5" , "embryo14.5" , "mouse_breast_cancer", "MERFISH_brain0.04", "MERFISH_brain0.09", "MERFISH_brain0.14", "MERFISH_brain0.19", "MERFISH_brain0.24","osmFISH"]
  1004. # simulated_dataset_names = ["simulated_kidney_cancer410" , "simulated_kidney_cancer411" , "simulated_kidney_cancer506" , "simulated_breast_cancerER+_CID4290" , "simulated_breast_cancerTNBC_CID44971" , "simulated_liver_cancerHCC-1L", "simulated_liver_cancerHCC-2L", "simulated_liver_cancerHCC-3L", "simulated_liver_cancerHCC-4L" , "simulated_breast_atlas" , "simulated_intestineA1", "simulated_intestineA2","simulated_chicken_heart", "simulated_prostate_cancer", "simulated_cerebellum"]
  1005. # method_names = ["SCANIT", "CCST" , "DeepST" , "GraphST" , "PROST" , "SpaSRL" , "STAGATE","SpatialPCA" , "banksy" , "giotto" , "DR_SC" , "ISC_MEB" , "BayesSpace" , "PRECAST" , "BayesCafe" , "BASS","SpaceFlow", "IRIS"]

Metrics.py at commit d3b2bec, under MIT · at the source

Overview

Authors: Ajita Shree1, V Aditya2, Tanush Kumar2, Hamim Zafar1,3,4
ORCID iDs: Hamim Zafar
  1. Department of Computer Science and Engineering, Indian Institute of Technology Kanpur
  2. Department of Mathematics and Statistics, Indian Institute of Technology Kanpur
  3. Department of Biological Sciences and Bioengineering, Indian Institute of Technology Kanpur
  4. Mehta Family Centre for Engineering in Medicine, Indian Institute of Technology Kanpur
Dates: published online 10 June 2026
Type: Preprint
License: CC BY
Identifiers: DOI 10.21203/rs.3.rs-9676637/v1 · OpenAlex W7164197983
Open access: green, a free copy (OpenAlex)
Status: code verified
Categories: genetics / omics (modality), methods / tools (subfield)
Methods: Connectivity, Machine learning
Keywords: Spatial transcriptomics, Spatial deconvolution, Spatial domain detection, Bivariate metrics
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Citations: not cited yet (Europe PMC); 107 references in the paper

Abstract

Spatial transcriptomic technologies enable high-resolution characterization of gene expression patterns and reconstruction of cellular architecture within tissue contexts. Two key computational problems have emerged for analyzing these datasets: spatial deconvolution, for disentangling celltype compositions at spatial locations, and spatial domain detection, for identifying spatially coherent regions within a tissue section. Although numerous methods have been developed for each task, a comprehensive and unified benchmarking study spanning diverse tissue types, spatial resolutions, and technological platforms remains lacking, hindering informed method selection by end users and impeding future methodological advancements. Here, we present spDDB (https://github.com/Zafar-Lab/spDDB), a comprehensive benchmarking framework for spatial deconvolution and domain detection methods across a large and diverse collection of datasets spanning multiple tissues, technologies, and biological conditions. We evaluated 21 deconvolution methods, including seven recently-developed approaches, across 37 datasets curated from brain, cancer, and organ tissues encompassing four distinct technologies. To enable rigorous evaluation, we introduced SynthST, a deep graph attention autoencoder-based simulator that generates realistic spatial cell-type distributions from spatial transcriptomic data, and employed a suite of spatial bivariate metrics including a novel bivariate Geary’s C metric, alongside rare cell-type, and cell-shape characterization metrics, for multidimensional performance assessment. While Cell2location, RCTD and SONAR emerged as top-performing methods for spatial deconvolution across tissue types, deconvolution performance varied substantially based on tissue architecture, spatial technology, dataset scale, and cell type diversity. For domain detection, we benchmarked 18 methods across 36 datasets spanning six spatial technologies, identifying PROST, BASS, and SpaceFlow as the leading approaches, while revealing notable limitations of existing methods in handling large-scale datasets. Finally, we provide practical guidelines to assist end users in selecting optimal methods for both tasks across diverse experimental settings.

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 20 matches between paragraphs and lines of code.

Zafar-Lab/spDDB

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: d3b2bec027726acb47fe5a57b0ea8bb1dffea380, 7 June 2026
Languages: Jupyter (60), Python (31), R (21)
Size: 164 files, 112 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (docs/requirements.txt), documentation, 35 notebooks
Not found: CITATION.cff, tests, continuous integration
Tools: pandas (81 files), NumPy (76 files), Matplotlib (66 files), Scanpy (62 files), seaborn (44 files), scikit-learn (43 files), SciPy (38 files), PyTorch (29 files), anndata (24 files), TensorFlow (12 files), Seurat (9 files), SingleCellExperiment (6 files), Squidpy (5 files), ggplot2 (4 files), UMAP (4 files), OpenCV (3 files), Pillow (3 files), PyTorch Geometric (3 files), rpy2 (3 files), tidyverse (3 files), h5py (2 files), cowplot (1 file), data.table (1 file), reticulate (1 file), scikit-image (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
114 files

Code availability

Our benchmarking pipeline for both the tasks is provided as reproducible pipeline at https://github.com/Zafar-Lab/spDDB. The simulator code is available at https://github.com/Zafar-Lab/spDDB/tree/main/SynthST and code for all the evaluation experiments at https://github.com/Zafar-Lab/spDDB/tree/main/Experiments.

Reproduced under the paper's license (CC BY), 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:

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

All datasets used in the study are publicly available. The simulated spatial gene expression datasets generated for spatial deconvolution and domain detection benchmarking have been deposited on figshare and are accessible through the project website: https://zafar-lab.github.io/spDDB_datasets.github.io/. Detailed information about all datasets used in this study is provided in Supplementary Tables 2 and 11.

For the spatial cell-type deconvolution task, the spatial and single cell gene expression datasets used to generate simulated spatial gene expression datasets are as follows. The spatial and scRNA dataset for DLPFC are available at ref (46) and (47) respectively, Mouse Brain at ref (39), Hippocampus at ref (24), Cerebellum at ref (24), Visium HD Mouse brain at ref (45) (https://www.10xgenomics.com/datasets, scRNA data at ref (39)), Kidney Cancer at ref (49), Breast Cancer at ref (50), Liver Cancer at ref (51) and (52) respectively, Prostate Cancer at ref (53), Visium HD Lung Cancer at ref (45) (https://www.10xgenomics.com/datasets)and (54), Developmental Lung at ref (56), Human Breast Atlas at ref (50), Intestine at ref (58), Mouse Liver Atlas at ref (59), Human Liver Atlas at ref (59), Kidney Atlas at ref (60) and Chicken heart at ref (57). For the deconvolution task on the simulation strategy 2 datasets, MERFISH preoptic hypothalamic brain dataset is available at ref (48), MERFISH Ileum is available at ref (61), MERFISH Breast Cancer at ref (55), and MERFISH Lung Cancer at ref (55).

For the domain detection task, Tthe spatial and single cell gene expression datasets for DLPFC is available at ref(46) (spatial:https://research.libd.org/spatialLIBD/ (107)) and (47) respectively, Mouse embryo at ref (88)https://db.cngb.org/stomics/mosta/ and (89) respectively, Mouse breast cancer at ref (3) and (90) respectively, osmFISH Mouse cortex at ref (91) (http://sdmbench.drai.cn/ (20)), and (92) respectively, and MERFISH brain at ref (48) (http://sdmbench.drai.cn/ (20)). All other data supporting the findings of this study are available within the article and its supplementary files. Any additional requests for information can be directed to, and will be fulfilled by, the corresponding author.

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, journal, dates, 4 authors, 4 keywords, 92 references.

Cite

This paper

Shree, A., Aditya, V., Kumar, T., & Zafar, H. (2026). A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic Technologies. Research Square (preprint). https://doi.org/10.21203/rs.3.rs-9676637/v1

BibTeX

@article{shree2026comprehensive,
author = {Shree, Ajita and Aditya, V and Kumar, Tanush and Zafar, Hamim},
title = {{A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic Technologies}},
journal = {Research Square (preprint)},
year = {2026},
month = jun,
publisher = {Research Square},
issn = {2693-5015},
doi = {10.21203/rs.3.rs-9676637/v1},
url = {https://doi.org/10.21203/rs.3.rs-9676637/v1}
}

RIS

TY - JOUR
AU - Shree, Ajita
AU - Aditya, V
AU - Kumar, Tanush
AU - Zafar, Hamim
TI - A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic Technologies
T2 - Research Square (preprint)
J2 - Res Sq
PY - 2026
DA - 2026/06/10
SN - 2693-5015
PB - Research Square
DO - 10.21203/rs.3.rs-9676637/v1
UR - https://doi.org/10.21203/rs.3.rs-9676637/v1
ER -

CSL-JSON

{
"id": "10.21203/rs.3.rs-9676637/v1",
"type": "article",
"title": "A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic Technologies",
"container-title": "Research Square (preprint)",
"author": [
{
"family": "Shree",
"given": "Ajita"
},
{
"family": "Aditya",
"given": "V"
},
{
"family": "Kumar",
"given": "Tanush"
},
{
"family": "Zafar",
"given": "Hamim"
}
],
"container-title-short": "Res Sq",
"DOI": "10.21203/rs.3.rs-9676637/v1",
"ISSN": "2693-5015",
"publisher": "Research Square",
"URL": "https://doi.org/10.21203/rs.3.rs-9676637/v1",
"issued": {
"date-parts": [
[
2026,
6,
10
]
]
}
}

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/s41592-026-03194-8 [code]
Beyond benchmarking: an expert-guided consensus approach to spatially aware clustering.
Journal: Nature methods
In common: Squidpy, rpy2, PyTorch Geometric, 17 other tools, methods / tools, genetics / omics, 14 references
[2] doi:10.1002/advs.77003 [code]
SemanticST: A Scalable Multi-Contextual Graph Learning Framework for Uncovering Spatial Niches and Robust Multi-Sample Integration in Spatial Transcriptomics.
Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)
In common: rpy2, PyTorch Geometric, anndata, 10 other tools, methods / tools, genetics / omics, 10 references
[3] doi:10.1016/j.isci.2026.117206 [code]
ReliST: A model-agnostic risk layer for spatial transcriptomics deconvolution.
Journal: iScience
In common: anndata, Scanpy, Pillow, 7 other tools, genetics / omics, 14 references
[4] doi:10.1093/bioinformatics/btag540 [code]
Deciphering spatial heterogeneity by multimodal spatial transcriptomics modelling with SpatialModal.
Journal: Bioinformatics (Oxford, England)
In common: Squidpy, rpy2, PyTorch Geometric, 11 other tools, methods / tools, genetics / omics, 6 references
[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, SingleCellExperiment, UMAP, 18 other tools, genetics / omics
[6] doi:10.1093/nar/gkag706 [code]
scDifformer: diffusion-based post-training for virtual cell modeling across large-scale single-cell data.
Journal: Nucleic acids research
In common: Squidpy, PyTorch Geometric, UMAP, 14 other tools, 3 references
[7] doi:10.1186/s13073-026-01704-z [code]
Gene expression profiling enables refined parcellation of cortical layers in the heterogeneous human cerebral cortex.
Journal: Genome medicine
In common: SingleCellExperiment, UMAP, anndata, 14 other tools, genetics / omics, 4 references
[8] doi:10.1093/bib/bbag404 [code]
Navigating cell maps by deep learning integration of single-cell and spatially resolved transcriptomics.
Journal: Briefings in bioinformatics
In common: PyTorch Geometric, anndata, Scanpy, 10 other tools, genetics / omics, 7 references
[9] doi:10.1093/bib/bbag298 [code]
Empowering multifaceted analysis of spatial transcriptomics data with RGAST.
Journal: Briefings in bioinformatics
In common: rpy2, PyTorch Geometric, anndata, 8 other tools, methods / tools, genetics / omics, 8 references
[10] doi:10.1093/bioinformatics/btag578 [code]
NicheDeSig: niche-aware deconvolution and adaptive signature analysis for spatial transcriptomics.
Journal: Bioinformatics (Oxford, England)
In common: anndata, Scanpy, PyTorch, 4 other tools, methods / tools, genetics / omics, 12 references

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

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.