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
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
#![allow(dead_code)]
use crate::candle::loss::nb_log_likelihood;
use crate::candle::nn::linear::*;
use crate::candle::traits::model::*;
use candle_core::{Result, Tensor};
use candle_nn::{Module, VarBuilder};
/////////////////////////
// Topic Model Decoder //
/////////////////////////
pub struct MultinomTopicDecoder {
n_features: usize,
n_topics: usize,
dictionary: SoftmaxLinear,
/// Per-feature multiplicative weight on the per-gene log-likelihood
/// term (NB-Fisher info, gathered/computed at the same `D` the
/// decoder operates on). Stored as a `[1, D]` non-trainable tensor
/// so it broadcasts over rows; `None` recovers the unweighted
/// `(x+1).log() · log_recon` form.
feature_weights: Option<Tensor>,
}
impl MultinomTopicDecoder {
/// Will create a new topic model decoder with the following parameters:
/// * `dictionary.weight`
pub fn new(n_features: usize, n_topics: usize, vs: VarBuilder) -> Result<Self> {
let dictionary = log_softmax_linear(n_topics, n_features, vs.pp("dictionary"))?;
Ok(Self {
n_features,
n_topics,
dictionary,
feature_weights: None,
})
}
pub fn dictionary(&self) -> &SoftmaxLinear {
&self.dictionary
}
/// Attach NB-Fisher (or any) per-feature weights `w_d ∈ (0, 1]`.
/// `weights.len()` must equal `self.n_features`. Stored as a
/// `[1, D]` tensor on `dev`. Pass after construction; weights are
/// not part of the trained varmap.
pub fn set_feature_weights(
&mut self,
weights: &[f32],
dev: &candle_core::Device,
) -> Result<()> {
debug_assert_eq!(weights.len(), self.n_features);
let t = Tensor::from_slice(weights, (1, self.n_features), dev)?;
self.feature_weights = Some(t);
Ok(())
}
}
impl NewDecoder for MultinomTopicDecoder {
fn new(n_features: usize, n_topics: usize, vs: VarBuilder) -> Result<Self> {
MultinomTopicDecoder::new(n_features, n_topics, vs)
}
}
impl DecoderModuleT for MultinomTopicDecoder {
/// Input z_nk is already on the probability simplex (from softmax/sparsemax)
fn forward(&self, z_nk: &Tensor) -> Result<Tensor> {
self.dictionary.forward(z_nk)
}
fn get_dictionary(&self) -> Result<Tensor> {
self.dictionary.weight_dk()
}
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 log_recon_nd = self.dictionary.forward_log(z_nk)?;
let recon_nd = log_recon_nd.exp()?;
// Multinomial log-likelihood with NB-Fisher per-gene weighting:
//
// llik = Σ_d w_d · x_d · log(p_d)
//
// The proper multinomial NLL is `Σ_d x_d · log(p_d)`, which lets
// high-count housekeeping dominate. NB-Fisher weights `w_d` (sub-
// linear in μ_d) carry the outlier adjustment that earlier code
// approximated with `log(x+1)` — but principled: `w_d` is
// calibrated by the fitted dispersion trend, not a fixed shape.
// When `feature_weights` is `None`, this reduces to the raw
// multinomial NLL.
let weighted_x = match &self.feature_weights {
Some(w) => x_nd.broadcast_mul(w)?,
None => x_nd.clone(),
};
let llik = weighted_x.mul(&log_recon_nd)?.sum(x_nd.rank() - 1)?;
Ok((recon_nd, llik))
}
fn llik_is_gene_chunked(&self) -> bool {
true
}
/// The same weighted multinomial as [`Self::forward_with_llik`], summed a
/// gene slice at a time so no `[N, D]` tensor is built.
///
/// One pass: `log_recon_d` is a `logsumexp` over TOPICS, so a gene's value
/// needs no other gene's, and the likelihood is a plain sum over genes.
fn llik_gene_chunked(&self, z_nk: &Tensor, x_nd: &Tensor, gene_chunk: usize) -> Result<Tensor> {
let last = x_nd.rank() - 1;
let log_w_kd = self.dictionary.log_weight_kd()?;
let mut llik: Option<Tensor> = None;
for (start, len) in super::gene_slices(self.n_features, gene_chunk) {
let log_recon = self
.dictionary
.forward_log_slice(z_nk, Some(&log_w_kd), start, len)?;
let x = x_nd.narrow(last, start, len)?;
let weighted = match &self.feature_weights {
Some(w) => x.broadcast_mul(&w.narrow(last, start, len)?)?,
None => x,
};
let part = weighted.mul(&log_recon)?.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_topics
}
fn attach_feature_weights(&mut self, weights: &[f32], dev: &candle_core::Device) -> Result<()> {
MultinomTopicDecoder::set_feature_weights(self, weights, dev)
}
}
/////////////////////////////////////
// Negative Binomial Topic Decoder //
/////////////////////////////////////
/// Topic decoder with negative binomial likelihood.
///
/// μ_gn = l_n · softmax(W · z_n)_g
/// x_gn ~ NB(μ_gn, φ_g)
///
/// where l_n is per-cell library size (total counts) and
/// φ_g is per-gene inverse dispersion (learned).
pub struct NbTopicDecoder {
n_features: usize,
n_topics: usize,
dictionary: SoftmaxLinear,
/// log(φ_g) per-gene inverse dispersion, [1, D]
log_phi_1d: Tensor,
}
impl NbTopicDecoder {
/// Create a NB topic decoder.
/// `log_phi` initialized to ln(2) ≈ 0.69 (moderate dispersion).
pub fn new(n_features: usize, n_topics: usize, vs: VarBuilder) -> Result<Self> {
let dictionary = log_softmax_linear(n_topics, n_features, vs.pp("dictionary"))?;
let init_val = candle_nn::Init::Const(0.693); // ln(2)
let log_phi_1d = vs.get_with_hints((1, n_features), "log_phi", init_val)?;
Ok(Self {
n_features,
n_topics,
dictionary,
log_phi_1d,
})
}
/// Return per-gene dispersion φ_g as [1, D]
pub fn phi(&self) -> Result<Tensor> {
self.log_phi_1d.exp()
}
/// Return log(φ_g) as [1, D]
pub fn log_phi(&self) -> &Tensor {
&self.log_phi_1d
}
}
impl NewDecoder for NbTopicDecoder {
fn new(n_features: usize, n_topics: usize, vs: VarBuilder) -> Result<Self> {
NbTopicDecoder::new(n_features, n_topics, vs)
}
}
impl DecoderModuleT for NbTopicDecoder {
fn forward(&self, z_nk: &Tensor) -> Result<Tensor> {
self.dictionary.forward(z_nk)
}
fn get_dictionary(&self) -> Result<Tensor> {
self.dictionary.weight_dk()
}
fn forward_with_llik<LlikFn>(
&self,
z_nk: &Tensor,
x_nd: &Tensor,
_llik: &LlikFn,
) -> Result<(Tensor, Tensor)>
where
LlikFn: Fn(&Tensor, &Tensor) -> Result<Tensor>,
{
// softmax(W · z) gives gene proportions per cell [N, D]
let log_recon_nd = self.dictionary.forward_log(z_nk)?;
let recon_nd = log_recon_nd.exp()?;
// Library size: l_n = Σ_g x_gn per cell [N, 1]
let lib_size = x_nd.sum(x_nd.rank() - 1)?.unsqueeze(1)?; // [N, 1]
// μ_gn = l_n · softmax(W · z_n)_g
let mu_nd = recon_nd.broadcast_mul(&lib_size)?;
// NB log-likelihood
let llik = nb_log_likelihood(x_nd, &mu_nd, &self.log_phi_1d)?;
// Return proportions (not mu) for dictionary extraction
Ok((log_recon_nd.exp()?, llik))
}
fn llik_is_gene_chunked(&self) -> bool {
true
}
/// The same NB likelihood as [`Self::forward_with_llik`], summed a gene
/// slice at a time so no `[N, D]` tensor is built.
///
/// One pass, for the same reason as the multinomial head: the rate is a
/// `logsumexp` over TOPICS. Only the library size is global, and it is
/// `[N, 1]`, read off `x` once and shared by every slice.
fn llik_gene_chunked(&self, z_nk: &Tensor, x_nd: &Tensor, gene_chunk: usize) -> Result<Tensor> {
let last = x_nd.rank() - 1;
let lib_size = x_nd.sum(last)?.unsqueeze(1)?; // [N, 1]
let log_w_kd = self.dictionary.log_weight_kd()?;
let mut llik: Option<Tensor> = None;
for (start, len) in super::gene_slices(self.n_features, gene_chunk) {
let recon = self
.dictionary
.forward_log_slice(z_nk, Some(&log_w_kd), start, len)?
.exp()?;
let mu = recon.broadcast_mul(&lib_size)?;
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_topics
}
}
#[cfg(test)]
#[path = "topic_tests.rs"]
mod tests;