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
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
//! **scVI-style NB decoder** for a Gaussian latent.
//!
//! Maps the unconstrained latent `z` to a gene distribution by a linear
//! decoder + gene-axis softmax, then a negative-binomial likelihood scaled by
//! the library size:
//!
//! ```text
//! logits_nd = z_nk · W + b # W: [n_latent, D]
//! π_nd = softmax_d(logits_nd) # gene distribution (sums to 1 over D)
//! μ_nd = library_n · π_nd
//! llik = NB(x_nd | μ_nd, φ_d)
//! ```
//!
//! Pairs with [`crate::candle::encoder::GaussianEncoder`]. The gene-axis softmax (not a
//! simplex mixture of per-topic distributions) is what makes a *Gaussian*,
//! unconstrained `z` valid here — unlike the topic decoders.
use crate::candle::loss::nb_log_likelihood;
use crate::candle::traits::model::*;
use candle_core::{Result, Tensor};
use candle_nn::{ops, Linear, Module, VarBuilder};
pub struct GaussianNbDecoder {
n_features: usize,
n_latent: usize,
/// Linear factor loadings `z → gene logits`; weight is `[D, n_latent]`.
decoder: Linear,
/// `[1, D]` log per-gene NB inverse dispersion.
log_phi_1d: Tensor,
}
impl GaussianNbDecoder {
pub fn new(n_features: usize, n_latent: usize, vs: VarBuilder) -> Result<Self> {
let decoder = candle_nn::linear(n_latent, n_features, vs.pp("gauss_decoder"))?;
let log_phi_1d =
vs.get_with_hints((1, n_features), "log_phi", candle_nn::Init::Const(0.693))?;
Ok(Self {
n_features,
n_latent,
decoder,
log_phi_1d,
})
}
/// The per-gene logit offset `b` in `π = softmax_d(z·W + b)`.
///
/// Exposed so a scorer can rebuild the same rate outside the module: the
/// loadings already ship as `dictionary.parquet`, but `b` lives only in
/// the checkpoint, and without it the reconstruction is off by a per-gene
/// factor.
#[must_use]
pub fn feature_bias(&self) -> Option<Tensor> {
self.decoder.bias().cloned()
}
/// `log π_nd = log_softmax_d(z·W + b)`.
fn log_pi(&self, z_nk: &Tensor) -> Result<Tensor> {
let logits_nd = self.decoder.forward(z_nk)?; // [N, D]
ops::log_softmax(&logits_nd, logits_nd.rank() - 1)
}
}
impl NewDecoder for GaussianNbDecoder {
fn new(n_features: usize, n_latent: usize, vs: VarBuilder) -> Result<Self> {
GaussianNbDecoder::new(n_features, n_latent, vs)
}
}
impl DecoderModuleT for GaussianNbDecoder {
fn forward(&self, z_nk: &Tensor) -> Result<Tensor> {
self.log_pi(z_nk)?.exp()
}
/// Factor loadings `[D, n_latent]` — the analogue of the topic dictionary.
fn get_dictionary(&self) -> Result<Tensor> {
Ok(self.decoder.weight().clone())
}
fn forward_with_llik<LlikFn>(
&self,
z_nk: &Tensor,
x_nd: &Tensor,
_llik: &LlikFn,
) -> Result<(Tensor, Tensor)>
where
LlikFn: Fn(&Tensor, &Tensor) -> Result<Tensor>,
{
let last = x_nd.rank() - 1;
// `softmax` directly — no need to `log_softmax` then `exp` back.
let logits_nd = self.decoder.forward(z_nk)?; // [N, D]
let pi_nd = ops::softmax(&logits_nd, last)?;
let lib_n1 = x_nd.sum_keepdim(last)?; // [N, 1]
let mu_nd = pi_nd.broadcast_mul(&lib_n1)?;
let llik = nb_log_likelihood(x_nd, &mu_nd, &self.log_phi_1d)?;
Ok((pi_nd, llik))
}
fn llik_is_gene_chunked(&self) -> bool {
true
}
/// Per-cell NB log-likelihood **without ever holding an `[N, D]` tensor**.
///
/// [`forward_with_llik`](DecoderModuleT::forward_with_llik) has to return
/// `π` itself, so training legitimately materialises `[N, D]`. Inference
/// only wants the scalar per cell, and the whole chain behind it —
/// logits, softmax, μ, and the dozen temporaries inside
/// [`crate::candle::loss::nb_log_likelihood_elem`] — is a sum over genes. So it
/// is taken in gene slices: the weight is narrowed to `[chunk, K]` and
/// nothing wider than `[N, chunk]` is ever allocated. At
/// whole-transcriptome D that is the difference between a couple of
/// hundred MB and tens of GB per block.
///
/// Two passes over the slices, because the softmax denominator is over
/// ALL genes: the first accumulates `log Σ_d exp(logit_d)` in the
/// streaming max/sumexp form, the second the likelihood terms. The
/// logits are recomputed rather than stored — that is the trade being
/// made, compute for memory.
fn llik_gene_chunked(&self, z_nk: &Tensor, x_nd: &Tensor, gene_chunk: usize) -> Result<Tensor> {
let chunk = gene_chunk.max(1).min(self.n_features);
let last = x_nd.rank() - 1;
let lib_n1 = x_nd.sum_keepdim(last)?; // [N, 1]
let w_dk = self.decoder.weight();
let bias_d = self.decoder.bias();
// Logits for one gene slice, as `[N, chunk]`.
let logits_of = |start: usize, len: usize| -> Result<Tensor> {
let w = w_dk.narrow(0, start, len)?; // [chunk, K]
let l = z_nk.matmul(&w.t()?)?; // [N, chunk]
match bias_d {
Some(b) => l.broadcast_add(&b.narrow(0, start, len)?.unsqueeze(0)?),
None => Ok(l),
}
};
// Pass 1: running max and Σexp, so the denominator never needs the
// full logit row in memory.
let mut running_max: Option<Tensor> = None;
let mut running_sum: Option<Tensor> = None;
for (start, len) in super::gene_slices(self.n_features, chunk) {
let logits = logits_of(start, len)?;
let m = logits.max_keepdim(last)?; // [N, 1]
let (new_max, sum) = match (running_max.take(), running_sum.take()) {
(Some(pm), Some(ps)) => {
let new_max = pm.maximum(&m)?;
let rescaled = ps.mul(&pm.sub(&new_max)?.exp()?)?;
let add = logits.broadcast_sub(&new_max)?.exp()?.sum_keepdim(last)?;
(new_max, rescaled.add(&add)?)
}
_ => {
let sum = logits.broadcast_sub(&m)?.exp()?.sum_keepdim(last)?;
(m, sum)
}
};
running_max = Some(new_max);
running_sum = Some(sum);
}
let max_n1 = running_max.expect("at least one gene slice");
let log_denom = running_sum
.expect("at least one gene slice")
.log()?
.add(&max_n1)?; // [N, 1]
// Pass 2: the likelihood terms, one slice at a time.
let mut llik: Option<Tensor> = None;
for (start, len) in super::gene_slices(self.n_features, chunk) {
let logits = logits_of(start, len)?;
let pi = logits.broadcast_sub(&log_denom)?.exp()?; // [N, chunk]
let mu = pi.broadcast_mul(&lib_n1)?;
let x = x_nd.narrow(last, start, len)?;
let log_phi = self
.log_phi_1d
.narrow(1, start, len)?
.broadcast_as(x.shape())?;
let part = crate::candle::loss::nb_log_likelihood_elem(&x, &mu, &log_phi)?.sum(last)?;
llik = Some(match llik.take() {
Some(acc) => acc.add(&part)?,
None => part,
});
}
llik.ok_or_else(|| candle_core::Error::Msg("no gene slices to score".into()))
}
fn dim_obs(&self) -> usize {
self.n_features
}
fn dim_latent(&self) -> usize {
self.n_latent
}
/// ESS log-likelihood closure for the Gaussian latent: multinomial
/// `Σ_d x_d · log π_d` with `π = softmax_d(z·W + b)` on detached weights.
/// Overrides the simplex-`θ` default (which would `softmax(z)` first).
fn build_ess_llik<'a>(
&'a self,
x_nd: &'a Tensor,
_topic_smoothing: f64,
) -> Result<EssLlikFn<'a>> {
// `[n_latent, D]`, contiguous — transposed once here, not per call.
let w_kd = self.decoder.weight().detach().t()?.contiguous()?;
let bias_d = self.decoder.bias().map(Tensor::detach);
let x_pos = x_nd.clamp(0.0, f64::INFINITY)?;
Ok(Box::new(move |z_nk: &Tensor| {
let logits = z_nk.matmul(&w_kd)?; // [N, D]
let logits = match &bias_d {
Some(b) => logits.broadcast_add(&b.unsqueeze(0)?)?,
None => logits,
};
let log_pi = ops::log_softmax(&logits, logits.rank() - 1)?;
x_pos.mul(&log_pi)?.sum(x_pos.rank() - 1)
}))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::candle::encoder::{GaussianEncoder, GaussianEncoderArgs};
use candle_core::{DType, Device};
use candle_nn::{VarBuilder, VarMap};
#[test]
fn test_gaussian_encoder_decoder_smoke() {
let dev = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &dev);
let (n, d, h, k) = (4usize, 6usize, 8usize, 3usize);
let enc = GaussianEncoder::new(
GaussianEncoderArgs {
n_features: d,
n_latent: k,
layers: &[h],
feature_mean: None,
},
&varmap,
vb.pp("enc"),
)
.unwrap();
let dec = GaussianNbDecoder::new(d, k, vb.pp("dec")).unwrap();
let x = Tensor::rand(0f32, 5f32, (n, d), &dev).unwrap();
let (z, kl) = enc.forward_t(&x, None, true).unwrap();
// Raw Gaussian latent — NOT projected to the simplex.
assert_eq!(z.dims(), &[n, k]);
assert_eq!(kl.dims(), &[n]);
let noop = |_a: &Tensor, _b: &Tensor| Ok(_a.clone());
let (pi, llik) = dec.forward_with_llik(&z, &x, &noop).unwrap();
assert_eq!(pi.dims(), &[n, d]);
assert_eq!(llik.dims(), &[n]);
// π is a gene distribution: each cell sums to 1 over genes.
for s in pi.sum(1).unwrap().to_vec1::<f32>().unwrap() {
assert!((s - 1.0).abs() < 1e-4, "pi row sum {s} != 1");
}
// NB likelihood is finite.
for v in llik.to_vec1::<f32>().unwrap() {
assert!(v.is_finite(), "llik {v} not finite");
}
}
}
#[cfg(test)]
#[path = "gaussian_nb_tests.rs"]
mod chunked_llik_tests;