OSCR

SlotDeconv: spatial transcriptomics deconvolution via diversity-constrained prototype learning and spatial refinement.

Code ↔ Paper

4 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 4 matches
  1. [1] § 2 Materials and methods › 2.3 Single-cell reference module › 2.3.4 Reference module optimization ↔ model/slot_model.py, lines 377–411 · score 0.76 · AdamW, weight decay, cosine annealing, clipped, training
  2. [2] § 2 Materials and methods › 2.5 Implementation details ↔ model/slot_model.py, lines 31–75 · score 0.76 · LayerNorm, ReLU, PyTorch, decoder, dropout, hidden
  3. [3] § 2 Materials and methods › 2.6 Evaluation metrics ↔ model/slot_utility.py, lines 23–37 · score 0.59 · Jensen Shannon, scipy, JSD, distance, Spot
  4. [4] § 3 Results › 3.1 Benchmarking on the mouse brain dataset ↔ model/visz.py, lines 53–159 · score 0.53 · clustering quality, Silhouette, ARI, PCA, embeddings, prototype

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 · 477 lines · 23 KB · MIT · 2 matches

  1. """
  2. SlotDeconv: Spatial Transcriptomics Deconvolution via Slot-based Reference Learning
  3. """
  4. import os
  5. import warnings
  6. import numpy as np
  7. import pandas as pd
  8. import torch
  9. import torch.nn as nn
  10. import torch.nn.functional as F
  11. import torch.optim as optim
  12. from scipy.spatial.distance import cdist
  13. from scipy.optimize import nnls
  14. warnings.filterwarnings("ignore")
  15. DEFAULT_CONFIG = {
  16. 'n_genes': 2500, # 3000
  17. 'b_epochs': 2500, #2000
  18. 'max_cells_per_type': 750,
  19. 'lambda_div': 5.0, # 4.0
  20. 'margin': 0.1,
  21. 'lambda_sp': 10.0, # 15.0
  22. 'sp_epochs': 500, # 1500
  23. 'sp_lr': 0.01,
  24. 'b_lr': 1e-3,
  25. 'pow_w': 0.8,
  26. 'knn': 15,
  27. 'd_slot':128,
  28. 'dec_hidden':(256,512),
  29. 'dec_dropout':0.1
  30. }
  31. class _SlotDecoder(nn.Module):
  32. def __init__(self, n_genes, n_types, d_slot=128, margin=0.1,hidden=(256,512),dropout=0.1):
  33. super().__init__()
  34. self.n_types = n_types
  35. self.n_genes = n_genes
  36. self.margin = margin
  37. self.slots = nn.Parameter(torch.randn(n_types, d_slot) * 0.1)
  38. dims=[d_slot]+list(hidden)+[n_genes]
  39. H=max(len(dims)-2,0)
  40. if isinstance(dropout,(list,tuple,np.ndarray)):
  41. drops=[float(x) for x in list(dropout)]
  42. if len(drops)<H: drops=drops+[0.0]*(H-len(drops))
  43. drops=drops[:H]
  44. else:
  45. d=float(dropout) if dropout is not None else 0.0
  46. drops=([d]+[0.0]*(H-1)) if H>0 else []
  47. layers=[]
  48. for i in range(H):
  49. layers.append(nn.Linear(dims[i],dims[i+1]))
  50. layers.append(nn.LayerNorm(dims[i+1]))
  51. layers.append(nn.ReLU())
  52. if drops[i]>0: layers.append(nn.Dropout(drops[i]))
  53. layers.append(nn.Linear(dims[-2],dims[-1]))
  54. self.decoder=nn.Sequential(*layers)
  55. self.alpha_raw = nn.Parameter(torch.log(torch.expm1(torch.ones(n_genes) * 5.0)))
  56. def alpha_disp(self):
  57. return F.softplus(self.alpha_raw) + 1e-6
  58. def forward(self, labels, size):
  59. logits = torch.clamp(self.decoder(self.slots[labels]), -20, 20)
  60. return (F.softplus(logits) * size).clamp(1e-8, 1e8)
  61. def get_reference_matrix(self):
  62. with torch.no_grad():
  63. logits = torch.clamp(self.decoder(self.slots), -20, 20)
  64. B = F.softplus(logits)
  65. B = B / (B.sum(dim=1, keepdim=True) + 1e-8)
  66. return B.cpu().numpy()
  67. def diversity_loss(self):
  68. logits = torch.clamp(self.decoder(self.slots), -20, 20)
  69. B = F.softplus(logits)
  70. B = B / (B.sum(dim=1, keepdim=True) + 1e-8)
  71. Bn = F.normalize(B, dim=-1)
  72. sim = Bn @ Bn.t()
  73. mask = torch.eye(self.n_types, device=sim.device)
  74. off_diag = sim * (1 - mask)
  75. return F.relu(off_diag - self.margin).max(), off_diag.max().item()
  76. def _nb_nll(y, mu, theta):
  77. mu = mu.clamp(1e-8, 1e8)
  78. theta = theta.clamp(1e-4, 1e4)
  79. t1 = torch.lgamma(y + theta) - torch.lgamma(theta) - torch.lgamma(y + 1.0)
  80. t2 = theta * torch.log(theta / (theta + mu) + 1e-12)
  81. t3 = y * torch.log(mu / (mu + theta) + 1e-12)
  82. return -(t1 + t2 + t3).mean()
  83. def set_seed(seed=42):
  84. import random
  85. os.environ["PYTHONHASHSEED"] = str(seed)
  86. random.seed(seed)
  87. np.random.seed(seed)
  88. torch.manual_seed(seed)
  89. torch.cuda.manual_seed_all(seed)
  90. torch.backends.cudnn.deterministic = True
  91. torch.backends.cudnn.benchmark = False
  92. def build_knn_graph(coords, knn=15):
  93. """Build KNN graph with RBF weights. O(NK) memory instead of O(N²)."""
  94. n = len(coords)
  95. knn = min(knn, n - 1)
  96. try:
  97. from sklearn.neighbors import NearestNeighbors
  98. nbrs = NearestNeighbors(n_neighbors=knn + 1, algorithm="auto").fit(coords)
  99. dist, idx = nbrs.kneighbors(coords)
  100. idx = idx[:, 1:]
  101. dist = dist[:, 1:]
  102. except Exception:
  103. D = cdist(coords, coords)
  104. idx = np.argsort(D, axis=1)[:, 1:knn + 1]
  105. dist = np.take_along_axis(D, idx, axis=1)
  106. sigma = np.median(dist[dist > 0]) + 1e-8
  107. w = np.exp(-(dist ** 2) / (2 * sigma ** 2)).astype(np.float32)
  108. w = w / (w.sum(1, keepdims=True) + 1e-8)
  109. return idx.astype(np.int64), w.astype(np.float32)
  110. def build_dense_graph(coords_norm):
  111. """Build dense spatial graph. Expects already normalized coords."""
  112. dist = cdist(coords_norm, coords_norm)
  113. sigma = np.median(dist[dist > 0])
  114. W = np.exp(-dist ** 2 / (2 * sigma ** 2))
  115. np.fill_diagonal(W, 0)
  116. W = W / (W.sum(axis=1, keepdims=True) + 1e-8)
  117. return W.astype(np.float32)
  118. class SlotDeconv:
  119. """
  120. SlotDeconv: Spatial-aware deconvolution with slot-based reference learning.
  121. Parameters
  122. ----------
  123. device : str, optional
  124. Device for computation. Default: auto-detect
  125. random_state : int, optional
  126. Random seed. Default: 42
  127. verbose : bool, optional
  128. Print progress. Default: True
  129. use_default_config : bool, optional
  130. Use optimized parameters from default configuration. Default: True
  131. """
  132. def __init__(self, device=None, random_state=42, verbose=True, use_default_config=True):
  133. if device is None:
  134. self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
  135. else:
  136. self.device = torch.device(device)
  137. self.random_state = random_state
  138. self.verbose = verbose
  139. self.use_default_config = use_default_config
  140. self.B_prob_ = None
  141. self.cell_types_ = None
  142. self.selected_genes_ = None
  143. self.gene_weights_ = None
  144. self._gene2idx = None
  145. self._last_fit_params = None
  146. self._last_transform_params = None
  147. set_seed(random_state)
  148. def _log(self, msg):
  149. if self.verbose:
  150. print(msg)
  151. def get_config(self):
  152. """Return current configuration including last used parameters."""
  153. d = dict(DEFAULT_CONFIG)
  154. d["use_default_config"] = bool(self.use_default_config)
  155. d["last_fit"] = self._last_fit_params
  156. d["last_transform"] = self._last_transform_params
  157. return d
  158. def fit(self, sc_count, sc_meta, cell_types, celltype_col='cellType',
  159. n_genes=None, max_cells_per_type=None, lambda_div=None, margin=None,
  160. b_epochs=None, b_lr=None,d_slot=None,dec_hidden=None,dec_dropout=None):
  161. """Learn reference matrix B from single-cell RNA-seq data."""
  162. set_seed(self.random_state)
  163. self.cell_types_ = list(cell_types)
  164. n_types = len(cell_types)
  165. if self.use_default_config:
  166. n_genes = n_genes if n_genes is not None else DEFAULT_CONFIG['n_genes']
  167. max_cells_per_type = max_cells_per_type if max_cells_per_type is not None else DEFAULT_CONFIG['max_cells_per_type']
  168. lambda_div = lambda_div if lambda_div is not None else DEFAULT_CONFIG['lambda_div']
  169. margin = margin if margin is not None else DEFAULT_CONFIG['margin']
  170. b_epochs = b_epochs if b_epochs is not None else DEFAULT_CONFIG['b_epochs']
  171. b_lr = b_lr if b_lr is not None else DEFAULT_CONFIG['b_lr']
  172. d_slot=d_slot if d_slot is not None else DEFAULT_CONFIG['d_slot']
  173. dec_hidden=dec_hidden if dec_hidden is not None else DEFAULT_CONFIG['dec_hidden']
  174. dec_dropout=dec_dropout if dec_dropout is not None else DEFAULT_CONFIG['dec_dropout']
  175. else:
  176. n_genes = n_genes if n_genes is not None else self._auto_n_genes(sc_count.shape[0], n_types)
  177. max_cells_per_type = max_cells_per_type if max_cells_per_type is not None else self._auto_max_cells(sc_meta, cell_types, celltype_col)
  178. lambda_div = lambda_div if lambda_div is not None else self._adaptive_diversity_weight(n_types)
  179. margin = margin if margin is not None else self._adaptive_margin(n_types)
  180. b_epochs = b_epochs if b_epochs is not None else 2000
  181. b_lr = b_lr if b_lr is not None else 1e-3
  182. d_slot=d_slot if d_slot is not None else DEFAULT_CONFIG['d_slot']
  183. dec_hidden=dec_hidden if dec_hidden is not None else DEFAULT_CONFIG['dec_hidden']
  184. dec_dropout=dec_dropout if dec_dropout is not None else DEFAULT_CONFIG['dec_dropout']
  185. if isinstance(dec_hidden,int): dec_hidden=(dec_hidden,)
  186. self._log(f"[SlotDeconv] Fitting: {n_types} types, {n_genes} genes, max {max_cells_per_type} cells/type")
  187. self._log(f"[SlotDeconv] Config: λ_div={lambda_div}, margin={margin}")
  188. self._last_fit_params = {"n_genes":n_genes,"max_cells_per_type":max_cells_per_type,"lambda_div":lambda_div,"margin":margin,"b_epochs":b_epochs,"b_lr":b_lr,"d_slot":int(d_slot),"dec_hidden":tuple(dec_hidden),"dec_dropout":float(dec_dropout)}
  189. sc_df, sc_meta_bal = self._balance_cells(sc_count, sc_meta, cell_types, celltype_col, max_cells_per_type)
  190. genes, self.gene_weights_ = self._select_discriminative_genes(sc_df, sc_meta_bal, cell_types, celltype_col, n_genes)
  191. self.selected_genes_ = genes
  192. self._gene2idx = {g: i for i, g in enumerate(genes)}
  193. sc_df = sc_df.loc[genes]
  194. self._log(f"[SlotDeconv] Training reference matrix...")
  195. self.B_prob_ = self._train_reference(sc_df,sc_meta_bal,cell_types,celltype_col,lambda_div=lambda_div,margin=margin,epochs=b_epochs,lr=b_lr,d_slot=d_slot,dec_hidden=dec_hidden,dec_dropout=dec_dropout)
  196. self._log(f"[SlotDeconv] Reference matrix learned")
  197. return self
  198. def transform(self, st_count, spatial_coords, use_spatial=True,
  199. lambda_sp=None, sp_epochs=None, sp_lr=None, pow_w=None,
  200. knn=None, spatial_backend='auto', use_gene_weight_in_kl=False,
  201. use_mixed_loss=False, mix_alpha=0.2, w_clip=(0.5, 2.0)):
  202. """
  203. Deconvolve spatial transcriptomics data.
  204. Parameters
  205. ----------
  206. st_count : pd.DataFrame
  207. Spatial count matrix, shape (genes, spots)
  208. spatial_coords : pd.DataFrame or np.ndarray
  209. Spatial coordinates
  210. use_spatial : bool, optional
  211. Use spatial refinement. Default: True
  212. lambda_sp : float, optional
  213. Spatial regularization weight. Default: from config
  214. sp_epochs : int, optional
  215. Spatial refinement epochs. Default: from config
  216. sp_lr : float, optional
  217. Learning rate. Default: from config
  218. pow_w : float, optional
  219. Gene weight power for NNLS. Default: from config
  220. knn : int, optional
  221. K for KNN graph. Default: from config
  222. spatial_backend : str, optional
  223. 'auto': dense if N<=8000, else KNN
  224. 'dense': force dense graph (may OOM for large data)
  225. 'knn': force KNN graph
  226. use_gene_weight_in_kl : bool, optional
  227. Use gene weights in KL loss (Stage3). Default: False
  228. use_mixed_loss : bool, optional
  229. Use mixed loss: (1-α)*KL + α*weighted_KL. Default: False
  230. Only effective when use_gene_weight_in_kl=True
  231. mix_alpha : float, optional
  232. Mixing weight for weighted KL (0-1). Default: 0.2
  233. w_clip : tuple, optional
  234. Clip range for gene weights in Stage3. Default: (0.5, 2.0)
  235. """
  236. if self.B_prob_ is None:
  237. raise ValueError("Model not fitted. Call fit() first.")
  238. if self.use_default_config:
  239. lambda_sp = lambda_sp if lambda_sp is not None else DEFAULT_CONFIG['lambda_sp']
  240. sp_epochs = sp_epochs if sp_epochs is not None else DEFAULT_CONFIG['sp_epochs']
  241. sp_lr = sp_lr if sp_lr is not None else DEFAULT_CONFIG['sp_lr']
  242. pow_w = pow_w if pow_w is not None else DEFAULT_CONFIG['pow_w']
  243. knn = knn if knn is not None else DEFAULT_CONFIG['knn']
  244. else:
  245. if isinstance(spatial_coords, pd.DataFrame):
  246. coords_tmp = spatial_coords[['x', 'y']].values.astype(np.float32)
  247. else:
  248. coords_tmp = np.asarray(spatial_coords, dtype=np.float32)
  249. lambda_sp = lambda_sp if lambda_sp is not None else self._adaptive_spatial_weight(coords_tmp)
  250. sp_epochs = sp_epochs if sp_epochs is not None else 1500
  251. sp_lr = sp_lr if sp_lr is not None else 0.01
  252. pow_w = pow_w if pow_w is not None else 0.8
  253. knn = knn if knn is not None else 15
  254. common_genes = [g for g in self.selected_genes_ if g in st_count.index]
  255. if len(common_genes) < len(self.selected_genes_) * 0.5:
  256. warnings.warn(f"Only {len(common_genes)}/{len(self.selected_genes_)} genes found")
  257. gene_idx = np.array([self._gene2idx[g] for g in common_genes])
  258. st_df = st_count.loc[common_genes]
  259. B_prob = self.B_prob_[:, gene_idx]
  260. B_prob = B_prob / (B_prob.sum(axis=1, keepdims=True) + 1e-8)
  261. w = self.gene_weights_[gene_idx]
  262. self._log(f"[SlotDeconv] NNLS deconvolution...")
  263. pred_nnls = self._deconv_nnls(st_df, B_prob, w, pow_w)
  264. if not use_spatial:
  265. return pred_nnls
  266. if isinstance(spatial_coords, pd.DataFrame):
  267. coords = spatial_coords.loc[st_df.columns][['x', 'y']].values.astype(np.float32)
  268. else:
  269. coords = np.asarray(spatial_coords, dtype=np.float32)
  270. n_spots = len(coords)
  271. if spatial_backend == 'auto':
  272. use_knn = n_spots > 8000
  273. elif spatial_backend == 'knn':
  274. use_knn = True
  275. else:
  276. use_knn = False
  277. gene_w = w if use_gene_weight_in_kl else None
  278. backend_str = 'knn' if use_knn else 'dense'
  279. self._last_transform_params = {"lambda_sp": lambda_sp, "sp_epochs": sp_epochs, "sp_lr": sp_lr, "pow_w": pow_w, "knn": knn, "spatial_backend_used": backend_str, "use_gene_weight_in_kl": bool(use_gene_weight_in_kl), "use_mixed_loss": bool(use_mixed_loss), "mix_alpha": float(mix_alpha), "w_clip": tuple(w_clip)}
  280. self._log(f"[SlotDeconv] Spatial refinement (λ_sp={lambda_sp}, backend={backend_str})...")
  281. pred_kl = self._deconv_spatial(st_df, B_prob, coords, pred_nnls, lambda_sp, sp_epochs, sp_lr,
  282. use_knn, knn, gene_w, pow_w, use_mixed_loss, mix_alpha, w_clip)
  283. self._log(f"[SlotDeconv] Done")
  284. return pred_kl
  285. def fit_transform(self, sc_count, sc_meta, cell_types, st_count, spatial_coords,
  286. celltype_col='cellType', use_spatial=True, **kwargs):
  287. fit_keys = ['n_genes','max_cells_per_type','lambda_div','margin','b_epochs','b_lr','d_slot','dec_hidden','dec_dropout']
  288. transform_keys = ['lambda_sp', 'sp_epochs', 'sp_lr', 'pow_w', 'knn', 'spatial_backend',
  289. 'use_gene_weight_in_kl', 'use_mixed_loss', 'mix_alpha', 'w_clip']
  290. fit_kwargs = {k: kwargs[k] for k in fit_keys if k in kwargs}
  291. transform_kwargs = {k: kwargs[k] for k in transform_keys if k in kwargs}
  292. self.fit(sc_count, sc_meta, cell_types, celltype_col, **fit_kwargs)
  293. return self.transform(st_count, spatial_coords, use_spatial, **transform_kwargs)
  294. def _auto_n_genes(self, total_genes, n_types):
  295. base = min(3000, total_genes)
  296. if n_types > 20:
  297. return min(int(base * 1.2), total_genes)
  298. elif n_types > 10:
  299. return min(base, total_genes)
  300. else:
  301. return min(int(base * 0.8), total_genes)
  302. def _auto_max_cells(self, sc_meta, cell_types, celltype_col):
  303. ct = sc_meta[celltype_col].astype(str).str.strip()
  304. counts = [sum(ct == t) for t in cell_types]
  305. median_count = np.median([c for c in counts if c > 0])
  306. return int(min(750, max(500, median_count)))
  307. def _adaptive_diversity_weight(self, n_types):
  308. if n_types <= 10:
  309. return 2.0
  310. elif n_types <= 20:
  311. return 3.0
  312. else:
  313. return 4.0
  314. def _adaptive_margin(self, n_types):
  315. if n_types <= 10:
  316. return 0.2
  317. elif n_types <= 20:
  318. return 0.15
  319. else:
  320. return 0.1
  321. def _adaptive_spatial_weight(self, coords):
  322. coords_norm = (coords - coords.mean(0)) / (coords.std(0) + 1e-8)
  323. dist = cdist(coords_norm, coords_norm)
  324. median_dist = np.median(dist[dist > 0])
  325. n_spots = len(coords)
  326. if n_spots < 1000:
  327. base_weight = 10.0
  328. elif n_spots < 5000:
  329. base_weight = 15.0
  330. else:
  331. base_weight = 20.0
  332. if median_dist < 0.5:
  333. return base_weight * 1.2
  334. elif median_dist > 1.5:
  335. return base_weight * 0.8
  336. return base_weight
  337. def _balance_cells(self, sc_df, sc_meta, cell_types, celltype_col, max_per_type):
  338. rng = np.random.RandomState(self.random_state)
  339. ct = sc_meta[celltype_col].astype(str).str.strip()
  340. keep = []
  341. for t in cell_types:
  342. idx = sc_meta.index[ct == t].astype(str).values
  343. if len(idx) == 0:
  344. continue
  345. if max_per_type and len(idx) > max_per_type:
  346. idx = rng.choice(idx, size=max_per_type, replace=False)
  347. keep.append(idx)
  348. keep = np.concatenate(keep).astype(str)
  349. keep = pd.Index(keep).intersection(sc_df.columns.astype(str)).intersection(sc_meta.index.astype(str))
  350. return sc_df.loc[:, keep], sc_meta.loc[keep]
  351. def _select_discriminative_genes(self, sc_df, sc_meta, cell_types, celltype_col, n_genes):
  352. X = sc_df.T.to_numpy(np.float32)
  353. ct = sc_meta[celltype_col].astype(str).values
  354. mp = {c: i for i, c in enumerate(cell_types)}
  355. idx = [i for i, x in enumerate(ct) if x in mp]
  356. X, ct = X[idx], ct[idx]
  357. K, G = len(cell_types), X.shape[1]
  358. mu = np.zeros((K, G), dtype=np.float32)
  359. var = np.zeros((K, G), dtype=np.float32)
  360. n = np.zeros(K, dtype=np.float32)
  361. for i, x in enumerate(ct):
  362. k = mp[x]
  363. mu[k] += X[i]
  364. var[k] += X[i] * X[i]
  365. n[k] += 1.0
  366. n = np.maximum(n, 1.0)
  367. mu = mu / n[:, None]
  368. var = np.maximum(var / n[:, None] - mu * mu, 0.0)
  369. mu_all = (mu * n[:, None]).sum(0) / n.sum()
  370. between_var = ((mu - mu_all[None, :]) ** 2 * n[:, None]).sum(0) / n.sum()
  371. within_var = (var * n[:, None]).sum(0) / n.sum()
  372. f_score = between_var / (within_var + 1e-6)
  373. top_idx = np.argsort(-f_score)[:min(n_genes, len(f_score))]
  374. return sc_df.index[top_idx], f_score[top_idx].astype(np.float32)
  375. def _train_reference(self, sc_df, sc_meta, cell_types, celltype_col, lambda_div, margin, epochs, lr, d_slot=None, dec_hidden=None, dec_dropout=None):
  376. if d_slot is None: d_slot=DEFAULT_CONFIG['d_slot']
  377. if dec_hidden is None: dec_hidden=DEFAULT_CONFIG['dec_hidden']
  378. if dec_dropout is None: dec_dropout=DEFAULT_CONFIG['dec_dropout']
  379. if isinstance(dec_hidden,int): dec_hidden=(dec_hidden,)
  380. cells=sc_meta.index.astype(str)
  381. X=sc_df.loc[:,cells].T.values.astype(np.float32)
  382. X[X<0]=0
  383. size=X.sum(axis=1,keepdims=True).astype(np.float32)
  384. size=np.maximum(size,1e-8)/np.median(size[size>0])
  385. mp={ct:i for i,ct in enumerate(cell_types)}
  386. ct=sc_meta.loc[cells,celltype_col].astype(str).str.strip().values
  387. mask=np.array([x in mp for x in ct],dtype=bool)
  388. X=X[mask]
  389. size=size[mask]
  390. labels=np.array([mp[x] for x in ct[mask]],dtype=np.int64)
  391. X_t=torch.tensor(X,device=self.device)
  392. S_t=torch.tensor(size,device=self.device)
  393. L_t=torch.tensor(labels,dtype=torch.long,device=self.device)
  394. model=_SlotDecoder(X.shape[1],len(cell_types),d_slot=int(d_slot),margin=margin,hidden=tuple(dec_hidden),dropout=float(dec_dropout)).to(self.device)
  395. opt=optim.AdamW(model.parameters(),lr=lr,weight_decay=1e-5)
  396. sch=optim.lr_scheduler.CosineAnnealingLR(opt,T_max=epochs,eta_min=1e-6)
  397. for ep in range(epochs):
  398. mu=model(L_t,S_t)
  399. recon=_nb_nll(X_t,mu,model.alpha_disp()[None,:])
  400. div,_=model.diversity_loss()
  401. loss=recon+lambda_div*div
  402. opt.zero_grad()
  403. loss.backward()
  404. nn.utils.clip_grad_norm_(model.parameters(),5.0)
  405. opt.step()
  406. sch.step()
  407. self.slots_=model.slots.detach().cpu().numpy()
  408. self.alpha_=model.alpha_disp().detach().cpu().numpy()
  409. return model.get_reference_matrix()
  410. def _deconv_nnls(self, st_df, B_prob, w, pow_w):
  411. X = st_df.T.values.astype(np.float32)
  412. X[X < 0] = 0
  413. size = X.sum(axis=1, keepdims=True).astype(np.float32)
  414. size = np.maximum(size, 1e-8) / np.median(size[size > 0])
  415. Xu = X / size
  416. ww = np.clip((w / (np.median(w) + 1e-8)) ** pow_w, 0.2, 5.0).astype(np.float64)
  417. sw = np.sqrt(ww)
  418. Bt = (B_prob * sw[None, :]).T.astype(np.float64)
  419. Xu_w = Xu * sw[None, :]
  420. V = np.zeros((Xu.shape[0], B_prob.shape[0]), dtype=np.float32)
  421. for i in range(Xu.shape[0]):
  422. v, _ = nnls(Bt, Xu_w[i].astype(np.float64))
  423. s = v.sum()
  424. V[i] = (v / s if s > 0 else np.ones(B_prob.shape[0]) / B_prob.shape[0]).astype(np.float32)
  425. return pd.DataFrame(V, index=st_df.columns, columns=self.cell_types_)
  426. def _deconv_spatial(self, st_df, B_prob, coords, V_init, lambda_sp, epochs, lr,
  427. use_knn, knn, gene_w, pow_w, use_mixed_loss, mix_alpha, w_clip):
  428. X = st_df.T.values.astype(np.float32)
  429. X[X < 0] = 0
  430. X_prob = X / (X.sum(axis=1, keepdims=True) + 1e-8)
  431. X_t = torch.tensor(X_prob, device=self.device)
  432. B_t = torch.tensor(B_prob, device=self.device)
  433. V_logits = nn.Parameter(torch.tensor(np.log(V_init.values + 1e-6), device=self.device))
  434. coords_norm = (coords - coords.mean(0)) / (coords.std(0) + 1e-8)
  435. if use_knn:
  436. idx, w = build_knn_graph(coords_norm, knn)
  437. idx_t = torch.tensor(idx, device=self.device, dtype=torch.long)
  438. w_t = torch.tensor(w, device=self.device)
  439. else:
  440. W = build_dense_graph(coords_norm)
  441. W_t = torch.tensor(W, device=self.device)
  442. if gene_w is not None:
  443. gw = np.asarray(gene_w, dtype=np.float64)
  444. gw = (gw / (np.median(gw) + 1e-12)) ** pow_w
  445. gw = np.clip(gw, w_clip[0], w_clip[1])
  446. gw = gw / (np.mean(gw) + 1e-12)
  447. gw_t = torch.tensor(gw, device=self.device, dtype=torch.float32)
  448. opt = optim.Adam([V_logits], lr=lr)
  449. sch = optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs, eta_min=1e-4)
  450. for ep in range(epochs):
  451. V = F.softmax(V_logits, dim=1)
  452. pred = V @ B_t
  453. pred_prob = pred / (pred.sum(dim=1, keepdim=True) + 1e-8)
  454. recon_base = (X_t * (torch.log(X_t + 1e-8) - torch.log(pred_prob + 1e-8))).sum(dim=1).mean()
  455. if gene_w is not None:
  456. recon_weighted = (gw_t[None, :] * X_t * (torch.log(X_t + 1e-8) - torch.log(pred_prob + 1e-8))).sum(dim=1).mean()
  457. if use_mixed_loss:
  458. recon = (1 - mix_alpha) * recon_base + mix_alpha * recon_weighted
  459. else:
  460. recon = recon_weighted
  461. else:
  462. recon = recon_base
  463. if use_knn:
  464. V_neighbor = (V[idx_t] * w_t[..., None]).sum(dim=1)
  465. else:
  466. V_neighbor = W_t @ V
  467. spatial = ((V - V_neighbor) ** 2).mean()
  468. loss = recon + lambda_sp * spatial
  469. opt.zero_grad()
  470. loss.backward()
  471. opt.step()
  472. sch.step()
  473. with torch.no_grad():
  474. V_final = F.softmax(V_logits, dim=1).cpu().numpy()
  475. return pd.DataFrame(V_final, index=st_df.columns, columns=self.cell_types_)

