import os
import sys
import numpy as np
import pandas as pd
import scanpy as sc
import torch
import torch.nn as nn
import torch.nn.functional as F
import pyro
import pyro.poutine as poutine
import pyro.distributions as dist
from ml_collections import ConfigDict
from torch_geometric.data import Data
from torch_geometric.loader import DataLoader
from torch_geometric.nn import GATConv, GCNConv
import torch_scatter
sys.path.append(os.path.dirname(os.path.realpath(__file__)))
from base_model import BaseModel
from module import Prior, StructuralPrior, ConvPrior
from module import Encoder, XtoZEncoder, ConvXtoZEncoder, XtoVEncoder, XtoOmegaCluEncoder, XtoKappaEncoder
from module import Decoder, ZtoSDecoder, StoXDecoder
from module import hsic
from dataset import XeniumDataset, HeteroDataset
EPS = 1e-8
[docs]
class HeteroAttnVGAE(BaseModel):
r"""Learning latent manifold w/ Conditional VGAE on hetero-graph
Generative path: DESI (u) -> Latent (z) -> Xenium (x)
"""
[docs]
def __init__(
self,
configs: ConfigDict,
device: torch.device = torch.device('cuda')
):
super().__init__(configs, device)
self.act = configs.act
# Parse node & edge types
self.ref = configs.ref
self.query = configs.query
self.r2q = (self.ref, 'to', self.query)
self.q2r = (self.query, 'to', self.ref)
self.r2r = (self.ref, 'to', self.ref)
# Whether to use conv. prior / posterior for `u` (i.e. histology patches)
self.patch_size = configs.patch_size if hasattr(configs, 'patch_size') else -1
self.prior = StructuralPrior(configs) if self.patch_size < 0 else ConvPrior(configs)
self.encode_z = XtoZEncoder(configs) if self.patch_size < 0 else ConvXtoZEncoder(configs)
self.encode_kappa = XtoKappaEncoder(configs)
self.encode_omega = XtoOmegaCluEncoder(configs)
self.kappa_mu = nn.Embedding(configs.n_cluster, configs.c_latent)
self.kappa_logvar = nn.Embedding(configs.n_cluster, configs.c_latent)
self.decode_s = ZtoSDecoder(configs)
self.decode_x = nn.Sequential(
nn.Linear(configs.c_latent, configs.c_hidden),
configs.act,
nn.Linear(configs.c_hidden, configs.c_in)
)
[docs]
def model(self, data):
pyro.module("VAE", self)
u = data[self.query].x
x = data[self.ref].x
clusters = data[self.ref].cluster
l = x.sum(axis=-1, keepdim=True)
# Reshape image patches if paired with histology
if self.patch_size > 0:
u = self._reshape_patches(u)
# Graph properties
edge_index_dict = data.edge_index_dict
edge_attr_dict = data.edge_attr_dict
edge_index = edge_index_dict[self.r2r]
edge_distances = edge_attr_dict[self.r2r]
src, dst = edge_index
# q2r edges for unpooling z -> cells
q2r_src, q2r_dst = edge_index_dict[self.q2r]
# --- Global parameters ---
# \theta: gene-specific dispersion
theta = pyro.param(
"theta",
torch.ones(self.configs.c_in, dtype=torch.float),
constraint=dist.constraints.positive
).to(self.device)
# Plates
cell_plate = pyro.plate("cell", x.size(0))
# -----------------------
# Sample z ~ p(z | u)
# -----------------------
with pyro.plate("patch", u.size(0)):
z_mu, z_logvar = self.prior(u, edge_index_dict)
z_dist = dist.Normal(z_mu, torch.exp(z_logvar/2))
z = pyro.sample("z", z_dist.to_event(1))
# --------------------------------------------
# Deterministic unpooling z -> z_cell
# --------------------------------------------
z_cell = torch_scatter.scatter_mean(
z[q2r_src], q2r_dst, dim=0, dim_size=x.size(0)
)
if self.configs.infer_cell_interaction:
# ---------------------------------------------
# Sample omega_raw ~ p(omega_raw | distance)
# ---------------------------------------------
# log_d = torch.log1p(edge_distances)
alpha = self.configs.alpha
scale = 1.0 / (edge_index.size(1) / x.size(0))
with pyro.plate("r2r_edges", edge_distances.size(0)):
with poutine.scale(scale=scale):
exponential_dist = dist.Exponential(
(1+edge_distances).pow(alpha)
)
omega = pyro.sample("omega", exponential_dist).squeeze(-1)
# --------------------------
# Sample kappa ~ p(kappa | cluster)
# (type-anchored intrinsic state; per-cell latent)
# --------------------------
with cell_plate:
C = int(clusters.max().item()) + 1
counts = torch.bincount(clusters, minlength=C).float().to(self.device)
w = 1.0 / (counts[clusters] + 1e-8)
kappa_mu = self.kappa_mu(clusters)
kappa_std = torch.ones(self.configs.c_latent, dtype=torch.float).to(self.device)
with poutine.scale(scale=w):
kappa = pyro.sample(
"kappa",
dist.Normal(kappa_mu, kappa_std).to_event(1)
)
# -----------------------------------------------------
# delta = z_cell - kappa (what neighbors transmit)
# -----------------------------------------------------
delta = z_cell - kappa
delta = delta - delta.mean(dim=0, keepdim=True)
msg = self._weighted_sum(edge_index, omega, delta)
s_prime = kappa + msg
else:
s_prime = z_cell
# --------------------------------------
# Reconstruct x from p(x | s_prime)
# --------------------------------------
with cell_plate:
mu = torch.softmax(self.decode_x(s_prime), dim=-1)
x_mu = l * mu
logits = (x_mu + EPS).log() - (theta + EPS).log()
nb_dist = dist.NegativeBinomial(total_count=theta, logits=logits)
pyro.sample("x", nb_dist.to_event(1), obs=x)
[docs]
def guide(self, data):
pyro.module("VAE", self)
x = data[self.ref].x
u = data[self.query].x
x = self.lognorm(x)
# Reshape image patches if paired with histology
if self.patch_size > 0:
u = self._reshape_patches(u)
edge_index_dict = data.edge_index_dict
edge_index = edge_index_dict[self.r2r]
src, dst = edge_index
# q2r edges for unpooling z -> cells
q2r_src, q2r_dst = edge_index_dict[self.q2r]
# Plates
cell_plate = pyro.plate("cell", x.size(0))
# -------------------------
# Sample z ~ q(z | x, u)
# -------------------------
with pyro.plate("patch", u.size(0)):
z_mu, z_logvar, _ = self.encode_z(x, u, edge_index_dict)
z_dist = dist.Normal(z_mu, torch.exp(z_logvar / 2))
z = pyro.sample("z", z_dist.to_event(1))
# deterministic unpool z -> z_cell
z_cell = torch_scatter.scatter_mean(
z[q2r_src], q2r_dst, dim=0, dim_size=x.size(0)
)
if self.configs.infer_cell_interaction:
clusters = data[self.ref].cluster
C = int(clusters.max().item()) + 1
# ----------------------------
# Sample kappa ~ q(kappa | x)
# ----------------------------
counts = torch.bincount(clusters, minlength=C).float().to(self.device)
w = 1.0 / (counts[clusters] + 1e-8)
with cell_plate:
with poutine.scale(scale=w):
kappa_mu, kappa_logvar = self.encode_kappa(x)
kappa = pyro.sample(
"kappa",
dist.Normal(kappa_mu, torch.exp(kappa_logvar / 2)).to_event(1)
)
# -------------------------------
# Sample omega_raw ~ q(omega_raw | x, z_cell)
# -------------------------------
omega_loc = self.encode_omega(x, z_cell, edge_index_dict).squeeze(-1)
scale = 1.0 / (edge_index.size(1) / x.size(0))
with pyro.plate("r2r_edges", omega_loc.size(0)):
with poutine.scale(scale=scale):
omega = pyro.sample("omega", dist.Delta(omega_loc)).squeeze(-1)
# -----------------------------------------------------
# delta = z_cell - kappa (transmittable component)
# -----------------------------------------------------
delta = z_cell - kappa
delta = delta - delta.mean(dim=0, keepdim=True) # global centering
# Regularization w/ HSIC btw kappa & m
# neighbor message = weighted neighborhood mean(delta)
msg = self._weighted_sum(edge_index, omega, delta)
k0 = kappa - kappa.mean(dim=0, keepdim=True)
m0 = msg - msg.mean(dim=0, keepdim=True)
hsic_loss = hsic(m0, k0.detach()) # detach kappa for stability
pyro.factor(
"hsic_indep",
1e-3 * hsic_loss,
has_rsample=True
)
[docs]
def predict(self, data, device):
with torch.no_grad():
data = data.to(device)
# Observed data
x = data[self.ref].x
l = x.sum(axis=-1, keepdim=True)
x = self.lognorm(x)
u = data[self.query].x
# Reshape image patches if paired with histology
if self.patch_size > 0:
u = self._reshape_patches(u)
edge_index_dict = data.edge_index_dict
edge_index = edge_index_dict[self.r2r]
n_edges = edge_index.size(1)
# q2r edges for unpooling patch z -> cell z
q2r_src, q2r_dst = edge_index_dict[self.q2r]
# ---------- p(z | u) ----------
pz_mu, _ = self.prior(u, edge_index_dict)
# ---------- q(z | x, u) ----------
qz_mu, _, _ = self.encode_z(x, u, edge_index_dict)
# ---------- unpool to cells ----------
z_cell = torch_scatter.scatter_mean(
qz_mu[q2r_src], q2r_dst, dim=0, dim_size=x.size(0)
)
infer_cci = self.configs.infer_cell_interaction
if infer_cci:
# ---------- omega (posterior mean / location) ----------
# TODO: directly use exponential dist.
q_omega = self.encode_omega(x, z_cell, edge_index_dict).squeeze(-1)
# ---------- kappa (posterior mean) ----------
q_kappa_mu, _ = self.encode_kappa(x)
kappa = q_kappa_mu
clusters = data[self.ref].cluster
C = int(clusters.max().item()) + 1
# ---------- delta = z_cell - kappa ----------
delta = z_cell - kappa
delta = delta - delta.mean(dim=0, keepdim=True)
# ---------- msg ----------
msg = self._weighted_sum(edge_index, q_omega, delta)
# ---------- intrinsic + extrinsic ----------
qs = z_cell
s_prime = kappa + msg
mu = torch.softmax(self.decode_x(s_prime), dim=-1)
else:
qs = z_cell
mu = torch.softmax(self.decode_x(qs), dim=-1)
q_omega = None
px = l * mu
return ConfigDict({
"qz": qz_mu,
"qs": qs,
"pz": pz_mu,
"px": px,
"omega": q_omega
})
[docs]
def fit(self, dataset, train_configs, DEBUG=False, log_wandb=False):
super().model_train(
self, dataset, train_configs, key=self.ref,
DEBUG=DEBUG, log_wandb=log_wandb
)
return None
[docs]
def evaluate(
self,
adata_ref: sc.AnnData,
adata_query: sc.AnnData,
graph_data: HeteroDataset,
n_subgraphs: int = 1,
device: torch.device = torch.device('cuda')
):
self.eval()
self.device = device
self.to(device)
n_cells, n_features = adata_ref.shape
n_pixels, _ = adata_query.shape
n_clusters = graph_data.num_clusters
full_graph_data = HeteroDataset(
adatas_ref=adata_ref,
adatas_query=adata_query,
n_subgraphs=n_subgraphs,
k=graph_data.k, r=graph_data.r, alpha=self.configs.alpha,
cluster_key=graph_data.cluster_key,
num_clusters=n_clusters,
is_weighted=graph_data.is_weighted,
ref=graph_data.ref, ref_proj_key=graph_data.ref_proj_key,
query=graph_data.query, query_proj_key=graph_data.query_proj_key,
is_ref_grid=graph_data.is_ref_grid,
is_query_grid=graph_data.is_query_grid,
verbose=False
)
dataloader = DataLoader(full_graph_data, shuffle=False)
qz = np.zeros((n_pixels, self.configs.c_latent), dtype=np.float32) # lowres latent
qs = np.zeros((n_cells, self.configs.c_latent), dtype=np.float32) # hires latent
pz = np.zeros_like(qz)
px = np.zeros((n_cells, n_features), dtype=np.float32)
# Summarized cell-type specific attention scores
qomega_scores = np.zeros((n_cells, n_clusters), dtype=np.float32)
# assume always one batch
data = next(iter(dataloader))
res = self.predict(data, device)
batch_pz = res.pz.detach().cpu().numpy()
batch_qz = res.qz.detach().cpu().numpy()
batch_qs = res.qs.detach().cpu().numpy()
batch_px = res.px.detach().cpu().numpy()
query_indices = data[self.query].idx.numpy()
qz[query_indices] = batch_qz
pz[query_indices] = batch_pz
ref_indices = data[self.ref].idx.numpy()
qs[ref_indices] = batch_qs
px[ref_indices] = batch_px
if self.configs.infer_cell_interaction:
eps = 1e-8
batch_omega = res.omega.detach().cpu().numpy() # (E,)
src, dst = data.edge_index_dict[self.r2r].cpu().numpy()
clusters = data[self.ref].cluster.cpu().numpy() # (N_subgraph,)
cell_indices = data[self.ref].idx[dst].cpu().numpy() # global target cell ids, (E,)
# -------------------------------------------------------
# den[i] = sum_{j->i} omega_{j->i} (per target cell)
# -------------------------------------------------------
den = np.zeros((n_cells,), dtype=np.float32)
np.add.at(den, cell_indices, batch_omega)
# -------------------------------------------------------
# normalized omega per edge (what the model actually uses)
# -------------------------------------------------------
omega_norm = batch_omega / (den[cell_indices] + eps) # (E,)
# -------------------------------------------------------
# MEAN aggregation of normalized omega by (cell, source-type)
# -------------------------------------------------------
attn_sum = np.zeros((n_cells, n_clusters), dtype=np.float32)
attn_count = np.zeros((n_cells, n_clusters), dtype=np.int32)
np.add.at(attn_sum, (cell_indices, clusters[src]), omega_norm)
np.add.at(attn_count, (cell_indices, clusters[src]), 1)
qomega_scores = np.divide(
attn_sum, attn_count,
out=np.zeros_like(attn_sum),
where=attn_count != 0
).astype(np.float32)
# ---------
# Abundance null per target cell (probabilities)
# ---------
edge_dist = data[self.r2r].edge_attr
edge_dist = edge_dist.squeeze(-1) if edge_dist.dim() > 1 else edge_dist
edge_dist = edge_dist.detach().cpu().numpy()
alpha = float(self.configs.alpha)
w_abun = (1.0 + edge_dist).astype(np.float32) ** (-alpha) # (E,)
abun_den = np.zeros((n_cells,), dtype=np.float32)
np.add.at(abun_den, cell_indices, w_abun)
abun_norm = w_abun / (abun_den[cell_indices] + eps) # (E,)
abundance_sum = np.zeros((n_cells, n_clusters), dtype=np.float32)
abundance_count = np.zeros((n_cells, n_clusters), dtype=np.int32)
np.add.at(abundance_sum, (cell_indices, clusters[src]), abun_norm)
np.add.at(abundance_count, (cell_indices, clusters[src]), 1)
abundance_count = np.divide(
abundance_sum, abundance_count,
out=np.zeros_like(abundance_sum),
where=abundance_count != 0
).astype(np.float32)
adata_query.obsm['X_z'] = qz.astype(np.float32) # Latent (z) for patches
adata_ref.obsm['X_z'] = qs.astype(np.float32) # Latent (z) for cells
# Save edge index & weights for visualization
if self.configs.infer_cell_interaction:
adata_ref.obsm['omega'] = qomega_scores # Attention scores summarized per cell
adata_ref.obsm['abundance'] = abundance_count # Cell-type abundance per cell
adata_ref.uns['omega'] = batch_omega
adata_ref.uns['edge_index'] = data.edge_index_dict[self.r2r].cpu().numpy()
return ConfigDict({
'qz': qz,
'qs': qs,
'pz': pz,
'px': px,
})
def _reshape_patches(self, u):
r"""Reshape flattened patches to proper image format"""
batch_size = u.shape[0]
expected_size = 3 * self.patch_size * self.patch_size
if u.shape[1] != expected_size:
raise ValueError(f"Expected flattened patch size {expected_size}, got {u.shape[1]}")
u_reshaped = u.view(batch_size, 3, self.patch_size, self.patch_size)
return u_reshaped
@staticmethod
def _weighted_sum(edge_index, edge_weights, x):
r"""Compute weighted neighboring scores per node"""
src, dst = edge_index
N = x.size(0)
num = torch_scatter.scatter_add(
edge_weights.unsqueeze(-1) * x[src],
dst,
dim=0,
dim_size=N,
)
den = torch_scatter.scatter_add(
edge_weights,
dst,
dim=0,
dim_size=N,
).unsqueeze(-1)
return num / (den + 1e-8)