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),
    )