slot_model.py at commit 5819b15, under MIT · at the source

Overview

Authors: Hanzhang Fang1, Cong Qi1, Yuanjie Zou1, Yeqing Chen1, Zhi Wei1
ORCID iDs: Cong Qi
  1. Department of Computer Science, New Jersey Institute of Technology, University Heights, Newark, New Jersey 07102, United States
Institutions: New Jersey Institute of Technology (United States)
Journal: Bioinformatics (Oxford, England), volume 42, issue Suppl 2, article btag424
Dates: published online 21 August 2026; in print August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1093/bioinformatics/btag424 · PMID 42635200 · PMCID PMC13501309 · OpenAlex W7204103832
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: genetics / omics (modality), human (organism), mouse (organism), methods / tools (subfield)
Methods: Statistics, Preprocessing, Connectivity, fMRI & imaging, Machine learning
MeSH: Computational Biology*, Gene Expression Profiling*, Software*, Transcriptome*, Algorithms, Animals, Humans, Machine Learning, Mice, Spatial Transcriptomics (* major topic)
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: NIH HHS (R35GM158529, R15HG012087); National Institutes of Health (R35GM158529, R15HG012087); NIGMS NIH HHS (R35 GM158529); NHGRI NIH HHS (R15 HG012087)
Citations: not cited yet (Europe PMC); 15 references in the paper

