use crate::candle::encoder::{dense_pool, scatter_pool};
use crate::candle::feature_embedding::FeatureEmbedding;
use crate::candle::nn::batch_norm;
use crate::candle::nn::layers::*;
use candle_core::{Result, Tensor};
use candle_nn::{ops, Linear, ModuleT, VarBuilder, VarMap};
pub struct PooledGeneEncoderArgs<'a> {
pub layers: &'a [usize],
pub out_dim: usize,
pub attn_pool: bool,
pub in_dim_extra: usize,
}
pub struct PooledGeneEncoder {
features: std::sync::Arc<FeatureEmbedding>,
attn_query: Option<Tensor>,
fc: StackLayers<Linear>,
bn_z: batch_norm::BatchNorm,
head: Linear,
}
impl PooledGeneEncoder {
pub fn new(
features: std::sync::Arc<FeatureEmbedding>,
args: PooledGeneEncoderArgs,
varmap: &VarMap,
vb: VarBuilder,
) -> Result<Self> {
debug_assert!(!args.layers.is_empty());
let embedding_dim = features.embedding_dim();
let bn_config = batch_norm::BatchNormConfig::default();
let fc_dims = args.layers[..args.layers.len() - 1].to_vec();
let in_dim = embedding_dim + args.in_dim_extra;
let out_dim = *args.layers.last().unwrap();
let fc = stack_relu_linear(in_dim, out_dim, &fc_dims, vb.pp("nn.enc.fc"))?;
let bn_z = batch_norm::batch_norm(out_dim, bn_config, varmap, vb.pp("nn.enc.bn_z"))?;
let head = candle_nn::linear(out_dim, args.out_dim, vb.pp("nn.enc.z.mean"))?;
let attn_query = if args.attn_pool {
Some(vb.get_with_hints(
(1, embedding_dim),
"attn.query",
candle_nn::init::DEFAULT_KAIMING_NORMAL,
)?)
} else {
None
};
Ok(Self {
features,
attn_query,
fc,
bn_z,
head,
})
}
pub fn features(&self) -> &FeatureEmbedding {
&self.features
}
#[must_use]
pub fn features_shared(&self) -> std::sync::Arc<FeatureEmbedding> {
std::sync::Arc::clone(&self.features)
}
pub fn attn_query(&self) -> Option<&Tensor> {
self.attn_query.as_ref()
}
pub fn pool(
&self,
x_nd: &Tensor,
x0_nd: Option<&Tensor>,
mean_1d: Option<&Tensor>,
visible_nd: Option<&Tensor>,
) -> Result<Tensor> {
if self.features.n_modules() > 0 {
candle_core::bail!(
"the dense pooled-gene read has no gene-module branch: modules pool a cell by \
membership over its context slots, and without a context window there are no \
slots. Build with n_gene_modules = 0."
);
}
let Some(attn_query) = self.attn_query.as_ref() else {
candle_core::bail!("pooling needs an attention query: build with attn_pool = true");
};
let scale = 1.0 / (self.features.embedding_dim() as f64).sqrt();
let a_nd = crate::candle::value_transform::anscombe_residual(x_nd, x0_nd, mean_1d)?; let rq_d = scatter_pool::query_over_features(&self.features, attn_query)?; let scores_nd = dense_pool::attention_scores_dense(&a_nd, &rq_d, visible_nd, scale)?;
let attn_nd = ops::softmax(&scores_nd, 1)?; let pooled_nh = dense_pool::pool_dense(&attn_nd, &a_nd, &self.features)?; match visible_nd {
Some(v) => {
let has_visible_n1 = v.sum_keepdim(1)?.gt(0.0)?.to_dtype(pooled_nh.dtype())?;
pooled_nh.broadcast_mul(&has_visible_n1)
}
None => Ok(pooled_nh),
}
}
pub fn trunk(&self, pooled: &Tensor, train: bool) -> Result<Tensor> {
let fc_nl = self.fc.forward_t(pooled, train)?;
self.bn_z.forward_t(&fc_nl, train)
}
pub fn head(&self, bn_nl: &Tensor, train: bool) -> Result<Tensor> {
self.head.forward_t(bn_nl, train)
}
pub fn forward(
&self,
x_nd: &Tensor,
x0_nd: Option<&Tensor>,
mean_1d: Option<&Tensor>,
visible_nd: Option<&Tensor>,
train: bool,
) -> Result<Tensor> {
let pooled = self.pool(x_nd, x0_nd, mean_1d, visible_nd)?;
self.head(&self.trunk(&pooled, train)?, train)
}
}
#[cfg(test)]
#[path = "pooled_tests.rs"]
mod pooled_tests;