src/sliding_pdps.rs

changeset 72
e9a460a0e638
parent 68
00d0881f89a6
child 75
677a5fd1b014
--- a/src/sliding_pdps.rs	Fri May 15 14:40:02 2026 -0500
+++ b/src/sliding_pdps.rs	Sun Jul 19 07:34:39 2026 +0200
@@ -8,10 +8,11 @@
 use crate::measures::merging::SpikeMerging;
 use crate::measures::{DiscreteMeasure, RNDM};
 use crate::plot::Plotter;
-use crate::prox_penalty::{ProxPenalty, StepLengthBoundPair};
+use crate::prox_penalty::StepLengthBoundPair;
 use crate::regularisation::SlidingRegTerm;
-use crate::sliding_fb::{SlidingFBConfig, Transport, TransportConfig, TransportStepLength};
+use crate::sliding_fb::{SlidingFBConfig, Transport, TransportConfig, TransportProxPenalty};
 use crate::types::*;
+use alg_tools::bounds::MinMaxMapping;
 use alg_tools::convex::{Conjugable, Prox, Zero};
 use alg_tools::direct_product::Pair;
 use alg_tools::error::DynResult;
@@ -100,12 +101,12 @@
     I: AlgIteratorFactory<IterInfo<F>>,
     Dat: DifferentiableMapping<MeasureZ<F, Z, N>, Codomain = F, DerivativeDomain = Pair<S, Z>>
         + BoundedCurvature<F>,
-    S: DifferentiableRealMapping<N, F> + ClosedMul<F>,
+    S: DifferentiableRealMapping<N, F> + ClosedMul<F> + MinMaxMapping<Loc<N, F>, F>,
     for<'a> Pair<&'a P, &'a IdOp<Z>>: StepLengthBoundPair<F, Dat>,
     //Pair<S, Z>: ClosedMul<F>,
     RNDM<N, F>: SpikeMerging<F>,
     Reg: SlidingRegTerm<Loc<N, F>, F>,
-    P: ProxPenalty<Loc<N, F>, S, Reg, F>,
+    P: TransportProxPenalty<Loc<N, F>, S, Reg, F>,
     // KOpM : Linear<RNDM<N, F>, Codomain=Y>
     //     + GEMV<F, RNDM<N, F>>
     //     + Preadjointable<
@@ -122,7 +123,7 @@
     KOpZ::SimpleAdjoint: GEMV<F, Y>,
     Y: ClosedEuclidean<F>,
     for<'b> &'b Y: Instance<Y>,
-    Z: ClosedEuclidean<F>,
+    Z: ClosedEuclidean<F> + AXPY,
     for<'b> &'b Z: Instance<Z>,
     R: Prox<Z, Codomain = F>,
     H: Conjugable<Y, F, Codomain = F>,
@@ -153,7 +154,6 @@
     let bigM = 0.0; //opKμ.adjoint_product_bound(&op𝒟).unwrap().sqrt();
     let nKz = opKz.opnorm_bound(L2, L2)?;
     let is_fb = nKz == 0.0;
-    let ℓ = 0.0;
     let idOpZ = IdOp::new();
     let opKz_adj = opKz.adjoint();
     let (l, l_z) = Pair(prox_penalty, &idOpZ).step_length_bound_pair(&f)?;
@@ -183,27 +183,13 @@
     //  The factor two in the manuscript disappears due to the definition of 𝚹 being
     // for ‖x-y‖₂² instead of c_2(x, y)=‖x-y‖₂²/2.
 
-    let mut θ_or_adaptive = match f.curvature_bound_components(config.guess) {
-        (_, Err(_)) => TransportStepLength::Fixed(config.transport.θ0),
-        (maybe_ℓ_F, Ok(transport_lip)) => {
-            let calculate_θτ = move |ℓ_F, max_transport| {
-                let ℓ_r = transport_lip * max_transport;
-                config.transport.θ0 / ((ℓ + ℓ_F + ℓ_r) + κ * bigθ * max_transport / τ)
-            };
-            match maybe_ℓ_F {
-                Ok(ℓ_F) => TransportStepLength::AdaptiveMax {
-                    l: ℓ_F, // TODO: could estimate computing the real reesidual
-                    max_transport: 0.0,
-                    g: calculate_θτ,
-                },
-                Err(_) => TransportStepLength::FullyAdaptive {
-                    l: F::EPSILON, // Start with something very small to estimate differentials
-                    max_transport: 0.0,
-                    g: calculate_θτ,
-                },
-            }
-        }
-    };
+    let mut τθ_or_adaptive = prox_penalty.get_transport_steplength(
+        f.curvature_bound_components(config.guess),
+        &config.transport,
+        0.0,
+        κ * bigθ / τ, // = 0 currently
+    );
+
     // Acceleration is not currently supported
     // let γ = dataterm.factor_of_strong_convexity();
     let ω = 1.0;
