CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning.
The 11 matches
- [1] § Methods › Spatial proximity graph construction ↔ cellniche/utils.py, lines 282–354 · score 0.81 · NearestNeighbors, multi slice, spatial graph, undirected, squidpy, COO
- [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] § Methods › Loss function ↔ cellniche/model.py, lines 8–101 · score 0.70 · learnable parameter, temperature parameter, node embeddings, module, matrix, batch
- [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] § Methods › Evaluation ↔ cellniche/utils.py, lines 561–640 · score 0.65 · clustering evaluation metrics, silhouette score, AMI, ARI, Macro
- [6] § Methods › Loss function ↔ cellniche/model.py, lines 8–101 · score 0.59 · log softmax, log probabilities, row, Loss
- [7] § Methods › Training module ↔ cellniche/model.py, lines 104–175 · score 0.57 · graph convolutional, node features, encoder, matrix, training, module
- [8] § Methods › Cellular identity representation ↔ cellniche/utils.py, lines 132–172 · score 0.54 · feature matrix, expression matrix, hot, phenotypic, gene, CellNiche
- [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] § Methods › Training module ↔ cellniche/model.py, lines 104–175 · score 0.53 · node embedding, node features, layer, edge, graph, module
- [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
- import os
- import random
- import numpy as np
- import pandas as pd
- import torch
- import scipy
- import scipy.sparse as sps
- import logging
- import torch.nn.functional as F
- import scanpy as sc
- from natsort import natsorted
- from scipy.sparse import coo_matrix
- from sklearn.preprocessing import LabelEncoder
- from sklearn.neighbors import NearestNeighbors
- from sklearn.cluster import KMeans
- from sklearn.metrics import (
- adjusted_rand_score,
- normalized_mutual_info_score,
- adjusted_mutual_info_score,
- f1_score,
- silhouette_score,
- homogeneity_score,
- v_measure_score,
- fowlkes_mallows_score
- )
- from torch_sparse import SparseTensor
- from torch_geometric.utils import to_undirected
- from typing import Optional, Union, Any, Tuple, List
- from numpy.typing import ArrayLike
- from scipy.stats import entropy
- from sklearn.metrics import silhouette_samples
- def to_float_tensor(arr):
- """Safely convert numpy array / tensor -> float32 tensor w/o grad."""
- if isinstance(arr, torch.Tensor):
- return arr.clone().detach().float()
- else: # numpy / list
- return torch.as_tensor(arr, dtype=torch.float)
- def _resolve_squidpy_connectivity_key(adata, connectivity_key: str = "spatial") -> str:
- """
- Resolve Squidpy connectivity key in adata.obsp.
- Squidpy usually stores the spatial graph as:
- adata.obsp["spatial_connectivities"]
- If users pass connectivity_key="spatial", this function will look for
- "spatial_connectivities". If users directly pass "spatial_connectivities",
- it will also work.
- """
- if connectivity_key is None:
- connectivity_key = "connectivities"
- # Case 1: user directly provides the obsp key
- if connectivity_key in adata.obsp:
- return connectivity_key
- # Case 2: user provides Squidpy key_added prefix, e.g. "spatial"
- candidate = f"{connectivity_key}_connectivities"
- if candidate in adata.obsp:
- return candidate
- raise KeyError(
- f"Cannot find Squidpy spatial connectivity graph in adata.obsp. "
- f"Tried '{connectivity_key}' and '{candidate}'. "
- f"Available obsp keys: {list(adata.obsp.keys())}. "
- f"For multi-slice data, please precompute the graph using "
- f"squidpy.gr.spatial_neighbors(..., library_key=..., key_added='{connectivity_key}')."
- )
- def _edge_index_from_squidpy_obsp(
- adata,
- connectivity_key: str = "spatial",
- ) -> torch.LongTensor:
- """
- Build edge_index from a Squidpy-generated spatial connectivity matrix.
- Expected input:
- adata.obsp["spatial_connectivities"]
- or:
- adata.obsp[f"{connectivity_key}_connectivities"]
- Returns:
- edge_index: torch.LongTensor with shape [2, num_edges]
- """
- obsp_key = _resolve_squidpy_connectivity_key(adata, connectivity_key)
- adj = adata.obsp[obsp_key]
- # Ensure sparse matrix
- if not sps.issparse(adj):
- adj = sps.csr_matrix(adj)
- # Convert to boolean adjacency
- adj = adj.astype(bool)
- # Convert to LIL for efficient diagonal modification
- if isinstance(adj, sps.csr_matrix):
- adj = adj.tolil()
- else:
- adj = adj.tolil()
- # Remove self-loops
- adj.setdiag(0)
- # Back to CSR and remove explicit zeros
- adj = adj.tocsr()
- adj.eliminate_zeros()
- # Extract nonzero indices
- row, col = adj.nonzero()
- edge_index = torch.tensor(
- np.array([row, col]),
- dtype=torch.long,
- )
- logging.info(
- f"Using precomputed Squidpy graph from adata.obsp['{obsp_key}']: "
- f"{adata.n_obs} nodes, {edge_index.shape[1]} edges."
- )
- return edge_index
- def load_data(
- data_path: str,
- dataset: str,
- phenoLabels: str,
- nicheLabels: Optional[str],
- embedding_type: str,
- radius: Optional[float],
- k_neighborhood: int,
- hvg: bool,
- n_hvg: int,
- multi_slice: bool = False,
- connectivity_key: str = "spatial",
- embedding_key: Optional[str] = None,
- ) -> Tuple[torch.FloatTensor, torch.LongTensor, np.ndarray, Any, int, Optional[torch.FloatTensor]]:
- """
- Load AnnData from .h5ad, build node features, edge index, and labels.
- Args:
- data_path: Directory containing the .h5ad files.
- dataset: Filename (without extension) to load.
- phenoLabels: Column name in `adata.obs` for phenotype labels (one-hot).
- nicheLabels: Column name in `adata.obs` for true labels (or None to skip).
- embedding_type: One of 'pheno_expr', 'pheno', 'expr'.
- radius: Radius threshold (if not None) for radius graph.
- k_neighborhood: Number of neighbors for kNN graph if radius is None.
- hvg: Whether to select highly variable genes for expression.
- multi_slice: bool = False,
- connectivity_key: str = "spatial",
- embedding_key: Optional[str] = None,
- Returns:
- x: Node feature matrix (one-hot or expr) as FloatTensor.
- edge_index: LongTensor[2, E] of graph edges.
- y: True label array of shape [N].
- adata: The AnnData object.
- n_classes: Number of unique labels in y.
- expr: Expression matrix as FloatTensor if needed, else None.
- """
- # 1) Read AnnData
- path = os.path.join(data_path, f"{dataset}.h5ad")
- adata = sc.read_h5ad(path).copy()
- # 2) Build phenotype one-hot features
- valid_embedding_types = ["pheno_expr", "pheno", "expr", "embedding"]
- if embedding_type not in valid_embedding_types:
- raise ValueError(
- f"Unknown embedding_type: {embedding_type}. "
- f"Expected one of {valid_embedding_types}."
- )
- # 3) Build phenotype one-hot features only when needed
- onehot = None
- if embedding_type in ["pheno", "pheno_expr"]:
- if phenoLabels is None:
- raise ValueError(
- "phenoLabels is required when embedding_type is "
- f"'{embedding_type}', but got phenoLabels=None."
- )
- if phenoLabels not in adata.obs:
- raise KeyError(
- f"phenoLabels='{phenoLabels}' was not found in adata.obs. "
- f"Available obs columns: {list(adata.obs.columns)}"
- )
- pheno = adata.obs[phenoLabels].astype(str)
- ph_le = LabelEncoder().fit(pheno)
- ph_idx = ph_le.transform(pheno)
- onehot = torch.nn.functional.one_hot(
- torch.from_numpy(ph_idx),
- num_classes=len(ph_le.classes_)
- ).float()
- # 4) Prepare expression matrix if needed
- expr = None
- if embedding_type in ["pheno_expr", "expr"]:
- temp_adata = adata.copy()
- if hvg:
- # # way 1
- # logging.info(f"hvg way1")
- # sc.pp.highly_variable_genes(temp_adata, n_top_genes=n_hvg, flavor='seurat_v3')
- # temp_adata = temp_adata[:,temp_adata.var.highly_variable]
- # temp_adata.raw=adata
- # sc.pp.normalize_total(temp_adata, target_sum=1e4,inplace=True)
- # sc.pp.log1p(temp_adata)
- # sc.pp.scale(temp_adata)
- # way 2
- sc.pp.highly_variable_genes(
- temp_adata,
- flavor="seurat_v3",
- n_top_genes=n_hvg,
- )
- sc.pp.normalize_total(temp_adata, target_sum=1e4)
- sc.pp.log1p(temp_adata)
- temp_adata = temp_adata[:, temp_adata.var["highly_variable"]].copy()
- mat = temp_adata.X
- if scipy.sparse.isspmatrix(mat):
- arr = mat.toarray()
- else:
- arr = np.asarray(mat)
- expr = torch.from_numpy(arr).float()
- del temp_adata
- # 5) For embedding_type="embedding", load features from adata.obsm[embedding_key]
- embedding = None
- if embedding_type == "embedding":
- if embedding_key is None:
- raise ValueError(
- "embedding_key is required when embedding_type='embedding'. "
- "Please provide the key of the embedding stored in adata.obsm, "
- "for example embedding_key='X_scVI' or embedding_key='X_pca'."
- )
- if embedding_key not in adata.obsm:
- raise KeyError(
- f"embedding_key='{embedding_key}' was not found in adata.obsm. "
- f"Available obsm keys: {list(adata.obsm.keys())}"
- )
- emb_arr = adata.obsm[embedding_key]
- if scipy.sparse.isspmatrix(emb_arr):
- emb_arr = emb_arr.toarray()
- else:
- emb_arr = np.asarray(emb_arr)
- if emb_arr.ndim != 2:
- raise ValueError(
- f"adata.obsm['{embedding_key}'] must be a 2D matrix, "
- f"but got shape {emb_arr.shape}."
- )
- if emb_arr.shape[0] != adata.n_obs:
- raise ValueError(
- f"adata.obsm['{embedding_key}'] has {emb_arr.shape[0]} rows, "
- f"but adata has {adata.n_obs} cells."
- )
- embedding = torch.from_numpy(emb_arr).float()
- # 6) Build graph edge_index
- if multi_slice:
- logging.warning(
- "multi_slice=True detected. "
- "For multi-slice or multi-sample data, CellNiche expects a precomputed "
- "sample-aware spatial graph generated by Squidpy"
- )
- edge_index = _edge_index_from_squidpy_obsp(
- adata,
- connectivity_key=connectivity_key,
- )
- else:
- # Original single-slice logic
- if "edgeList" in adata.uns:
- edge_np = np.array(adata.uns["edgeList"])
- # Support both shape [2, E] and [E, 2]
- if edge_np.ndim != 2:
- raise ValueError("adata.uns['edgeList'] must be a 2D array.")
- if edge_np.shape[0] == 2:
- edge_index = torch.from_numpy(edge_np).long()
- elif edge_np.shape[1] == 2:
- edge_index = torch.from_numpy(edge_np.T).long()
- else:
- raise ValueError(
- "adata.uns['edgeList'] should have shape [2, E] or [E, 2]."
- )
- edge_index = to_undirected(edge_index)
- else:
- # choose coords
- if "spatial" in adata.obsm and adata.obsm["spatial"] is not None:
- coords = adata.obsm["spatial"]
- else:
- coords = adata.obs[["x", "y"]].to_numpy()
- if radius is not None:
- nbrs = NearestNeighbors(radius=radius).fit(coords)
- _, idxs = nbrs.radius_neighbors(coords)
- rows_list, cols_list = [], []
- for i, neighbors in enumerate(idxs):
- # remove self-loop
- neighbors = neighbors[neighbors != i]
- rows_list.extend([i] * len(neighbors))
- cols_list.extend(neighbors.tolist())
- rows = np.asarray(rows_list)
- cols = np.asarray(cols_list)
- else:
- nbrs = NearestNeighbors(n_neighbors=k_neighborhood + 1).fit(coords)
- _, idxs = nbrs.kneighbors(coords)
- rows = np.repeat(np.arange(coords.shape[0]), k_neighborhood)
- cols = idxs[:, 1:].flatten()
- mat = coo_matrix(
- (np.ones_like(rows), (rows, cols)),
- shape=(coords.shape[0], coords.shape[0]),
- )
- mat = mat + mat.T # make undirected
- edge_index = torch.from_numpy(
- np.vstack(mat.nonzero()).astype(np.int64)
- )
- neighbors_count = np.array([len(neighbors) for neighbors in idxs])
- average_neighbors = neighbors_count.mean()
- logging.info(f"Average number of neighbors per node: {average_neighbors}")
- # 7) Encode true labels from nicheLabels if provided
- if nicheLabels is not None and nicheLabels in adata.obs:
- true_vals = adata.obs[nicheLabels].astype(str)
- nl_encoder = LabelEncoder().fit(true_vals)
- y = nl_encoder.transform(true_vals)
- n_classes = len(nl_encoder.classes_)
- else:
- # default dummy labels: all-zero, one class
- y = np.zeros(adata.n_obs, dtype=int)
- n_classes = 1
- # 8) Return according to embedding type
- if embedding_type == "pheno_expr":
- logging.info(
- f"Loaded {dataset}: "
- f"{onehot.shape[0]} nodes, {edge_index.shape[1]} edges, "
- f"{onehot.shape[-1]} phenotype features, "
- f"{expr.shape[-1]} expression features"
- )
- return to_float_tensor(onehot), edge_index, y, adata, n_classes, to_float_tensor(expr)
- elif embedding_type == "pheno":
- logging.info(
- f"Loaded {dataset}: "
- f"{onehot.shape[0]} nodes, {edge_index.shape[1]} edges, "
- f"{onehot.shape[-1]} phenotype features"
- )
- return to_float_tensor(onehot), edge_index, y, adata, n_classes, None
- elif embedding_type == "expr":
- logging.info(
- f"Loaded {dataset}: "
- f"{expr.shape[0]} nodes, {edge_index.shape[1]} edges, "
- f"{expr.shape[-1]} expression features"
- )
- return to_float_tensor(expr), edge_index, y, adata, n_classes, None
- elif embedding_type == "embedding":
- logging.info(
- f"Loaded {dataset}: "
- f"{embedding.shape[0]} nodes, {edge_index.shape[1]} edges, "
- f"{embedding.shape[-1]} embedding features from adata.obsm['{embedding_key}']"
- )
- return to_float_tensor(embedding), edge_index, y, adata, n_classes, None
- def setup_seed(seed: int) -> None:
- """
- Set random seed for reproducibility across Python, NumPy, Torch, and CUDA.
- Args:
- seed (int): The seed to set.
- """
- random.seed(seed)
- np.random.seed(seed)
- torch.manual_seed(seed)
- torch.cuda.manual_seed_all(seed)
- # torch.backends.cudnn.deterministic = True
- # torch.backends.cudnn.benchmark = False
- def create_sparse_tensor_from_edges(
- rows: list[int],
- cols: list[int],
- sparse_size: tuple[int, int],
- device: torch.device = torch.device("cpu"),
- ) -> SparseTensor:
- """
- Create a SparseTensor from row/col indices.
- Args:
- rows (list[int]): Row indices.
- cols (list[int]): Column indices.
- sparse_size (tuple[int,int]): Matrix size.
- device (torch.device): Device for tensor.
- Returns:
- SparseTensor
- """
- vals = torch.ones(len(rows), device=device)
- return SparseTensor(
- row=torch.tensor(rows, device=device),
- col=torch.tensor(cols, device=device),
- value=vals,
- sparse_sizes=sparse_size,
- )
- def sparse_intersection_and_union(
- adj1: SparseTensor,
- adj2: SparseTensor,
- strategy: str = "and",
- ) -> SparseTensor:
- """
- Compute intersection or union of two sparse adjacency matrices.
- Args:
- adj1, adj2 (SparseTensor): Input graphs.
- strategy (str): 'and' or 'or'.
- Returns:
- SparseTensor
- """
- rows1, cols1 = adj1.storage.row(), adj1.storage.col()
- rows2, cols2 = adj2.storage.row(), adj2.storage.col()
- device = rows1.device
- set1 = set(zip(rows1.tolist(), cols1.tolist()))
- set2 = set(zip(rows2.tolist(), cols2.tolist()))
- if strategy == "and":
- common = set1 & set2
- else:
- common = set1 | set2
- if not common:
- return SparseTensor(sparse_sizes=adj1.sparse_sizes())
- rows, cols = zip(*common)
- return create_sparse_tensor_from_edges(rows, cols, adj1.sparse_sizes(), device)
- def get_positivePairs(
- subAdj: SparseTensor,
- features: Optional[torch.Tensor] = None,
- strategy: str = "freq",
- ) -> SparseTensor:
- """
- Generate positive pair adjacency based on strategy.
- Args:
- subAdj (SparseTensor): Stochastic subgraph adjacency.
- features (Tensor): Node features.
- strategy (str): 'freq', 'sim', 'and', 'or'.
- Returns:
- SparseTensor: positive adjacency.
- """
- row, col, val = subAdj.storage.row(), subAdj.storage.col(), subAdj.storage.value()
- # Frequency-based mask
- freq_thresh = subAdj.sum(dim=1) / subAdj.storage.colptr()[1:]
- mask = val > freq_thresh[row]
- rows, cols, vals = row[mask], col[mask], val[mask]
- freq_adj = SparseTensor(row=rows, col=cols, value=vals, sparse_sizes=subAdj.sparse_sizes())
- if strategy == "freq":
- return freq_adj
- if features is None:
- raise ValueError("Features required for non-freq strategy.")
- # Similarity-based mask
- f_row, f_col = features[row], features[col]
- sim_vals = F.cosine_similarity(f_row, f_col, dim=1)
- sim_vals[row == col] = 0
- sim_adj = SparseTensor(row=row, col=col, value=sim_vals, sparse_sizes=subAdj.sparse_sizes())
- sim_thresh = sim_adj.sum(dim=1) / sim_adj.storage.colptr()[1:]
- sim_mask = sim_vals > sim_thresh[row]
- sim_rows, sim_cols, sim_vals = row[sim_mask], col[sim_mask], sim_vals[sim_mask]
- sim_based = SparseTensor(row=sim_rows, col=sim_cols, value=sim_vals, sparse_sizes=subAdj.sparse_sizes())
- if strategy == "sim":
- return sim_based
- # AND / OR combination
- return sparse_intersection_and_union(freq_adj, sim_based, strategy)
- def match_labels(true_labels, predicted_labels, n_classes):
- from scipy.optimize import linear_sum_assignment as linear_assignment
- cost_matrix = np.zeros((n_classes, n_classes))
- for i in range(n_classes):
- for j in range(n_classes):
- cost_matrix[i, j] = np.sum((true_labels == i) & (predicted_labels == j))
- row_ind, col_ind = linear_assignment(-cost_matrix)
- new_labels = np.copy(predicted_labels)
- for i, j in zip(row_ind, col_ind):
- new_labels[predicted_labels == j] = i
- return new_labels
- def refine_spatial_domains(y_pred, coord, n_neighbors=6):
- nbrs = NearestNeighbors(n_neighbors=n_neighbors + 1).fit(coord)
- distances, indices = nbrs.kneighbors(coord)
- indices = indices[:, 1:]
- y_refined = pd.Series(index=y_pred.index, dtype='object')
- for i in range(y_pred.shape[0]):
- y_pred_count = y_pred[indices[i, :]].value_counts()
- if y_pred[i] in y_pred_count.index:
- if (y_pred_count.loc[y_pred[i]] < n_neighbors / 2) and (y_pred_count.max() > n_neighbors / 2):
- y_refined[i] = y_pred_count.idxmax()
- else:
- # y_refined[i] = y_pred[i] # waring
- y_refined.iloc[i] = y_pred[i]
- else:
- y_refined.iloc[i] = y_pred[i]
- y_refined = pd.Categorical(
- values=y_refined.astype('U'),
- categories=natsorted(map(str, y_refined.unique())),
- )
- return y_refined
- def clustering_st(
- adata: Any,
- n_clusters: int,
- features: Optional[Union[torch.Tensor, np.ndarray]] = None,
- true_labels: Optional[np.ndarray] = None,
- refine: bool = False,
- ) -> Tuple[Any, dict]:
- """
- Perform KMeans clustering and compute evaluation metrics.
- Args:
- adata (AnnData): Annotated data object.
- features (Tensor or ndarray): Embeddings.
- n_clusters (int): Number of clusters.
- true_labels (ndarray): Ground-truth labels.
- Returns:
- adata (AnnData): Updated with 'kmeans' clusters.
- metrics (dict): Cluster evaluation metrics.
- """
- # 1) Convert to numpy
- if torch.is_tensor(features):
- feats = features.cpu().numpy()
- else:
- feats = features
- # 2) Run KMeans
- km = KMeans(n_clusters=n_clusters, max_iter=5000, n_init=10)
- # km = KMeans(n_clusters=n_clusters, max_iter=10000, n_init=20)
- raw_labels = km.fit_predict(feats).astype(int)
- # 3) Store raw labels
- adata.obs['kmeans'] = pd.Categorical(raw_labels)
- clustering_results = {'kmeans': raw_labels}
- # 4) Optional spatial refinement
- if refine:
- # spatial coords must exist
- coords = adata.obsm.get('spatial')
- if coords is None:
- raise ValueError("adata.obsm['spatial'] needed for refinement")
- for method, labels in list(clustering_results.items()):
- refined = refine_spatial_domains(pd.Series(labels), coords)
- refined = refined.astype(int)
- col = f"{method}_refined"
- adata.obs[col] = pd.Categorical(refined)
- clustering_results[col] = refined
- # 5) Compute metrics
- metrics_results: dict = {}
- if true_labels is not None:
- for method, labels in clustering_results.items():
- # align predicted → true
- aligned = match_labels(true_labels, labels, n_clusters)
- acc = (aligned == true_labels).mean()
- nmi = normalized_mutual_info_score(true_labels, aligned)
- ari = adjusted_rand_score(true_labels, aligned)
- ami = adjusted_mutual_info_score(true_labels, aligned)
- f1m = f1_score(true_labels, aligned, average='macro')
- f1i = f1_score(true_labels, aligned, average='micro')
- sil = silhouette_score(feats, true_labels)
- FMI = fowlkes_mallows_score(true_labels, aligned)
- v_measure = v_measure_score(true_labels, aligned)
- homogeneity = homogeneity_score(true_labels, aligned)
- metrics_results[method] = {
- 'Acc': acc,
- 'NMI': nmi,
- 'AMI': ami,
- 'ARI': ari,
- 'F1 Macro': f1m,
- 'F1 Micro': f1i,
- 'Silhouette': sil,
- 'Fowlkes-Mallows': FMI,
- 'V-Measure': v_measure,
- 'Homogeneity': homogeneity
- }
- return adata, metrics_results
- def _rng(random_state: Optional[int] = None) -> np.random.Generator:
- """Return a NumPy Generator with the requested seed (or global RNG)."""
- return np.random.default_rng(random_state)
- def _nearest_neighbors(
- x: np.ndarray, k: int, **kwargs
- ) -> np.ndarray:
- """
- Return indices of the *k* nearest neighbours for every point in *x*.
- The first neighbour returned by ``sklearn`` is the query point itself,
- so we discard it.
- """
- nn = NearestNeighbors(n_neighbors=k + 1, **kwargs).fit(x)
- indices = nn.kneighbors(x, return_distance=False)[:, 1:] # drop self‑index
- return indices
- def _encode_labels(labels: ArrayLike) -> np.ndarray:
- """
- Map arbitrary label values to consecutive integers starting from 0.
- This simplifies downstream use of ``np.bincount`` and avoids
- large sparse counts when label values are not contiguous.
- """
- labels = np.asarray(labels)
- _, encoded = np.unique(labels, return_inverse=True)
- return encoded
- # ---------------------------------------------------------------------
- # 1. Entropy of Batch Mixing
- # ---------------------------------------------------------------------
- def compute_entropy_batch_mixing(
- embeddings: np.ndarray,
- batch_labels: ArrayLike,
- k: int = 50,
- normalize: bool = True,
- **nn_kwargs,
- ) -> float:
- """
- Average entropy of batch labels in the *k*‑NN neighbourhood of each cell.
- Parameters
- ----------
- embeddings
- Low‑dimensional representation of shape *(n_cells, n_dims)*.
- batch_labels
- Iterable of length *n_cells* with one label per cell.
- k
- Number of neighbours (*excluding* the query cell) to consider.
- normalize
- If ``True`` (default) divide by the maximal entropy
- ``log(n_batches)``, yielding values in ``[0, 1]``.
- **nn_kwargs
- Additional arguments forwarded to :class:`sklearn.neighbors.NearestNeighbors`.
- Returns
- -------
- float
- Mean entropy across all cells.
- """
- if k < 1:
- raise ValueError("k must be ≥ 1")
- batch_labels = _encode_labels(batch_labels)
- indices = _nearest_neighbors(embeddings, k=k, **nn_kwargs)
- n_batches = int(batch_labels.max()) + 1
- max_ent = np.log(n_batches) if normalize else 1.0
- entropies = []
- for nbr in indices:
- counts = np.bincount(batch_labels[nbr], minlength=n_batches)
- probs = counts / k # guaranteed non‑negative, summing to 1
- ent = entropy(probs) / max_ent if max_ent > 0 else 0.0
- entropies.append(ent)
- return float(np.mean(entropies))
- # ---------------------------------------------------------------------
- # 2. iLISI
- # ---------------------------------------------------------------------
- def compute_ilisi(
- embeddings: np.ndarray,
- batch_labels: ArrayLike,
- k: int = 90,
- **nn_kwargs,
- ) -> np.ndarray:
- """
- Compute the *inverse* Local Inverse Simpson’s Index (iLISI).
- For each cell *i*::
- iLISI_i = 1 − (# neighbours from same batch) / k
- Thus, 0 indicates perfect batch isolation; 1 indicates perfect mixing.
- Parameters
- ----------
- embeddings
- *(n_cells, n_dims)* array.
- batch_labels
- Iterable with one batch label per cell.
- k
- Number of neighbours (*excluding* the query cell) to consider.
- **nn_kwargs
- Extra arguments for :class:`sklearn.neighbors.NearestNeighbors`.
- Returns
- -------
- np.ndarray
- Vector of length *n_cells* with iLISI scores.
- """
- batch_labels = _encode_labels(batch_labels)
- indices = _nearest_neighbors(embeddings, k=k, **nn_kwargs)
- same_batch = (batch_labels[indices] == batch_labels[:, None]).sum(axis=1)
- ilisi = 1.0 - (same_batch / k)
- return ilisi
- # ---------------------------------------------------------------------
- # 3. Seurat Alignment Score (SAS)
- # ---------------------------------------------------------------------
- def compute_seurat_alignment_score(
- embeddings: np.ndarray,
- batch_labels: ArrayLike,
- neighbor_frac: float = 0.01,
- n_repeats: int = 3,
- random_state: Optional[int] = None,
- **nn_kwargs,
- ) -> float:
- """
- Seurat Alignment Score (Butler *et al.*, Cell 2018).
- The score estimates how well batches mix after integration by repeatedly
- down‑sampling to equal batch sizes and measuring the proportion of
- cross‑batch neighbours.
- Parameters
- ----------
- embeddings
- *(n_cells, n_dims)* array.
- batch_labels
- Iterable with one batch label per cell.
- neighbor_frac
- Fraction of the (sub‑sampled) cells to use as *k* in *k*‑NN.
- Must be in ``(0, 1]``.
- n_repeats
- Number of random sub‑samples to average.
- random_state
- Seed for reproducibility.
- **nn_kwargs
- Extra arguments for :class:`sklearn.neighbors.NearestNeighbors`.
- Returns
- -------
- float
- Mean SAS across repeats (1 = perfect mixing, 0 = no mixing).
- """
- if not (0 < neighbor_frac <= 1):
- raise ValueError("neighbor_frac must be in (0, 1]")
- rng = _rng(random_state)
- batch_labels = _encode_labels(batch_labels)
- batch_indices = [np.where(batch_labels == b)[0] for b in np.unique(batch_labels)]
- min_size = min(len(idx) for idx in batch_indices)
- n_batches = len(batch_indices)
- scores = []
- for _ in range(n_repeats):
- # balanced subsample
- sel = np.concatenate([rng.choice(idx, min_size, replace=False) for idx in batch_indices])
- x_sub, y_sub = embeddings[sel], batch_labels[sel]
- k = max(int(round(len(sel) * neighbor_frac)), 1)
- indices = _nearest_neighbors(x_sub, k=k, **nn_kwargs)
- same_batch = (y_sub[indices] == y_sub[:, None]).sum(axis=1).mean()
- score = (k - same_batch) * n_batches / (k * (n_batches - 1))
- scores.append(min(score, 1.0)) # numerical guard
- return float(np.mean(scores))
- # ---------------------------------------------------------------------
- # 4. ASW‑batch
- # ---------------------------------------------------------------------
- def compute_avg_silhouette_width_batch(
- embeddings: np.ndarray,
- batch_labels: ArrayLike,
- cell_types: ArrayLike,
- min_cells: int = 3,
- **silhouette_kwargs,
- ) -> float:
- """
- Average Silhouette Width computed per cell‑type, then averaged.
- Parameters
- ----------
- ...
- min_cells
- Minimum number of cells a cell‑type must have to be included.
- """
- x = embeddings
- y = _encode_labels(batch_labels)
- ct = _encode_labels(cell_types)
- scores = []
- for t in np.unique(ct):
- mask = ct == t
- n = mask.sum()
- n_labels = np.unique(y[mask]).size
- # Skip if too few cells OR labels≈cells (invalid for silhouette)
- if n < min_cells or n_labels < 2 or n_labels >= n:
- scores.append(0.0)
- continue
- try:
- s = silhouette_samples(x[mask], y[mask], **silhouette_kwargs)
- scores.append((1.0 - np.abs(s)).mean())
- except ValueError:
- scores.append(0.0)
- return float(np.mean(scores))
utils.py at commit af58974, under MIT · at the source
Overview
- 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
- State Key Laboratory of Genome and Multi-omics Technologies, BGI Research,Hangzhou, China
- Key Laboratory of Spatial Omics of Zhejiang Province, BGI Research,Hangzhou, China
- 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
- State Key Laboratory of Mathematical Sciences, Academy of Mathematics and Systems Science, Chinese Academy of Sciences,Beijing, China
- School of Mathematics, University of Chinese Academy of Sciences, Chinese Academy of Sciences,Beijing, China
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
Availability: 1 check, the latest on 29 September 2026: the link is dead
- 29 September 2026: the link is dead
Super-LzzZ/CellNiche
af58974ded7cf57299a9f8952d4cc6dffee39c6f, 4 May 2026Availability: 1 check, the latest on 29 September 2026: the link answers
- 29 September 2026: the link answers
12 files
- cellniche/
__init__.py , Python, 61 lines, 1 match - cellniche/
main.py , Python, 44 lines - cellniche/
model.py , Python, 294 lines, 4 matches - cellniche/
sampler.py , Python, 271 lines - cellniche/
trainer.py , Python, 218 lines - cellniche/
utils.py , Python, 871 lines, 6 matches - tutorial/
CosMxMouseBrain.ipynb , Jupyter, 85 lines - tutorial/
NSCLC.ipynb , Jupyter, 127 lines - tutorial/
brain_STARmap.ipynb , Jupyter, 135 lines - tutorial/
cortex.ipynb , Jupyter, 137 lines - repository limit reached (2,000 files or 30 MB): the rest is at the source (7 files)
- LICENSE, License, 22 lines
- README.md, Text, 105 lines
Zenodo 19143524
Availability: 1 check, the latest on 29 September 2026: the link answers (HTTP 200)
- 29 September 2026: the link answers (HTTP 200)
Code availability
The software package implementing the CellNiche algorithm has been deposited at GitHub https://
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.mendeley.com/
datasets/ , at Mendeley Data; found in “Data availability”mpjzbtfgfr - data.mendeley.com/
datasets/ , at Mendeley Data; found in “Data availability”zjnpwh8m5b - github.com/
drieslab/ , at github.com; found in “Data availability”spatial-datasets - github.com/
hubioinfo/ , at github.com; found in the text, “Evaluation”cytocommunity
Data availability
The osmFISH dataset of mouse somatosensory cortex is available at https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 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://
BibTeX
@article{liang2026cellni
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/
url = {https://
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/
VL - 17
IS - 1
SP - 5547
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"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":
"volume": "17",
"issue": "1",
"page": "5547",
"DOI": "10.1038/
"PMID": "42020427",
"PMCID": "PMC13287597",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://
"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 methodsIn 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 researchIn 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 biologyIn 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 communicationsIn 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 bioinformaticsIn 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: iScienceIn 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 bioinformaticsIn 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 bioinformaticsIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 3 repositories of the authors' code, each at its verified commit and with its license, 10 scripts, and 11 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:b72c40cb4ca9bf13…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
