src/forward_pdps.rs

changeset 72
e9a460a0e638
parent 66
fe47ad484deb
child 75
677a5fd1b014
--- a/src/forward_pdps.rs	Fri May 15 14:40:02 2026 -0500
+++ b/src/forward_pdps.rs	Sun Jul 19 07:34:39 2026 +0200
@@ -128,7 +128,7 @@
     KOpZ::SimpleAdjoint: GEMV<F, Y, Z>,
     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>,
@@ -200,7 +200,8 @@
     // Run the algorithm
     for state in iterator.iter_init(|| full_stats(&μ, &z, ε, stats.clone())) {
         // Calculate initial transport
-        let Pair(mut τv, τz) = f.differential(Pair(&μ, &z)) * τ;
+        let Pair(v, mut z_tmp) = f.differential(Pair(&μ, &z));
+        let mut τv = v * τ;
         let μ_base = μ.clone();
 
         // Construct μ^{k+1} by solving finite-dimensional subproblems and insert new spikes.
@@ -233,7 +234,7 @@
                 ε,
                 ins,
                 &reg,
-                is_fb.then_some(|μ̃: &RNDM<N, F>| f.apply(Pair(μ̃, &z))),
+                Some(|μ̃: &RNDM<N, F>| f.apply(Pair(μ̃, &z))),
             );
         }
 
@@ -241,14 +242,14 @@
         stats.pruned += prune_with_stats(&mut μ);
 
         // Do z variable primal update
-        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_tmp, &y);
+        z_tmp.axpy(1.0, &z, -σ_p);
+        let z_new = fnR.prox(σ_p, z_tmp);
+        //let z_new = fnR.prox(σ_p, z_tmp * (-σ_p) + &z);
         // 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;
 

mercurial