src/sliding_fb.rs

changeset 72
e9a460a0e638
parent 68
00d0881f89a6
child 75
677a5fd1b014
--- a/src/sliding_fb.rs	Fri May 15 14:40:02 2026 -0500
+++ b/src/sliding_fb.rs	Sun Jul 19 07:34:39 2026 +0200
@@ -16,13 +16,16 @@
 use crate::measures::merging::SpikeMerging;
 use crate::measures::{DeltaMeasure, DiscreteMeasure, Radon, RNDM};
 use crate::plot::Plotter;
-use crate::prox_penalty::{ProxPenalty, StepLengthBound};
+use crate::prox_penalty::{ProxPenalty, RadonSquared, StepLengthBound};
 use crate::regularisation::SlidingRegTerm;
+use crate::seminorms::DiscreteMeasureOp;
 use crate::types::*;
+use alg_tools::bounds::{Bounds, MinMaxMapping};
 use alg_tools::error::DynResult;
 use alg_tools::euclidean::Euclidean;
+use alg_tools::instance::Space;
 use alg_tools::iterate::AlgIteratorFactory;
-use alg_tools::mapping::{DifferentiableMapping, DifferentiableRealMapping};
+use alg_tools::mapping::{DifferentiableMapping, DifferentiableRealMapping, Mapping, RealMapping};
 use alg_tools::nalgebra_support::ToNalgebraRealField;
 use alg_tools::norms::Norm;
 use anyhow::ensure;
@@ -34,14 +37,24 @@
 pub struct TransportConfig<F: Float> {
     /// Transport step length $θ$ normalised to $(0, 1)$.
     pub θ0: F,
+    /// Unnormalised transport step length $θ$. Overrides θ0.
+    pub τθ: Option<F>,
     /// Factor in $(0, 1)$ for decreasing transport to adapt to tolerance.
     pub adaptation: F,
-    /// A posteriori transport tolerance multiplier (C_pos)
-    pub tolerance_mult_con: F,
+    /// A posteriori transport tolerance multiplier
+    pub tolerance_mult: F,
+    /// Multiplier for rough estimate of ℓ_{∇v}. Should be ≥ 1.
+    /// If explicit τθ is given, ℓ_{∇v} is one divided by this number times τθ.
+    /// Otherwise, if  τθ relative to an estimate of ℓ_{∇v}, this number multiplies that estimate.
+    pub ℓ_gradv_mult: F,
     /// maximum number of adaptation iterations, until cancelling transport.
     pub max_attempts: usize,
     /// Maximum number of failed transportations for a single source point
     pub max_fail: usize,
+    /// Allow points to be transported partially.
+    pub allow_partial_transport: bool,
+    /// Use an alternative remainder control rule.
+    pub alt_remainder_control: bool,
 }
 
 #[replace_float_literals(F::cast_from(literal))]
@@ -50,7 +63,7 @@
     pub fn check(&self) -> DynResult<()> {
         ensure!(self.θ0 > 0.0);
         ensure!(0.0 < self.adaptation && self.adaptation < 1.0);
-        ensure!(self.tolerance_mult_con > 0.0);
+        ensure!(self.tolerance_mult > 0.0);
         Ok(())
     }
 }
