OSCR

CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning.

Code ↔ Paper

11 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 11 matches
  1. [1] § Methods › Spatial proximity graph construction ↔ cellniche/utils.py, lines 282–354 · score 0.81 · NearestNeighbors, multi slice, spatial graph, undirected, squidpy, COO
  2. [2] § Results › CellNiche integrates cross-platform spatial maps and constructs a unified virtual tissue atlas ↔ cellniche/__init__.py, lines 1–44 · score 0.76 · Seurat alignment score, silhouette width, batch mixing, entropy, iLISI, CellNiche
  3. [3] § Methods › Loss function ↔ cellniche/model.py, lines 8–101 · score 0.70 · learnable parameter, temperature parameter, node embeddings, module, matrix, batch
  4. [4] § Results › CellNiche integrates cross-platform spatial maps and constructs a unified virtual tissue atlas ↔ cellniche/utils.py, lines 768–829 · score 0.65 · Seurat alignment score, batch mixing, cross, CellNiche
  5. [5] § Methods › Evaluation ↔ cellniche/utils.py, lines 561–640 · score 0.65 · clustering evaluation metrics, silhouette score, AMI, ARI, Macro
  6. [6] § Methods › Loss function ↔ cellniche/model.py, lines 8–101 · score 0.59 · log softmax, log probabilities, row, Loss
  7. [7] § Methods › Training module ↔ cellniche/model.py, lines 104–175 · score 0.57 · graph convolutional, node features, encoder, matrix, training, module
  8. [8] § Methods › Cellular identity representation ↔ cellniche/utils.py, lines 132–172 · score 0.54 · feature matrix, expression matrix, hot, phenotypic, gene, CellNiche
  9. [9] § Results › CellNiche shows accurate and robust performance on spatial transcriptomics data ↔ cellniche/utils.py, lines 561–640 · score 0.54 · silhouette score, spatial domain, AMI, ARI, Macro, Clustering
  10. [10] § Methods › Training module ↔ cellniche/model.py, lines 104–175 · score 0.53 · node embedding, node features, layer, edge, graph, module
  11. [11] § Methods › Differential gene expression analysis ↔ cellniche/utils.py, lines 175–221 · score 0.52 · expression matrices, Scanpy, sum, Log, genes, CellNiche

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 · 871 lines · 28 KB · MIT · 6 matches

  1. import os
  2. import random
  3. import numpy as np
  4. import pandas as pd
  5. import torch
  6. import scipy
  7. import scipy.sparse as sps
  8. import logging
  9. import torch.nn.functional as F
  10. import scanpy as sc
  11. from natsort import natsorted
  12. from scipy.sparse import coo_matrix
  13. from sklearn.preprocessing import LabelEncoder
  14. from sklearn.neighbors import NearestNeighbors
  15. from sklearn.cluster import KMeans
  16. from sklearn.metrics import (
  17. adjusted_rand_score,
  18. normalized_mutual_info_score,
  19. adjusted_mutual_info_score,
  20. f1_score,
  21. silhouette_score,
  22. homogeneity_score,
  23. v_measure_score,
  24. fowlkes_mallows_score
  25. )
  26. from torch_sparse import SparseTensor
  27. from torch_geometric.utils import to_undirected
  28. from typing import Optional, Union, Any, Tuple, List
  29. from numpy.typing import ArrayLike
  30. from scipy.stats import entropy
  31. from sklearn.metrics import silhouette_samples
  32. def to_float_tensor(arr):
  33. """Safely convert numpy array / tensor -> float32 tensor w/o grad."""
  34. if isinstance(arr, torch.Tensor):
  35. return arr.clone().detach().float()
  36. else: # numpy / list
  37. return torch.as_tensor(arr, dtype=torch.float)
  38. def _resolve_squidpy_connectivity_key(adata, connectivity_key: str = "spatial") -> str:
  39. """
  40. Resolve Squidpy connectivity key in adata.obsp.
  41. Squidpy usually stores the spatial graph as:
  42. adata.obsp["spatial_connectivities"]
  43. If users pass connectivity_key="spatial", this function will look for
  44. "spatial_connectivities". If users directly pass "spatial_connectivities",
  45. it will also work.
  46. """
  47. if connectivity_key is None:
  48. connectivity_key = "connectivities"
  49. # Case 1: user directly provides the obsp key
  50. if connectivity_key in adata.obsp:
  51. return connectivity_key
  52. # Case 2: user provides Squidpy key_added prefix, e.g. "spatial"
  53. candidate = f"{connectivity_key}_connectivities"
  54. if candidate in adata.obsp:
  55. return candidate
  56. raise KeyError(
  57. f"Cannot find Squidpy spatial connectivity graph in adata.obsp. "
  58. f"Tried '{connectivity_key}' and '{candidate}'. "
  59. f"Available obsp keys: {list(adata.obsp.keys())}. "
  60. f"For multi-slice data, please precompute the graph using "
  61. f"squidpy.gr.spatial_neighbors(..., library_key=..., key_added='{connectivity_key}')."
  62. )
  63. def _edge_index_from_squidpy_obsp(
  64. adata,
  65. connectivity_key: str = "spatial",
  66. ) -> torch.LongTensor:
  67. """
  68. Build edge_index from a Squidpy-generated spatial connectivity matrix.
  69. Expected input:
  70. adata.obsp["spatial_connectivities"]
  71. or:
  72. adata.obsp[f"{connectivity_key}_connectivities"]
  73. Returns:
  74. edge_index: torch.LongTensor with shape [2, num_edges]
  75. """
  76. obsp_key = _resolve_squidpy_connectivity_key(adata, connectivity_key)
  77. adj = adata.obsp[obsp_key]
  78. # Ensure sparse matrix
  79. if not sps.issparse(adj):
  80. adj = sps.csr_matrix(adj)
  81. # Convert to boolean adjacency
  82. adj = adj.astype(bool)
  83. # Convert to LIL for efficient diagonal modification
  84. if isinstance(adj, sps.csr_matrix):
  85. adj = adj.tolil()
  86. else:
  87. adj = adj.tolil()
  88. # Remove self-loops
  89. adj.setdiag(0)
  90. # Back to CSR and remove explicit zeros
  91. adj = adj.tocsr()
  92. adj.eliminate_zeros()
  93. # Extract nonzero indices
  94. row, col = adj.nonzero()
  95. edge_index = torch.tensor(
  96. np.array([row, col]),
  97. dtype=torch.long,
  98. )
  99. logging.info(
  100. f"Using precomputed Squidpy graph from adata.obsp['{obsp_key}']: "
  101. f"{adata.n_obs} nodes, {edge_index.shape[1]} edges."
  102. )
  103. return edge_index
  104. def load_data(
  105. data_path: str,
  106. dataset: str,
  107. phenoLabels: str,
  108. nicheLabels: Optional[str],
  109. embedding_type: str,
  110. radius: Optional[float],
  111. k_neighborhood: int,
  112. hvg: bool,
  113. n_hvg: int,
  114. multi_slice: bool = False,
  115. connectivity_key: str = "spatial",
  116. embedding_key: Optional[str] = None,
  117. ) -> Tuple[torch.FloatTensor, torch.LongTensor, np.ndarray, Any, int, Optional[torch.FloatTensor]]:
  118. """
  119. Load AnnData from .h5ad, build node features, edge index, and labels.
  120. Args:
  121. data_path: Directory containing the .h5ad files.
  122. dataset: Filename (without extension) to load.
  123. phenoLabels: Column name in `adata.obs` for phenotype labels (one-hot).
  124. nicheLabels: Column name in `adata.obs` for true labels (or None to skip).
  125. embedding_type: One of 'pheno_expr', 'pheno', 'expr'.
  126. radius: Radius threshold (if not None) for radius graph.
  127. k_neighborhood: Number of neighbors for kNN graph if radius is None.
  128. hvg: Whether to select highly variable genes for expression.
  129. multi_slice: bool = False,
  130. connectivity_key: str = "spatial",
  131. embedding_key: Optional[str] = None,
  132. Returns:
  133. x: Node feature matrix (one-hot or expr) as FloatTensor.
  134. edge_index: LongTensor[2, E] of graph edges.
  135. y: True label array of shape [N].
  136. adata: The AnnData object.
  137. n_classes: Number of unique labels in y.
  138. expr: Expression matrix as FloatTensor if needed, else None.
  139. """
  140. # 1) Read AnnData
  141. path = os.path.join(data_path, f"{dataset}.h5ad")
  142. adata = sc.read_h5ad(path).copy()
  143. # 2) Build phenotype one-hot features
  144. valid_embedding_types = ["pheno_expr", "pheno", "expr", "embedding"]
  145. if embedding_type not in valid_embedding_types:
  146. raise ValueError(
  147. f"Unknown embedding_type: {embedding_type}. "
  148. f"Expected one of {valid_embedding_types}."
  149. )
  150. # 3) Build phenotype one-hot features only when needed
  151. onehot = None
  152. if embedding_type in ["pheno", "pheno_expr"]:
  153. if phenoLabels is None:
  154. raise ValueError(
  155. "phenoLabels is required when embedding_type is "
  156. f"'{embedding_type}', but got phenoLabels=None."
  157. )
  158. if phenoLabels not in adata.obs:
  159. raise KeyError(
  160. f"phenoLabels='{phenoLabels}' was not found in adata.obs. "
  161. f"Available obs columns: {list(adata.obs.columns)}"
  162. )
  163. pheno = adata.obs[phenoLabels].astype(str)
  164. ph_le = LabelEncoder().fit(pheno)
  165. ph_idx = ph_le.transform(pheno)
  166. onehot = torch.nn.functional.one_hot(
  167. torch.from_numpy(ph_idx),
  168. num_classes=len(ph_le.classes_)
  169. ).float()
  170. # 4) Prepare expression matrix if needed
  171. expr = None
  172. if embedding_type in ["pheno_expr", "expr"]:
  173. temp_adata = adata.copy()
  174. if hvg:
  175. # # way 1
  176. # logging.info(f"hvg way1")
  177. # sc.pp.highly_variable_genes(temp_adata, n_top_genes=n_hvg, flavor='seurat_v3')
  178. # temp_adata = temp_adata[:,temp_adata.var.highly_variable]
  179. # temp_adata.raw=adata
  180. # sc.pp.normalize_total(temp_adata, target_sum=1e4,inplace=True)
  181. # sc.pp.log1p(temp_adata)
  182. # sc.pp.scale(temp_adata)
  183. # way 2
  184. sc.pp.highly_variable_genes(
  185. temp_adata,
  186. flavor="seurat_v3",
  187. n_top_genes=n_hvg,
  188. )
  189. sc.pp.normalize_total(temp_adata, target_sum=1e4)
  190. sc.pp.log1p(temp_adata)
  191. temp_adata = temp_adata[:, temp_adata.var["highly_variable"]].copy()
  192. mat = temp_adata.X
  193. if scipy.sparse.isspmatrix(mat):
  194. arr = mat.toarray()
  195. else:
  196. arr = np.asarray(mat)
  197. expr = torch.from_numpy(arr).float()
  198. del temp_adata
  199. # 5) For embedding_type="embedding", load features from adata.obsm[embedding_key]
  200. embedding = None
  201. if embedding_type == "embedding":
  202. if embedding_key is None:
  203. raise ValueError(
  204. "embedding_key is required when embedding_type='embedding'. "
  205. "Please provide the key of the embedding stored in adata.obsm, "
  206. "for example embedding_key='X_scVI' or embedding_key='X_pca'."
  207. )
  208. if embedding_key not in adata.obsm:
  209. raise KeyError(
  210. f"embedding_key='{embedding_key}' was not found in adata.obsm. "
  211. f"Available obsm keys: {list(adata.obsm.keys())}"
  212. )
  213. emb_arr = adata.obsm[embedding_key]
  214. if scipy.sparse.isspmatrix(emb_arr):
  215. emb_arr = emb_arr.toarray()
  216. else:
  217. emb_arr = np.asarray(emb_arr)
  218. if emb_arr.ndim != 2:
  219. raise ValueError(
  220. f"adata.obsm['{embedding_key}'] must be a 2D matrix, "
  221. f"but got shape {emb_arr.shape}."
  222. )
  223. if emb_arr.shape[0] != adata.n_obs:
  224. raise ValueError(
  225. f"adata.obsm['{embedding_key}'] has {emb_arr.shape[0]} rows, "
  226. f"but adata has {adata.n_obs} cells."
  227. )
  228. embedding = torch.from_numpy(emb_arr).float()
  229. # 6) Build graph edge_index
  230. if multi_slice:
  231. logging.warning(
  232. "multi_slice=True detected. "
  233. "For multi-slice or multi-sample data, CellNiche expects a precomputed "
  234. "sample-aware spatial graph generated by Squidpy"
  235. )
  236. edge_index = _edge_index_from_squidpy_obsp(
  237. adata,
  238. connectivity_key=connectivity_key,
  239. )
  240. else:
  241. # Original single-slice logic
  242. if "edgeList" in adata.uns:
  243. edge_np = np.array(adata.uns["edgeList"])
  244. # Support both shape [2, E] and [E, 2]
  245. if edge_np.ndim != 2:
  246. raise ValueError("adata.uns['edgeList'] must be a 2D array.")
  247. if edge_np.shape[0] == 2:
  248. edge_index = torch.from_numpy(edge_np).long()
  249. elif edge_np.shape[1] == 2:
  250. edge_index = torch.from_numpy(edge_np.T).long()
  251. else:
  252. raise ValueError(
  253. "adata.uns['edgeList'] should have shape [2, E] or [E, 2]."
  254. )
  255. edge_index = to_undirected(edge_index)
  256. else:
  257. # choose coords
  258. if "spatial" in adata.obsm and adata.obsm["spatial"] is not None:
  259. coords = adata.obsm["spatial"]
  260. else:
  261. coords = adata.obs[["x", "y"]].to_numpy()
  262. if radius is not None:
  263. nbrs = NearestNeighbors(radius=radius).fit(coords)
  264. _, idxs = nbrs.radius_neighbors(coords)
  265. rows_list, cols_list = [], []
  266. for i, neighbors in enumerate(idxs):
  267. # remove self-loop
  268. neighbors = neighbors[neighbors != i]
  269. rows_list.extend([i] * len(neighbors))
  270. cols_list.extend(neighbors.tolist())
  271. rows = np.asarray(rows_list)
  272. cols = np.asarray(cols_list)
  273. else:
  274. nbrs = NearestNeighbors(n_neighbors=k_neighborhood + 1).fit(coords)
  275. _, idxs = nbrs.kneighbors(coords)
  276. rows = np.repeat(np.arange(coords.shape[0]), k_neighborhood)
  277. cols = idxs[:, 1:].flatten()
  278. mat = coo_matrix(
  279. (np.ones_like(rows), (rows, cols)),
  280. shape=(coords.shape[0], coords.shape[0]),
  281. )
  282. mat = mat + mat.T # make undirected
  283. edge_index = torch.from_numpy(
  284. np.vstack(mat.nonzero()).astype(np.int64)
  285. )
  286. neighbors_count = np.array([len(neighbors) for neighbors in idxs])
  287. average_neighbors = neighbors_count.mean()
  288. logging.info(f"Average number of neighbors per node: {average_neighbors}")
  289. # 7) Encode true labels from nicheLabels if provided
  290. if nicheLabels is not None and nicheLabels in adata.obs:
  291. true_vals = adata.obs[nicheLabels].astype(str)
  292. nl_encoder = LabelEncoder().fit(true_vals)
  293. y = nl_encoder.transform(true_vals)
  294. n_classes = len(nl_encoder.classes_)
  295. else:
  296. # default dummy labels: all-zero, one class
  297. y = np.zeros(adata.n_obs, dtype=int)
  298. n_classes = 1
  299. # 8) Return according to embedding type
  300. if embedding_type == "pheno_expr":
  301. logging.info(
  302. f"Loaded {dataset}: "
  303. f"{onehot.shape[0]} nodes, {edge_index.shape[1]} edges, "
  304. f"{onehot.shape[-1]} phenotype features, "
  305. f"{expr.shape[-1]} expression features"
  306. )
  307. return to_float_tensor(onehot), edge_index, y, adata, n_classes, to_float_tensor(expr)
  308. elif embedding_type == "pheno":
  309. logging.info(
  310. f"Loaded {dataset}: "
  311. f"{onehot.shape[0]} nodes, {edge_index.shape[1]} edges, "
  312. f"{onehot.shape[-1]} phenotype features"
  313. )
  314. return to_float_tensor(onehot), edge_index, y, adata, n_classes, None
  315. elif embedding_type == "expr":
  316. logging.info(
  317. f"Loaded {dataset}: "
  318. f"{expr.shape[0]} nodes, {edge_index.shape[1]} edges, "
  319. f"{expr.shape[-1]} expression features"
  320. )
  321. return to_float_tensor(expr), edge_index, y, adata, n_classes, None
  322. elif embedding_type == "embedding":
  323. logging.info(
  324. f"Loaded {dataset}: "
  325. f"{embedding.shape[0]} nodes, {edge_index.shape[1]} edges, "
  326. f"{embedding.shape[-1]} embedding features from adata.obsm['{embedding_key}']"
  327. )
  328. return to_float_tensor(embedding), edge_index, y, adata, n_classes, None
  329. def setup_seed(seed: int) -> None:
  330. """
  331. Set random seed for reproducibility across Python, NumPy, Torch, and CUDA.
  332. Args:
  333. seed (int): The seed to set.
  334. """
  335. random.seed(seed)
  336. np.random.seed(seed)
  337. torch.manual_seed(seed)
  338. torch.cuda.manual_seed_all(seed)
  339. # torch.backends.cudnn.deterministic = True
  340. # torch.backends.cudnn.benchmark = False
  341. def create_sparse_tensor_from_edges(
  342. rows: list[int],
  343. cols: list[int],
  344. sparse_size: tuple[int, int],
  345. device: torch.device = torch.device("cpu"),
  346. ) -> SparseTensor:
  347. """
  348. Create a SparseTensor from row/col indices.
  349. Args:
  350. rows (list[int]): Row indices.
  351. cols (list[int]): Column indices.
  352. sparse_size (tuple[int,int]): Matrix size.
  353. device (torch.device): Device for tensor.
  354. Returns:
  355. SparseTensor
  356. """
  357. vals = torch.ones(len(rows), device=device)
  358. return SparseTensor(
  359. row=torch.tensor(rows, device=device),
  360. col=torch.tensor(cols, device=device),
  361. value=vals,
  362. sparse_sizes=sparse_size,
  363. )
  364. def sparse_intersection_and_union(
  365. adj1: SparseTensor,
  366. adj2: SparseTensor,
  367. strategy: str = "and",
  368. ) -> SparseTensor:
  369. """
  370. Compute intersection or union of two sparse adjacency matrices.
  371. Args:
  372. adj1, adj2 (SparseTensor): Input graphs.
  373. strategy (str): 'and' or 'or'.
  374. Returns:
  375. SparseTensor
  376. """
  377. rows1, cols1 = adj1.storage.row(), adj1.storage.col()
  378. rows2, cols2 = adj2.storage.row(), adj2.storage.col()
  379. device = rows1.device
  380. set1 = set(zip(rows1.tolist(), cols1.tolist()))
  381. set2 = set(zip(rows2.tolist(), cols2.tolist()))
  382. if strategy == "and":
  383. common = set1 & set2
  384. else:
  385. common = set1 | set2
  386. if not common:
  387. return SparseTensor(sparse_sizes=adj1.sparse_sizes())
  388. rows, cols = zip(*common)
  389. return create_sparse_tensor_from_edges(rows, cols, adj1.sparse_sizes(), device)
  390. def get_positivePairs(
  391. subAdj: SparseTensor,
  392. features: Optional[torch.Tensor] = None,
  393. strategy: str = "freq",
  394. ) -> SparseTensor:
  395. """
  396. Generate positive pair adjacency based on strategy.
  397. Args:
  398. subAdj (SparseTensor): Stochastic subgraph adjacency.
  399. features (Tensor): Node features.
  400. strategy (str): 'freq', 'sim', 'and', 'or'.
  401. Returns:
  402. SparseTensor: positive adjacency.
  403. """
  404. row, col, val = subAdj.storage.row(), subAdj.storage.col(), subAdj.storage.value()
  405. # Frequency-based mask
  406. freq_thresh = subAdj.sum(dim=1) / subAdj.storage.colptr()[1:]
  407. mask = val > freq_thresh[row]
  408. rows, cols, vals = row[mask], col[mask], val[mask]
  409. freq_adj = SparseTensor(row=rows, col=cols, value=vals, sparse_sizes=subAdj.sparse_sizes())
  410. if strategy == "freq":
  411. return freq_adj
  412. if features is None:
  413. raise ValueError("Features required for non-freq strategy.")
  414. # Similarity-based mask
  415. f_row, f_col = features[row], features[col]
  416. sim_vals = F.cosine_similarity(f_row, f_col, dim=1)
  417. sim_vals[row == col] = 0
  418. sim_adj = SparseTensor(row=row, col=col, value=sim_vals, sparse_sizes=subAdj.sparse_sizes())
  419. sim_thresh = sim_adj.sum(dim=1) / sim_adj.storage.colptr()[1:]
  420. sim_mask = sim_vals > sim_thresh[row]
  421. sim_rows, sim_cols, sim_vals = row[sim_mask], col[sim_mask], sim_vals[sim_mask]
  422. sim_based = SparseTensor(row=sim_rows, col=sim_cols, value=sim_vals, sparse_sizes=subAdj.sparse_sizes())
  423. if strategy == "sim":
  424. return sim_based
  425. # AND / OR combination
  426. return sparse_intersection_and_union(freq_adj, sim_based, strategy)
  427. def match_labels(true_labels, predicted_labels, n_classes):
  428. from scipy.optimize import linear_sum_assignment as linear_assignment
  429. cost_matrix = np.zeros((n_classes, n_classes))
  430. for i in range(n_classes):
  431. for j in range(n_classes):
  432. cost_matrix[i, j] = np.sum((true_labels == i) & (predicted_labels == j))
  433. row_ind, col_ind = linear_assignment(-cost_matrix)
  434. new_labels = np.copy(predicted_labels)
  435. for i, j in zip(row_ind, col_ind):
  436. new_labels[predicted_labels == j] = i
  437. return new_labels
  438. def refine_spatial_domains(y_pred, coord, n_neighbors=6):
  439. nbrs = NearestNeighbors(n_neighbors=n_neighbors + 1).fit(coord)
  440. distances, indices = nbrs.kneighbors(coord)
  441. indices = indices[:, 1:]
  442. y_refined = pd.Series(index=y_pred.index, dtype='object')
  443. for i in range(y_pred.shape[0]):
  444. y_pred_count = y_pred[indices[i, :]].value_counts()
  445. if y_pred[i] in y_pred_count.index:
  446. if (y_pred_count.loc[y_pred[i]] < n_neighbors / 2) and (y_pred_count.max() > n_neighbors / 2):
  447. y_refined[i] = y_pred_count.idxmax()
  448. else:
  449. # y_refined[i] = y_pred[i] # waring
  450. y_refined.iloc[i] = y_pred[i]
  451. else:
  452. y_refined.iloc[i] = y_pred[i]
  453. y_refined = pd.Categorical(
  454. values=y_refined.astype('U'),
  455. categories=natsorted(map(str, y_refined.unique())),
  456. )
  457. return y_refined
  458. def clustering_st(
  459. adata: Any,
  460. n_clusters: int,
  461. features: Optional[Union[torch.Tensor, np.ndarray]] = None,
  462. true_labels: Optional[np.ndarray] = None,
  463. refine: bool = False,
  464. ) -> Tuple[Any, dict]:
  465. """
  466. Perform KMeans clustering and compute evaluation metrics.
  467. Args:
  468. adata (AnnData): Annotated data object.
  469. features (Tensor or ndarray): Embeddings.
  470. n_clusters (int): Number of clusters.
  471. true_labels (ndarray): Ground-truth labels.
  472. Returns:
  473. adata (AnnData): Updated with 'kmeans' clusters.
  474. metrics (dict): Cluster evaluation metrics.
  475. """
  476. # 1) Convert to numpy
  477. if torch.is_tensor(features):
  478. feats = features.cpu().numpy()
  479. else:
  480. feats = features
  481. # 2) Run KMeans
  482. km = KMeans(n_clusters=n_clusters, max_iter=5000, n_init=10)
  483. # km = KMeans(n_clusters=n_clusters, max_iter=10000, n_init=20)
  484. raw_labels = km.fit_predict(feats).astype(int)
  485. # 3) Store raw labels
  486. adata.obs['kmeans'] = pd.Categorical(raw_labels)
  487. clustering_results = {'kmeans': raw_labels}
  488. # 4) Optional spatial refinement
  489. if refine:
  490. # spatial coords must exist
  491. coords = adata.obsm.get('spatial')
  492. if coords is None:
  493. raise ValueError("adata.obsm['spatial'] needed for refinement")
  494. for method, labels in list(clustering_results.items()):
  495. refined = refine_spatial_domains(pd.Series(labels), coords)
  496. refined = refined.astype(int)
  497. col = f"{method}_refined"
  498. adata.obs[col] = pd.Categorical(refined)
  499. clustering_results[col] = refined
  500. # 5) Compute metrics
  501. metrics_results: dict = {}
  502. if true_labels is not None:
  503. for method, labels in clustering_results.items():
  504. # align predicted → true
  505. aligned = match_labels(true_labels, labels, n_clusters)
  506. acc = (aligned == true_labels).mean()
  507. nmi = normalized_mutual_info_score(true_labels, aligned)
  508. ari = adjusted_rand_score(true_labels, aligned)
  509. ami = adjusted_mutual_info_score(true_labels, aligned)
  510. f1m = f1_score(true_labels, aligned, average='macro')
  511. f1i = f1_score(true_labels, aligned, average='micro')
  512. sil = silhouette_score(feats, true_labels)
  513. FMI = fowlkes_mallows_score(true_labels, aligned)
  514. v_measure = v_measure_score(true_labels, aligned)
  515. homogeneity = homogeneity_score(true_labels, aligned)
  516. metrics_results[method] = {
  517. 'Acc': acc,
  518. 'NMI': nmi,
  519. 'AMI': ami,
  520. 'ARI': ari,
  521. 'F1 Macro': f1m,
  522. 'F1 Micro': f1i,
  523. 'Silhouette': sil,
  524. 'Fowlkes-Mallows': FMI,
  525. 'V-Measure': v_measure,
  526. 'Homogeneity': homogeneity
  527. }
  528. return adata, metrics_results
  529. def _rng(random_state: Optional[int] = None) -> np.random.Generator:
  530. """Return a NumPy Generator with the requested seed (or global RNG)."""
  531. return np.random.default_rng(random_state)
  532. def _nearest_neighbors(
  533. x: np.ndarray, k: int, **kwargs
  534. ) -> np.ndarray:
  535. """
  536. Return indices of the *k* nearest neighbours for every point in *x*.
  537. The first neighbour returned by ``sklearn`` is the query point itself,
  538. so we discard it.
  539. """
  540. nn = NearestNeighbors(n_neighbors=k + 1, **kwargs).fit(x)
  541. indices = nn.kneighbors(x, return_distance=False)[:, 1:] # drop self‑index
  542. return indices
  543. def _encode_labels(labels: ArrayLike) -> np.ndarray:
  544. """
  545. Map arbitrary label values to consecutive integers starting from 0.
  546. This simplifies downstream use of ``np.bincount`` and avoids
  547. large sparse counts when label values are not contiguous.
  548. """
  549. labels = np.asarray(labels)
  550. _, encoded = np.unique(labels, return_inverse=True)
  551. return encoded
  552. # ---------------------------------------------------------------------
  553. # 1. Entropy of Batch Mixing
  554. # ---------------------------------------------------------------------
  555. def compute_entropy_batch_mixing(
  556. embeddings: np.ndarray,
  557. batch_labels: ArrayLike,
  558. k: int = 50,
  559. normalize: bool = True,
  560. **nn_kwargs,
  561. ) -> float:
  562. """
  563. Average entropy of batch labels in the *k*‑NN neighbourhood of each cell.
  564. Parameters
  565. ----------
  566. embeddings
  567. Low‑dimensional representation of shape *(n_cells, n_dims)*.
  568. batch_labels
  569. Iterable of length *n_cells* with one label per cell.
  570. k
  571. Number of neighbours (*excluding* the query cell) to consider.
  572. normalize
  573. If ``True`` (default) divide by the maximal entropy
  574. ``log(n_batches)``, yielding values in ``[0, 1]``.
  575. **nn_kwargs
  576. Additional arguments forwarded to :class:`sklearn.neighbors.NearestNeighbors`.
  577. Returns
  578. -------
  579. float
  580. Mean entropy across all cells.
  581. """
  582. if k < 1:
  583. raise ValueError("k must be ≥ 1")
  584. batch_labels = _encode_labels(batch_labels)
  585. indices = _nearest_neighbors(embeddings, k=k, **nn_kwargs)
  586. n_batches = int(batch_labels.max()) + 1
  587. max_ent = np.log(n_batches) if normalize else 1.0
  588. entropies = []
  589. for nbr in indices:
  590. counts = np.bincount(batch_labels[nbr], minlength=n_batches)
  591. probs = counts / k # guaranteed non‑negative, summing to 1
  592. ent = entropy(probs) / max_ent if max_ent > 0 else 0.0
  593. entropies.append(ent)
  594. return float(np.mean(entropies))
  595. # ---------------------------------------------------------------------
  596. # 2. iLISI
  597. # ---------------------------------------------------------------------
  598. def compute_ilisi(
  599. embeddings: np.ndarray,
  600. batch_labels: ArrayLike,
  601. k: int = 90,
  602. **nn_kwargs,
  603. ) -> np.ndarray:
  604. """
  605. Compute the *inverse* Local Inverse Simpson’s Index (iLISI).
  606. For each cell *i*::
  607. iLISI_i = 1 − (# neighbours from same batch) / k
  608. Thus, 0 indicates perfect batch isolation; 1 indicates perfect mixing.
  609. Parameters
  610. ----------
  611. embeddings
  612. *(n_cells, n_dims)* array.
  613. batch_labels
  614. Iterable with one batch label per cell.
  615. k
  616. Number of neighbours (*excluding* the query cell) to consider.
  617. **nn_kwargs
  618. Extra arguments for :class:`sklearn.neighbors.NearestNeighbors`.
  619. Returns
  620. -------
  621. np.ndarray
  622. Vector of length *n_cells* with iLISI scores.
  623. """
  624. batch_labels = _encode_labels(batch_labels)
  625. indices = _nearest_neighbors(embeddings, k=k, **nn_kwargs)
  626. same_batch = (batch_labels[indices] == batch_labels[:, None]).sum(axis=1)
  627. ilisi = 1.0 - (same_batch / k)
  628. return ilisi
  629. # ---------------------------------------------------------------------
  630. # 3. Seurat Alignment Score (SAS)
  631. # ---------------------------------------------------------------------
  632. def compute_seurat_alignment_score(
  633. embeddings: np.ndarray,
  634. batch_labels: ArrayLike,
  635. neighbor_frac: float = 0.01,
  636. n_repeats: int = 3,
  637. random_state: Optional[int] = None,
  638. **nn_kwargs,
  639. ) -> float:
  640. """
  641. Seurat Alignment Score (Butler *et al.*, Cell 2018).
  642. The score estimates how well batches mix after integration by repeatedly
  643. down‑sampling to equal batch sizes and measuring the proportion of
  644. cross‑batch neighbours.
  645. Parameters
  646. ----------
  647. embeddings
  648. *(n_cells, n_dims)* array.
  649. batch_labels
  650. Iterable with one batch label per cell.
  651. neighbor_frac
  652. Fraction of the (sub‑sampled) cells to use as *k* in *k*‑NN.
  653. Must be in ``(0, 1]``.
  654. n_repeats
  655. Number of random sub‑samples to average.
  656. random_state
  657. Seed for reproducibility.
  658. **nn_kwargs
  659. Extra arguments for :class:`sklearn.neighbors.NearestNeighbors`.
  660. Returns
  661. -------
  662. float
  663. Mean SAS across repeats (1 = perfect mixing, 0 = no mixing).
  664. """
  665. if not (0 < neighbor_frac <= 1):
  666. raise ValueError("neighbor_frac must be in (0, 1]")
  667. rng = _rng(random_state)
  668. batch_labels = _encode_labels(batch_labels)
  669. batch_indices = [np.where(batch_labels == b)[0] for b in np.unique(batch_labels)]
  670. min_size = min(len(idx) for idx in batch_indices)
  671. n_batches = len(batch_indices)
  672. scores = []
  673. for _ in range(n_repeats):
  674. # balanced subsample
  675. sel = np.concatenate([rng.choice(idx, min_size, replace=False) for idx in batch_indices])
  676. x_sub, y_sub = embeddings[sel], batch_labels[sel]
  677. k = max(int(round(len(sel) * neighbor_frac)), 1)
  678. indices = _nearest_neighbors(x_sub, k=k, **nn_kwargs)
  679. same_batch = (y_sub[indices] == y_sub[:, None]).sum(axis=1).mean()
  680. score = (k - same_batch) * n_batches / (k * (n_batches - 1))
  681. scores.append(min(score, 1.0)) # numerical guard
  682. return float(np.mean(scores))
  683. # ---------------------------------------------------------------------
  684. # 4. ASW‑batch
  685. # ---------------------------------------------------------------------
  686. def compute_avg_silhouette_width_batch(
  687. embeddings: np.ndarray,
  688. batch_labels: ArrayLike,
  689. cell_types: ArrayLike,
  690. min_cells: int = 3,
  691. **silhouette_kwargs,
  692. ) -> float:
  693. """
  694. Average Silhouette Width computed per cell‑type, then averaged.
  695. Parameters
  696. ----------
  697. ...
  698. min_cells
  699. Minimum number of cells a cell‑type must have to be included.
  700. """
  701. x = embeddings
  702. y = _encode_labels(batch_labels)
  703. ct = _encode_labels(cell_types)
  704. scores = []
  705. for t in np.unique(ct):
  706. mask = ct == t
  707. n = mask.sum()
  708. n_labels = np.unique(y[mask]).size
  709. # Skip if too few cells OR labels≈cells (invalid for silhouette)
  710. if n < min_cells or n_labels < 2 or n_labels >= n:
  711. scores.append(0.0)
  712. continue
  713. try:
  714. s = silhouette_samples(x[mask], y[mask], **silhouette_kwargs)
  715. scores.append((1.0 - np.abs(s)).mean())
  716. except ValueError:
  717. scores.append(0.0)
  718. return float(np.mean(scores))

