Skip to main content

gam_terms/analytic_penalties/
total_variation.rs

1use super::*;
2
3// ---------------------------------------------------------------------------
4// Total variation penalty
5// ---------------------------------------------------------------------------
6
7/// Shape of the first-difference operator used by [`TotalVariationPenalty`].
8#[derive(Debug, Clone)]
9pub enum DifferenceOpKind {
10    /// Path graph with rows connected as `(0, 1), (1, 2), ...`.
11    ForwardDiff1D,
12    /// Explicit adjacency list; each edge row has `-1` at `from`, `+1` at `to`.
13    GraphEdges(Vec<(usize, usize)>),
14}
15
16/// Coordinatewise/anisotropic smoothed-L¹ total variation on a row-major
17/// `(n_eff, d)` latent block.
18///
19/// Uses the differentiable Huber-style kernel `φ(x)=sqrt(x²+ε²)-ε` separately
20/// for each edge and latent axis. This is not vector-norm/isotropic edge TV:
21/// the Hessian intentionally has no cross-axis terms. The difference operator
22/// defines the prior shape: forward 1-D differences for ordered context
23/// windows, or graph edges for adjacency-structured atoms. Pair TV with
24/// Orthogonality when piecewise-constant atoms need a gauge-fixed basis.
25#[derive(Debug, Clone)]
26pub struct TotalVariationPenalty {
27    /// Base strength. If `learnable_weight` is true, the resolved strength is
28    /// `weight * exp(rho[rho_index])`; otherwise it is fixed at `weight`.
29    pub weight: f64,
30    /// Number of rows in the row-major latent coefficient block.
31    pub n_eff: usize,
32    pub difference_op: DifferenceOpKind,
33    pub smoothing_eps: f64,
34    pub learnable_weight: bool,
35    pub rho_index: usize,
36    pub weight_schedule: Option<ScalarWeightSchedule>,
37}
38
39impl TotalVariationPenalty {
40    #[must_use = "build error must be handled"]
41    pub fn new(
42        weight: f64,
43        n_eff: usize,
44        difference_op: DifferenceOpKind,
45        smoothing_eps: f64,
46        learnable_weight: bool,
47    ) -> Result<Self, String> {
48        if !(weight.is_finite() && weight > 0.0) {
49            return Err(format!(
50                "TotalVariationPenalty::new requires finite weight > 0, got {weight}"
51            ));
52        }
53        if n_eff == 0 {
54            return Err("TotalVariationPenalty::new requires n_eff > 0".to_string());
55        }
56        if !(smoothing_eps.is_finite() && smoothing_eps > 0.0) {
57            return Err(format!(
58                "TotalVariationPenalty::new requires finite smoothing_eps > 0, got {smoothing_eps}"
59            ));
60        }
61        if let DifferenceOpKind::GraphEdges(edges) = &difference_op {
62            if edges.is_empty() {
63                return Err(
64                    "TotalVariationPenalty::new GraphEdges requires at least one edge".to_string(),
65                );
66            }
67            for &(a, b) in edges {
68                if a >= n_eff || b >= n_eff {
69                    return Err(format!(
70                        "TotalVariationPenalty::new graph edge ({a}, {b}) exceeds n_eff {n_eff}"
71                    ));
72                }
73                if a == b {
74                    return Err(format!(
75                        "TotalVariationPenalty::new graph edge ({a}, {b}) is self-referential"
76                    ));
77                }
78            }
79        }
80        Ok(Self {
81            weight,
82            n_eff,
83            difference_op,
84            smoothing_eps,
85            learnable_weight,
86            rho_index: 0,
87            weight_schedule: None,
88        })
89    }
90
91    impl_with_weight_schedule!(weight);
92
93    fn resolved_weight(&self, rho: ArrayView1<'_, f64>) -> f64 {
94        if self.learnable_weight {
95            validated_learnable_weight(self.weight, rho[self.rho_index])
96        } else {
97            self.weight
98        }
99    }
100
101    fn latent_dim(&self, target_len: usize) -> Option<usize> {
102        if self.n_eff == 0 || !target_len.is_multiple_of(self.n_eff) {
103            assert_eq!(
104                target_len % self.n_eff.max(1),
105                0,
106                "target length must be divisible by n_eff"
107            );
108            return None;
109        }
110        Some(target_len / self.n_eff)
111    }
112
113    fn edge_count(&self) -> usize {
114        match &self.difference_op {
115            DifferenceOpKind::ForwardDiff1D => self.n_eff.saturating_sub(1),
116            DifferenceOpKind::GraphEdges(edges) => edges.len(),
117        }
118    }
119
120    fn add_edge_hvp(
121        &self,
122        target: ArrayView1<'_, f64>,
123        v: ArrayView1<'_, f64>,
124        out: &mut Array1<f64>,
125        d: usize,
126        a: usize,
127        b: usize,
128        weight: f64,
129    ) {
130        let eps2 = self.smoothing_eps * self.smoothing_eps;
131        for j in 0..d {
132            let ia = a * d + j;
133            let ib = b * d + j;
134            let diff = target[ib] - target[ia];
135            let r = (diff * diff + eps2).sqrt();
136            let curvature = eps2 / (r * r * r);
137            let dv = v[ib] - v[ia];
138            let h = weight * curvature * dv;
139            out[ia] -= h;
140            out[ib] += h;
141        }
142    }
143
144    fn add_edge_grad(
145        &self,
146        target: ArrayView1<'_, f64>,
147        out: &mut Array1<f64>,
148        d: usize,
149        a: usize,
150        b: usize,
151        weight: f64,
152    ) {
153        let eps2 = self.smoothing_eps * self.smoothing_eps;
154        for j in 0..d {
155            let ia = a * d + j;
156            let ib = b * d + j;
157            let diff = target[ib] - target[ia];
158            let smooth_sign = diff / (diff * diff + eps2).sqrt();
159            let g = weight * smooth_sign;
160            out[ia] -= g;
161            out[ib] += g;
162        }
163    }
164
165    fn add_edge_diag(
166        &self,
167        target: ArrayView1<'_, f64>,
168        out: &mut Array1<f64>,
169        d: usize,
170        a: usize,
171        b: usize,
172        weight: f64,
173    ) {
174        let eps2 = self.smoothing_eps * self.smoothing_eps;
175        for j in 0..d {
176            let ia = a * d + j;
177            let ib = b * d + j;
178            let diff = target[ib] - target[ia];
179            let r = (diff * diff + eps2).sqrt();
180            let curvature = weight * eps2 / (r * r * r);
181            out[ia] += curvature;
182            out[ib] += curvature;
183        }
184    }
185
186    fn add_edge_dense(
187        &self,
188        target: ArrayView1<'_, f64>,
189        out: &mut Array2<f64>,
190        d: usize,
191        a: usize,
192        b: usize,
193        weight: f64,
194    ) {
195        let eps2 = self.smoothing_eps * self.smoothing_eps;
196        for j in 0..d {
197            let ia = a * d + j;
198            let ib = b * d + j;
199            let diff = target[ib] - target[ia];
200            let r = (diff * diff + eps2).sqrt();
201            let curvature = weight * eps2 / (r * r * r);
202            out[[ia, ia]] += curvature;
203            out[[ib, ib]] += curvature;
204            out[[ia, ib]] -= curvature;
205            out[[ib, ia]] -= curvature;
206        }
207    }
208
209    pub fn diag_target(
210        &self,
211        target: ArrayView1<'_, f64>,
212        rho: ArrayView1<'_, f64>,
213    ) -> Array1<f64> {
214        let Some(d) = self.latent_dim(target.len()) else {
215            return Array1::<f64>::zeros(target.len());
216        };
217        let weight = self.resolved_weight(rho);
218        let mut out = Array1::<f64>::zeros(target.len());
219        match &self.difference_op {
220            DifferenceOpKind::ForwardDiff1D => {
221                for a in 0..self.n_eff.saturating_sub(1) {
222                    self.add_edge_diag(target, &mut out, d, a, a + 1, weight);
223                }
224            }
225            DifferenceOpKind::GraphEdges(edges) => {
226                for &(a, b) in edges {
227                    self.add_edge_diag(target, &mut out, d, a, b, weight);
228                }
229            }
230        }
231        out
232    }
233
234    /// Materialize `Dᵀ diag(φ''(D T)) D` for diagnostics and small graph cases.
235    pub fn as_dense(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array2<f64> {
236        let n = target.len();
237        let Some(d) = self.latent_dim(n) else {
238            return Array2::<f64>::zeros((n, n));
239        };
240        let weight = self.resolved_weight(rho);
241        let mut out = Array2::<f64>::zeros((n, n));
242        match &self.difference_op {
243            DifferenceOpKind::ForwardDiff1D => {
244                for a in 0..self.n_eff.saturating_sub(1) {
245                    self.add_edge_dense(target, &mut out, d, a, a + 1, weight);
246                }
247            }
248            DifferenceOpKind::GraphEdges(edges) => {
249                for &(a, b) in edges {
250                    self.add_edge_dense(target, &mut out, d, a, b, weight);
251                }
252            }
253        }
254        out
255    }
256
257    pub fn log_det_plus_lambda_i_forward_1d(
258        &self,
259        target: ArrayView1<'_, f64>,
260        rho: ArrayView1<'_, f64>,
261        lambda: f64,
262    ) -> Result<f64, String> {
263        if !matches!(&self.difference_op, DifferenceOpKind::ForwardDiff1D) {
264            return Err(
265                "TotalVariationPenalty::log_det_plus_lambda_i_forward_1d requires ForwardDiff1D"
266                    .to_string(),
267            );
268        }
269        let Some(d) = self.latent_dim(target.len()) else {
270            return Err(format!(
271                "TotalVariationPenalty target length {} is not divisible by n_eff {}",
272                target.len(),
273                self.n_eff
274            ));
275        };
276        if !(lambda.is_finite() && lambda > 0.0) {
277            return Err(format!(
278                "TotalVariationPenalty::log_det_plus_lambda_i_forward_1d requires finite λ > 0; got {lambda}"
279            ));
280        }
281        let n = self.n_eff;
282        if n == 1 {
283            return Ok((d as f64) * lambda.ln());
284        }
285        let weight = self.resolved_weight(rho);
286        let eps2 = self.smoothing_eps * self.smoothing_eps;
287        let mut total = 0.0;
288        for j in 0..d {
289            let mut edge_w = vec![0.0; n - 1];
290            for a in 0..n - 1 {
291                let diff = target[(a + 1) * d + j] - target[a * d + j];
292                let r = (diff * diff + eps2).sqrt();
293                edge_w[a] = weight * eps2 / (r * r * r);
294            }
295
296            let mut prev_pivot = lambda + edge_w[0];
297            if !prev_pivot.is_finite() || prev_pivot <= 0.0 {
298                return Err(format!(
299                    "TotalVariationPenalty log-det encountered non-positive pivot {prev_pivot:.3e}"
300                ));
301            }
302            total += prev_pivot.ln();
303            for row in 1..n {
304                let left = edge_w[row - 1];
305                let right = if row + 1 < n { edge_w[row] } else { 0.0 };
306                let diag = lambda + left + right;
307                let pivot = diag - left * left / prev_pivot;
308                if !pivot.is_finite() || pivot <= 0.0 {
309                    return Err(format!(
310                        "TotalVariationPenalty log-det encountered non-positive pivot {pivot:.3e}"
311                    ));
312                }
313                total += pivot.ln();
314                prev_pivot = pivot;
315            }
316        }
317        Ok(total)
318    }
319}
320
321impl AnalyticPenalty for TotalVariationPenalty {
322    fn tier(&self) -> PenaltyTier {
323        PenaltyTier::Psi
324    }
325
326    fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
327        let Some(d) = self.latent_dim(target.len()) else {
328            return 0.0;
329        };
330        if self.edge_count() == 0 {
331            return 0.0;
332        }
333        let weight = self.resolved_weight(rho);
334        let eps = self.smoothing_eps;
335        let eps2 = eps * eps;
336        let mut acc = 0.0;
337        match &self.difference_op {
338            DifferenceOpKind::ForwardDiff1D => {
339                for a in 0..self.n_eff.saturating_sub(1) {
340                    let b = a + 1;
341                    for j in 0..d {
342                        let diff = target[b * d + j] - target[a * d + j];
343                        acc += (diff * diff + eps2).sqrt() - eps;
344                    }
345                }
346            }
347            DifferenceOpKind::GraphEdges(edges) => {
348                for &(a, b) in edges {
349                    for j in 0..d {
350                        let diff = target[b * d + j] - target[a * d + j];
351                        acc += (diff * diff + eps2).sqrt() - eps;
352                    }
353                }
354            }
355        }
356        weight * acc
357    }
358
359    fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
360        let Some(d) = self.latent_dim(target.len()) else {
361            return Array1::<f64>::zeros(target.len());
362        };
363        let weight = self.resolved_weight(rho);
364        let mut out = Array1::<f64>::zeros(target.len());
365        match &self.difference_op {
366            DifferenceOpKind::ForwardDiff1D => {
367                for a in 0..self.n_eff.saturating_sub(1) {
368                    self.add_edge_grad(target, &mut out, d, a, a + 1, weight);
369                }
370            }
371            DifferenceOpKind::GraphEdges(edges) => {
372                for &(a, b) in edges {
373                    self.add_edge_grad(target, &mut out, d, a, b, weight);
374                }
375            }
376        }
377        out
378    }
379
380    fn hvp(
381        &self,
382        target: ArrayView1<'_, f64>,
383        rho: ArrayView1<'_, f64>,
384        v: ArrayView1<'_, f64>,
385    ) -> Array1<f64> {
386        assert_eq!(target.len(), v.len(), "hvp dimension mismatch");
387        if target.len() != v.len() {
388            return Array1::<f64>::zeros(target.len());
389        }
390        let Some(d) = self.latent_dim(target.len()) else {
391            return Array1::<f64>::zeros(target.len());
392        };
393        let weight = self.resolved_weight(rho);
394        let mut out = Array1::<f64>::zeros(target.len());
395        match &self.difference_op {
396            DifferenceOpKind::ForwardDiff1D => {
397                for a in 0..self.n_eff.saturating_sub(1) {
398                    self.add_edge_hvp(target, v, &mut out, d, a, a + 1, weight);
399                }
400            }
401            DifferenceOpKind::GraphEdges(edges) => {
402                for &(a, b) in edges {
403                    self.add_edge_hvp(target, v, &mut out, d, a, b, weight);
404                }
405            }
406        }
407        out
408    }
409
410    impl_learnable_weight_grad_rho!();
411
412    impl_learnable_weight_rho_count!();
413    impl_learnable_weight_domain!(weight);
414
415    fn name(&self) -> &str {
416        "total_variation"
417    }
418
419    impl_scalar_apply_schedule!(weight);
420}
421
422// ---------------------------------------------------------------------------
423// Monotonicity penalty (1D shape constraint)
424// ---------------------------------------------------------------------------
425
426/// Soft monotonicity penalty over a row-major `(n_eff, d)` latent block.
427///
428/// For each adjacent pair `(a, a+1)` along the leading axis and each output
429/// column `j`, the penalty contribution is
430///
431/// ```text
432/// softplus(-direction * (target[a+1, j] - target[a, j]) / smoothing_eps)
433/// * smoothing_eps
434/// ```
435///
436/// which is the smoothed hinge that hits zero when the slope agrees with
437/// `direction` (+1 ⇒ non-decreasing, -1 ⇒ non-increasing) and grows
438/// approximately linearly when it disagrees. The Hessian is positive
439/// semidefinite (softplus is convex) so the penalty composes cleanly with
440/// PIRLS/REML.
441///
442/// `n_eff` is the number of latent rows along the constrained axis; the
443/// remaining `target.len() / n_eff` columns are penalized independently and
444/// summed.
445#[derive(Debug, Clone)]
446pub struct ShapeMonotonicityPenalty {
447    pub weight: f64,
448    pub n_eff: usize,
449    /// `+1.0` for non-decreasing, `-1.0` for non-increasing along the leading axis.
450    pub direction: f64,
451    pub smoothing_eps: f64,
452    pub learnable_weight: bool,
453    pub rho_index: usize,
454    pub weight_schedule: Option<ScalarWeightSchedule>,
455}
456
457impl ShapeMonotonicityPenalty {
458    #[must_use = "build error must be handled"]
459    pub fn new(
460        weight: f64,
461        n_eff: usize,
462        direction: f64,
463        smoothing_eps: f64,
464        learnable_weight: bool,
465    ) -> Result<Self, String> {
466        if !(weight.is_finite() && weight > 0.0) {
467            return Err(format!(
468                "ShapeMonotonicityPenalty::new requires finite weight > 0, got {weight}"
469            ));
470        }
471        if n_eff == 0 {
472            return Err("ShapeMonotonicityPenalty::new requires n_eff > 0".to_string());
473        }
474        if !(direction.is_finite() && direction.abs() > 0.0) {
475            return Err(format!(
476                "ShapeMonotonicityPenalty::new requires finite non-zero direction (+1 or -1), got {direction}"
477            ));
478        }
479        if !(smoothing_eps.is_finite() && smoothing_eps > 0.0) {
480            return Err(format!(
481                "ShapeMonotonicityPenalty::new requires finite smoothing_eps > 0, got {smoothing_eps}"
482            ));
483        }
484        Ok(Self {
485            weight,
486            n_eff,
487            direction: direction.signum(),
488            smoothing_eps,
489            learnable_weight,
490            rho_index: 0,
491            weight_schedule: None,
492        })
493    }
494
495    impl_with_weight_schedule!(weight);
496
497    fn resolved_weight(&self, rho: ArrayView1<'_, f64>) -> f64 {
498        if self.learnable_weight {
499            validated_learnable_weight(self.weight, rho[self.rho_index])
500        } else {
501            self.weight
502        }
503    }
504
505    fn latent_dim(&self, target_len: usize) -> Option<usize> {
506        if self.n_eff == 0 || !target_len.is_multiple_of(self.n_eff) {
507            return None;
508        }
509        Some(target_len / self.n_eff)
510    }
511
512    /// Smoothed-hinge contribution for a single edge `(a, b)` and column `j`.
513    fn edge_value(&self, target: ArrayView1<'_, f64>, d: usize, a: usize, b: usize) -> f64 {
514        let eps = self.smoothing_eps;
515        let mut acc = 0.0;
516        for j in 0..d {
517            let slope = target[b * d + j] - target[a * d + j];
518            let z = -self.direction * slope / eps;
519            // softplus(z) * eps, computed in a numerically stable form.
520            let sp = if z > 0.0 {
521                z + (-z).exp().ln_1p()
522            } else {
523                z.exp().ln_1p()
524            };
525            acc += sp * eps;
526        }
527        acc
528    }
529
530    /// d softplus(-dir * slope / eps) * eps / d target = -dir * sigma(-dir*slope/eps).
531    fn edge_grad(
532        &self,
533        target: ArrayView1<'_, f64>,
534        out: &mut Array1<f64>,
535        d: usize,
536        a: usize,
537        b: usize,
538        weight: f64,
539    ) {
540        let eps = self.smoothing_eps;
541        for j in 0..d {
542            let slope = target[b * d + j] - target[a * d + j];
543            let z = -self.direction * slope / eps;
544            // Stable sigmoid(z).
545            let sigma = if z > 0.0 {
546                1.0 / (1.0 + (-z).exp())
547            } else {
548                let ez = z.exp();
549                ez / (1.0 + ez)
550            };
551            let g = weight * (-self.direction) * sigma;
552            out[a * d + j] -= g;
553            out[b * d + j] += g;
554        }
555    }
556}
557
558impl AnalyticPenalty for ShapeMonotonicityPenalty {
559    fn tier(&self) -> PenaltyTier {
560        PenaltyTier::Psi
561    }
562
563    fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
564        let Some(d) = self.latent_dim(target.len()) else {
565            return 0.0;
566        };
567        if self.n_eff < 2 {
568            return 0.0;
569        }
570        let weight = self.resolved_weight(rho);
571        let mut acc = 0.0;
572        for a in 0..self.n_eff.saturating_sub(1) {
573            acc += self.edge_value(target, d, a, a + 1);
574        }
575        weight * acc
576    }
577
578    fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
579        let Some(d) = self.latent_dim(target.len()) else {
580            return Array1::<f64>::zeros(target.len());
581        };
582        let weight = self.resolved_weight(rho);
583        let mut out = Array1::<f64>::zeros(target.len());
584        for a in 0..self.n_eff.saturating_sub(1) {
585            self.edge_grad(target, &mut out, d, a, a + 1, weight);
586        }
587        out
588    }
589
590    fn hvp(
591        &self,
592        target: ArrayView1<'_, f64>,
593        rho: ArrayView1<'_, f64>,
594        v: ArrayView1<'_, f64>,
595    ) -> Array1<f64> {
596        assert_eq!(target.len(), v.len(), "hvp dimension mismatch");
597        let Some(d) = self.latent_dim(target.len()) else {
598            return Array1::<f64>::zeros(target.len());
599        };
600        let weight = self.resolved_weight(rho);
601        let eps = self.smoothing_eps;
602        let mut out = Array1::<f64>::zeros(target.len());
603        for a in 0..self.n_eff.saturating_sub(1) {
604            let b = a + 1;
605            for j in 0..d {
606                let slope = target[b * d + j] - target[a * d + j];
607                let z = -self.direction * slope / eps;
608                let sigma = if z > 0.0 {
609                    1.0 / (1.0 + (-z).exp())
610                } else {
611                    let ez = z.exp();
612                    ez / (1.0 + ez)
613                };
614                // d²P/d(target_a)d(target_b) follows from the chain rule on
615                // z = -dir * (target_b - target_a) / eps. The penalty value is
616                // `softplus(z) * eps` (note the outer eps from `edge_value`).
617                // softplus''(z) = sigma(z)(1 - sigma(z)) and the (dz/dtarget)²
618                // factor is 1/eps², but the value's outer `* eps` cancels one of
619                // those, leaving `sigma(1 - sigma) / eps` — exactly the eps power
620                // that keeps `hvp` consistent with the finite difference of
621                // `grad_target` (whose own eps already cancelled). Off-diagonal
622                // entries carry an extra minus sign from the difference.
623                let h = weight * sigma * (1.0 - sigma) / eps;
624                let dv = v[b * d + j] - v[a * d + j];
625                out[a * d + j] -= h * dv;
626                out[b * d + j] += h * dv;
627            }
628        }
629        out
630    }
631
632    impl_learnable_weight_grad_rho!();
633
634    impl_learnable_weight_rho_count!();
635    impl_learnable_weight_domain!(weight);
636
637    fn name(&self) -> &str {
638        "monotonicity"
639    }
640
641    impl_scalar_apply_schedule!(weight);
642}