@@ -59,9 +72,13 @@
 impl<F: Float> Default for TransportConfig<F> {
     fn default() -> Self {
         TransportConfig {
-            θ0: 0.9,
+            θ0: 0.99,
+            τθ: None,
             adaptation: 0.9,
-            tolerance_mult_con: 100.0,
+            allow_partial_transport: true,
+            alt_remainder_control: false,
+            tolerance_mult: 1e1,
+            ℓ_gradv_mult: 3.0,
             max_attempts: 2,
             max_fail: usize::MAX,
         }
@@ -98,63 +115,514 @@
 }
 
 /// Internal type of adaptive transport step length calculation
-pub(crate) enum TransportStepLength<F: Float, G: Fn(F, F) -> F> {
+#[derive(Clone, Debug, Serialize, Deserialize)]
+pub enum TransportStepLength<F: Float> {
     /// Fixed, known step length
     #[allow(dead_code)]
-    Fixed(F),
+    Fixed { τθ: F, ℓ_gradv: F },
+    /// Simple step lengths that do not depend on maximum transport
+    Simple { ℓ_gradv: F, τθ0: F, ℓ_base: F },
     /// Adaptive step length, only wrt. maximum transport.
-    /// Content of `l` depends on use case, while `g` calculates the step length from `l`.
-    AdaptiveMax { l: F, max_transport: F, g: G },
+    AdaptiveMax {
+        ℓ_gradv: F,
+        adaptive_max_transport: F,
+        τθ0: F,
+        ℓ_base: F,
+        ℓ_base_max_transport: F,
+    },
     /// Adaptive step length.
-    /// Content of `l` depends on use case, while `g` calculates the step length from `l`.
-    FullyAdaptive { l: F, max_transport: F, g: G },
+    FullyAdaptive {
+        adaptive_ℓ_gradv: F,
+        adaptive_max_transport: F,
+        τθ0: F,
+        ℓ_base: F,
+        ℓ_base_max_transport: F,
+    },
+}
+
+#[replace_float_literals(F::cast_from(literal))]
+impl<F: Float> TransportStepLength<F> {
+    fn get_ℓ_gradv(&self) -> F {
+        use TransportStepLength::*;
+        match *self {
+            Fixed { ℓ_gradv, .. } => ℓ_gradv,
+            Simple { ℓ_gradv, .. } => ℓ_gradv,
+            AdaptiveMax { ℓ_gradv, .. } => ℓ_gradv,
+            FullyAdaptive { adaptive_ℓ_gradv, .. } => adaptive_ℓ_gradv,
+        }
+    }
+
+    fn new(
+        maybe_ℓ_gradv_est: DynResult<F>,
+        tconfig: &TransportConfig<F>,
+        ℓ_base: F,
+        ℓ_base_max_transport: F,
+    ) -> Self {
+        if let Some(τθ) = tconfig.τθ {
+            TransportStepLength::Fixed { τθ, ℓ_gradv: 1.0 / (τθ * tconfig.ℓ_gradv_mult) }
+        } else {
+            match maybe_ℓ_gradv_est {
+                Ok(ℓ_gradv_est) => {
+                    if ℓ_base_max_transport == 0.0 {
+                        TransportStepLength::Simple {
+                            ℓ_gradv: tconfig.ℓ_gradv_mult * ℓ_gradv_est,
+                            τθ0: tconfig.θ0,
+                            ℓ_base,
+                        }
+                    } else {
+                        TransportStepLength::AdaptiveMax {
+                            ℓ_gradv: tconfig.ℓ_gradv_mult * ℓ_gradv_est,
+                            adaptive_max_transport: 0.0,
+                            τθ0: tconfig.θ0,
+                            ℓ_base,
+                            ℓ_base_max_transport,
+                        }
+                    }
+                }
+                Err(_) => TransportStepLength::FullyAdaptive {
+                    adaptive_ℓ_gradv: 10.0 * F::EPSILON, // Start with something very small to estimate differentials
+                    adaptive_max_transport: 0.0,
+                    τθ0: tconfig.θ0,
+                    ℓ_base,
+                    ℓ_base_max_transport,
+                },
+            }
+        }
+    }
 }
 
 #[derive(Clone, Debug, Serialize)]
-pub struct SingleTransport<const N: usize, F: Float> {
+pub struct SingleTransport<Domain, F: Float> {
     /// Source point
-    x: Loc<N, F>,
+    x: Domain,
     /// Target point
-    y: Loc<N, F>,
+    y: Domain,
     /// Original mass
     α_μ_orig: F,
     /// Transported mass
     α_γ: F,
     /// Helper for pruning
-    prune: bool,
+    retain: bool,
     /// Fail count
     fail_count: usize,
+    /// Contribution to remainder (temporary variable)
+    excess: F,
 }
 
 #[derive(Clone, Debug, Serialize)]
-pub struct Transport<const N: usize, F: Float> {
-    vec: Vec<SingleTransport<N, F>>,
+pub struct Transport<Domain, F: Float> {
+    vec: Vec<SingleTransport<Domain, F>>,
 }
 
-/// Whether partiall transported points are allowed.
+/// Whether partially transported points are allowed.
 ///
 /// Partial transport can cause spike count explosion, so full or zero
 /// transport is generally preferred. If this is set to `true`, different
 /// transport adaptation heuristics will be used.
-const ALLOW_PARTIAL_TRANSPORT: bool = true;
 const MINIMAL_PARTIAL_TRANSPORT: bool = true;
+const NEW_APPROACH: bool = true;
 
-impl<const N: usize, F: Float> Transport<N, F> {
+pub trait TransportProxPenalty<Domain, PreadjointCodomain, Reg, F = f64>:
+    ProxPenalty<Domain, PreadjointCodomain, Reg, F>
+where
+    F: Float + ToNalgebraRealField,
+    Reg: SlidingRegTerm<Domain, F>,
+    Domain: Space + Clone,
+{
+    type TransportStepLength;
+
+    /// Constrution of initial transport `γ1` from initial measure `μ` and `v=F'(μ)`
+    /// with step lengh τ and transport step length `θ_or_adaptive`.
+    fn initial_transport(
+        &self,
+        γ: &mut Transport<Domain, F>,
+        μ: &DiscreteMeasure<Domain, F>,
+        ε: F,
+        τ: F,
+        τθ_or_adaptive: &mut Self::TransportStepLength,
+        v: &PreadjointCodomain,
+        tconfig: &TransportConfig<F>,
+    );
+
+    /// A posteriori transport adaptation.
+    fn aposteriori_transport(
+        &self,
+        γ: &mut Transport<Domain, F>,
+        μ: &DiscreteMeasure<Domain, F>,
+        μ̆: &DiscreteMeasure<Domain, F>,
+        τv̆: &mut PreadjointCodomain,
+        v: &mut PreadjointCodomain,
+        extra: Option<F>,
+        ε: F,
+        τ: F,
+        τθ_or_adaptive: &Self::TransportStepLength,
+        reg: &Reg,
+        tconfig: &TransportConfig<F>,
+        rconfig: &RefinementSettings<F>,
+        attempts: &mut usize,
+    ) -> bool;
+
+    fn get_transport_steplength(
+        &self,
+        lips: (DynResult<F>, DynResult<F>, DynResult<F>),
+        tconfig: &TransportConfig<F>,
+        ℓ_base: F,
+        ℓ_base_max_transport: F,
+    ) -> Self::TransportStepLength;
+}
+
+#[replace_float_literals(F::cast_from(literal))]
+impl<F, M, Reg, const N: usize> TransportProxPenalty<Loc<N, F>, M, Reg, F> for RadonSquared
+where
+    RadonSquared: ProxPenalty<Loc<N, F>, M, Reg, F>,
+    F: Float + ToNalgebraRealField,
+    M: MinMaxMapping<Loc<N, F>, F> + DifferentiableRealMapping<N, F>,
+    Reg: SlidingRegTerm<Loc<N, F>, F>,
+    RNDM<N, F>: SpikeMerging<F>,
+{
+    type TransportStepLength = TransportStepLength<F>;
+
+    fn get_transport_steplength(
+        &self,
+        (ℓ_F, maybe_ℓ_gradv, _maybe_transport_lip): (DynResult<F>, DynResult<F>, DynResult<F>),
+        tconfig: &TransportConfig<F>,
+        ℓ: F,
+        ℓ_base_max_transport: F,
+    ) -> Self::TransportStepLength {
+        TransportStepLength::new(
+            maybe_ℓ_gradv,
+            tconfig,
+            ℓ + ℓ_F.unwrap_or(0.0),
+            ℓ_base_max_transport,
+        )
+    }
+
+    fn initial_transport(
+        &self,
+        γ: &mut Transport<Loc<N, F>, F>,
+        μ: &RNDM<N, F>,
+        _ε: F,
+        _τ: F,
+        τθ_or_adaptive: &mut TransportStepLength<F>,
+        v: &M,
+        tconfig: &TransportConfig<F>,
+    ) {
+        γ.do_init_transport(v, μ, τθ_or_adaptive, tconfig);
+    }
+
+    fn aposteriori_transport(
+        &self,
+        γ: &mut Transport<Loc<N, F>, F>,
+        μ: &RNDM<N, F>,
+        μ̆: &RNDM<N, F>,
+        τv̆: &mut M,
+        v: &mut M,
+        _extra: Option<F>,
+        ε: F,
+        τ: F,
+        τθ_or_adaptive: &TransportStepLength<F>,
+        reg: &Reg,
+        tconfig: &TransportConfig<F>,
+        _rconfig: &RefinementSettings<F>,
+        attempts: &mut usize,
+    ) -> bool {
+        *attempts += 1;
+
+        if *attempts > tconfig.max_attempts {
+            // Previous round has set transport to zero
+            return true;
+        }
+
+        let nΔ = μ.dist_matching(&μ̆);
+        let all_ok = γ.do_new_aposteriori_transport(
+            μ,
+            τv̆,
+            v,
+            ε,
+            τ,
+            τθ_or_adaptive,
+            reg,
+            tconfig,
+            |_, mass_zero, negated| {
+                use std::cmp::Ordering::*;
+                match (mass_zero, negated) {
+                    (Less, _) => -nΔ,
+                    (Greater, _) => nΔ,
+                    (Equal, false) => nΔ, // pessimistic estimated for ω
+                    (Equal, true) => -nΔ, // pessimistic estimated for -ω
+                }
+            },
+        );
+
+        if !all_ok {
+            if *attempts >= tconfig.max_attempts {
+                for ρ in γ.iter_mut() {
+                    ρ.α_γ = 0.0;
+                }
+            }
+        } else {
+            for ρ in γ.iter_mut() {
+                if ρ.α_γ == 0.0 {
+                    ρ.fail_count += 1;
+                } else {
+                    ρ.fail_count = 0;
+                }
+            }
+        }
+        all_ok
+    }
+}
+
+#[replace_float_literals(F::cast_from(literal))]
+impl<F, M, Reg, 𝒟, O, const N: usize> TransportProxPenalty<Loc<N, F>, M, Reg, F> for 𝒟
+where
+    F: Float + ToNalgebraRealField,
+    𝒟: DiscreteMeasureOp<Loc<N, F>, F>,
+    𝒟::Codomain: RealMapping<N, F>,
+    M: MinMaxMapping<Loc<N, F>, F> + DifferentiableRealMapping<N, F>,
+    for<'a> &'a M: std::ops::Add<𝒟::PreCodomain, Output = O>,
+    O: MinMaxMapping<Loc<N, F>, F>,
+    Reg: SlidingRegTerm<Loc<N, F>, F>,
+    Self: ProxPenalty<Loc<N, F>, M, Reg, F>,
+    RNDM<N, F>: SpikeMerging<F>,
+{
+    type TransportStepLength = TransportStepLength<F>;
+
+    fn get_transport_steplength(
+        &self,
+        (ℓ_F, maybe_ℓ_gradv, _maybe_transport_lip): (DynResult<F>, DynResult<F>, DynResult<F>),
+        tconfig: &TransportConfig<F>,
+        ℓ_base: F,
+        ℓ_base_max_transport: F,
+    ) -> Self::TransportStepLength {
+        TransportStepLength::new(
+            maybe_ℓ_gradv,
+            tconfig,
+            ℓ_base + ℓ_F.unwrap_or(0.0),
+            ℓ_base_max_transport,
+        )
+    }
+
+    fn initial_transport(
+        &self,
+        γ: &mut Transport<Loc<N, F>, F>,
+        μ: &RNDM<N, F>,
+        _ε: F,
+        _τ: F,
+        τθ_or_adaptive: &mut TransportStepLength<F>,
+        v: &M,
+        tconfig: &TransportConfig<F>,
+    ) {
+        γ.do_init_transport(v, μ, τθ_or_adaptive, tconfig);
+    }
+
+    fn aposteriori_transport(
+        &self,
+        γ: &mut Transport<Loc<N, F>, F>,
+        μ: &RNDM<N, F>,
+        μ̆: &RNDM<N, F>,
+        τv̆: &mut M,
+        v: &mut M,
+        extra: Option<F>,
+        ε: F,
+        τ: F,
+        τθ_or_adaptive: &TransportStepLength<F>,
+        reg: &Reg,
+        tconfig: &TransportConfig<F>,
+        _rconfig: &RefinementSettings<F>,
+        attempts: &mut usize,
+    ) -> bool {
+        *attempts += 1;
+
+        if *attempts > tconfig.max_attempts {
+            // Previous round has set transport to zero
+            return true;
+        }
+
+        let all_ok = if NEW_APPROACH {
+            let ω = self.apply(μ.sub_matching(&μ̆));
+            γ.do_new_aposteriori_transport(
+                μ,
+                τv̆,
+                v,
+                ε,
+                τ,
+                τθ_or_adaptive,
+                reg,
+                tconfig,
+                |x, _, _| ω.apply(x),
+            )
+        } else {
+            let mut all_ok0 = true;
+
+            // 1. If π_♯^1γ^{k+1} = γ1 has non-zero mass at some point y, but μ = μ^{k+1} does not,
+            // then the ansatz ∇w̃_x(y) = w^{k+1}(y) may not be satisfied. So set the mass of γ1
+            // at that point to zero, and retry.
+            for (δ, ρ) in izip!(μ.iter_spikes(), γ.iter_mut()) {
+                if δ.α == 0.0 && ρ.α_γ != 0.0 {
+                    all_ok0 = false;
+                    ρ.α_γ = 0.0;
+                }
+                // TODO: sign
+            }
+
+            // 2. Through bounding ∫ B_ω(y, z) dλ(x, y, z).
+            //    through the estimate ≤ C ‖Δ‖‖γ^{k+1}‖ for Δ := μ^{k+1}-μ̆^k
+            //    which holds for some some C if the convolution kernel in 𝒟 has Lipschitz gradient.
+
+            let nγ = γ.norm(Radon);
+            let nΔ = μ.dist_matching(&μ̆) + extra.unwrap_or(0.0);
+            let t = ε * tconfig.tolerance_mult;
+            if nγ * nΔ > t && *attempts >= tconfig.max_attempts {
+                all_ok0 = false;
+            } else if nγ * nΔ > t {
+                // Since t/(nγ*nΔ)<1, and the constant tconfig.adaptation < 1,
+                // this will guarantee that eventually ‖γ‖ decreases sufficiently that we
+                // will not enter here.
+                //*γ *= tconfig.adaptation * t / (nγ * nΔ);
+
+                // We want a consistent behaviour that has the potential to set many weights to zero.
+                // Therefore, we find the smallest uniform reduction `chg_one`, subtracted
+                // from all weights, that achieves total `adapt` adaptation.
+                let adapt_to = tconfig.adaptation * t / nΔ;
+                let reduction_target = nγ - adapt_to;
+                assert!(reduction_target > 0.0);
+                if tconfig.allow_partial_transport {
+                    if MINIMAL_PARTIAL_TRANSPORT {
+                        // This reduces weights of transport, starting from … until `adapt` is
+                        // exhausted. It will, therefore, only ever cause one extrap point insertion
+                        // at the sources, unlike “full” partial transport.
+                        //let refs = γ.vec.iter_mut().collect::<Vec<_>>();
+                        //refs.sort_by(|ρ1, ρ2| ρ1.α_γ.abs().partial_cmp(&ρ2.α_γ.abs()).unwrap());
+                        // let mut it = refs.into_iter();
+                        //
+                        // Maybe sort by differential norm
+                        // let mut refs = γ
+                        //     .vec
+                        //     .iter_mut()
+                        //     .map(|ρ| {
+                        //         let val = v.differential(&ρ.x).norm2_squared();
+                        //         (ρ, val)
+                        //     })
+                        //     .collect::<Vec<_>>();
+                        // refs.sort_by(|(_, v1), (_, v2)| v2.partial_cmp(&v1).unwrap());
+                        // let mut it = refs.into_iter().map(|(ρ, _)| ρ);
+                        let mut it = γ.vec.iter_mut().rev();
+                        let _unused = it.try_fold(reduction_target, |left, ρ| {
+                            let w = ρ.α_γ.abs();
+                            if left <= w {
+                                ρ.α_γ = ρ.α_γ.signum() * (w - left);
+                                ControlFlow::Break(())
+                            } else {
+                                ρ.α_γ = 0.0;
+                                ControlFlow::Continue(left - w)
+                            }
+                        });
+                    } else {
+                        // This version equally reduces all weights. It causes partial transport, which
+                        // has the problem that that we need to then adapt weights in both start and
+                        // end points, in insert_and_reweigh, somtimes causing the number of spikes μ
+                        // to explode.
+                        let mut abs_weights = γ
+                            .vec
+                            .iter()
+                            .map(|ρ| ρ.α_γ.abs())
+                            .filter(|t| *t > F::EPSILON)
+                            .collect::<Vec<F>>();
+                        abs_weights.sort_by(|a, b| a.total_cmp(b));
+                        let n = abs_weights.len();
+                        // Cannot have partial transport; can cause spike count explosion
+                        let chg = abs_weights.into_iter().zip((1..=n).rev()).try_fold(
+                            0.0,
+                            |smaller_total, (w, m)| {
+                                let mf = F::cast_from(m);
+                                let reduction = w * mf + smaller_total;
+                                if reduction >= reduction_target {
+                                    ControlFlow::Break((reduction_target - smaller_total) / mf)
+                                } else {
+                                    ControlFlow::Continue(smaller_total + w)
+                                }
+                            },
+                        );
+                        match chg {
+                            ControlFlow::Continue(_) => γ.vec.iter_mut().for_each(|δ| δ.α_γ = 0.0),
+                            ControlFlow::Break(chg_one) => γ.vec.iter_mut().for_each(|ρ| {
+                                let t = ρ.α_γ.abs();
+                                if t > 0.0 {
+                                    if tconfig.allow_partial_transport {
+                                        let new = (t - chg_one).max(0.0);
+                                        ρ.α_γ = ρ.α_γ.signum() * new;
+                                    }
+                                }
+                            }),
+                        }
+                    }
+                } else {
+                    // This version zeroes smallest weights, avoiding partial transport.
+                    let mut abs_weights_idx = γ
+                        .vec
+                        .iter()
+                        .map(|ρ| ρ.α_γ.abs())
+                        .zip(0..)
+                        .filter(|(w, _)| *w >= 0.0)
+                        .collect::<Vec<(F, usize)>>();
+                    abs_weights_idx.sort_by(|(a, _), (b, _)| a.total_cmp(b));
+
+                    let mut left = reduction_target;
+
+                    for (w, i) in abs_weights_idx {
+                        left -= w;
+                        let ρ = &mut γ.vec[i];
+                        ρ.α_γ = 0.0;
+                        if left < 0.0 {
+                            break;
+                        }
+                    }
+                }
+
+                all_ok0 = false
+            }
+            all_ok0
+        };
+
+        if !all_ok {
+            if *attempts >= tconfig.max_attempts {
+                for ρ in γ.iter_mut() {
+                    ρ.α_γ = 0.0;
+                }
+            }
+        } else {
+            for ρ in γ.iter_mut() {
+                if ρ.α_γ == 0.0 {
+                    ρ.fail_count += 1;
+                } else {
+                    ρ.fail_count = 0;
+                }
+            }
+        }
+        all_ok
+    }
+}
+
+#[replace_float_literals(F::cast_from(literal))]
+impl<const N: usize, F: Float> Transport<Loc<N, F>, F> {
     pub(crate) fn new() -> Self {
         Transport { vec: Vec::new() }
     }
 
-    pub(crate) fn iter(&self) -> impl Iterator<Item = &'_ SingleTransport<N, F>> {
+    pub(crate) fn iter(&self) -> impl Iterator<Item = &'_ SingleTransport<Loc<N, F>, F>> {
         self.vec.iter()
     }
 
-    pub(crate) fn iter_mut(&mut self) -> impl Iterator<Item = &'_ mut SingleTransport<N, F>> {
+    pub(crate) fn iter_mut(
+        &mut self,
+    ) -> impl Iterator<Item = &'_ mut SingleTransport<Loc<N, F>, F>> {
         self.vec.iter_mut()
     }
 
     pub(crate) fn extend<I>(&mut self, it: I)
     where
-        I: IntoIterator<Item = SingleTransport<N, F>>,
+        I: IntoIterator<Item = SingleTransport<Loc<N, F>, F>>,
     {
         self.vec.extend(it)
     }
@@ -171,7 +639,6 @@
     // }
 
     /// Construct `μ̆`, replacing the contents of `μ`.
-    #[replace_float_literals(F::cast_from(literal))]
     pub(crate) fn μ̆_into(&self, μ: &mut RNDM<N, F>) {
         assert!(self.len() <= μ.len());
 
@@ -190,287 +657,28 @@
 
         // Then source points with partial transport
         let mut i = self.len();
-        if ALLOW_PARTIAL_TRANSPORT {
-            // This can cause the number of points to explode, so cannot have partial transport.
-            for ρ in self.iter() {
-                let α = ρ.α_μ_orig - ρ.α_γ;
-                if ρ.α_γ.abs() > F::EPSILON && α != 0.0 {
-                    let δ = DeltaMeasure { α, x: ρ.x };
-                    if i < μ.len() {
-                        μ[i] = δ;
-                    } else {
-                        μ.push(δ)
-                    }
-                    i += 1;
+        // This can cause the number of points to explode, so cannot have partial transport.
+        for ρ in self.iter() {
+            let α = ρ.α_μ_orig - ρ.α_γ;
+            if ρ.α_γ.abs() > F::EPSILON && α.abs() > F::EPSILON {
+                let δ = DeltaMeasure { α, x: ρ.x };
+                if i < μ.len() {
+                    μ[i] = δ;
+                } else {
+                    μ.push(δ)
                 }
+                i += 1;
             }
         }
         μ.truncate(i);
     }
 
-    /// Constrution of initial transport `γ1` from initial measure `μ` and `v=F'(μ)`
-    /// with step lengh τ and transport step length `θ_or_adaptive`.
-    #[replace_float_literals(F::cast_from(literal))]
-    pub(crate) fn initial_transport<G, D>(
-        &mut self,
-        μ: &RNDM<N, F>,
-        _τ: F,
-        τθ_or_adaptive: &mut TransportStepLength<F, G>,
-        v: D,
-        tconfig: &TransportConfig<F>,
-    ) where
-        G: Fn(F, F) -> F,
-        D: DifferentiableRealMapping<N, F>,
-    {
-        use TransportStepLength::*;
-
-        // Initialise transport structure weights
-        for (δ, ρ) in izip!(μ.iter_spikes(), self.iter_mut()) {
-            ρ.α_μ_orig = δ.α;
-            ρ.x = δ.x;
-            if ρ.fail_count > tconfig.max_fail {
-                ρ.α_γ = 0.0
-            } else {
-                // If old transport has opposing sign, the new transport will be none.
-                ρ.α_γ = if (ρ.α_γ > 0.0 && δ.α < 0.0) || (ρ.α_γ < 0.0 && δ.α > 0.0) {
-                    0.0
-                } else {
-                    δ.α
-                }
-            }
-        }
-
-        let γ_prev_len = self.len();
-        assert!(μ.len() >= γ_prev_len);
-        self.extend(μ[γ_prev_len..].iter().map(|δ| SingleTransport {
-            x: δ.x,
-            y: δ.x, // Just something, will be filled properly in the next phase
-            α_μ_orig: δ.α,
-            α_γ: δ.α,
-            prune: false,
-            fail_count: 0,
-        }));
-
-        // Calculate transport rays.
-        match *τθ_or_adaptive {
-            Fixed(θ) => {
-                for ρ in self.iter_mut() {
-                    if ρ.fail_count <= tconfig.max_fail {
-                        ρ.y = ρ.x - v.differential(&ρ.x) * (ρ.α_γ.signum() * θ);
-                    }
-                }
-            }
-            AdaptiveMax { l: ℓ_F, ref mut max_transport, g: ref calculate_θτ } => {
-                *max_transport = max_transport.max(self.norm(Radon));
-                let θτ = calculate_θτ(ℓ_F, *max_transport);
-                for ρ in self.iter_mut() {
-                    if ρ.fail_count <= tconfig.max_fail {
-                        ρ.y = ρ.x - v.differential(&ρ.x) * (ρ.α_γ.signum() * θτ);
-                    }
-                }
-            }
-            FullyAdaptive {
-                l: ref mut adaptive_ℓ_F,
-                ref mut max_transport,
-                g: ref calculate_θτ,
-            } => {
-                *max_transport = max_transport.max(self.norm(Radon));
-                let mut θτ = calculate_θτ(*adaptive_ℓ_F, *max_transport);
-                // Do two runs through the spikes to update θ, breaking if first run did not cause
-                // a change.
-                for _i in 0..=1 {
-                    let mut changes = false;
-                    for ρ in self.iter_mut() {
-                        if ρ.fail_count < tconfig.max_fail {
-                            let dv_x = v.differential(&ρ.x);
-                            let g = &dv_x * (ρ.α_γ.signum() * θτ);
-                            ρ.y = ρ.x - g;
-                            let n = g.norm2();
-                            if n >= F::EPSILON {
-                                // Estimate Lipschitz factor of ∇v
-                                let this_ℓ_F = (dv_x - v.differential(&ρ.y)).norm2() / n;
-                                *adaptive_ℓ_F = adaptive_ℓ_F.max(this_ℓ_F);
-                                θτ = calculate_θτ(*adaptive_ℓ_F, *max_transport);
-                                changes = true
-                            }
-                        }
-                    }
-                    if !changes {
-                        break;
-                    }
-                }
-            }
-        }
-    }
-
-    /// A posteriori transport adaptation.
-    #[replace_float_literals(F::cast_from(literal))]
-    pub(crate) fn aposteriori_transport<D>(
-        &mut self,
-        μ: &RNDM<N, F>,
-        μ̆: &RNDM<N, F>,
-        _v: &mut D,
-        extra: Option<F>,
-        ε: F,
-        tconfig: &TransportConfig<F>,
-        attempts: &mut usize,
-    ) -> bool
-    where
-        D: DifferentiableRealMapping<N, F>,
-    {
-        *attempts += 1;
-
-        // 1. If π_♯^1γ^{k+1} = γ1 has non-zero mass at some point y, but μ = μ^{k+1} does not,
-        // then the ansatz ∇w̃_x(y) = w^{k+1}(y) may not be satisfied. So set the mass of γ1
-        // at that point to zero, and retry.
-        let mut all_ok = true;
-        for (δ, ρ) in izip!(μ.iter_spikes(), self.iter_mut()) {
-            if δ.α == 0.0 && ρ.α_γ != 0.0 {
-                all_ok = false;
-                ρ.α_γ = 0.0;
-            }
-        }
-
-        // 2. Through bounding ∫ B_ω(y, z) dλ(x, y, z).
-        //    through the estimate ≤ C ‖Δ‖‖γ^{k+1}‖ for Δ := μ^{k+1}-μ̆^k
-        //    which holds for some some C if the convolution kernel in 𝒟 has Lipschitz gradient.
-        let nγ = self.norm(Radon);
-        let nΔ = μ.dist_matching(&μ̆) + extra.unwrap_or(0.0);
-        let t = ε * tconfig.tolerance_mult_con;
-        if nγ * nΔ > t && *attempts >= tconfig.max_attempts {
-            all_ok = false;
-        } else if nγ * nΔ > t {
-            // Since t/(nγ*nΔ)<1, and the constant tconfig.adaptation < 1,
-            // this will guarantee that eventually ‖γ‖ decreases sufficiently that we
-            // will not enter here.
-            //*self *= tconfig.adaptation * t / (nγ * nΔ);
-
-            // We want a consistent behaviour that has the potential to set many weights to zero.
-            // Therefore, we find the smallest uniform reduction `chg_one`, subtracted
-            // from all weights, that achieves total `adapt` adaptation.
-            let adapt_to = tconfig.adaptation * t / nΔ;
-            let reduction_target = nγ - adapt_to;
-            assert!(reduction_target > 0.0);
-            if ALLOW_PARTIAL_TRANSPORT {
-                if MINIMAL_PARTIAL_TRANSPORT {
-                    // This reduces weights of transport, starting from … until `adapt` is
-                    // exhausted. It will, therefore, only ever cause one extrap point insertion
-                    // at the sources, unlike “full” partial transport.
-                    //let refs = self.vec.iter_mut().collect::<Vec<_>>();
-                    //refs.sort_by(|ρ1, ρ2| ρ1.α_γ.abs().partial_cmp(&ρ2.α_γ.abs()).unwrap());
-                    // let mut it = refs.into_iter();
-                    //
-                    // Maybe sort by differential norm
-                    // let mut refs = self
-                    //     .vec
-                    //     .iter_mut()
-                    //     .map(|ρ| {
-                    //         let val = v.differential(&ρ.x).norm2_squared();
-                    //         (ρ, val)
-                    //     })
-                    //     .collect::<Vec<_>>();
-                    // refs.sort_by(|(_, v1), (_, v2)| v2.partial_cmp(&v1).unwrap());
-                    // let mut it = refs.into_iter().map(|(ρ, _)| ρ);
-                    let mut it = self.vec.iter_mut().rev();
-                    let _unused = it.try_fold(reduction_target, |left, ρ| {
-                        let w = ρ.α_γ.abs();
-                        if left <= w {
-                            ρ.α_γ = ρ.α_γ.signum() * (w - left);
-                            ControlFlow::Break(())
-                        } else {
-                            ρ.α_γ = 0.0;
-                            ControlFlow::Continue(left - w)
-                        }
-                    });
-                } else {
-                    // This version equally reduces all weights. It causes partial transport, which
-                    // has the problem that that we need to then adapt weights in both start and
-                    // end points, in insert_and_reweigh, somtimes causing the number of spikes μ
-                    // to explode.
-                    let mut abs_weights = self
-                        .vec
-                        .iter()
-                        .map(|ρ| ρ.α_γ.abs())
-                        .filter(|t| *t > F::EPSILON)
-                        .collect::<Vec<F>>();
-                    abs_weights.sort_by(|a, b| a.partial_cmp(b).unwrap());
-                    let n = abs_weights.len();
-                    // Cannot have partial transport; can cause spike count explosion
-                    let chg = abs_weights.into_iter().zip((1..=n).rev()).try_fold(
-                        0.0,
-                        |smaller_total, (w, m)| {
-                            let mf = F::cast_from(m);
-                            let reduction = w * mf + smaller_total;
-                            if reduction >= reduction_target {
-                                ControlFlow::Break((reduction_target - smaller_total) / mf)
-                            } else {
-                                ControlFlow::Continue(smaller_total + w)
-                            }
-                        },
-                    );
-                    match chg {
-                        ControlFlow::Continue(_) => self.vec.iter_mut().for_each(|δ| δ.α_γ = 0.0),
-                        ControlFlow::Break(chg_one) => self.vec.iter_mut().for_each(|ρ| {
-                            let t = ρ.α_γ.abs();
-                            if t > 0.0 {
-                                if ALLOW_PARTIAL_TRANSPORT {
-                                    let new = (t - chg_one).max(0.0);
-                                    ρ.α_γ = ρ.α_γ.signum() * new;
-                                }
-                            }
-                        }),
-                    }
-                }
-            } else {
-                // This version zeroes smallest weights, avoiding partial transport.
-                let mut abs_weights_idx = self
-                    .vec
-                    .iter()
-                    .map(|ρ| ρ.α_γ.abs())
-                    .zip(0..)
-                    .filter(|(w, _)| *w >= 0.0)
-                    .collect::<Vec<(F, usize)>>();
-                abs_weights_idx.sort_by(|(a, _), (b, _)| a.partial_cmp(b).unwrap());
-
-                let mut left = reduction_target;
-
-                for (w, i) in abs_weights_idx {
-                    left -= w;
-                    let ρ = &mut self.vec[i];
-                    ρ.α_γ = 0.0;
-                    if left < 0.0 {
-                        break;
-                    }
-                }
-            }
-
-            all_ok = false
-        }
-
-        if !all_ok && *attempts >= tconfig.max_attempts {
-            for ρ in self.iter_mut() {
-                ρ.α_γ = 0.0;
-            }
-        }
-
-        for ρ in self.iter_mut() {
-            if ρ.α_γ == 0.0 {
-                ρ.fail_count += 1;
-            } else if all_ok {
-                ρ.fail_count = 0;
-            }
-        }
-
-        all_ok
-    }
-
     /// Returns $‖μ\^k - π\_♯\^0γ\^{k+1}‖$
     pub(crate) fn μ0_minus_γ0_radon(&self) -> F {
         self.vec.iter().map(|ρ| (ρ.α_μ_orig - ρ.α_γ).abs()).sum()
     }
 
     /// Returns $∫ c_2 d|γ|$
-    #[replace_float_literals(F::cast_from(literal))]
     pub(crate) fn c2integral(&self) -> F {
         self.vec
             .iter()
@@ -478,7 +686,6 @@
             .sum()
     }
 
-    #[replace_float_literals(F::cast_from(literal))]
     pub(crate) fn get_transport_stats(&self, stats: &mut IterInfo<F>, μ: &RNDM<N, F>) {
         // TODO: This doesn't take into account μ[i].α becoming zero in the latest tranport
         // attempt, for i < self.len(), when a corresponding source term also exists with index
@@ -522,22 +729,229 @@
         assert!(self.vec.len() <= μ.len());
         let old_len = μ.len();
         for (ρ, δ) in self.vec.iter_mut().zip(μ.iter_spikes()) {
-            ρ.prune = !(δ.α.abs() > F::EPSILON);
+            ρ.retain = δ.α.abs() > F::EPSILON;
         }
         μ.prune_by(|δ| δ.α.abs() > F::EPSILON);
         stats.pruned += old_len - μ.len();
-        self.vec.retain(|ρ| !ρ.prune);
+        self.vec.retain(|ρ| ρ.retain);
         assert!(self.vec.len() <= μ.len());
     }
+
+    /// Helper for initial transport. Called from [`TransportProxPenalty::initial_transport`].
+    pub(super) fn do_init_transport<D>(
+        &mut self,
+        v: &D,
+        μ: &RNDM<N, F>,
+        τθ_or_adaptive: &mut TransportStepLength<F>,
+        tconfig: &TransportConfig<F>,
+    ) where
+        D: DifferentiableRealMapping<N, F>,
+    {
+        use TransportStepLength::*;
+
+        // Initialise transport structure weights
+        for (δ, ρ) in izip!(μ.iter_spikes(), self.iter_mut()) {
+            ρ.α_μ_orig = δ.α;
+            ρ.x = δ.x;
+            ρ.y = δ.x; // Later updated if no fails.
+            ρ.α_γ = if ρ.fail_count > tconfig.max_fail {
+                0.0
+            } else {
+                // If old transport has opposing sign, the new transport will be none.
+                if (ρ.α_γ > 0.0 && δ.α < 0.0) || (ρ.α_γ < 0.0 && δ.α > 0.0) {
+                    0.0
+                } else {
+                    δ.α
+                }
+            };
+        }
+
+        let γ_prev_len = self.len();
+        assert!(μ.len() >= γ_prev_len);
+        self.extend(μ[γ_prev_len..].iter().map(|δ| SingleTransport {
+            x: δ.x,
+            y: δ.x, // Just something, will be filled properly in the next phase
+            α_μ_orig: δ.α,
+            α_γ: δ.α,
+            retain: true,
+            fail_count: 0,
+            excess: 0.0,
+        }));
+
+        // Calculate transport rays.
+        let simple_τθ = match *τθ_or_adaptive {
+            Fixed { τθ, .. } => Some(τθ),
+            Simple { ℓ_gradv, τθ0, ℓ_base } => Some(τθ0 / (ℓ_gradv + ℓ_base)),
+            AdaptiveMax {
+                ℓ_gradv,
+                ref mut adaptive_max_transport,
+                τθ0,
+                ℓ_base,
+                ℓ_base_max_transport,
+            } => {
+                *adaptive_max_transport = adaptive_max_transport.max(self.norm(Radon));
+                Some(τθ0 / (ℓ_gradv + ℓ_base + ℓ_base_max_transport * *adaptive_max_transport))
+            }
+            FullyAdaptive {
+                ref mut adaptive_ℓ_gradv,
+                ref mut adaptive_max_transport,
+                τθ0,
+                ℓ_base,
+                ℓ_base_max_transport,
+            } => {
+                *adaptive_max_transport = adaptive_max_transport.max(self.norm(Radon));
+                let mut τθ = τθ0
+                    / (*adaptive_ℓ_gradv + ℓ_base + ℓ_base_max_transport * *adaptive_max_transport);
+                // Do two runs through the spikes to update θ, breaking if first run did not cause
+                // a change.
+                for _i in 0..=1 {
+                    let mut changes = false;
+                    for ρ in self.iter_mut() {
+                        if ρ.fail_count < tconfig.max_fail {
+                            let dv_x = v.differential(&ρ.x);
+                            let g = &dv_x * (ρ.α_γ.signum() * τθ);
+                            ρ.y = ρ.x - g;
+                            let n = g.norm2();
+                            if n >= F::EPSILON {
+                                // Estimate Lipschitz factor of ∇v
+                                let this_ℓ_gradv = (dv_x - v.differential(&ρ.y)).norm2() / n;
+                                *adaptive_ℓ_gradv = adaptive_ℓ_gradv.max(this_ℓ_gradv);
+                                τθ = τθ0
+                                    / (*adaptive_ℓ_gradv
+                                        + ℓ_base
+                                        + ℓ_base_max_transport * *adaptive_max_transport);
+                                changes = true
+                            }
+                        }
+                    }
+                    if !changes {
+                        break;
+                    }
+                }
+                None
+            }
+        };
+
+        if let Some(τθ) = simple_τθ {
+            for ρ in self.iter_mut() {
+                if ρ.fail_count <= tconfig.max_fail {
+                    ρ.y = ρ.x - v.differential(&ρ.x) * (ρ.α_γ.signum() * τθ);
+                }
+            }
+        }
+    }
+
+    /// Helper for a posteriori transport error control.
+    /// Called from [`TransportProxPenalty::aposteriori_transport`].
+    fn do_new_aposteriori_transport<Reg, M>(
+        &mut self,
+        μ: &RNDM<N, F>,
+        τv̆: &mut M,
+        v: &mut M,
+        ε: F,
+        τ: F,
+        τθ_or_adaptive: &TransportStepLength<F>,
+        reg: &Reg,
+        tconfig: &TransportConfig<F>,
+        ω: impl Fn(&Loc<N, F>, std::cmp::Ordering, bool) -> F,
+    ) -> bool
+    where
+        Reg: SlidingRegTerm<Loc<N, F>, F>,
+        F: ToNalgebraRealField,
+        M: MinMaxMapping<Loc<N, F>, F> + DifferentiableRealMapping<N, F>,
+    {
+        let τℓ_gradv2 = τ * τθ_or_adaptive.get_ℓ_gradv() / 2.0;
+        let Bounds(α_lower, α_upper) = reg.subdiff_range();
+
+        let (all_ok, m) = izip!(self.vec.iter_mut(), μ.iter_spikes()).fold(
+            (true, 0.0),
+            |(all_ok_so_far, total_excess), (ρ, δ)| {
+                use std::cmp::Ordering::*;
+                // NOTE: The tolerances ε are commented out, because they will in any case
+                // be consumed by `t` below by suitably large choice of `tolerance_mult`.
+                // Hence, we simply implicitly adapt the `tolerance_mult` be commenting out
+                // the `ε` here. That way, `d` is in its entirely multiplied by τ, as is `t`,
+                // making all factors independent of τ.
+                let maybe_excess = match ρ.α_γ.total_cmp(&0.0) {
+                    Greater => (δ.α >= 0.0).then(|| {
+                        let gvx = v.differential(&ρ.x);
+                        let d = if tconfig.alt_remainder_control {
+                            let τv̆y = τv̆.apply(&ρ.y);
+                            (/*ε +*/τ * α_upper + ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), false))
+                                + (τv̆y + τ * gvx.dot(&ρ.x - &ρ.y))
+                                - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
+                        } else {
+                            (/*2.0 * ε*/-ω(&ρ.y, δ.α.total_cmp(&ρ.α_γ), true)
+                                + ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), false))
+                                + τ * gvx.dot(&ρ.x - &ρ.y)
+                                - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
+                        };
+                        d * ρ.α_γ
+                    }),
+                    Less => (δ.α <= 0.0).then(|| {
+                        let gvx = v.differential(&ρ.x);
+                        let d = if tconfig.alt_remainder_control {
+                            let τv̆y = τv̆.apply(&ρ.y);
+                            (/*ε*/-τ * α_lower - ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), true))
+                                - (τv̆y + τ * gvx.dot(&ρ.x - &ρ.y))
+                                - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
+                        } else {
+                            (/*2.0 * ε +*/ω(&ρ.y, δ.α.total_cmp(&ρ.α_γ), false)
+                                - ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), true))
+                                - τ * gvx.dot(&ρ.x - &ρ.y)
+                                - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
+                        };
+                        d * (-ρ.α_γ)
+                    }),
+                    Equal => Some(0.0),
+                };
+                match maybe_excess {
+                    None => {
+                        ρ.α_γ = 0.0;
+                        ρ.excess = 0.0;
+                        (false, total_excess)
+                    }
+                    Some(e) => {
+                        ρ.excess = e;
+                        (all_ok_so_far, total_excess + e)
+                    }
+                }
+            },
+        );
+
+        let t = τ * ε * tconfig.tolerance_mult;
+
+        if m > t {
+            let mut it = self.vec.iter_mut().rev().filter(|ρ| ρ.excess > 0.0);
+            let reduction_target = m - tconfig.adaptation * t;
+            let _unused = it.try_fold(reduction_target, |left, ρ| {
+                let d = ρ.excess;
+                if d >= left {
+                    if tconfig.allow_partial_transport {
+                        ρ.α_γ *= (d - left) / d;
+                    } else {
+                        ρ.α_γ = 0.0;
+                    }
+                    ControlFlow::Break(())
+                } else {
+                    ρ.α_γ = 0.0;
+                    ControlFlow::Continue(left - d)
+                }
+            });
+            false
+        } else {
+            all_ok
+        }
+    }
 }
 