Abstract

Motivation: Spatial transcriptomics (ST) measures gene expression in intact tissues. In spot-based ST assays, each spot can contain mixtures of multiple cell types. Deconvolution is particularly challenging when closely related cell subtypes share highly similar expression profiles and when spatial context is underutilized during proportion estimation.

Results: We present SlotDeconv, a method for ST deconvolution consisting of a single-cell reference module and a spatial inference module. The reference module learns discriminative cell-type signatures using slot-based prototype vectors decoded into a reference matrix, trained with a negative binomial reconstruction loss and a max-margin diversity constraint that discourages similar cell-type signatures. Ablation studies confirm that both components are essential: removing the diversity constraint reduces spot-wise Pearson correlation by 41%, and replacing learned prototypes with cell-type mean expression reduces it to near zero. The spatial inference module initializes spot-level proportions via gene-weighted nonnegative least squares (NNLS), then refines them by minimizing Kullback–Leibler (KL) divergence between observed and reconstructed spot expression under a spatial neighborhood consistency regularizer. Benchmarked against CARD, RCTD, Cell2location, and Spotiphy on a 27 cell type mouse brain dataset, SlotDeconv achieves the highest spot wise Pearson correlation (approximately 0.56) and cosine similarity (0.633), outperforming competing methods in spot-wise correlation, with particularly strong gains on transcriptionally similar cortical neuronal subtypes. Biological validation on human pancreatic cancer and mouse olfactory bulb datasets further confirms spatial specificity.

