File size: 7,302 Bytes
07fcdfe | 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 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 | """Population-level encoder: [B, N, G] β [B, H]."""
from typing import Optional
import torch
import torch.nn as nn
from ..metrics.distribution_metrics import masked_mean, masked_var
class PopulationEncoder(nn.Module):
"""Encode a variable-size cell population into a fixed-size vector.
Uses a **four-way pooling** strategy to capture population structure:
1. **Mean pooling**: per-gene mean expression across cells
2. **Variance pooling**: per-gene variance (captures subpopulation heterogeneity)
3. **Attention pooling**: learnable query attends over per-cell embeddings,
upweighting cells near the population center
4. **Skewness pooling** (optional): third moment for tail behavior
The four streams are projected to a common dimension, concatenated,
and processed through an MLP to produce the final embedding.
Parameters
----------
num_genes : input feature dimension G
hidden_dim : width of MLP layers (default 256)
output_dim : size of output embedding H (default 128)
n_layers : number of MLP layers (default 2)
use_var : if True, include variance pooling
use_attention : if True, include attention pooling
use_skewness : if True, include skewness pooling (clamped for stability)
cell_encoder_hidden : hidden dim for per-cell encoder (default 128)
"""
def __init__(
self,
num_genes: int,
hidden_dim: int = 256,
output_dim: int = 128,
n_layers: int = 2,
use_var: bool = True,
use_attention: bool = True,
use_skewness: bool = False,
cell_encoder_hidden: int = 128,
) -> None:
super().__init__()
self.num_genes = num_genes
self.output_dim = output_dim
self.use_var = use_var
self.use_attention = use_attention
self.use_skewness = use_skewness
# ββ Per-cell encoder (for attention pooling) ββββββββββββββββββ
# Maps each cell's gene expression to a compact embedding
if use_attention:
enc_layers = []
in_dim = num_genes
for i in range(2): # 2-layer small encoder
out_dim = cell_encoder_hidden if i == 0 else cell_encoder_hidden
enc_layers += [nn.Linear(in_dim, out_dim), nn.LayerNorm(out_dim), nn.GELU()]
in_dim = out_dim
self.cell_encoder = nn.Sequential(*enc_layers)
self.cell_emb_dim = cell_encoder_hidden
# Learnable query for attention pooling [1, H_cell]
self.query = nn.Parameter(torch.randn(1, cell_encoder_hidden) * 0.02)
# Projection for mean/var to cell_emb_dim for concatenation
self.mean_proj = nn.Linear(num_genes, cell_encoder_hidden)
if use_var:
self.var_proj = nn.Linear(num_genes, cell_encoder_hidden)
if use_skewness:
self.skew_proj = nn.Linear(num_genes, cell_encoder_hidden)
else:
self.cell_encoder = None
self.cell_emb_dim = 0
# ββ Output MLP ββββββββββββββββββββββββββββββββββββββββββββββββ
# Compute total input dimension
mlp_in = 0
if use_var:
mlp_in += num_genes * 2 # mean + var
else:
mlp_in += num_genes # mean only
if use_attention:
mlp_in += self.cell_emb_dim # attention pooling
if use_skewness:
mlp_in += num_genes # skewness
layers = []
in_dim = mlp_in
for _ in range(n_layers):
layers += [nn.Linear(in_dim, hidden_dim), nn.LayerNorm(hidden_dim), nn.GELU()]
in_dim = hidden_dim
layers.append(nn.Linear(hidden_dim, output_dim))
self.mlp = nn.Sequential(*layers)
def forward(
self,
cells: torch.Tensor,
mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Encode a cell population into a fixed-size vector.
Parameters
----------
cells : [B, N, G] raw gene expression (log1p normalized)
mask : [B, N] optional boolean mask (True = valid cell)
Returns
-------
z : [B, output_dim] population embedding
"""
B, N, G = cells.shape
device = cells.device
# ββ Stream 1: Mean + Var βββββββββββββββββββββββββββββββββββββ
mu = masked_mean(cells, mask) # [B, G]
streams = [mu]
if self.use_var:
var = masked_var(cells, mask) # [B, G]
streams.append(var)
# ββ Stream 2: Attention pooling ββββββββββββββββββββββββββββββ
if self.use_attention and self.cell_encoder is not None:
# Encode each cell independently
flat_cells = cells.reshape(B * N, G)
cell_emb = self.cell_encoder(flat_cells) # [B*N, H_cell]
cell_emb = cell_emb.reshape(B, N, -1) # [B, N, H_cell]
# Learnable query attends over all cells
# query: [1, H_cell] β [B, 1, H_cell]
q = self.query.unsqueeze(0).expand(B, -1, -1) # [B, 1, H_cell]
scores = torch.bmm(q, cell_emb.transpose(1, 2)) # [B, 1, N]
scores = scores.squeeze(1) # [B, N]
# Apply mask: set invalid cells to -inf before softmax
if mask is not None:
scores = scores.masked_fill(~mask, float("-inf"))
weights = torch.softmax(scores, dim=-1) # [B, N]
attn_pooled = torch.bmm(weights.unsqueeze(1), cell_emb).squeeze(1) # [B, H_cell]
streams.append(attn_pooled)
# ββ Stream 3: Skewness (optional) ββββββββββββββββββββββββββββ
if self.use_skewness:
# Third central moment, clamped for numerical stability
if mask is not None:
masked_cells = cells * mask.unsqueeze(-1) # [B, N, G]
n_valid = mask.sum(dim=-1, keepdim=True).clamp(min=1) # [B, 1]
mean_expanded = mu.unsqueeze(1) # [B, 1, G]
diff = masked_cells - mean_expanded # [B, N, G]
skew_num = (diff ** 3).sum(dim=1) # [B, G]
var_val = masked_var(cells, mask).clamp(min=1e-6) # [B, G]
skewness = skew_num / (n_valid * var_val ** 1.5) # [B, G]
skewness = skewness.clamp(-3.0, 3.0) # Clamp for stability
else:
mean_expanded = mu.unsqueeze(1)
diff = cells - mean_expanded
var_val = cells.var(dim=1, keepdim=False).clamp(min=1e-6)
skew_num = (diff ** 3).mean(dim=1)
skewness = (skew_num / (var_val ** 1.5)).clamp(-3.0, 3.0)
streams.append(skewness)
# ββ Concatenate & project ββββββββββββββββββββββββββββββββββββ
pooled = torch.cat(streams, dim=-1) # [B, mlp_in]
return self.mlp(pooled) # [B, output_dim]
|