utils.py at commit af58974, under MIT · at the source

Overview

  1. Key Laboratory of Systems Health Science of Zhejiang Province, School of Life Science, Hangzhou Institute for Advanced Study, University of Chinese Academy of Sciences,Hangzhou, China
  2. State Key Laboratory of Genome and Multi-omics Technologies, BGI Research,Hangzhou, China
  3. Key Laboratory of Spatial Omics of Zhejiang Province, BGI Research,Hangzhou, China
  4. SJTU-Yale Joint Center for Biostatistics and Data Science, Department of Bioinformatics and Biostatistics, School of Life Sciences and Biotechnology, Shanghai Jiao Tong University,Shanghai, China
  5. State Key Laboratory of Mathematical Sciences, Academy of Mathematics and Systems Science, Chinese Academy of Sciences,Beijing, China
  6. School of Mathematics, University of Chinese Academy of Sciences, Chinese Academy of Sciences,Beijing, China
Journal: Nature communications, volume 17, issue 1, article 5547
Dates: received 12 August 2025; accepted 24 March 2026; published online 22 April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41467-026-71759-4 · PMID 42020427 · PMCID PMC13287597 · OpenAlex W7155187838
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), human (organism), mouse (organism), other condition (population), cellular / molecular (subfield)
Methods: Statistics, Smoothing, state filtering, decompositions, Machine learning
Keywords: Computational models, Machine learning, Computer modelling, Computer science
MeSH: Carcinoma, Non-Small-Cell Lung*, Cellular Microenvironment*, Lung Neoplasms*, Tumor Microenvironment*, Animals, Brain, Clustering Algorithms, Humans, Machine Learning, Mice (* major topic)
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: National Key Research and Development Program of China (2022YFA1004800); Strategic Priority Research Program of the Chinese Academy of Sciences (XDB1350000); CAS Project for Young Scientists in Basic Research (YSBR-077); National Natural Science Foundation of China (12025107, 12571550, 12326610)
Citations: cited by 1 paper (Europe PMC); 65 references in the paper
Research resources: RRID:SCR_024672

