File size: 4,140 Bytes
925ee3b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 | """Posterior sampling and extraction for DeepPTR."""
from __future__ import annotations
from typing import TYPE_CHECKING
import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, TensorDataset
if TYPE_CHECKING:
from anndata import AnnData
from ._model import DeepPTR
def posterior_gamma(
model: "DeepPTR",
adata: "AnnData",
n_samples: int = 50,
batch_size: int = 512,
device: str | None = None,
) -> tuple[np.ndarray, np.ndarray]:
"""Compute posterior mean and variance of gamma via MC sampling.
Parameters
----------
model
Trained :class:`DeepPTR` model.
adata
AnnData with ``layers['spliced']`` and ``layers['unspliced']``.
n_samples
Number of MC samples from the posterior.
batch_size
Inference batch size.
device
Torch device.
Returns
-------
gamma_mean, gamma_var : np.ndarray
Each shape ``(n_obs, n_genes)``, dtype float32.
"""
from scipy.sparse import issparse
if device is None:
device = next(model.parameters()).device
else:
device = torch.device(device)
model.eval()
def _dense(mat):
if issparse(mat):
return np.asarray(mat.todense())
return np.asarray(mat)
s_np = _dense(adata.layers["spliced"]).astype(np.float32)
u_np = _dense(adata.layers["unspliced"]).astype(np.float32)
s_t = torch.from_numpy(s_np)
u_t = torch.from_numpy(u_np)
ds = TensorDataset(s_t, u_t)
loader = DataLoader(ds, batch_size=batch_size, shuffle=False)
n_obs = adata.n_obs
n_genes = adata.n_vars
# Accumulators (Welford online mean/var)
mean_acc = np.zeros((n_obs, n_genes), dtype=np.float64)
m2_acc = np.zeros((n_obs, n_genes), dtype=np.float64)
with torch.no_grad():
for k in range(n_samples):
gamma_list: list[np.ndarray] = []
for s_b, u_b in loader:
s_b = s_b.to(device)
u_b = u_b.to(device)
mu_T, logvar_T, mu_PT, logvar_PT = model.encoder(s_b, u_b)
z_PT = model.reparameterize(mu_PT, logvar_PT)
gamma_b = F.softplus(model.decoder.f_gamma(z_PT))
gamma_list.append(gamma_b.cpu().numpy())
gamma_k = np.concatenate(gamma_list, axis=0) # (n_obs, n_genes)
# Welford update
delta = gamma_k - mean_acc
mean_acc += delta / (k + 1)
delta2 = gamma_k - mean_acc
m2_acc += delta * delta2
gamma_mean = mean_acc.astype(np.float32)
gamma_var = (m2_acc / max(n_samples - 1, 1)).astype(np.float32)
return gamma_mean, gamma_var
def extract_latent(
model: "DeepPTR",
adata: "AnnData",
batch_size: int = 512,
device: str | None = None,
) -> tuple[np.ndarray, np.ndarray]:
"""Extract posterior mean of z_T and z_PT for all cells.
Returns
-------
z_T, z_PT : np.ndarray
Shapes ``(n_obs, d_T)`` and ``(n_obs, d_PT)``.
"""
from scipy.sparse import issparse
if device is None:
device = next(model.parameters()).device
else:
device = torch.device(device)
model.eval()
def _dense(mat):
if issparse(mat):
return np.asarray(mat.todense())
return np.asarray(mat)
s_np = _dense(adata.layers["spliced"]).astype(np.float32)
u_np = _dense(adata.layers["unspliced"]).astype(np.float32)
s_t = torch.from_numpy(s_np)
u_t = torch.from_numpy(u_np)
ds = TensorDataset(s_t, u_t)
loader = DataLoader(ds, batch_size=batch_size, shuffle=False)
z_T_list: list[np.ndarray] = []
z_PT_list: list[np.ndarray] = []
with torch.no_grad():
for s_b, u_b in loader:
s_b = s_b.to(device)
u_b = u_b.to(device)
mu_T, _, mu_PT, _ = model.encoder(s_b, u_b)
z_T_list.append(mu_T.cpu().numpy())
z_PT_list.append(mu_PT.cpu().numpy())
return (
np.concatenate(z_T_list, axis=0).astype(np.float32),
np.concatenate(z_PT_list, axis=0).astype(np.float32),
)
|