@@ -230,8 +216,8 @@
 
     // Run the algorithm
     for state in iterator.iter_init(|| full_stats(&μ, &z, ε, stats.clone())) {
-        // Calculate initial transport
-        let Pair(v, _) = f.differential(Pair(&μ, &z));
+        let Pair(mut v, _) = f.differential(Pair(&μ, &z));
+
         //opKμ.preadjoint().apply_add(&mut v, y);
         // We want to proceed as in Example 4.12 but with v and v̆ as in §5.
         // With A(ν, z) = A_μ ν + A_z z, following Example 5.1, we have
@@ -242,7 +228,15 @@
 
         //dbg!(&μ);
 
-        γ.initial_transport(&μ, τ, &mut θ_or_adaptive, v, &config.transport);
+        prox_penalty.initial_transport(
+            &mut γ,
+            &μ,
+            ε,
+            τ,
+            &mut τθ_or_adaptive,
+            &v,
+            &config.transport,
+        );
 
         let mut attempts = 0;
 
@@ -254,7 +248,8 @@
             let μ̆ = μ.clone();
 
             // Calculate τv̆ = τA_*(A[μ_transported + μ_transported_base]-b)
-            let Pair(mut τv̆, τz̆) = f.differential(Pair(&μ̆, &z)) * τ;
+            let Pair(v̆, mut z̆) = f.differential(Pair(&μ̆, &z));
+            let mut τv̆ = v̆ * τ;
             // opKμ.preadjoint().gemv(&mut τv̆, τ, y, 1.0);
 
             // Construct μ^{k+1} by solving finite-dimensional subproblems and insert new spikes.
@@ -270,18 +265,25 @@
             )?;
 
             // Do z variable primal update here to able to estimate B_{v̆^k-v^{k+1}}
-            let mut z_new = τz̆;
-            opKz_adj.gemv(&mut z_new, -σ_p, &y, -σ_p / τ);
-            z_new = fnR.prox(σ_p, z_new + &z);
+            opKz_adj.apply_add(&mut z̆, &y);
+            z̆.axpy(1.0, &z, -σ_p);
+            let z_new = fnR.prox(σ_p, z̆);
+            //let z_new = fnR.prox(σ_p, z̆ * (-σ_p) + &z);
 
             // A posteriori transport adaptation.
-            if γ.aposteriori_transport(
+            if prox_penalty.aposteriori_transport(
+                &mut γ,
                 &μ,
                 &μ̆,
                 &mut τv̆,
-                Some(z_new.dist2(&z)),
+                &mut v,
+                None, //Some(z_new.dist2(&z)),
                 ε,
+                τ,
+                &τθ_or_adaptive,
+                reg,
                 &config.transport,
+                &config.insertion.refinement,
                 &mut attempts,
             ) {
                 break 'adapt_transport (maybe_d, within_tolerances, τv̆, z_new, μ̆);
@@ -294,7 +296,7 @@
         // This crucially expects the merge routine to be stable with respect to spike locations,
         // and not to performing any pruning. That is be to done below simultaneously for γ.
         if config.insertion.merge_now(&state) {
-            stats.merged += prox_penalty.merge_spikes(
+            let m = prox_penalty.merge_spikes(
                 &mut μ,
                 &mut τv̆,
                 &μ̆,
@@ -302,17 +304,20 @@
                 ε,
                 &config.insertion,
                 &reg,
-                is_fb.then_some(|μ̃: &RNDM<N, F>| f.apply(Pair(μ̃, &z))),
+                Some(|μ̃: &RNDM<N, F>| f.apply(Pair(μ̃, &z))),
             );
+            if m > 0 {
+                stats.merged += m;
+            }
         }
 
         γ.prune_compat(&mut μ, &mut stats);
 
         // Do dual update
         // opKμ.gemv(&mut y, σ_d*(1.0 + ω), &μ, 1.0);    // y = y + σ_d K[(1+ω)(μ,z)^{k+1}]
-        opKz.gemv(&mut y, σ_d * (1.0 + ω), &z_new, 1.0);
-        // opKμ.gemv(&mut y, -σ_d*ω, μ_base, 1.0);// y = y + σ_d K[(1+ω)(μ,z)^{k+1} - ω (μ,z)^k]-b
-        opKz.gemv(&mut y, -σ_d * ω, z, 1.0); // y = y + σ_d K[(1+ω)(μ,z)^{k+1} - ω (μ,z)^k]-b
+        z.axpy(1.0 + ω, &z_new, -ω);
+        //z = (z - &z_new) * (-ω) + &z_new;
+        opKz.gemv(&mut y, σ_d, z, 1.0);
         y = starH.prox(σ_d, y);
         z = z_new;
 
@@ -364,10 +369,10 @@
     I: AlgIteratorFactory<IterInfo<F>>,
     Dat: DifferentiableMapping<MeasureZ<F, Z, N>, Codomain = F, DerivativeDomain = Pair<S, Z>>
         + BoundedCurvature<F>,
-    S: DifferentiableRealMapping<N, F> + ClosedMul<F>,
+    S: DifferentiableRealMapping<N, F> + ClosedMul<F> + MinMaxMapping<Loc<N, F>, F>,
     RNDM<N, F>: SpikeMerging<F>,
     Reg: SlidingRegTerm<Loc<N, F>, F>,
-    P: ProxPenalty<Loc<N, F>, S, Reg, F>,
+    P: TransportProxPenalty<Loc<N, F>, S, Reg, F>,
     for<'a> Pair<&'a P, &'a IdOp<Z>>: StepLengthBoundPair<F, Dat>,
     Z: ClosedEuclidean<F> + AXPY + Clone,
     for<'b> &'b Z: Instance<Z>,

mercurial