Abstract

Deciphering cellular microenvironments at atlas scale remains challenging because molecular identity, spatial context, and platform heterogeneity are tightly coupled. Here we present CellNiche, a scalable contrastive-learning framework that identifies and characterizes cellular microenvironments from spatial omics data using cell-centric spatial-proximity subgraphs. CellNiche combines spatial co-localization and molecular co-expression cues to learn microenvironment-aware embeddings. Across spatial omics datasets from multiple platforms (>10 million cells in total), scaling experiments show improved representations with more training data and competitive clustering and embedding-quality performance with efficient computation. In a multi-sample human non-small-cell lung cancer (NSCLC) cohort, CellNiche identifies conserved and sample-specific tumor and immune microenvironments and captures localized spatial transitions. In four independent mouse brain atlases, CellNiche integrates 293 slices into a unified virtual brain map for cross-atlas annotation transfer and spatial refinement.

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

Repositories

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

QIFEIDKN/STGATE_pyG

License: none: the authors keep all their rights
State: the link is dead, verified on 29 September 2026
Evidence: found in the paper
Software Heritage: not archived
Found in: the text, “STAGATE (v1.0.0)”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 29 September 2026: the link is dead
  • 29 September 2026: the link is dead

Super-LzzZ/CellNiche

License: MIT
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: af58974ded7cf57299a9f8952d4cc6dffee39c6f, 4 May 2026
Languages: Jupyter (11), Python (6)
Size: 29 files, 17 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (pyproject.toml), 11 notebooks
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (8 files), NumPy (7 files), pandas (5 files), Scanpy (5 files), anndata (4 files), Matplotlib (4 files), PyTorch Geometric (4 files), scikit-learn (4 files), seaborn (4 files), SciPy (2 files), NetworkX (1 file), Squidpy (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
12 files

Zenodo 19143524

License: MIT
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: the references
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 29 September 2026: the link answers (HTTP 200)
  • 29 September 2026: the link answers (HTTP 200)
At the source:

Code availability

The software package implementing the CellNiche algorithm has been deposited at GitHub https://github.com/Super-LzzZ/CellNiche under the MIT license. The version associated with this study has been archived at Zenodo (10.5281/zenodo.19143524)65.

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:

  • 3 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 10 scripts, each with its path and the digest of its content;
  • 11 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data availability

The osmFISH dataset of mouse somatosensory cortex is available at https://github.com/drieslab/spatial-datasets (https://github.com/drieslab/spatial-datasets.)28. The mouse spleen CODEX dataset is available at https://data.mendeley.com/datasets/zjnpwh8m5b/130. The STARMap dataset of mouse brain is available at https://singlecell.broadinstitute.org/single_cell/study/SCP183032. The human CRC CODEX dataset is available at https://data.mendeley.com/datasets/mpjzbtfgfr/137. The NSCLC CosMx dataset is available at https://nanostring.com/products/cosmx-spatial-molecular-imager/nsclc-ffpe-dataset/27. The spatial transcriptomics atlases of mouse brain are available at https://singlecell.broadinstitute.org/single_cell/study/SCP1830 (Atlas 1)32, https://doi.brainimagelibrary.org/doi/10.35077/g.610 (Atlas 2)31, https://doi.brainimagelibrary.org/doi/10.35077/act-bag (Atlas 3)33, https://info.vizgen.com/mouse-brain-map (Atlas 4)34. The mouse E16.5 whole embryo Stereo-seq data is available at https://db.cngb.org/stomics/mosta/download/36. Source data are provided in this paper.

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, 29 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 4 keywords, 10 MeSH terms, 1 funder, 57 references, 1 RRID.

Cite

This paper

Liang, Z., Zhong, B., Jiao, M., Wang, Y., & Liu, S. (2026). CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning. Nature communications, 17(1), 5547. https://doi.org/10.1038/s41467-026-71759-4

BibTeX

@article{liang2026cellniche,
author = {Liang, Zhongming and Zhong, Bingxu and Jiao, Mingqi and Wang, Yong and Liu, Shiping},
title = {{CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning}},
journal = {Nature communications},
year = {2026},
month = apr,
volume = {17},
number = {1},
pages = {5547},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-71759-4},
url = {https://doi.org/10.1038/s41467-026-71759-4},
pmid = {42020427},
pmcid = {PMC13287597}
}

RIS

TY - JOUR
AU - Liang, Zhongming
AU - Zhong, Bingxu
AU - Jiao, Mingqi
AU - Wang, Yong
AU - Liu, Shiping
TI - CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/04/22
VL - 17
IS - 1
SP - 5547
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-71759-4
UR - https://doi.org/10.1038/s41467-026-71759-4
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-71759-4",
"type": "article-journal",
"title": "CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning",
"container-title": "Nature communications",
"author": [
{
"family": "Liang",
"given": "Zhongming"
},
{
"family": "Zhong",
"given": "Bingxu"
},
{
"family": "Jiao",
"given": "Mingqi"
},
{
"family": "Wang",
"given": "Yong"
},
{
"family": "Liu",
"given": "Shiping"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "5547",
"DOI": "10.1038/s41467-026-71759-4",
"PMID": "42020427",
"PMCID": "PMC13287597",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-71759-4",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
22
]
]
}
}

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, PyTorch Geometric, anndata, 8 other tools, 12 references
[2] 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, anndata, 9 other tools, 6 references
[3] 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: PyTorch Geometric, anndata, Scanpy, 7 other tools, 10 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, PyTorch Geometric, anndata, 8 other tools, mouse, 7 references
[5] doi:10.1038/s42003-026-10462-y [code]
SpaDC enables sequence-based integrative analysis and regulatory inference of spatial chromatin accessibility data.
Journal: Communications biology
In common: Squidpy, anndata, Scanpy, 7 other tools, mouse, 5 references
[6] doi:10.1038/s41467-026-68596-w [code]
Spatial cartography of human thymus enables the geopositioning of lineage transcription factors in rare mimetic thymic epithelial cells.
Journal: Nature communications
In common: Squidpy, anndata, Scanpy, 8 other tools, cellular / molecular, 4 references
[7] doi:10.1186/s12859-026-06490-4 [code]
Tissueformer: extending single-cell foundation models to predict population-level phenotypes.
Journal: BMC bioinformatics
In common: anndata, PyTorch, seaborn, 5 other tools, other condition, mouse, cellular / molecular, 7 references
[8] doi:10.1016/j.isci.2026.116906 [code]
Evaluating exon skipping in the central nervous system in Duchenne muscular dystrophy using spatial transcriptomics.
Journal: iScience
In common: Squidpy, anndata, Scanpy, 5 other tools, other condition, mouse, 6 references
[9] 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, 6 other tools, other condition, mouse, cellular / molecular, 5 references
[10] doi:10.1093/bib/bbag298 [code]
Empowering multifaceted analysis of spatial transcriptomics data with RGAST.
Journal: Briefings in bioinformatics
In common: PyTorch Geometric, anndata, Scanpy, 7 other tools, mouse, 5 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.