-impl<const N: usize, F: Float> Norm<Radon, F> for Transport<N, F> {
+impl<const N: usize, F: Float> Norm<Radon, F> for Transport<Loc<N, F>, F> {
     fn norm(&self, _: Radon) -> F {
         self.iter().map(|ρ| ρ.α_γ.abs()).sum()
     }
 }
 
-impl<const N: usize, F: Float> MulAssign<F> for Transport<N, F> {
+impl<const N: usize, F: Float> MulAssign<F> for Transport<Loc<N, F>, F> {
     fn mul_assign(&mut self, factor: F) {
         for ρ in self.iter_mut() {
             ρ.α_γ *= factor;
@@ -564,11 +978,12 @@
     F: Float + ToNalgebraRealField,
     I: AlgIteratorFactory<IterInfo<F>>,
     Dat: DifferentiableMapping<RNDM<N, F>, Codomain = F> + BoundedCurvature<F>,
-    Dat::DerivativeDomain: DifferentiableRealMapping<N, F> + ClosedMul<F>,
+    Dat::DerivativeDomain:
+        DifferentiableRealMapping<N, F> + ClosedMul<F> + MinMaxMapping<Loc<N, F>, F>,
     //for<'a> Dat::Differential<'a>: Lipschitz<&'a P, FloatType = F>,
     RNDM<N, F>: SpikeMerging<F>,
     Reg: SlidingRegTerm<Loc<N, F>, F>,
-    P: ProxPenalty<Loc<N, F>, Dat::DerivativeDomain, Reg, F> + StepLengthBound<F, Dat>,
+    P: TransportProxPenalty<Loc<N, F>, Dat::DerivativeDomain, Reg, F> + StepLengthBound<F, Dat>,
     Plot: Plotter<P::ReturnMapping, Dat::DerivativeDomain, RNDM<N, F>>,
 {
     // Check parameters
@@ -584,30 +999,14 @@
     //let max_transport = config.max_transport.scale
     //                    * reg.radon_norm_bound(b.norm2_squared() / 2.0);
     //let ℓ = opA.transport.lipschitz_factor(L2Squared) * max_transport;
-    let ℓ = 0.0;
     let τ = config.τ0 / prox_penalty.step_length_bound(&f)?;
 
-    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)
-            };
-            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: 10.0 * 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,
+        0.0,
+    );
     // We multiply tolerance by τ for FB since our subproblems depending on tolerances are scaled
     // by τ compared to the conditional gradient approach.
     let tolerance = config.insertion.tolerance * τ * reg.tolerance_scaling();
@@ -625,9 +1024,18 @@
 
     // Run the algorithm
     for state in iterator.iter_init(|| full_stats(&μ, ε, stats.clone())) {
+        let mut v = f.differential(&μ);
+
         // Calculate initial transport
-        let v = f.differential(&μ);
-        γ.initial_transport(&μ, τ, &mut θ_or_adaptive, v, &config.transport);
+        prox_penalty.initial_transport(
+            &mut γ,
+            &μ,
+            ε,
+            τ,
+            &mut τθ_or_adaptive,
+            &v,
+            &config.transport,
+        );
 
         let mut attempts = 0;
 
@@ -658,8 +1066,21 @@
             )?;
 
             // A posteriori transport adaptation.
-            if γ.aposteriori_transport(&μ, &μ̆, &mut τv̆, None, ε, &config.transport, &mut attempts)
-            {
+            if prox_penalty.aposteriori_transport(
+                &mut γ,
+                &μ,
+                &μ̆,
+                &mut τv̆,
+                &mut v,
+                None,
+                ε,
+                τ,
+                &τθ_or_adaptive,
+                reg,
+                &config.transport,
+                &config.insertion.refinement,
+                &mut attempts,
+            ) {
                 break 'adapt_transport (maybe_d, within_tolerances, τv̆, μ̆);
             }
 
@@ -672,7 +1093,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̆,
                 &μ̆,
@@ -682,6 +1103,10 @@
                 &reg,
                 Some(|μ̃: &RNDM<N, F>| f.apply(μ̃)),
             );
+            if m > 0 {
+                stats.merged += m;
+                v = f.differential(&μ);
+            }
         }
 
         γ.prune_compat(&mut μ, &mut stats);

mercurial