OSCR

ProtoCloud: A prototypical self-explaining model for single-cell analysis.

Code ↔ Paper

21 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 21 matches
  1. [1] § STAR★Methods › Method details › Reliability and uncertainty estimation ↔ src/ProtoCloud/model/calibrator.py, lines 10–26 · score 0.72 · global isotonic regression, global calibration, calibrated probability, fit, mapping, scores
  2. [2] § Design ↔ src/ProtoCloud/model/model.py, lines 68–101 · score 0.72 · deep generative model, latent space organized, low dimensional, embeds, ProtoCloud, encoder
  3. [3] § STAR★Methods › Method details › The ProtoCloud model ↔ src/ProtoCloud/model/model.py, lines 68–101 · score 0.70 · low dimensional, generative model, prototype vector, latent space, embeds, bias
  4. [4] § Design ↔ src/ProtoCloud/model/model.py, lines 272–340 · score 0.69 · cross entropy, orthogonal loss, atomic loss, multinomial, components, identity
  5. [5] § STAR★Methods › Method details › Ablation study › Two-stage vs. single-stage curriculum ↔ src/ProtoCloud/model/train.py, lines 19–84 · score 0.68 · stage curriculum, orthogonal loss, atomic loss, stage training, ProtoCloud, model
  6. [6] § STAR★Methods › Method details › Evaluation metrics ↔ src/ProtoCloud/model/calibrator.py, lines 209–236 · score 0.68 · expected calibration error, predicted probability, conf, bins, ECE, accuracy
  7. [7] § Design ↔ src/ProtoCloud/prp/lrp.py, lines 295–377 · score 0.67 · prototypical relevance propagation, gene relevance, relevance score, backpropagating, gene expression, explanatory
  8. [8] § STAR★Methods › Method details › Shaping latent spaces through a two-stage curriculum and latent decomposition ↔ src/ProtoCloud/model/model.py, lines 486–512 · score 0.65 · L1 regularization, layer weight, sparsity, encourages, matrix, encoder
  9. [9] § Results › Batch-separated informative latent space ↔ src/ProtoCloud/model/train.py, lines 19–84 · score 0.64 · orthogonal loss, stage curriculum, atomic loss, stage training, ProtoCloud, batch
  10. [10] § STAR★Methods › Method details › Prototypical relevance propagation for gene-level interpretation ↔ src/ProtoCloud/model/model.py, lines 21–47 · score 0.63 · fully connected layer, ReLU, activation, linear
  11. [11] § STAR★Methods › Method details › Model architecture and training ↔ main.py, lines 371–438 · score 0.63 · AdamW, batch normalization, optimizer, ReLU, activation, bias
  12. [12] § Results › Batch-separated informative latent space ↔ src/ProtoCloud/model/model.py, lines 272–340 · score 0.61 · loss component, orthogonal loss, atomic loss, entropy, batch, latent
  13. [13] § STAR★Methods › Method details › The ProtoCloud model ↔ src/ProtoCloud/model/model.py, lines 386–429 · score 0.59 · negative binomial, dispersion parameter, likelihood, reconstructing, NB, latent
  14. [14] § STAR★Methods › Method details › Data collection and processing › Patch-seq RGC ↔ src/ProtoCloud/data/scRNAdata.py, lines 33–91 · score 0.58 · single cell RNA, model trained, Seq, raw, genes
  15. [15] § Results › Similarity-guided annotation correction with justification ↔ src/ProtoCloud/model/calibrator.py, lines 157–206 · score 0.56 · expected calibration error, Brier score, ECE, metrics, prediction, model
  16. [16] § Results › Similarity-guided annotation correction with justification ↔ src/ProtoCloud/model/calibrator.py, lines 99–131 · score 0.55 · expected calibration error, Brier score, ECE, prediction, model
  17. [17] § STAR★Methods › Method details › The orthogonal loss to help learn diverse prototypes for each class ↔ src/ProtoCloud/model/model.py, lines 179–213 · score 0.53 · cross entropy, orthogonal loss, class, prototype, model
  18. [18] § STAR★Methods › Method details › Evaluation metrics ↔ src/ProtoCloud/utils/utils.py, lines 461–484 · score 0.51 · nearest neighbors, batch entropy
  19. [19] § STAR★Methods › Method details › Evaluation metrics ↔ src/ProtoCloud/utils/utils.py, lines 441–464 · score 0.51 · nearest neighbors, batch entropy
  20. [20] § Results › Similarity-guided annotation correction with justification › Type 2: Ambiguous misprediction ↔ src/ProtoCloud/model/calibrator.py, lines 157–206 · score 0.51 · Brier score, calibrated score, ECE, metrics, probability
  21. [21] § STAR★Methods › Method details › Prototypical relevance propagation for gene-level interpretation ↔ src/ProtoCloud/prp/lrp.py, lines 295–377 · score 0.51 · prototypical relevance propagation, backpropagated, LRP, PRP, dimensional, protoCloud

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 · 867 lines · 32 KB · MIT · 6 matches

  1. from typing import Any, Iterable, Mapping, Sequence, Tuple, Union, Optional, Callable, Literal, List
  2. import torch
  3. import torch.nn as nn
  4. import torch.nn.functional as F
  5. from torch.distributions import Distribution, Gamma, Poisson
  6. import numpy as np
  7. import pandas as pd
  8. import ProtoCloud.glo as glo
  9. glo.set_value('EPS', 1e-16)
  10. glo.set_value('LRP_FILTER_TOP_K', 0.1)
  11. from ..utils import seed_torch, log_likelihood_nb, one_hot_encoder
  12. device = 'cuda' if torch.cuda.is_available() else 'cpu'
  13. num_workers = 4 if torch.cuda.is_available() else 0
  14. EPS = glo.get_value('EPS')
  15. def form_block(in_dim, out_dim,
  16. use_bn = True, activation = 'relu',
  17. bias = False, dropout = 0,
  18. ):
  19. """Construct a sequential block of Linear, BatchNorm, activation, and dropout.
  20. Parameters
  21. ----------
  22. in_dim : int
  23. Input feature dimension.
  24. out_dim : int
  25. Output feature dimension.
  26. use_bn : bool, optional
  27. Whether to include BatchNorm1d, by default True.
  28. activation : {'relu', 'leakyrelu'}, optional
  29. Activation function type, by default 'relu'.
  30. bias : bool, optional
  31. Whether to add bias to the Linear layer, by default False.
  32. dropout : float, optional
  33. Dropout rate. Set to 0 to disable, by default 0.
  34. Returns
  35. -------
  36. torch.nn.Sequential
  37. The composed layer block.
  38. Raises
  39. ------
  40. ValueError
  41. If ``activation`` is not recognized.
  42. """
  43. layers = [nn.Linear(in_dim, out_dim, bias=bias)]
  44. if use_bn:
  45. layers.append(nn.BatchNorm1d(out_dim))
  46. # activation
  47. if activation == 'relu':
  48. layers.append(nn.ReLU())
  49. elif activation == 'leakyrelu':
  50. layers.append(nn.LeakyReLU())
  51. else:
  52. raise ValueError('activation not recognized')
  53. if dropout != 0:
  54. layers.append(nn.Dropout(dropout))
  55. return nn.Sequential(*layers)
  56. class protoCloud(nn.Module):
  57. """
  58. ProtoCloud, a self-explaining deep generative model trained end-to-end to embed cells into a structured, low-dimensional space organized around cell type-specific prototypes.
  59. Parameters
  60. ----------
  61. input_dim : int
  62. Number of input genes.
  63. num_prototypes_per_class : int
  64. Number of prototype vectors per cell type.
  65. num_classes : int
  66. Number of cell types.
  67. latent_dim : int
  68. Dimensionality of the latent space.
  69. raw_input : int
  70. 1 if input is raw counts (applies log1p), 0 if log-normalized.
  71. encoder_layer_sizes : list of int, optional
  72. Hidden layer sizes for the encoder, by default [1024, 512, 256].
  73. decoder_layer_sizes : list of int, optional
  74. Hidden layer sizes for the decoder, by default [512, 1024].
  75. activation : {'relu', 'leakyrelu'}, optional
  76. Activation function, by default 'relu'.
  77. use_bias : bool, optional
  78. Whether to use bias in Linear layers, by default False.
  79. use_dropout : float, optional
  80. Dropout rate, by default 0.
  81. use_bn : bool, optional
  82. Whether to use batch normalization, by default True.
  83. obs_dist : {'nb', 'normal'}, optional
  84. Observation distribution for reconstruction, by default 'nb'.
  85. nb_dispersion : {'celltype_target', 'celltype_pred', 'gene'}, optional
  86. How to model negative binomial dispersion, by default
  87. 'celltype_target'.
  88. """
  89. def __init__(self, input_dim:int,
  90. num_classes:int,
  91. num_prototypes_per_class: int = 6,
  92. latent_dim: int = 20,
  93. raw_input:int = 1,
  94. encoder_layer_sizes: Optional[list] = None,
  95. decoder_layer_sizes: Optional[list] = None,
  96. activation: Literal['relu', 'leakyrelu'] = 'relu',
  97. use_bias:bool = False,
  98. use_dropout:float = 0,
  99. use_bn:bool = True,
  100. obs_dist: Literal['nb', 'normal'] = 'nb',
  101. nb_dispersion: Literal['celltype_target', 'celltype_pred', 'gene'] = "celltype_target",
  102. # n_batch:int = 1,
  103. ):
  104. super(protoCloud, self).__init__()
  105. self.input_dim = input_dim
  106. self.latent_dim = latent_dim
  107. self.num_prototypes_per_class = num_prototypes_per_class
  108. self.num_classes = num_classes
  109. self.raw_input = raw_input
  110. self.activation = activation
  111. self.use_bn = use_bn
  112. self.use_bias = bool(use_bias)
  113. self.use_dropout = use_dropout
  114. self.obs_dist = None if not self.raw_input else obs_dist
  115. self.nb_dispersion = nb_dispersion
  116. self.epsilon = EPS
  117. # self.n_batch = n_batch
  118. # prototype-class labeled matrix
  119. self.num_prototypes = self.num_prototypes_per_class * self.num_classes
  120. identity = torch.zeros(self.num_prototypes, self.num_classes)
  121. for j in range(self.num_prototypes):
  122. identity[j, j // self.num_prototypes_per_class] = 1
  123. self.register_buffer('prototype_class_identity', identity) # register buffer to move to GPU automatically
  124. # prototype vectors
  125. prototype_shape = (self.num_prototypes, self.latent_dim)
  126. self.prototype_vectors = nn.Parameter(torch.randn(prototype_shape), requires_grad = True)
  127. # mask
  128. self.scale = nn.Parameter(torch.ones(1) * 1.0)
  129. #######################################################
  130. if encoder_layer_sizes is None:
  131. self.encoder_layer_sizes = [self.input_dim] + [1024, 512, 256] # [128, 64, 32]
  132. else:
  133. self.encoder_layer_sizes = [self.input_dim] + encoder_layer_sizes
  134. if decoder_layer_sizes is None:
  135. # self.latent_dim += self.n_batch
  136. self.decoder_layer_sizes = [self.latent_dim] + [512, 1024] # [32, 128]
  137. else:
  138. self.decoder_layer_sizes = [self.latent_dim] + decoder_layer_sizes
  139. # self.decoder_layer_sizes[0] += self.n_batch
  140. # Encoder
  141. self.encoder = nn.Sequential()
  142. for i, (in_dim, out_dim) in enumerate(zip(self.encoder_layer_sizes[:-1], self.encoder_layer_sizes[1:])):
  143. self.encoder.add_module(str(i), form_block(in_dim, out_dim,
  144. self.use_bn, self.activation, self.use_bias, self.use_dropout))
  145. self.z_mean = nn.Linear(self.encoder_layer_sizes[-1], latent_dim, bias = True)
  146. self.z_log_var = nn.Linear(self.encoder_layer_sizes[-1], latent_dim, bias = True)
  147. # Decoder
  148. self.decoder = nn.Sequential()
  149. for i, (in_dim, out_dim) in enumerate(zip(self.decoder_layer_sizes[:-1], self.decoder_layer_sizes[1:])):
  150. self.decoder.add_module(str(i), form_block(in_dim, out_dim,
  151. self.use_bn, self.activation, self.use_bias, self.use_dropout))
  152. self.px_mean = nn.Linear(self.decoder_layer_sizes[-1], input_dim, bias = True)
  153. # likelihood
  154. self.softmax = nn.Softmax(dim = -1)
  155. # nb dispersion: gene-specific
  156. self.px_theta = nn.Sequential(
  157. nn.Linear(self.decoder_layer_sizes[-1], input_dim, bias = True),
  158. nn.Softplus()) # output always positive
  159. # nb dispersion: celltype-specific
  160. self.theta = nn.Parameter(torch.randn(self.input_dim, self.num_classes))
  161. # Classifier
  162. self.classifier = nn.Linear(self.num_prototypes, self.num_classes, bias = False)
  163. self._initialize_weights()
  164. def forward(self, x, batch_id=None):
  165. """Forward pass through the full model.
  166. Parameters
  167. ----------
  168. x : torch.Tensor
  169. Input gene expression of shape ``(batch_size, input_dim)``.
  170. batch_id : optional
  171. Batch information (currently unused).
  172. Returns
  173. -------
  174. pred : torch.Tensor
  175. Classification logits, shape ``(batch_size, num_classes)``.
  176. px_mu : torch.Tensor
  177. Reconstruction mean, shape ``(batch_size, input_dim)``.
  178. px_t : torch.Tensor or None
  179. Dispersion parameters for NB distribution, or None.
  180. z_mu : torch.Tensor
  181. Latent mean, shape ``(batch_size, latent_dim)``.
  182. z_logVar : torch.Tensor
  183. Latent log-variance, shape ``(batch_size, latent_dim)``.
  184. sim_scores : torch.Tensor
  185. Prototype similarity scores, shape ``(batch_size, num_prototypes)``.
  186. """
  187. self.lib_size = torch.sum(x, 1, True)
  188. if self.raw_input: # raw: 1
  189. x = torch.log1p(x)
  190. encode = self.encoder(x)
  191. z_mu = self.z_mean(encode)
  192. z_logVar = self.z_log_var(encode)
  193. z = self.reparameterize(z_mu, z_logVar)
  194. sim_scores = self.calc_sim_scores(z)
  195. pred = self.classifier(sim_scores)
  196. px = self.decoder(z)
  197. px_mu = self.px_mean(px)
  198. if self.obs_dist == 'nb':
  199. px_mu = self.softmax(px_mu) * self.lib_size
  200. if self.nb_dispersion.startswith('celltype') :
  201. px_t = self.theta
  202. elif self.nb_dispersion == 'gene':
  203. px_t = self.px_theta(px)
  204. px_t = torch.mean(px_t, 0, True)
  205. else:
  206. raise NotImplementedError
  207. px_t = torch.clamp(px_t, min = EPS)
  208. else:
  209. px_t = None
  210. return pred, px_mu, px_t, z_mu, z_logVar, sim_scores
  211. def reparameterize(self, mu, logvar):
  212. """Sample from the latent distribution using the reparameterization trick.
  213. Computes ``z = mu + exp(logvar / 2) * epsilon`` where epsilon is
  214. sampled from a standard normal.
  215. Parameters
  216. ----------
  217. mu : torch.Tensor
  218. Mean of the latent distribution.
  219. logvar : torch.Tensor
  220. Log-variance of the latent distribution.
  221. Returns
  222. -------
  223. torch.Tensor
  224. Sampled latent vector, same shape as ``mu``.
  225. """
  226. std = torch.exp(logvar / 2)
  227. eps = torch.randn_like(std)
  228. return mu + std * eps
  229. def loss_function(self, x, target, pred,
  230. px_mu, px_theta,
  231. z_mu, z_logVar,
  232. sim_scores):
  233. """Compute all loss components for training.
  234. Parameters
  235. ----------
  236. x : torch.Tensor
  237. Input gene expression, shape ``(batch_size, input_dim)``.
  238. target : torch.Tensor
  239. Ground truth class labels, shape ``(batch_size,)``.
  240. pred : torch.Tensor
  241. Classification logits from the forward pass.
  242. px_mu : torch.Tensor
  243. Reconstruction mean from the forward pass.
  244. px_theta : torch.Tensor
  245. Dispersion parameters from the forward pass.
  246. z_mu : torch.Tensor
  247. Latent mean from the forward pass.
  248. z_logVar : torch.Tensor
  249. Latent log-variance from the forward pass.
  250. sim_scores : torch.Tensor
  251. Prototype similarity scores from the forward pass.
  252. Returns
  253. -------
  254. recon_loss : torch.Tensor
  255. Reconstruction loss (NB log-likelihood or MSE).
  256. kl_loss : torch.Tensor
  257. KL divergence to class prototypes.
  258. classify_loss : torch.Tensor
  259. Cross-entropy classification loss.
  260. ortho_loss : torch.Tensor
  261. Orthogonality loss for prototype separation.
  262. atomic_loss : torch.Tensor
  263. Attraction/repulsion loss for prototype assignment.
  264. """
  265. if target.max() >= self.prototype_class_identity.shape[0]:
  266. print("Max target:", target.max())
  267. print("Shape of prototype_class_identity:", self.prototype_class_identity.shape)
  268. raise IndexError("Target index is out of bounds.")
  269. # Reconstruction loss
  270. if self.nb_dispersion == 'celltype_target' or self.nb_dispersion == 'gene':
  271. recon_loss, _ = self.recon_loss(x, target, px_mu, px_theta)
  272. else: # self.nb_dispersion == 'celltype_pred'
  273. softmax_pred = F.softmax(pred, dim=1)
  274. max_index = torch.multinomial(softmax_pred, 1)
  275. recon_loss, _ = self.recon_loss(x, max_index, px_mu, px_theta)
  276. prototypes_of_correct_class = self.prototype_class_identity[:, target].t()
  277. index_prototypes_of_correct_class = (prototypes_of_correct_class == 1).nonzero(as_tuple = True)[1]
  278. # class-corresponding prototypes' index for each sample in the batch
  279. index_prototypes_of_correct_class = index_prototypes_of_correct_class.view(x.shape[0],
  280. self.num_prototypes_per_class)
  281. # KL divergence loss
  282. kl_loss, mask = self.kl_divergence_nearest(z_mu, z_logVar, index_prototypes_of_correct_class, sim_scores)
  283. # Classification loss
  284. classify_loss = F.cross_entropy(pred, target)
  285. # Orthogonality loss
  286. ortho_loss = self.orthogonal_loss()
  287. # Atomic loss
  288. atomic_loss = self.atomic_loss(sim_scores, prototypes_of_correct_class)
  289. return recon_loss, kl_loss, classify_loss, ortho_loss, atomic_loss
  290. def calc_sim_scores(self, z):
  291. """Compute similarity scores between latent embeddings and prototypes.
  292. Uses the first half of latent dimensions to compute pairwise
  293. Euclidean distances, then converts to similarities via the
  294. Cauchy kernel.
  295. Parameters
  296. ----------
  297. z : torch.Tensor
  298. Latent embeddings, shape ``(batch_size, latent_dim)``.
  299. Returns
  300. -------
  301. torch.Tensor
  302. Similarity scores, shape ``(batch_size, num_prototypes)``.
  303. """
  304. # pairwise Euclidean distances between z and prototype vectors
  305. d = torch.cdist(z[:, :self.latent_dim // 2],
  306. self.prototype_vectors[:, :self.latent_dim // 2], p = 2) ## Batch size x num_prototypes
  307. sim_scores = self.distance_2_similarity(d)
  308. return sim_scores
  309. def distance_2_similarity(self, distances):
  310. """Convert distances to similarities using a Cauchy kernel.
  311. Computes ``1 / (scale^2 * d^2 + 1)``, yielding values in [0, 1].
  312. Parameters
  313. ----------
  314. distances : torch.Tensor
  315. Pairwise distance values.
  316. Returns
  317. -------
  318. torch.Tensor
  319. Similarity scores in [0, 1], same shape as input.
  320. """
  321. # return torch.log((distances + 1) / (distances + self.epsilon))
  322. return 1.0 / (torch.square(distances * self.scale) + 1.0) # heavy tail
  323. def recon_loss(self, x, target, px_mu, px_t):
  324. """Compute reconstruction loss.
  325. Uses negative binomial log-likelihood for count data or MSE
  326. for normalized data.
  327. Parameters
  328. ----------
  329. x : torch.Tensor
  330. Ground truth expression, shape ``(batch_size, input_dim)``.
  331. target : torch.Tensor
  332. Class labels, used for cell-type-specific dispersion.
  333. px_mu : torch.Tensor
  334. Predicted reconstruction mean.
  335. px_t : torch.Tensor or None
  336. Dispersion parameters.
  337. Returns
  338. -------
  339. loss : torch.Tensor
  340. Scalar reconstruction loss.
  341. dispersion : torch.Tensor or None
  342. Computed dispersion values, or None for normal distribution.
  343. """
  344. if self.obs_dist == 'nb':
  345. if self.nb_dispersion.startswith('celltype'):
  346. dispersion = F.linear(one_hot_encoder(target, self.num_classes), self.theta)
  347. dispersion = torch.exp(dispersion)
  348. elif self.nb_dispersion == 'gene':
  349. dispersion = px_t
  350. else:
  351. raise NotImplementedError
  352. ll = -log_likelihood_nb(x, px_mu, dispersion)
  353. recon_loss = torch.mean(torch.sum(ll, dim = -1))
  354. recon_loss = recon_loss / self.input_dim * self.latent_dim / 2.0 # scale nb loss down
  355. else:
  356. # x = F.normalize(x, dim = 0)
  357. recon_loss = torch.nn.functional.mse_loss(px_mu, x, reduction = "mean")
  358. dispersion = None
  359. return recon_loss, dispersion
  360. def kl_divergence_nearest(self, mu, logVar, nearest_pt, sim_scores):
  361. """Compute KL divergence to nearest class prototypes.
  362. The first half of the latent space is regularized toward prototype
  363. distributions (weight 5); the second half toward a standard
  364. normal (weight 1). Losses are weighted by similarity scores.
  365. Parameters
  366. ----------
  367. mu : torch.Tensor
  368. Latent means, shape ``(batch_size, latent_dim)``.
  369. logVar : torch.Tensor
  370. Latent log-variances, shape ``(batch_size, latent_dim)``.
  371. nearest_pt : torch.Tensor
  372. Prototype indices for correct class, shape
  373. ``(batch_size, num_prototypes_per_class)``.
  374. sim_scores : torch.Tensor
  375. Similarity scores, shape ``(batch_size, num_prototypes)``.
  376. Returns
  377. -------
  378. kl_loss : torch.Tensor
  379. Scalar KL divergence loss.
  380. mask : torch.Tensor
  381. Boolean mask indicating contributing prototypes.
  382. """
  383. kl_loss = torch.zeros(sim_scores.shape).to(device)
  384. half_latent = self.latent_dim // 2
  385. for i in range(self.num_prototypes_per_class):
  386. p_v = self.prototype_vectors[nearest_pt[:, i], :] # all class prototype i vector
  387. kl1 = torch.distributions.kl.kl_divergence(
  388. torch.distributions.Normal(mu[:, :half_latent], torch.exp(logVar[:, :half_latent] / 2)),
  389. torch.distributions.Normal(p_v[:, :half_latent], torch.ones(p_v[:, :half_latent].shape).to(device))
  390. )
  391. kl2 = torch.distributions.kl.kl_divergence(
  392. torch.distributions.Normal(mu[:, half_latent:], torch.exp(logVar[:, half_latent:] / 2)),
  393. # torch.distributions.Normal(torch.zeros_like(p_v[:, half_latent:]).to(device), torch.ones_like(p_v[:, half_latent:]).to(device))
  394. torch.distributions.Normal(torch.zeros_like(p_v[:, half_latent:]), torch.ones_like(p_v[:, half_latent:]).to(device))
  395. )
  396. kl = torch.mean(kl1 * 5 + kl2, dim=-1)
  397. kl_loss[np.arange(sim_scores.shape[0]), nearest_pt[:, i]] = kl
  398. kl_loss = kl_loss * sim_scores # element-wise scale by similarity scores
  399. mask = kl_loss > 0 # prototypes contributes
  400. kl_loss = torch.sum(kl_loss, dim = -1) / (torch.sum(sim_scores * mask, dim = -1))
  401. kl_loss = torch.mean(kl_loss)
  402. return kl_loss, mask
  403. def orthogonal_loss(self):
  404. """Compute orthogonality loss for prototype diversity.
  405. Encourages prototypes within the same class to be orthogonal
  406. in the first half of the latent space. Also includes L1
  407. sparsity regularization on the first encoder layer weights.
  408. Returns
  409. -------
  410. torch.Tensor
  411. Scalar orthogonality + sparsity loss.
  412. """
  413. s_loss = 0
  414. for k in range(self.num_classes):
  415. # p_k = self.prototype_vectors[k*self.num_prototypes_per_class : (k+1)*self.num_prototypes_per_class, :]
  416. p_k = self.prototype_vectors[k*self.num_prototypes_per_class : (k+1)*self.num_prototypes_per_class, :self.latent_dim//2]
  417. p_k_mean = torch.mean(p_k, dim = 0)
  418. p_k_2 = p_k - p_k_mean
  419. p_k_dot = p_k_2 @ p_k_2.T
  420. s_matrix = p_k_dot - (torch.eye(p_k.shape[0]).to(device))
  421. s_loss += torch.norm(s_matrix, p = 2)
  422. # # L1 regularization
  423. sparsity = 1.0 / torch.numel(self.encoder[0][0].weight) * torch.norm(self.encoder[0][0].weight, 1)
  424. return s_loss / self.num_classes + sparsity
  425. def atomic_loss(self, sim_scores, mask):
  426. """Compute attraction/repulsion loss for prototype assignment.
  427. Encourages high similarity to correct-class prototypes
  428. (attraction) and low similarity to incorrect-class prototypes
  429. (repulsion).
  430. Parameters
  431. ----------
  432. sim_scores : torch.Tensor
  433. Prototype similarities, shape ``(batch_size, num_prototypes)``.
  434. mask : torch.Tensor
  435. Boolean mask for correct-class prototypes.
  436. Returns
  437. -------
  438. torch.Tensor
  439. Scalar loss (repulsion - attraction).
  440. """
  441. attraction = torch.mean(torch.max(sim_scores * mask, 1).values)
  442. repulsion = torch.mean(torch.max(sim_scores * torch.logical_not(mask), 1).values)
  443. # repulsion = torch.sum(torch.mean(sim_scores * torch.logical_not(mask), 1).values)
  444. return repulsion - attraction
  445. def set_last_layer_incorrect_connection(self, incorrect_strength):
  446. """Initialize classifier weights based on prototype-class identity.
  447. Sets weights to 1 for correct class connections and
  448. ``incorrect_strength`` for incorrect class connections.
  449. Parameters
  450. ----------
  451. incorrect_strength : float
  452. Weight for incorrect class connections (e.g., -0.5).
  453. """
  454. positive_one_weights_locations = torch.t(self.prototype_class_identity)
  455. negative_one_weights_locations = 1 - positive_one_weights_locations
  456. correct_class_connection = 1
  457. incorrect_class_connection = incorrect_strength
  458. self.classifier.weight.data.copy_(
  459. correct_class_connection * positive_one_weights_locations
  460. + incorrect_class_connection * negative_one_weights_locations)
  461. def _initialize_weights(self):
  462. '''
  463. initialize weights for vae
  464. '''
  465. for m in self.encoder.modules():
  466. if isinstance(m, nn.Linear):
  467. nn.init.uniform_(m.weight, -0.08, 0.08)
  468. if m.bias is not None:
  469. nn.init.constant_(m.bias, 0)
  470. elif isinstance(m, nn.BatchNorm1d):
  471. nn.init.constant_(m.weight, 1)
  472. nn.init.constant_(m.bias, 0.001)
  473. for m in self.decoder.modules():
  474. if isinstance(m, nn.Linear):
  475. nn.init.uniform_(m.weight, -0.08, 0.08)
  476. if m.bias is not None:
  477. nn.init.constant_(m.bias, 0)
  478. elif isinstance(m, nn.BatchNorm1d):
  479. nn.init.constant_(m.weight, 1)
  480. nn.init.constant_(m.bias, 0.001)
  481. self.set_last_layer_incorrect_connection(incorrect_strength = -0.5)
  482. # get results helper functions
  483. #######################################################
  484. @property
  485. def get_prototypes(self):
  486. """torch.Tensor : Prototype vectors of shape ``(num_prototypes, latent_dim)``."""
  487. return self.prototype_vectors
  488. def get_prototype_cells(self):
  489. """Decode prototype vectors to gene expression space.
  490. Samples 100 times from the reconstruction distribution for each
  491. prototype and returns the mean counts.
  492. Returns
  493. -------
  494. torch.Tensor
  495. Prototype cells in gene space, shape
  496. ``(num_prototypes, input_dim)``.
  497. """
  498. px_mu, px_theta = self.get_latent_decode(self.prototype_vectors)
  499. # sample 100 and take avg for each
  500. proto_cells = torch.zeros(self.num_prototypes, self.input_dim)
  501. for i in range(self.num_classes):
  502. x_mu = px_mu[i*self.num_prototypes_per_class : (i+1)*self.num_prototypes_per_class, :]
  503. for j in range(self.num_prototypes_per_class):
  504. t = px_theta[:,i]
  505. mu = x_mu[j]
  506. proto_cells[i*self.num_prototypes_per_class + j, :] = torch.mean(self.sample_recon(mu, t, 100), axis=0)
  507. return proto_cells
  508. def max_sim_score(self, sim_scores):
  509. """Find the maximum similarity score per class and the best-matching prototype.
  510. Parameters
  511. ----------
  512. sim_scores : torch.Tensor
  513. Similarity scores, shape ``(batch_size, num_prototypes)``.
  514. Returns
  515. -------
  516. max_sim : torch.Tensor
  517. Maximum similarity to nearest prototype, shape ``(batch_size,)``.
  518. nearest_proto_idx : torch.Tensor
  519. Index of the nearest prototype within its class, shape ``(batch_size,)``.
  520. """
  521. sim_reshaped = sim_scores.view(-1, self.num_classes, self.num_prototypes_per_class)
  522. # max sim to a prototype for each class
  523. max_sim_per_cls, max_proto_indices_per_class = torch.max(sim_reshaped, dim=2)
  524. # Find the nearest class for each cell
  525. max_sim, nearest_cls_idx = torch.max(max_sim_per_cls, dim=1)
  526. # The indices of the nearest prototypes
  527. nearest_proto_idx = max_proto_indices_per_class[range(sim_reshaped.shape[0]), nearest_cls_idx]
  528. return max_sim, nearest_proto_idx
  529. def get_pred(self, x, test=False):
  530. """Get classification predictions.
  531. Parameters
  532. ----------
  533. x : torch.Tensor
  534. Input gene expression, shape ``(batch_size, input_dim)``.
  535. test : bool, optional
  536. If True, use latent mean only (deterministic); otherwise
  537. sample via reparameterization, by default False.
  538. Returns
  539. -------
  540. pred : torch.Tensor
  541. Classification logits, shape ``(batch_size, num_classes)``.
  542. max_sim : torch.Tensor
  543. Maximum similarity to nearest prototype.
  544. proto_idx : torch.Tensor
  545. Index of the nearest prototype.
  546. """
  547. self.eval()
  548. if self.raw_input: # raw: 1
  549. x = torch.log1p(x)
  550. encode = self.encoder(x)
  551. z_mu = self.z_mean(encode)
  552. if test:
  553. z = z_mu
  554. else:
  555. z_logVar = self.z_log_var(encode)
  556. z = self.reparameterize(z_mu, z_logVar)
  557. sim_scores = self.calc_sim_scores(z)
  558. pred = self.classifier(sim_scores)
  559. max_sim, proto_idx = self.max_sim_score(sim_scores)
  560. return pred, max_sim, proto_idx
  561. def get_latent(self, x):
  562. """Encode input to latent mean (deterministic, no sampling).
  563. Parameters
  564. ----------
  565. x : torch.Tensor
  566. Input gene expression, shape ``(batch_size, input_dim)``.
  567. Returns
  568. -------
  569. torch.Tensor
  570. Latent means, shape ``(batch_size, latent_dim)``.
  571. """
  572. self.eval()
  573. if self.raw_input:
  574. x = torch.log1p(x)
  575. encode = self.encoder(x)
  576. z_mu = self.z_mean(encode)
  577. return z_mu
  578. def get_latent_decode(self, z):
  579. """Decode latent vectors to gene expression space.
  580. Parameters
  581. ----------
  582. z : torch.Tensor
  583. Latent embeddings, shape ``(batch_size, latent_dim)``.
  584. Returns
  585. -------
  586. px_mu : torch.Tensor
  587. Reconstruction mean, shape ``(batch_size, input_dim)``.
  588. px_t : torch.Tensor or None
  589. Dispersion parameters, or None for normal distribution.
  590. """
  591. px = self.decoder(z)
  592. px_mu = self.px_mean(px)
  593. px_t = self.px_theta(px)
  594. if self.obs_dist == 'nb':
  595. px_mu = self.softmax(px_mu) * self.num_prototypes
  596. if self.nb_dispersion.startswith('celltype'):
  597. px_t = self.theta
  598. elif self.nb_dispersion == 'gene':
  599. px_t = self.px_theta(px)
  600. px_t = torch.mean(px_t, 0, True)
  601. else:
  602. raise NotImplementedError
  603. px_t = torch.clamp(px_t, min = EPS)
  604. else:
  605. px_t = None
  606. return px_mu, px_t
  607. def get_recon(self, x):
  608. """Reconstruct input via full encode-decode pipeline.
  609. Parameters
  610. ----------
  611. x : torch.Tensor
  612. Input gene expression, shape ``(batch_size, input_dim)``.
  613. Returns
  614. -------
  615. px_mu : torch.Tensor
  616. Reconstruction mean.
  617. px_t : torch.Tensor or None
  618. Dispersion parameters, or None for normal distribution.
  619. """
  620. self.eval()
  621. if self.raw_input: # raw: 1
  622. x = torch.log1p(x)
  623. encode = self.encoder(x)
  624. z_mu = self.z_mean(encode)
  625. z_logVar = self.z_log_var(encode)
  626. z = self.reparameterize(z_mu, z_logVar)
  627. px_mu, px_t = self.get_latent_decode(z)
  628. return px_mu, px_t
  629. def get_log_likelihood(self, input, target=None):
  630. """Compute negative log-likelihood averaged over multiple samples.
  631. Samples 5 times with different random seeds and averages the
  632. negative binomial log-likelihood across samples.
  633. Parameters
  634. ----------
  635. input : torch.Tensor
  636. Input gene expression, shape ``(batch_size, input_dim)``.
  637. target : torch.Tensor, optional
  638. Ground truth labels, required when ``nb_dispersion`` is
  639. ``'celltype_target'``.
  640. Returns
  641. -------
  642. torch.Tensor
  643. Mean negative log-likelihood per cell, shape ``(batch_size,)``.
  644. Raises
  645. ------
  646. NotImplementedError
  647. If ``obs_dist`` is not ``'nb'`` or if ``target`` is required
  648. but not provided.
  649. """
  650. if self.obs_dist != 'nb':
  651. raise NotImplementedError
  652. elif self.nb_dispersion == 'celltype_target' and target is None:
  653. print("Provide target label for log-likelihood calculation due to your choice of dispersion")
  654. raise NotImplementedError
  655. self.eval()
  656. n_sample = 5
  657. ll_value = 0
  658. for i in range(n_sample):
  659. with torch.no_grad():
  660. seed_torch(torch.device(device), seed = i, msg=False)
  661. pred, px_mu, px_t, _, _, _ = self.forward(input)
  662. if self.nb_dispersion == 'celltype_target':
  663. # data target
  664. assert target is not None
  665. dispersion = F.linear(one_hot_encoder(target, self.num_classes), self.theta)
  666. dispersion = torch.exp(dispersion)
  667. elif self.nb_dispersion == 'celltype_pred':
  668. # pred target
  669. softmax_pred = F.softmax(pred, dim=1)
  670. max_index = torch.multinomial(softmax_pred, 1)
  671. dispersion = F.linear(one_hot_encoder(max_index, self.num_classes), self.theta)
  672. dispersion = torch.exp(dispersion)
  673. elif self.nb_dispersion == 'gene':
  674. dispersion = px_t
  675. else:
  676. raise NotImplementedError
  677. ll = -log_likelihood_nb(input, px_mu, dispersion)
  678. ll = torch.sum(ll, dim = -1)
  679. ll_value += ll
  680. del px_mu, px_t, dispersion, ll
  681. torch.cuda.empty_cache()
  682. return ll_value / n_sample
  683. def sample_recon(self, px_mu, px_t, sample_size):
  684. """Sample from the negative binomial reconstruction distribution.
  685. Uses a Gamma-Poisson compound to generate count samples.
  686. Parameters
  687. ----------
  688. px_mu : torch.Tensor
  689. Reconstruction mean.
  690. px_t : torch.Tensor
  691. Dispersion (concentration) parameter.
  692. sample_size : int
  693. Number of samples to draw.
  694. Returns
  695. -------
  696. torch.Tensor
  697. Sampled counts, shape ``(sample_size, ..., input_dim)``.
  698. """
  699. concentration = px_t
  700. rate = px_t / px_mu
  701. # Gamma(alpha, beta: rate = 1/scale)
  702. gamma_d = Gamma(concentration=concentration, rate=rate)
  703. p_means = gamma_d.rsample((sample_size,))
  704. l_train = torch.clamp(p_means, max=1e8)
  705. counts = Poisson(l_train).sample() # (n_samples, n_cells, n_vars)
  706. return counts

model.py at commit 90c2678, under MIT · at the source

Overview

  1. Department of Computer Science, University of British Columbia, Vancouver, BC V6T 1Z4, Canada
Institutions: University of British Columbia (Canada)
Journal: Cell genomics, volume 6, issue 6, article 101217
Dates: received 8 May 2025; accepted 17 March 2026; published online 16 April 2026; in print June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1016/j.xgen.2026.101217 · PMID 41997134 · PMCID PMC13261663 · OpenAlex W7154597936
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: genetics / omics (modality), human (organism), cellular / molecular (subfield)
Methods: Statistics, Smoothing, state filtering, decompositions, Machine learning, Connectivity
Keywords: single-cell RNA sequencing, cell state, rare cell type, prototypical network, self-explaining model, disentanglement, layer-wise relevance propagation, prototypical relevance propagation, deep generative models, variational autoencoder
MeSH: Single-Cell Gene Expression Analysis*, Animals, Genomics, Humans (* major topic)
Journal subjects: Technology
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: Natural Sciences and Engineering Research Council of Canada; University of British Columbia; Canadian Institutes of Health Research; Canada Research Chair Program; Canada Foundation for Innovation & John. R. Evans Leader Fund
Citations: not cited yet (Europe PMC); 96 references in the paper

Abstract

Cell type annotation is a fundamental task in single-cell genomics. Although various methods have been developed for automatic annotation, they often function as black-box models lacking explainability, proper uncertainty estimation, and robustness for rare cell types. We introduce ProtoCloud, a self-explanatory deep generative model that embeds cells into a structured, low-dimensional space organized around cell-type-specific prototypes. ProtoCloud matches or outperforms existing methods across 11 large-scale datasets, particularly for rare cell types. Its built-in uncertainty quantification mechanism, based on cell-prototype similarity, identifies and re-annotates misannotated training cells. By backpropagating cell prototype similarities to the gene space, ProtoCloud identifies key genes driving its classifications, facilitating the discovery of both known and novel marker genes. Applied to a time-course dataset of post-injury retinal neurons, ProtoCloud successfully annotates previously unassigned cells; in an esophageal cell atlas, it identifies rare but potentially important cell populations and their marker genes associated with esophageal inflammation.

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

Ding-Group/ProtoCloud

License: MIT
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: 90c2678f9ce54e1e9911bd1aa5086c63a7798d89, 11 August 2026
Languages: Python (17), Jupyter (1)
Size: 50 files, 18 scripts
Software Heritage: not archived
Found in: “Data and code availability”
Holds: README, license file, environment (pyproject.toml), 1 notebook
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (10 files), PyTorch (10 files), pandas (9 files), scikit-learn (6 files), anndata (5 files), Scanpy (4 files), SciPy (4 files), Matplotlib (2 files), seaborn (2 files), h5py (1 file), UMAP (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
20 files

Zenodo 18882740

License: MIT
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: “Data and code availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (8 files), pandas (8 files), PyTorch (8 files), scikit-learn (5 files), anndata (4 files), Scanpy (3 files), SciPy (3 files), Matplotlib (2 files), seaborn (2 files), h5py (1 file), UMAP (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers (HTTP 200)
  • 29 September 2026: the link answers (HTTP 200)
18 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 34 scripts, each with its path and the digest of its content;
  • 21 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 and code availability

• All data used in this study are publicly available. The datasets used are available through their sources in the key resources table. • All original code has been deposited at https://github.com/Ding-Group/ProtoCloud and achieved in Zenodo (https://doi.org/10.5281/zenodo.18882740). • Any additional information required to reanalyze the data reported in this paper is available from the lead contact upon request.

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, 2 authors, 10 keywords, 4 MeSH terms, 5 funders, 86 references.

Cite

This paper

Guo, K., & Ding, J. (2026). ProtoCloud: A prototypical self-explaining model for single-cell analysis. Cell genomics, 6(6), 101217. https://doi.org/10.1016/j.xgen.2026.101217

BibTeX

@article{guo2026protocloud,
author = {Guo, Kaiyun and Ding, Jiarui},
title = {{ProtoCloud: A prototypical self-explaining model for single-cell analysis}},
journal = {Cell genomics},
year = {2026},
month = apr,
volume = {6},
number = {6},
pages = {101217},
publisher = {Elsevier},
issn = {2666-979X},
doi = {10.1016/j.xgen.2026.101217},
url = {https://doi.org/10.1016/j.xgen.2026.101217},
pmid = {41997134},
pmcid = {PMC13261663}
}

RIS

TY - JOUR
AU - Guo, Kaiyun
AU - Ding, Jiarui
TI - ProtoCloud: A prototypical self-explaining model for single-cell analysis
T2 - Cell genomics
J2 - Cell Genom
PY - 2026
DA - 2026/04/16
VL - 6
IS - 6
SP - 101217
SN - 2666-979X
PB - Elsevier
DO - 10.1016/j.xgen.2026.101217
UR - https://doi.org/10.1016/j.xgen.2026.101217
LA - en
ER -

CSL-JSON

{
"id": "10.1016/j.xgen.2026.101217",
"type": "article-journal",
"title": "ProtoCloud: A prototypical self-explaining model for single-cell analysis",
"container-title": "Cell genomics",
"author": [
{
"family": "Guo",
"given": "Kaiyun"
},
{
"family": "Ding",
"given": "Jiarui"
}
],
"container-title-short": "Cell Genom",
"volume": "6",
"issue": "6",
"page": "101217",
"DOI": "10.1016/j.xgen.2026.101217",
"PMID": "41997134",
"PMCID": "PMC13261663",
"ISSN": "2666-979X",
"publisher": "Elsevier",
"URL": "https://doi.org/10.1016/j.xgen.2026.101217",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
16
]
]
}
}

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.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, 8 other tools, 9 references
[2] doi:10.1038/s44320-026-00208-7 [code]
Interpretable deep generative ensemble learning for single-cell omics with Hydra.
Journal: Molecular systems biology
In common: UMAP, anndata, Scanpy, 8 other tools, cellular / molecular, 6 references
[3] doi:10.64898/2026.03.30.714220 [code]
An integrated single cell and spatial omics atlas of human prenatal development
Journal: bioRxiv (preprint)
In common: UMAP, anndata, Scanpy, 8 other tools, 5 references
[4] doi:10.1371/journal.pcbi.1014327 [code]
Supervised deep learning with gene functional annotation for cell classification.
Journal: PLoS computational biology
In common: anndata, Scanpy, h5py, 7 other tools, genetics / omics, cellular / molecular, 4 references
[5] doi:10.1038/s41467-026-71759-4 [code]
CellNiche represents cellular microenvironments in atlas-scale spatial omics data with contrastive learning.
Journal: Nature communications
In common: anndata, Scanpy, PyTorch, 6 other tools, cellular / molecular, 6 references
[6] doi:10.1016/j.isci.2026.116055 [code]
Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.
Journal: iScience
In common: UMAP, anndata, Scanpy, 8 other tools, genetics / omics, 3 references
[7] doi:10.1093/bib/bbag490 [code]
Systematic benchmarking and optimal strategy selection of cross-species integration methods.
Journal: Briefings in bioinformatics
In common: anndata, Scanpy, seaborn, 5 other tools, genetics / omics, 6 references
[8] doi:10.1093/bioinformatics/btag652 [code]
mmVelo: a deep generative model for estimating cell state-dependent dynamics across multiple modalities.
Journal: Bioinformatics (Oxford, England)
In common: UMAP, anndata, Scanpy, 7 other tools, genetics / omics, 4 references
[9] doi:10.1016/j.xcrm.2026.102766 [code]
A longitudinal single-cell and spatial multiomic atlas of pediatric high-grade glioma.
Journal: Cell reports. Medicine
In common: UMAP, anndata, Scanpy, 7 other tools, genetics / omics, cellular / molecular, 3 references
[10] 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: anndata, Scanpy, h5py, 7 other tools, genetics / omics, 4 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.