Availability: Source code is available at https://github.com/HannahNJIT/SlotDeconv

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

HannahNJIT/SlotDeconv

License: MIT
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 5819b15a6acc9e753a16c1c781ca0b155edfc7fa, 18 May 2026
Languages: Python (5), Jupyter (1)
Size: 12 files, 6 scripts
Software Heritage: not archived
Found in: “Data availability”
Holds: README, license file, environment (environment.yml, requirements.txt), 1 notebook
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (5 files), pandas (5 files), scikit-learn (4 files), SciPy (4 files), PyTorch (3 files), Matplotlib (2 files), Scanpy (2 files), anndata (1 file), seaborn (1 file), UMAP (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
8 files

Availability

Source code is available at https://github.com/HannahNJIT/SlotDeconv

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;
  • 6 scripts, each with its path and the digest of its content;
  • 4 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

No dataset and no data link were found in the paper.

Data availability

Source code, scripts, and accession links needed to reproduce the analyses are available at https://github.com/HannahNJIT/SlotDeconv. Public datasets used in this study are described in Supplementary Note S1, available as supplementary data at Bioinformatics online.

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

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 10 MeSH terms, 4 funders, 15 references.

Cite

This paper

Fang, H., Qi, C., Zou, Y., Chen, Y., & Wei, Z. (2026). SlotDeconv: spatial transcriptomics deconvolution via diversity-constrained prototype learning and spatial refinement. Bioinformatics (Oxford, England), 42(Suppl 2), btag424. https://doi.org/10.1093/bioinformatics/btag424

BibTeX

@article{fang2026slotdeconv,
author = {Fang, Hanzhang and Qi, Cong and Zou, Yuanjie and Chen, Yeqing and Wei, Zhi},
title = {{SlotDeconv: spatial transcriptomics deconvolution via diversity-constrained prototype learning and spatial refinement}},
journal = {Bioinformatics (Oxford, England)},
year = {2026},
month = aug,
volume = {42},
number = {Suppl 2},
pages = {btag424},
publisher = {Oxford University Press},
issn = {1367-4803},
doi = {10.1093/bioinformatics/btag424},
url = {https://doi.org/10.1093/bioinformatics/btag424},
pmid = {42635200},
pmcid = {PMC13501309}
}

RIS

TY - JOUR
AU - Fang, Hanzhang
AU - Qi, Cong
AU - Zou, Yuanjie
AU - Chen, Yeqing
AU - Wei, Zhi
TI - SlotDeconv: spatial transcriptomics deconvolution via diversity-constrained prototype learning and spatial refinement
T2 - Bioinformatics (Oxford, England)
J2 - Bioinformatics
PY - 2026
DA - 2026/08/01
VL - 42
IS - Suppl 2
SP - btag424
SN - 1367-4803
PB - Oxford University Press
DO - 10.1093/bioinformatics/btag424
UR - https://doi.org/10.1093/bioinformatics/btag424
LA - en
ER -

CSL-JSON

{
"id": "10.1093/bioinformatics/btag424",
"type": "article-journal",
"title": "SlotDeconv: spatial transcriptomics deconvolution via diversity-constrained prototype learning and spatial refinement",
"container-title": "Bioinformatics (Oxford, England)",
"author": [
{
"family": "Fang",
"given": "Hanzhang"
},
{
"family": "Qi",
"given": "Cong"
},
{
"family": "Zou",
"given": "Yuanjie"
},
{
"family": "Chen",
"given": "Yeqing"
},
{
"family": "Wei",
"given": "Zhi"
}
],
"container-title-short": "Bioinformatics",
"volume": "42",
"issue": "Suppl 2",
"page": "btag424",
"DOI": "10.1093/bioinformatics/btag424",
"PMID": "42635200",
"PMCID": "PMC13501309",
"ISSN": "1367-4803",
"publisher": "Oxford University Press",
"URL": "https://doi.org/10.1093/bioinformatics/btag424",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
1
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.21203/rs.3.rs-9676637/v1 [code]
A Comprehensive Benchmarking of Spatial Deconvolution and Domain Detection Methods across Diverse Tissues and Spatial Transcriptomic Technologies
Journal: Research Square (preprint)
In common: UMAP, anndata, Scanpy, 7 other tools, methods / tools, genetics / omics, 7 references
[2] doi:10.1016/j.isci.2026.117206 [code]
ReliST: A model-agnostic risk layer for spatial transcriptomics deconvolution.
Journal: iScience
In common: anndata, Scanpy, PyTorch, 6 other tools, genetics / omics, mouse, 6 references
[3] 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, 6 references
[4] doi:10.1038/s42003-026-10259-z [code]
Spatial transcriptomic profiling of developing mouse hearts reveals a spatially patterned signaling environment.
Journal: Communications biology
In common: UMAP, anndata, Scanpy, 7 other tools, genetics / omics, mouse, 3 references
[5] 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: anndata, Scanpy, PyTorch, 5 other tools, genetics / omics, mouse, 3 references
[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: UMAP, anndata, Scanpy, 7 other tools, 2 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: UMAP, anndata, Scanpy, 7 other tools, genetics / omics, mouse, 1 reference
[8] doi:10.1038/s41593-026-02293-1 [code]
Optics-free spatial genomics for mapping mammalian brain aging by IRISeq.
Journal: Nature neuroscience
In common: UMAP, Scanpy, seaborn, 5 other tools, genetics / omics, mouse, 3 references
[9] doi:10.1038/s41592-026-03057-2 [code]
CREsted: modeling genomic and synthetic cell-type-specific enhancers across tissues and species.
Journal: Nature methods
In common: UMAP, anndata, Scanpy, 7 other tools, methods / tools, genetics / omics, mouse
[10] doi:10.1093/bioinformatics/btag515 [code]
PRISM: Prior-enhanced Inference for Spatial Transcriptomic Cell Type Mapping.
Journal: Bioinformatics (Oxford, England)
In common: anndata, Scanpy, PyTorch, 5 other tools, genetics / omics, 3 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.