src/sliding_fb.rs

changeset 72
e9a460a0e638
parent 68
00d0881f89a6
child 75
677a5fd1b014
equal deleted inserted replaced
71:e2953ffd4e0b 72:e9a460a0e638
14 use crate::fb::*; 14 use crate::fb::*;
15 use crate::forward_model::{BoundedCurvature, BoundedCurvatureGuess}; 15 use crate::forward_model::{BoundedCurvature, BoundedCurvatureGuess};
16 use crate::measures::merging::SpikeMerging; 16 use crate::measures::merging::SpikeMerging;
17 use crate::measures::{DeltaMeasure, DiscreteMeasure, Radon, RNDM}; 17 use crate::measures::{DeltaMeasure, DiscreteMeasure, Radon, RNDM};
18 use crate::plot::Plotter; 18 use crate::plot::Plotter;
19 use crate::prox_penalty::{ProxPenalty, StepLengthBound}; 19 use crate::prox_penalty::{ProxPenalty, RadonSquared, StepLengthBound};
20 use crate::regularisation::SlidingRegTerm; 20 use crate::regularisation::SlidingRegTerm;
21 use crate::seminorms::DiscreteMeasureOp;
21 use crate::types::*; 22 use crate::types::*;
23 use alg_tools::bounds::{Bounds, MinMaxMapping};
22 use alg_tools::error::DynResult; 24 use alg_tools::error::DynResult;
23 use alg_tools::euclidean::Euclidean; 25 use alg_tools::euclidean::Euclidean;
26 use alg_tools::instance::Space;
24 use alg_tools::iterate::AlgIteratorFactory; 27 use alg_tools::iterate::AlgIteratorFactory;
25 use alg_tools::mapping::{DifferentiableMapping, DifferentiableRealMapping}; 28 use alg_tools::mapping::{DifferentiableMapping, DifferentiableRealMapping, Mapping, RealMapping};
26 use alg_tools::nalgebra_support::ToNalgebraRealField; 29 use alg_tools::nalgebra_support::ToNalgebraRealField;
27 use alg_tools::norms::Norm; 30 use alg_tools::norms::Norm;
28 use anyhow::ensure; 31 use anyhow::ensure;
29 use std::ops::ControlFlow; 32 use std::ops::ControlFlow;
30 33
32 #[derive(Clone, Copy, Eq, PartialEq, Serialize, Deserialize, Debug)] 35 #[derive(Clone, Copy, Eq, PartialEq, Serialize, Deserialize, Debug)]
33 #[serde(default)] 36 #[serde(default)]
34 pub struct TransportConfig<F: Float> { 37 pub struct TransportConfig<F: Float> {
35 /// Transport step length $θ$ normalised to $(0, 1)$. 38 /// Transport step length $θ$ normalised to $(0, 1)$.
36 pub θ0: F, 39 pub θ0: F,
40 /// Unnormalised transport step length $θ$. Overrides θ0.
41 pub τθ: Option<F>,
37 /// Factor in $(0, 1)$ for decreasing transport to adapt to tolerance. 42 /// Factor in $(0, 1)$ for decreasing transport to adapt to tolerance.
38 pub adaptation: F, 43 pub adaptation: F,
39 /// A posteriori transport tolerance multiplier (C_pos) 44 /// A posteriori transport tolerance multiplier
40 pub tolerance_mult_con: F, 45 pub tolerance_mult: F,
46 /// Multiplier for rough estimate of ℓ_{∇v}. Should be ≥ 1.
47 /// If explicit τθ is given, ℓ_{∇v} is one divided by this number times τθ.
48 /// Otherwise, if τθ relative to an estimate of ℓ_{∇v}, this number multiplies that estimate.
49 pub ℓ_gradv_mult: F,
41 /// maximum number of adaptation iterations, until cancelling transport. 50 /// maximum number of adaptation iterations, until cancelling transport.
42 pub max_attempts: usize, 51 pub max_attempts: usize,
43 /// Maximum number of failed transportations for a single source point 52 /// Maximum number of failed transportations for a single source point
44 pub max_fail: usize, 53 pub max_fail: usize,
54 /// Allow points to be transported partially.
55 pub allow_partial_transport: bool,
56 /// Use an alternative remainder control rule.
57 pub alt_remainder_control: bool,
45 } 58 }
46 59
47 #[replace_float_literals(F::cast_from(literal))] 60 #[replace_float_literals(F::cast_from(literal))]
48 impl<F: Float> TransportConfig<F> { 61 impl<F: Float> TransportConfig<F> {
49 /// Check that the parameters are ok. Panics if not. 62 /// Check that the parameters are ok. Panics if not.
50 pub fn check(&self) -> DynResult<()> { 63 pub fn check(&self) -> DynResult<()> {
51 ensure!(self.θ0 > 0.0); 64 ensure!(self.θ0 > 0.0);
52 ensure!(0.0 < self.adaptation && self.adaptation < 1.0); 65 ensure!(0.0 < self.adaptation && self.adaptation < 1.0);
53 ensure!(self.tolerance_mult_con > 0.0); 66 ensure!(self.tolerance_mult > 0.0);
54 Ok(()) 67 Ok(())
55 } 68 }
56 } 69 }
57 70
58 #[replace_float_literals(F::cast_from(literal))] 71 #[replace_float_literals(F::cast_from(literal))]
59 impl<F: Float> Default for TransportConfig<F> { 72 impl<F: Float> Default for TransportConfig<F> {
60 fn default() -> Self { 73 fn default() -> Self {
61 TransportConfig { 74 TransportConfig {
62 θ0: 0.9, 75 θ0: 0.99,
76 τθ: None,
63 adaptation: 0.9, 77 adaptation: 0.9,
64 tolerance_mult_con: 100.0, 78 allow_partial_transport: true,
79 alt_remainder_control: false,
80 tolerance_mult: 1e1,
81 ℓ_gradv_mult: 3.0,
65 max_attempts: 2, 82 max_attempts: 2,
66 max_fail: usize::MAX, 83 max_fail: usize::MAX,
67 } 84 }
68 } 85 }
69 } 86 }
96 } 113 }
97 } 114 }
98 } 115 }
99 116
100 /// Internal type of adaptive transport step length calculation 117 /// Internal type of adaptive transport step length calculation
101 pub(crate) enum TransportStepLength<F: Float, G: Fn(F, F) -> F> { 118 #[derive(Clone, Debug, Serialize, Deserialize)]
119 pub enum TransportStepLength<F: Float> {
102 /// Fixed, known step length 120 /// Fixed, known step length
103 #[allow(dead_code)] 121 #[allow(dead_code)]
104 Fixed(F), 122 Fixed { τθ: F, ℓ_gradv: F },
123 /// Simple step lengths that do not depend on maximum transport
124 Simple { ℓ_gradv: F, τθ0: F, ℓ_base: F },
105 /// Adaptive step length, only wrt. maximum transport. 125 /// Adaptive step length, only wrt. maximum transport.
106 /// Content of `l` depends on use case, while `g` calculates the step length from `l`. 126 AdaptiveMax {
107 AdaptiveMax { l: F, max_transport: F, g: G }, 127 ℓ_gradv: F,
128 adaptive_max_transport: F,
129 τθ0: F,
130 ℓ_base: F,
131 ℓ_base_max_transport: F,
132 },
108 /// Adaptive step length. 133 /// Adaptive step length.
109 /// Content of `l` depends on use case, while `g` calculates the step length from `l`. 134 FullyAdaptive {
110 FullyAdaptive { l: F, max_transport: F, g: G }, 135 adaptive_ℓ_gradv: F,
136 adaptive_max_transport: F,
137 τθ0: F,
138 ℓ_base: F,
139 ℓ_base_max_transport: F,
140 },
141 }
142
143 #[replace_float_literals(F::cast_from(literal))]
144 impl<F: Float> TransportStepLength<F> {
145 fn get_ℓ_gradv(&self) -> F {
146 use TransportStepLength::*;
147 match *self {
148 Fixed { ℓ_gradv, .. } => ℓ_gradv,
149 Simple { ℓ_gradv, .. } => ℓ_gradv,
150 AdaptiveMax { ℓ_gradv, .. } => ℓ_gradv,
151 FullyAdaptive { adaptive_ℓ_gradv, .. } => adaptive_ℓ_gradv,
152 }
153 }
154
155 fn new(
156 maybe_ℓ_gradv_est: DynResult<F>,
157 tconfig: &TransportConfig<F>,
158 ℓ_base: F,
159 ℓ_base_max_transport: F,
160 ) -> Self {
161 if let Some(τθ) = tconfig.τθ {
162 TransportStepLength::Fixed { τθ, ℓ_gradv: 1.0 / (τθ * tconfig.ℓ_gradv_mult) }
163 } else {
164 match maybe_ℓ_gradv_est {
165 Ok(ℓ_gradv_est) => {
166 if ℓ_base_max_transport == 0.0 {
167 TransportStepLength::Simple {
168 ℓ_gradv: tconfig.ℓ_gradv_mult * ℓ_gradv_est,
169 τθ0: tconfig.θ0,
170 ℓ_base,
171 }
172 } else {
173 TransportStepLength::AdaptiveMax {
174 ℓ_gradv: tconfig.ℓ_gradv_mult * ℓ_gradv_est,
175 adaptive_max_transport: 0.0,
176 τθ0: tconfig.θ0,
177 ℓ_base,
178 ℓ_base_max_transport,
179 }
180 }
181 }
182 Err(_) => TransportStepLength::FullyAdaptive {
183 adaptive_ℓ_gradv: 10.0 * F::EPSILON, // Start with something very small to estimate differentials
184 adaptive_max_transport: 0.0,
185 τθ0: tconfig.θ0,
186 ℓ_base,
187 ℓ_base_max_transport,
188 },
189 }
190 }
191 }
111 } 192 }
112 193
113 #[derive(Clone, Debug, Serialize)] 194 #[derive(Clone, Debug, Serialize)]
114 pub struct SingleTransport<const N: usize, F: Float> { 195 pub struct SingleTransport<Domain, F: Float> {
115 /// Source point 196 /// Source point
116 x: Loc<N, F>, 197 x: Domain,
117 /// Target point 198 /// Target point
118 y: Loc<N, F>, 199 y: Domain,
119 /// Original mass 200 /// Original mass
120 α_μ_orig: F, 201 α_μ_orig: F,
121 /// Transported mass 202 /// Transported mass
122 α_γ: F, 203 α_γ: F,
123 /// Helper for pruning 204 /// Helper for pruning
124 prune: bool, 205 retain: bool,
125 /// Fail count 206 /// Fail count
126 fail_count: usize, 207 fail_count: usize,
208 /// Contribution to remainder (temporary variable)
209 excess: F,
127 } 210 }
128 211
129 #[derive(Clone, Debug, Serialize)] 212 #[derive(Clone, Debug, Serialize)]
130 pub struct Transport<const N: usize, F: Float> { 213 pub struct Transport<Domain, F: Float> {
131 vec: Vec<SingleTransport<N, F>>, 214 vec: Vec<SingleTransport<Domain, F>>,
132 } 215 }
133 216
134 /// Whether partiall transported points are allowed. 217 /// Whether partially transported points are allowed.
135 /// 218 ///
136 /// Partial transport can cause spike count explosion, so full or zero 219 /// Partial transport can cause spike count explosion, so full or zero
137 /// transport is generally preferred. If this is set to `true`, different 220 /// transport is generally preferred. If this is set to `true`, different
138 /// transport adaptation heuristics will be used. 221 /// transport adaptation heuristics will be used.
139 const ALLOW_PARTIAL_TRANSPORT: bool = true;
140 const MINIMAL_PARTIAL_TRANSPORT: bool = true; 222 const MINIMAL_PARTIAL_TRANSPORT: bool = true;
141 223 const NEW_APPROACH: bool = true;
142 impl<const N: usize, F: Float> Transport<N, F> { 224
225 pub trait TransportProxPenalty<Domain, PreadjointCodomain, Reg, F = f64>:
226 ProxPenalty<Domain, PreadjointCodomain, Reg, F>
227 where
228 F: Float + ToNalgebraRealField,
229 Reg: SlidingRegTerm<Domain, F>,
230 Domain: Space + Clone,
231 {
232 type TransportStepLength;
233
234 /// Constrution of initial transport `γ1` from initial measure `μ` and `v=F'(μ)`
235 /// with step lengh τ and transport step length `θ_or_adaptive`.
236 fn initial_transport(
237 &self,
238 γ: &mut Transport<Domain, F>,
239 μ: &DiscreteMeasure<Domain, F>,
240 ε: F,
241 τ: F,
242 τθ_or_adaptive: &mut Self::TransportStepLength,
243 v: &PreadjointCodomain,
244 tconfig: &TransportConfig<F>,
245 );
246
247 /// A posteriori transport adaptation.
248 fn aposteriori_transport(
249 &self,
250 γ: &mut Transport<Domain, F>,
251 μ: &DiscreteMeasure<Domain, F>,
252 μ̆: &DiscreteMeasure<Domain, F>,
253 τv̆: &mut PreadjointCodomain,
254 v: &mut PreadjointCodomain,
255 extra: Option<F>,
256 ε: F,
257 τ: F,
258 τθ_or_adaptive: &Self::TransportStepLength,
259 reg: &Reg,
260 tconfig: &TransportConfig<F>,
261 rconfig: &RefinementSettings<F>,
262 attempts: &mut usize,
263 ) -> bool;
264
265 fn get_transport_steplength(
266 &self,
267 lips: (DynResult<F>, DynResult<F>, DynResult<F>),
268 tconfig: &TransportConfig<F>,
269 ℓ_base: F,
270 ℓ_base_max_transport: F,
271 ) -> Self::TransportStepLength;
272 }
273
274 #[replace_float_literals(F::cast_from(literal))]
275 impl<F, M, Reg, const N: usize> TransportProxPenalty<Loc<N, F>, M, Reg, F> for RadonSquared
276 where
277 RadonSquared: ProxPenalty<Loc<N, F>, M, Reg, F>,
278 F: Float + ToNalgebraRealField,
279 M: MinMaxMapping<Loc<N, F>, F> + DifferentiableRealMapping<N, F>,
280 Reg: SlidingRegTerm<Loc<N, F>, F>,
281 RNDM<N, F>: SpikeMerging<F>,
282 {
283 type TransportStepLength = TransportStepLength<F>;
284
285 fn get_transport_steplength(
286 &self,
287 (ℓ_F, maybe_ℓ_gradv, _maybe_transport_lip): (DynResult<F>, DynResult<F>, DynResult<F>),
288 tconfig: &TransportConfig<F>,
289 ℓ: F,
290 ℓ_base_max_transport: F,
291 ) -> Self::TransportStepLength {
292 TransportStepLength::new(
293 maybe_ℓ_gradv,
294 tconfig,
295 ℓ + ℓ_F.unwrap_or(0.0),
296 ℓ_base_max_transport,
297 )
298 }
299
300 fn initial_transport(
301 &self,
302 γ: &mut Transport<Loc<N, F>, F>,
303 μ: &RNDM<N, F>,
304 _ε: F,
305 _τ: F,
306 τθ_or_adaptive: &mut TransportStepLength<F>,
307 v: &M,
308 tconfig: &TransportConfig<F>,
309 ) {
310 γ.do_init_transport(v, μ, τθ_or_adaptive, tconfig);
311 }
312
313 fn aposteriori_transport(
314 &self,
315 γ: &mut Transport<Loc<N, F>, F>,
316 μ: &RNDM<N, F>,
317 μ̆: &RNDM<N, F>,
318 τv̆: &mut M,
319 v: &mut M,
320 _extra: Option<F>,
321 ε: F,
322 τ: F,
323 τθ_or_adaptive: &TransportStepLength<F>,
324 reg: &Reg,
325 tconfig: &TransportConfig<F>,
326 _rconfig: &RefinementSettings<F>,
327 attempts: &mut usize,
328 ) -> bool {
329 *attempts += 1;
330
331 if *attempts > tconfig.max_attempts {
332 // Previous round has set transport to zero
333 return true;
334 }
335
336 let nΔ = μ.dist_matching(&μ̆);
337 let all_ok = γ.do_new_aposteriori_transport(
338 μ,
339 τv̆,
340 v,
341 ε,
342 τ,
343 τθ_or_adaptive,
344 reg,
345 tconfig,
346 |_, mass_zero, negated| {
347 use std::cmp::Ordering::*;
348 match (mass_zero, negated) {
349 (Less, _) => -nΔ,
350 (Greater, _) => nΔ,
351 (Equal, false) => nΔ, // pessimistic estimated for ω
352 (Equal, true) => -nΔ, // pessimistic estimated for -ω
353 }
354 },
355 );
356
357 if !all_ok {
358 if *attempts >= tconfig.max_attempts {
359 for ρ in γ.iter_mut() {
360 ρ.α_γ = 0.0;
361 }
362 }
363 } else {
364 for ρ in γ.iter_mut() {
365 if ρ.α_γ == 0.0 {
366 ρ.fail_count += 1;
367 } else {
368 ρ.fail_count = 0;
369 }
370 }
371 }
372 all_ok
373 }
374 }
375
376 #[replace_float_literals(F::cast_from(literal))]
377 impl<F, M, Reg, 𝒟, O, const N: usize> TransportProxPenalty<Loc<N, F>, M, Reg, F> for 𝒟
378 where
379 F: Float + ToNalgebraRealField,
380 𝒟: DiscreteMeasureOp<Loc<N, F>, F>,
381 𝒟::Codomain: RealMapping<N, F>,
382 M: MinMaxMapping<Loc<N, F>, F> + DifferentiableRealMapping<N, F>,
383 for<'a> &'a M: std::ops::Add<𝒟::PreCodomain, Output = O>,
384 O: MinMaxMapping<Loc<N, F>, F>,
385 Reg: SlidingRegTerm<Loc<N, F>, F>,
386 Self: ProxPenalty<Loc<N, F>, M, Reg, F>,
387 RNDM<N, F>: SpikeMerging<F>,
388 {
389 type TransportStepLength = TransportStepLength<F>;
390
391 fn get_transport_steplength(
392 &self,
393 (ℓ_F, maybe_ℓ_gradv, _maybe_transport_lip): (DynResult<F>, DynResult<F>, DynResult<F>),
394 tconfig: &TransportConfig<F>,
395 ℓ_base: F,
396 ℓ_base_max_transport: F,
397 ) -> Self::TransportStepLength {
398 TransportStepLength::new(
399 maybe_ℓ_gradv,
400 tconfig,
401 ℓ_base + ℓ_F.unwrap_or(0.0),
402 ℓ_base_max_transport,
403 )
404 }
405
406 fn initial_transport(
407 &self,
408 γ: &mut Transport<Loc<N, F>, F>,
409 μ: &RNDM<N, F>,
410 _ε: F,
411 _τ: F,
412 τθ_or_adaptive: &mut TransportStepLength<F>,
413 v: &M,
414 tconfig: &TransportConfig<F>,
415 ) {
416 γ.do_init_transport(v, μ, τθ_or_adaptive, tconfig);
417 }
418
419 fn aposteriori_transport(
420 &self,
421 γ: &mut Transport<Loc<N, F>, F>,
422 μ: &RNDM<N, F>,
423 μ̆: &RNDM<N, F>,
424 τv̆: &mut M,
425 v: &mut M,
426 extra: Option<F>,
427 ε: F,
428 τ: F,
429 τθ_or_adaptive: &TransportStepLength<F>,
430 reg: &Reg,
431 tconfig: &TransportConfig<F>,
432 _rconfig: &RefinementSettings<F>,
433 attempts: &mut usize,
434 ) -> bool {
435 *attempts += 1;
436
437 if *attempts > tconfig.max_attempts {
438 // Previous round has set transport to zero
439 return true;
440 }
441
442 let all_ok = if NEW_APPROACH {
443 let ω = self.apply(μ.sub_matching(&μ̆));
444 γ.do_new_aposteriori_transport(
445 μ,
446 τv̆,
447 v,
448 ε,
449 τ,
450 τθ_or_adaptive,
451 reg,
452 tconfig,
453 |x, _, _| ω.apply(x),
454 )
455 } else {
456 let mut all_ok0 = true;
457
458 // 1. If π_♯^1γ^{k+1} = γ1 has non-zero mass at some point y, but μ = μ^{k+1} does not,
459 // then the ansatz ∇w̃_x(y) = w^{k+1}(y) may not be satisfied. So set the mass of γ1
460 // at that point to zero, and retry.
461 for (δ, ρ) in izip!(μ.iter_spikes(), γ.iter_mut()) {
462 if δ.α == 0.0 && ρ.α_γ != 0.0 {
463 all_ok0 = false;
464 ρ.α_γ = 0.0;
465 }
466 // TODO: sign
467 }
468
469 // 2. Through bounding ∫ B_ω(y, z) dλ(x, y, z).
470 // through the estimate ≤ C ‖Δ‖‖γ^{k+1}‖ for Δ := μ^{k+1}-μ̆^k
471 // which holds for some some C if the convolution kernel in 𝒟 has Lipschitz gradient.
472
473 let nγ = γ.norm(Radon);
474 let nΔ = μ.dist_matching(&μ̆) + extra.unwrap_or(0.0);
475 let t = ε * tconfig.tolerance_mult;
476 if nγ * nΔ > t && *attempts >= tconfig.max_attempts {
477 all_ok0 = false;
478 } else if nγ * nΔ > t {
479 // Since t/(nγ*nΔ)<1, and the constant tconfig.adaptation < 1,
480 // this will guarantee that eventually ‖γ‖ decreases sufficiently that we
481 // will not enter here.
482 //*γ *= tconfig.adaptation * t / (nγ * nΔ);
483
484 // We want a consistent behaviour that has the potential to set many weights to zero.
485 // Therefore, we find the smallest uniform reduction `chg_one`, subtracted
486 // from all weights, that achieves total `adapt` adaptation.
487 let adapt_to = tconfig.adaptation * t / nΔ;
488 let reduction_target = nγ - adapt_to;
489 assert!(reduction_target > 0.0);
490 if tconfig.allow_partial_transport {
491 if MINIMAL_PARTIAL_TRANSPORT {
492 // This reduces weights of transport, starting from … until `adapt` is
493 // exhausted. It will, therefore, only ever cause one extrap point insertion
494 // at the sources, unlike “full” partial transport.
495 //let refs = γ.vec.iter_mut().collect::<Vec<_>>();
496 //refs.sort_by(|ρ1, ρ2| ρ1.α_γ.abs().partial_cmp(&ρ2.α_γ.abs()).unwrap());
497 // let mut it = refs.into_iter();
498 //
499 // Maybe sort by differential norm
500 // let mut refs = γ
501 // .vec
502 // .iter_mut()
503 // .map(|ρ| {
504 // let val = v.differential(&ρ.x).norm2_squared();
505 // (ρ, val)
506 // })
507 // .collect::<Vec<_>>();
508 // refs.sort_by(|(_, v1), (_, v2)| v2.partial_cmp(&v1).unwrap());
509 // let mut it = refs.into_iter().map(|(ρ, _)| ρ);
510 let mut it = γ.vec.iter_mut().rev();
511 let _unused = it.try_fold(reduction_target, |left, ρ| {
512 let w = ρ.α_γ.abs();
513 if left <= w {
514 ρ.α_γ = ρ.α_γ.signum() * (w - left);
515 ControlFlow::Break(())
516 } else {
517 ρ.α_γ = 0.0;
518 ControlFlow::Continue(left - w)
519 }
520 });
521 } else {
522 // This version equally reduces all weights. It causes partial transport, which
523 // has the problem that that we need to then adapt weights in both start and
524 // end points, in insert_and_reweigh, somtimes causing the number of spikes μ
525 // to explode.
526 let mut abs_weights = γ
527 .vec
528 .iter()
529 .map(|ρ| ρ.α_γ.abs())
530 .filter(|t| *t > F::EPSILON)
531 .collect::<Vec<F>>();
532 abs_weights.sort_by(|a, b| a.total_cmp(b));
533 let n = abs_weights.len();
534 // Cannot have partial transport; can cause spike count explosion
535 let chg = abs_weights.into_iter().zip((1..=n).rev()).try_fold(
536 0.0,
537 |smaller_total, (w, m)| {
538 let mf = F::cast_from(m);
539 let reduction = w * mf + smaller_total;
540 if reduction >= reduction_target {
541 ControlFlow::Break((reduction_target - smaller_total) / mf)
542 } else {
543 ControlFlow::Continue(smaller_total + w)
544 }
545 },
546 );
547 match chg {
548 ControlFlow::Continue(_) => γ.vec.iter_mut().for_each(|δ| δ.α_γ = 0.0),
549 ControlFlow::Break(chg_one) => γ.vec.iter_mut().for_each(|ρ| {
550 let t = ρ.α_γ.abs();
551 if t > 0.0 {
552 if tconfig.allow_partial_transport {
553 let new = (t - chg_one).max(0.0);
554 ρ.α_γ = ρ.α_γ.signum() * new;
555 }
556 }
557 }),
558 }
559 }
560 } else {
561 // This version zeroes smallest weights, avoiding partial transport.
562 let mut abs_weights_idx = γ
563 .vec
564 .iter()
565 .map(|ρ| ρ.α_γ.abs())
566 .zip(0..)
567 .filter(|(w, _)| *w >= 0.0)
568 .collect::<Vec<(F, usize)>>();
569 abs_weights_idx.sort_by(|(a, _), (b, _)| a.total_cmp(b));
570
571 let mut left = reduction_target;
572
573 for (w, i) in abs_weights_idx {
574 left -= w;
575 let ρ = &mut γ.vec[i];
576 ρ.α_γ = 0.0;
577 if left < 0.0 {
578 break;
579 }
580 }
581 }
582
583 all_ok0 = false
584 }
585 all_ok0
586 };
587
588 if !all_ok {
589 if *attempts >= tconfig.max_attempts {
590 for ρ in γ.iter_mut() {
591 ρ.α_γ = 0.0;
592 }
593 }
594 } else {
595 for ρ in γ.iter_mut() {
596 if ρ.α_γ == 0.0 {
597 ρ.fail_count += 1;
598 } else {
599 ρ.fail_count = 0;
600 }
601 }
602 }
603 all_ok
604 }
605 }
606
607 #[replace_float_literals(F::cast_from(literal))]
608 impl<const N: usize, F: Float> Transport<Loc<N, F>, F> {
143 pub(crate) fn new() -> Self { 609 pub(crate) fn new() -> Self {
144 Transport { vec: Vec::new() } 610 Transport { vec: Vec::new() }
145 } 611 }
146 612
147 pub(crate) fn iter(&self) -> impl Iterator<Item = &'_ SingleTransport<N, F>> { 613 pub(crate) fn iter(&self) -> impl Iterator<Item = &'_ SingleTransport<Loc<N, F>, F>> {
148 self.vec.iter() 614 self.vec.iter()
149 } 615 }
150 616
151 pub(crate) fn iter_mut(&mut self) -> impl Iterator<Item = &'_ mut SingleTransport<N, F>> { 617 pub(crate) fn iter_mut(
618 &mut self,
619 ) -> impl Iterator<Item = &'_ mut SingleTransport<Loc<N, F>, F>> {
152 self.vec.iter_mut() 620 self.vec.iter_mut()
153 } 621 }
154 622
155 pub(crate) fn extend<I>(&mut self, it: I) 623 pub(crate) fn extend<I>(&mut self, it: I)
156 where 624 where
157 I: IntoIterator<Item = SingleTransport<N, F>>, 625 I: IntoIterator<Item = SingleTransport<Loc<N, F>, F>>,
158 { 626 {
159 self.vec.extend(it) 627 self.vec.extend(it)
160 } 628 }
161 629
162 pub(crate) fn len(&self) -> usize { 630 pub(crate) fn len(&self) -> usize {
169 // .map(|(ρ, δ)| (ρ.α_γ - δ.α).abs()) 637 // .map(|(ρ, δ)| (ρ.α_γ - δ.α).abs())
170 // .sum() 638 // .sum()
171 // } 639 // }
172 640
173 /// Construct `μ̆`, replacing the contents of `μ`. 641 /// Construct `μ̆`, replacing the contents of `μ`.
174 #[replace_float_literals(F::cast_from(literal))]
175 pub(crate) fn μ̆_into(&self, μ: &mut RNDM<N, F>) { 642 pub(crate) fn μ̆_into(&self, μ: &mut RNDM<N, F>) {
176 assert!(self.len() <= μ.len()); 643 assert!(self.len() <= μ.len());
177 644
178 // First transported points 645 // First transported points
179 for (δ, ρ) in izip!(μ.iter_spikes_mut(), self.iter()) { 646 for (δ, ρ) in izip!(μ.iter_spikes_mut(), self.iter()) {
188 } 655 }
189 } 656 }
190 657
191 // Then source points with partial transport 658 // Then source points with partial transport
192 let mut i = self.len(); 659 let mut i = self.len();
193 if ALLOW_PARTIAL_TRANSPORT { 660 // This can cause the number of points to explode, so cannot have partial transport.
194 // This can cause the number of points to explode, so cannot have partial transport. 661 for ρ in self.iter() {
195 for ρ in self.iter() { 662 let α = ρ.α_μ_orig - ρ.α_γ;
196 let α = ρ.α_μ_orig - ρ.α_γ; 663 if ρ.α_γ.abs() > F::EPSILON && α.abs() > F::EPSILON {
197 if ρ.α_γ.abs() > F::EPSILON && α != 0.0 { 664 let δ = DeltaMeasure { α, x: ρ.x };
198 let δ = DeltaMeasure { α, x: ρ.x }; 665 if i < μ.len() {
199 if i < μ.len() { 666 μ[i] = δ;
200 μ[i] = δ; 667 } else {
201 } else { 668 μ.push(δ)
202 μ.push(δ) 669 }
203 } 670 i += 1;
204 i += 1;
205 }
206 } 671 }
207 } 672 }
208 μ.truncate(i); 673 μ.truncate(i);
209 }
210
211 /// Constrution of initial transport `γ1` from initial measure `μ` and `v=F'(μ)`
212 /// with step lengh τ and transport step length `θ_or_adaptive`.
213 #[replace_float_literals(F::cast_from(literal))]
214 pub(crate) fn initial_transport<G, D>(
215 &mut self,
216 μ: &RNDM<N, F>,
217 _τ: F,
218 τθ_or_adaptive: &mut TransportStepLength<F, G>,
219 v: D,
220 tconfig: &TransportConfig<F>,
221 ) where
222 G: Fn(F, F) -> F,
223 D: DifferentiableRealMapping<N, F>,
224 {
225 use TransportStepLength::*;
226
227 // Initialise transport structure weights
228 for (δ, ρ) in izip!(μ.iter_spikes(), self.iter_mut()) {
229 ρ.α_μ_orig = δ.α;
230 ρ.x = δ.x;
231 if ρ.fail_count > tconfig.max_fail {
232 ρ.α_γ = 0.0
233 } else {
234 // If old transport has opposing sign, the new transport will be none.
235 ρ.α_γ = if (ρ.α_γ > 0.0 && δ.α < 0.0) || (ρ.α_γ < 0.0 && δ.α > 0.0) {
236 0.0
237 } else {
238 δ.α
239 }
240 }
241 }
242
243 let γ_prev_len = self.len();
244 assert!(μ.len() >= γ_prev_len);
245 self.extend(μ[γ_prev_len..].iter().map(|δ| SingleTransport {
246 x: δ.x,
247 y: δ.x, // Just something, will be filled properly in the next phase
248 α_μ_orig: δ.α,
249 α_γ: δ.α,
250 prune: false,
251 fail_count: 0,
252 }));
253
254 // Calculate transport rays.
255 match *τθ_or_adaptive {
256 Fixed(θ) => {
257 for ρ in self.iter_mut() {
258 if ρ.fail_count <= tconfig.max_fail {
259 ρ.y = ρ.x - v.differential(&ρ.x) * (ρ.α_γ.signum() * θ);
260 }
261 }
262 }
263 AdaptiveMax { l: ℓ_F, ref mut max_transport, g: ref calculate_θτ } => {
264 *max_transport = max_transport.max(self.norm(Radon));
265 let θτ = calculate_θτ(ℓ_F, *max_transport);
266 for ρ in self.iter_mut() {
267 if ρ.fail_count <= tconfig.max_fail {
268 ρ.y = ρ.x - v.differential(&ρ.x) * (ρ.α_γ.signum() * θτ);
269 }
270 }
271 }
272 FullyAdaptive {
273 l: ref mut adaptive_ℓ_F,
274 ref mut max_transport,
275 g: ref calculate_θτ,
276 } => {
277 *max_transport = max_transport.max(self.norm(Radon));
278 let mut θτ = calculate_θτ(*adaptive_ℓ_F, *max_transport);
279 // Do two runs through the spikes to update θ, breaking if first run did not cause
280 // a change.
281 for _i in 0..=1 {
282 let mut changes = false;
283 for ρ in self.iter_mut() {
284 if ρ.fail_count < tconfig.max_fail {
285 let dv_x = v.differential(&ρ.x);
286 let g = &dv_x * (ρ.α_γ.signum() * θτ);
287 ρ.y = ρ.x - g;
288 let n = g.norm2();
289 if n >= F::EPSILON {
290 // Estimate Lipschitz factor of ∇v
291 let this_ℓ_F = (dv_x - v.differential(&ρ.y)).norm2() / n;
292 *adaptive_ℓ_F = adaptive_ℓ_F.max(this_ℓ_F);
293 θτ = calculate_θτ(*adaptive_ℓ_F, *max_transport);
294 changes = true
295 }
296 }
297 }
298 if !changes {
299 break;
300 }
301 }
302 }
303 }
304 }
305
306 /// A posteriori transport adaptation.
307 #[replace_float_literals(F::cast_from(literal))]
308 pub(crate) fn aposteriori_transport<D>(
309 &mut self,
310 μ: &RNDM<N, F>,
311 μ̆: &RNDM<N, F>,
312 _v: &mut D,
313 extra: Option<F>,
314 ε: F,
315 tconfig: &TransportConfig<F>,
316 attempts: &mut usize,
317 ) -> bool
318 where
319 D: DifferentiableRealMapping<N, F>,
320 {
321 *attempts += 1;
322
323 // 1. If π_♯^1γ^{k+1} = γ1 has non-zero mass at some point y, but μ = μ^{k+1} does not,
324 // then the ansatz ∇w̃_x(y) = w^{k+1}(y) may not be satisfied. So set the mass of γ1
325 // at that point to zero, and retry.
326 let mut all_ok = true;
327 for (δ, ρ) in izip!(μ.iter_spikes(), self.iter_mut()) {
328 if δ.α == 0.0 && ρ.α_γ != 0.0 {
329 all_ok = false;
330 ρ.α_γ = 0.0;
331 }
332 }
333
334 // 2. Through bounding ∫ B_ω(y, z) dλ(x, y, z).
335 // through the estimate ≤ C ‖Δ‖‖γ^{k+1}‖ for Δ := μ^{k+1}-μ̆^k
336 // which holds for some some C if the convolution kernel in 𝒟 has Lipschitz gradient.
337 let nγ = self.norm(Radon);
338 let nΔ = μ.dist_matching(&μ̆) + extra.unwrap_or(0.0);
339 let t = ε * tconfig.tolerance_mult_con;
340 if nγ * nΔ > t && *attempts >= tconfig.max_attempts {
341 all_ok = false;
342 } else if nγ * nΔ > t {
343 // Since t/(nγ*nΔ)<1, and the constant tconfig.adaptation < 1,
344 // this will guarantee that eventually ‖γ‖ decreases sufficiently that we
345 // will not enter here.
346 //*self *= tconfig.adaptation * t / (nγ * nΔ);
347
348 // We want a consistent behaviour that has the potential to set many weights to zero.
349 // Therefore, we find the smallest uniform reduction `chg_one`, subtracted
350 // from all weights, that achieves total `adapt` adaptation.
351 let adapt_to = tconfig.adaptation * t / nΔ;
352 let reduction_target = nγ - adapt_to;
353 assert!(reduction_target > 0.0);
354 if ALLOW_PARTIAL_TRANSPORT {
355 if MINIMAL_PARTIAL_TRANSPORT {
356 // This reduces weights of transport, starting from … until `adapt` is
357 // exhausted. It will, therefore, only ever cause one extrap point insertion
358 // at the sources, unlike “full” partial transport.
359 //let refs = self.vec.iter_mut().collect::<Vec<_>>();
360 //refs.sort_by(|ρ1, ρ2| ρ1.α_γ.abs().partial_cmp(&ρ2.α_γ.abs()).unwrap());
361 // let mut it = refs.into_iter();
362 //
363 // Maybe sort by differential norm
364 // let mut refs = self
365 // .vec
366 // .iter_mut()
367 // .map(|ρ| {
368 // let val = v.differential(&ρ.x).norm2_squared();
369 // (ρ, val)
370 // })
371 // .collect::<Vec<_>>();
372 // refs.sort_by(|(_, v1), (_, v2)| v2.partial_cmp(&v1).unwrap());
373 // let mut it = refs.into_iter().map(|(ρ, _)| ρ);
374 let mut it = self.vec.iter_mut().rev();
375 let _unused = it.try_fold(reduction_target, |left, ρ| {
376 let w = ρ.α_γ.abs();
377 if left <= w {
378 ρ.α_γ = ρ.α_γ.signum() * (w - left);
379 ControlFlow::Break(())
380 } else {
381 ρ.α_γ = 0.0;
382 ControlFlow::Continue(left - w)
383 }
384 });
385 } else {
386 // This version equally reduces all weights. It causes partial transport, which
387 // has the problem that that we need to then adapt weights in both start and
388 // end points, in insert_and_reweigh, somtimes causing the number of spikes μ
389 // to explode.
390 let mut abs_weights = self
391 .vec
392 .iter()
393 .map(|ρ| ρ.α_γ.abs())
394 .filter(|t| *t > F::EPSILON)
395 .collect::<Vec<F>>();
396 abs_weights.sort_by(|a, b| a.partial_cmp(b).unwrap());
397 let n = abs_weights.len();
398 // Cannot have partial transport; can cause spike count explosion
399 let chg = abs_weights.into_iter().zip((1..=n).rev()).try_fold(
400 0.0,
401 |smaller_total, (w, m)| {
402 let mf = F::cast_from(m);
403 let reduction = w * mf + smaller_total;
404 if reduction >= reduction_target {
405 ControlFlow::Break((reduction_target - smaller_total) / mf)
406 } else {
407 ControlFlow::Continue(smaller_total + w)
408 }
409 },
410 );
411 match chg {
412 ControlFlow::Continue(_) => self.vec.iter_mut().for_each(|δ| δ.α_γ = 0.0),
413 ControlFlow::Break(chg_one) => self.vec.iter_mut().for_each(|ρ| {
414 let t = ρ.α_γ.abs();
415 if t > 0.0 {
416 if ALLOW_PARTIAL_TRANSPORT {
417 let new = (t - chg_one).max(0.0);
418 ρ.α_γ = ρ.α_γ.signum() * new;
419 }
420 }
421 }),
422 }
423 }
424 } else {
425 // This version zeroes smallest weights, avoiding partial transport.
426 let mut abs_weights_idx = self
427 .vec
428 .iter()
429 .map(|ρ| ρ.α_γ.abs())
430 .zip(0..)
431 .filter(|(w, _)| *w >= 0.0)
432 .collect::<Vec<(F, usize)>>();
433 abs_weights_idx.sort_by(|(a, _), (b, _)| a.partial_cmp(b).unwrap());
434
435 let mut left = reduction_target;
436
437 for (w, i) in abs_weights_idx {
438 left -= w;
439 let ρ = &mut self.vec[i];
440 ρ.α_γ = 0.0;
441 if left < 0.0 {
442 break;
443 }
444 }
445 }
446
447 all_ok = false
448 }
449
450 if !all_ok && *attempts >= tconfig.max_attempts {
451 for ρ in self.iter_mut() {
452 ρ.α_γ = 0.0;
453 }
454 }
455
456 for ρ in self.iter_mut() {
457 if ρ.α_γ == 0.0 {
458 ρ.fail_count += 1;
459 } else if all_ok {
460 ρ.fail_count = 0;
461 }
462 }
463
464 all_ok
465 } 674 }
466 675
467 /// Returns $‖μ\^k - π\_♯\^0γ\^{k+1}‖$ 676 /// Returns $‖μ\^k - π\_♯\^0γ\^{k+1}‖$
468 pub(crate) fn μ0_minus_γ0_radon(&self) -> F { 677 pub(crate) fn μ0_minus_γ0_radon(&self) -> F {
469 self.vec.iter().map(|ρ| (ρ.α_μ_orig - ρ.α_γ).abs()).sum() 678 self.vec.iter().map(|ρ| (ρ.α_μ_orig - ρ.α_γ).abs()).sum()
470 } 679 }
471 680
472 /// Returns $∫ c_2 d|γ|$ 681 /// Returns $∫ c_2 d|γ|$
473 #[replace_float_literals(F::cast_from(literal))]
474 pub(crate) fn c2integral(&self) -> F { 682 pub(crate) fn c2integral(&self) -> F {
475 self.vec 683 self.vec
476 .iter() 684 .iter()
477 .map(|ρ| ρ.y.dist2_squared(&ρ.x) / 2.0 * ρ.α_γ.abs()) 685 .map(|ρ| ρ.y.dist2_squared(&ρ.x) / 2.0 * ρ.α_γ.abs())
478 .sum() 686 .sum()
479 } 687 }
480 688
481 #[replace_float_literals(F::cast_from(literal))]
482 pub(crate) fn get_transport_stats(&self, stats: &mut IterInfo<F>, μ: &RNDM<N, F>) { 689 pub(crate) fn get_transport_stats(&self, stats: &mut IterInfo<F>, μ: &RNDM<N, F>) {
483 // TODO: This doesn't take into account μ[i].α becoming zero in the latest tranport 690 // TODO: This doesn't take into account μ[i].α becoming zero in the latest tranport
484 // attempt, for i < self.len(), when a corresponding source term also exists with index 691 // attempt, for i < self.len(), when a corresponding source term also exists with index
485 // j ≥ self.len(). For now, we let that be reflected in the prune count. 692 // j ≥ self.len(). For now, we let that be reflected in the prune count.
486 stats.inserted += μ.len() - self.len(); 693 stats.inserted += μ.len() - self.len();
520 /// latter needs to be pruned when μ is. 727 /// latter needs to be pruned when μ is.
521 pub(crate) fn prune_compat(&mut self, μ: &mut RNDM<N, F>, stats: &mut IterInfo<F>) { 728 pub(crate) fn prune_compat(&mut self, μ: &mut RNDM<N, F>, stats: &mut IterInfo<F>) {
522 assert!(self.vec.len() <= μ.len()); 729 assert!(self.vec.len() <= μ.len());
523 let old_len = μ.len(); 730 let old_len = μ.len();
524 for (ρ, δ) in self.vec.iter_mut().zip(μ.iter_spikes()) { 731 for (ρ, δ) in self.vec.iter_mut().zip(μ.iter_spikes()) {
525 ρ.prune = !(δ.α.abs() > F::EPSILON); 732 ρ.retain = δ.α.abs() > F::EPSILON;
526 } 733 }
527 μ.prune_by(|δ| δ.α.abs() > F::EPSILON); 734 μ.prune_by(|δ| δ.α.abs() > F::EPSILON);
528 stats.pruned += old_len - μ.len(); 735 stats.pruned += old_len - μ.len();
529 self.vec.retain(|ρ| !ρ.prune); 736 self.vec.retain(|ρ| ρ.retain);
530 assert!(self.vec.len() <= μ.len()); 737 assert!(self.vec.len() <= μ.len());
531 } 738 }
532 } 739
533 740 /// Helper for initial transport. Called from [`TransportProxPenalty::initial_transport`].
534 impl<const N: usize, F: Float> Norm<Radon, F> for Transport<N, F> { 741 pub(super) fn do_init_transport<D>(
742 &mut self,
743 v: &D,
744 μ: &RNDM<N, F>,
745 τθ_or_adaptive: &mut TransportStepLength<F>,
746 tconfig: &TransportConfig<F>,
747 ) where
748 D: DifferentiableRealMapping<N, F>,
749 {
750 use TransportStepLength::*;
751
752 // Initialise transport structure weights
753 for (δ, ρ) in izip!(μ.iter_spikes(), self.iter_mut()) {
754 ρ.α_μ_orig = δ.α;
755 ρ.x = δ.x;
756 ρ.y = δ.x; // Later updated if no fails.
757 ρ.α_γ = if ρ.fail_count > tconfig.max_fail {
758 0.0
759 } else {
760 // If old transport has opposing sign, the new transport will be none.
761 if (ρ.α_γ > 0.0 && δ.α < 0.0) || (ρ.α_γ < 0.0 && δ.α > 0.0) {
762 0.0
763 } else {
764 δ.α
765 }
766 };
767 }
768
769 let γ_prev_len = self.len();
770 assert!(μ.len() >= γ_prev_len);
771 self.extend(μ[γ_prev_len..].iter().map(|δ| SingleTransport {
772 x: δ.x,
773 y: δ.x, // Just something, will be filled properly in the next phase
774 α_μ_orig: δ.α,
775 α_γ: δ.α,
776 retain: true,
777 fail_count: 0,
778 excess: 0.0,
779 }));
780
781 // Calculate transport rays.
782 let simple_τθ = match *τθ_or_adaptive {
783 Fixed { τθ, .. } => Some(τθ),
784 Simple { ℓ_gradv, τθ0, ℓ_base } => Some(τθ0 / (ℓ_gradv + ℓ_base)),
785 AdaptiveMax {
786 ℓ_gradv,
787 ref mut adaptive_max_transport,
788 τθ0,
789 ℓ_base,
790 ℓ_base_max_transport,
791 } => {
792 *adaptive_max_transport = adaptive_max_transport.max(self.norm(Radon));
793 Some(τθ0 / (ℓ_gradv + ℓ_base + ℓ_base_max_transport * *adaptive_max_transport))
794 }
795 FullyAdaptive {
796 ref mut adaptive_ℓ_gradv,
797 ref mut adaptive_max_transport,
798 τθ0,
799 ℓ_base,
800 ℓ_base_max_transport,
801 } => {
802 *adaptive_max_transport = adaptive_max_transport.max(self.norm(Radon));
803 let mut τθ = τθ0
804 / (*adaptive_ℓ_gradv + ℓ_base + ℓ_base_max_transport * *adaptive_max_transport);
805 // Do two runs through the spikes to update θ, breaking if first run did not cause
806 // a change.
807 for _i in 0..=1 {
808 let mut changes = false;
809 for ρ in self.iter_mut() {
810 if ρ.fail_count < tconfig.max_fail {
811 let dv_x = v.differential(&ρ.x);
812 let g = &dv_x * (ρ.α_γ.signum() * τθ);
813 ρ.y = ρ.x - g;
814 let n = g.norm2();
815 if n >= F::EPSILON {
816 // Estimate Lipschitz factor of ∇v
817 let this_ℓ_gradv = (dv_x - v.differential(&ρ.y)).norm2() / n;
818 *adaptive_ℓ_gradv = adaptive_ℓ_gradv.max(this_ℓ_gradv);
819 τθ = τθ0
820 / (*adaptive_ℓ_gradv
821 + ℓ_base
822 + ℓ_base_max_transport * *adaptive_max_transport);
823 changes = true
824 }
825 }
826 }
827 if !changes {
828 break;
829 }
830 }
831 None
832 }
833 };
834
835 if let Some(τθ) = simple_τθ {
836 for ρ in self.iter_mut() {
837 if ρ.fail_count <= tconfig.max_fail {
838 ρ.y = ρ.x - v.differential(&ρ.x) * (ρ.α_γ.signum() * τθ);
839 }
840 }
841 }
842 }
843
844 /// Helper for a posteriori transport error control.
845 /// Called from [`TransportProxPenalty::aposteriori_transport`].
846 fn do_new_aposteriori_transport<Reg, M>(
847 &mut self,
848 μ: &RNDM<N, F>,
849 τv̆: &mut M,
850 v: &mut M,
851 ε: F,
852 τ: F,
853 τθ_or_adaptive: &TransportStepLength<F>,
854 reg: &Reg,
855 tconfig: &TransportConfig<F>,
856 ω: impl Fn(&Loc<N, F>, std::cmp::Ordering, bool) -> F,
857 ) -> bool
858 where
859 Reg: SlidingRegTerm<Loc<N, F>, F>,
860 F: ToNalgebraRealField,
861 M: MinMaxMapping<Loc<N, F>, F> + DifferentiableRealMapping<N, F>,
862 {
863 let τℓ_gradv2 = τ * τθ_or_adaptive.get_ℓ_gradv() / 2.0;
864 let Bounds(α_lower, α_upper) = reg.subdiff_range();
865
866 let (all_ok, m) = izip!(self.vec.iter_mut(), μ.iter_spikes()).fold(
867 (true, 0.0),
868 |(all_ok_so_far, total_excess), (ρ, δ)| {
869 use std::cmp::Ordering::*;
870 // NOTE: The tolerances ε are commented out, because they will in any case
871 // be consumed by `t` below by suitably large choice of `tolerance_mult`.
872 // Hence, we simply implicitly adapt the `tolerance_mult` be commenting out
873 // the `ε` here. That way, `d` is in its entirely multiplied by τ, as is `t`,
874 // making all factors independent of τ.
875 let maybe_excess = match ρ.α_γ.total_cmp(&0.0) {
876 Greater => (δ.α >= 0.0).then(|| {
877 let gvx = v.differential(&ρ.x);
878 let d = if tconfig.alt_remainder_control {
879 let τv̆y = τv̆.apply(&ρ.y);
880 (/*ε +*/τ * α_upper + ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), false))
881 + (τv̆y + τ * gvx.dot(&ρ.x - &ρ.y))
882 - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
883 } else {
884 (/*2.0 * ε*/-ω(&ρ.y, δ.α.total_cmp(&ρ.α_γ), true)
885 + ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), false))
886 + τ * gvx.dot(&ρ.x - &ρ.y)
887 - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
888 };
889 d * ρ.α_γ
890 }),
891 Less => (δ.α <= 0.0).then(|| {
892 let gvx = v.differential(&ρ.x);
893 let d = if tconfig.alt_remainder_control {
894 let τv̆y = τv̆.apply(&ρ.y);
895 (/*ε*/-τ * α_lower - ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), true))
896 - (τv̆y + τ * gvx.dot(&ρ.x - &ρ.y))
897 - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
898 } else {
899 (/*2.0 * ε +*/ω(&ρ.y, δ.α.total_cmp(&ρ.α_γ), false)
900 - ω(&ρ.x, ρ.α_γ.total_cmp(&ρ.α_μ_orig), true))
901 - τ * gvx.dot(&ρ.x - &ρ.y)
902 - τℓ_gradv2 * ρ.y.dist2_squared(&ρ.x)
903 };
904 d * (-ρ.α_γ)
905 }),
906 Equal => Some(0.0),
907 };
908 match maybe_excess {
909 None => {
910 ρ.α_γ = 0.0;
911 ρ.excess = 0.0;
912 (false, total_excess)
913 }
914 Some(e) => {
915 ρ.excess = e;
916 (all_ok_so_far, total_excess + e)
917 }
918 }
919 },
920 );
921
922 let t = τ * ε * tconfig.tolerance_mult;
923
924 if m > t {
925 let mut it = self.vec.iter_mut().rev().filter(|ρ| ρ.excess > 0.0);
926 let reduction_target = m - tconfig.adaptation * t;
927 let _unused = it.try_fold(reduction_target, |left, ρ| {
928 let d = ρ.excess;
929 if d >= left {
930 if tconfig.allow_partial_transport {
931 ρ.α_γ *= (d - left) / d;
932 } else {
933 ρ.α_γ = 0.0;
934 }
935 ControlFlow::Break(())
936 } else {
937 ρ.α_γ = 0.0;
938 ControlFlow::Continue(left - d)
939 }
940 });
941 false
942 } else {
943 all_ok
944 }
945 }
946 }
947
948 impl<const N: usize, F: Float> Norm<Radon, F> for Transport<Loc<N, F>, F> {
535 fn norm(&self, _: Radon) -> F { 949 fn norm(&self, _: Radon) -> F {
536 self.iter().map(|ρ| ρ.α_γ.abs()).sum() 950 self.iter().map(|ρ| ρ.α_γ.abs()).sum()
537 } 951 }
538 } 952 }
539 953
540 impl<const N: usize, F: Float> MulAssign<F> for Transport<N, F> { 954 impl<const N: usize, F: Float> MulAssign<F> for Transport<Loc<N, F>, F> {
541 fn mul_assign(&mut self, factor: F) { 955 fn mul_assign(&mut self, factor: F) {
542 for ρ in self.iter_mut() { 956 for ρ in self.iter_mut() {
543 ρ.α_γ *= factor; 957 ρ.α_γ *= factor;
544 } 958 }
545 } 959 }
562 ) -> DynResult<RNDM<N, F>> 976 ) -> DynResult<RNDM<N, F>>
563 where 977 where
564 F: Float + ToNalgebraRealField, 978 F: Float + ToNalgebraRealField,
565 I: AlgIteratorFactory<IterInfo<F>>, 979 I: AlgIteratorFactory<IterInfo<F>>,
566 Dat: DifferentiableMapping<RNDM<N, F>, Codomain = F> + BoundedCurvature<F>, 980 Dat: DifferentiableMapping<RNDM<N, F>, Codomain = F> + BoundedCurvature<F>,
567 Dat::DerivativeDomain: DifferentiableRealMapping<N, F> + ClosedMul<F>, 981 Dat::DerivativeDomain:
982 DifferentiableRealMapping<N, F> + ClosedMul<F> + MinMaxMapping<Loc<N, F>, F>,
568 //for<'a> Dat::Differential<'a>: Lipschitz<&'a P, FloatType = F>, 983 //for<'a> Dat::Differential<'a>: Lipschitz<&'a P, FloatType = F>,
569 RNDM<N, F>: SpikeMerging<F>, 984 RNDM<N, F>: SpikeMerging<F>,
570 Reg: SlidingRegTerm<Loc<N, F>, F>, 985 Reg: SlidingRegTerm<Loc<N, F>, F>,
571 P: ProxPenalty<Loc<N, F>, Dat::DerivativeDomain, Reg, F> + StepLengthBound<F, Dat>, 986 P: TransportProxPenalty<Loc<N, F>, Dat::DerivativeDomain, Reg, F> + StepLengthBound<F, Dat>,
572 Plot: Plotter<P::ReturnMapping, Dat::DerivativeDomain, RNDM<N, F>>, 987 Plot: Plotter<P::ReturnMapping, Dat::DerivativeDomain, RNDM<N, F>>,
573 { 988 {
574 // Check parameters 989 // Check parameters
575 ensure!(config.τ0 > 0.0, "Invalid step length parameter"); 990 ensure!(config.τ0 > 0.0, "Invalid step length parameter");
576 config.transport.check()?; 991 config.transport.check()?;
582 // Set up parameters 997 // Set up parameters
583 // let opAnorm = opA.opnorm_bound(Radon, L2); 998 // let opAnorm = opA.opnorm_bound(Radon, L2);
584 //let max_transport = config.max_transport.scale 999 //let max_transport = config.max_transport.scale
585 // * reg.radon_norm_bound(b.norm2_squared() / 2.0); 1000 // * reg.radon_norm_bound(b.norm2_squared() / 2.0);
586 //let ℓ = opA.transport.lipschitz_factor(L2Squared) * max_transport; 1001 //let ℓ = opA.transport.lipschitz_factor(L2Squared) * max_transport;
587 let ℓ = 0.0;
588 let τ = config.τ0 / prox_penalty.step_length_bound(&f)?; 1002 let τ = config.τ0 / prox_penalty.step_length_bound(&f)?;
589 1003
590 let mut θ_or_adaptive = match f.curvature_bound_components(config.guess) { 1004 let mut τθ_or_adaptive = prox_penalty.get_transport_steplength(
591 (_, Err(_)) => TransportStepLength::Fixed(config.transport.θ0), 1005 f.curvature_bound_components(config.guess),
592 (maybe_ℓ_F, Ok(transport_lip)) => { 1006 &config.transport,
593 let calculate_θτ = move |ℓ_F, max_transport| { 1007 0.0,
594 let ℓ_r = transport_lip * max_transport; 1008 0.0,
595 config.transport.θ0 / (ℓ + ℓ_F + ℓ_r) 1009 );
596 };
597 match maybe_ℓ_F {
598 Ok(ℓ_F) => TransportStepLength::AdaptiveMax {
599 l: ℓ_F, // TODO: could estimate computing the real reesidual
600 max_transport: 0.0,
601 g: calculate_θτ,
602 },
603 Err(_) => TransportStepLength::FullyAdaptive {
604 l: 10.0 * F::EPSILON, // Start with something very small to estimate differentials
605 max_transport: 0.0,
606 g: calculate_θτ,
607 },
608 }
609 }
610 };
611 // We multiply tolerance by τ for FB since our subproblems depending on tolerances are scaled 1010 // We multiply tolerance by τ for FB since our subproblems depending on tolerances are scaled
612 // by τ compared to the conditional gradient approach. 1011 // by τ compared to the conditional gradient approach.
613 let tolerance = config.insertion.tolerance * τ * reg.tolerance_scaling(); 1012 let tolerance = config.insertion.tolerance * τ * reg.tolerance_scaling();
614 let mut ε = tolerance.initial(); 1013 let mut ε = tolerance.initial();
615 1014
623 }; 1022 };
624 let mut stats = IterInfo::new(); 1023 let mut stats = IterInfo::new();
625 1024
626 // Run the algorithm 1025 // Run the algorithm
627 for state in iterator.iter_init(|| full_stats(&μ, ε, stats.clone())) { 1026 for state in iterator.iter_init(|| full_stats(&μ, ε, stats.clone())) {
1027 let mut v = f.differential(&μ);
1028
628 // Calculate initial transport 1029 // Calculate initial transport
629 let v = f.differential(&μ); 1030 prox_penalty.initial_transport(
630 γ.initial_transport(&μ, τ, &mut θ_or_adaptive, v, &config.transport); 1031 &mut γ,
1032 &μ,
1033 ε,
1034 τ,
1035 &mut τθ_or_adaptive,
1036 &v,
1037 &config.transport,
1038 );
631 1039
632 let mut attempts = 0; 1040 let mut attempts = 0;
633 1041
634 // Solve finite-dimensional subproblem several times until the dual variable for the 1042 // Solve finite-dimensional subproblem several times until the dual variable for the
635 // regularisation term conforms to the assumptions made for the transport above. 1043 // regularisation term conforms to the assumptions made for the transport above.
656 &state, 1064 &state,
657 &mut stats, 1065 &mut stats,
658 )?; 1066 )?;
659 1067
660 // A posteriori transport adaptation. 1068 // A posteriori transport adaptation.
661 if γ.aposteriori_transport(&μ, &μ̆, &mut τv̆, None, ε, &config.transport, &mut attempts) 1069 if prox_penalty.aposteriori_transport(
662 { 1070 &mut γ,
1071 &μ,
1072 &μ̆,
1073 &mut τv̆,
1074 &mut v,
1075 None,
1076 ε,
1077 τ,
1078 &τθ_or_adaptive,
1079 reg,
1080 &config.transport,
1081 &config.insertion.refinement,
1082 &mut attempts,
1083 ) {
663 break 'adapt_transport (maybe_d, within_tolerances, τv̆, μ̆); 1084 break 'adapt_transport (maybe_d, within_tolerances, τv̆, μ̆);
664 } 1085 }
665 1086
666 stats.get_transport_mut().readjustment_iters += 1; 1087 stats.get_transport_mut().readjustment_iters += 1;
667 }; 1088 };
670 1091
671 // Merge spikes. 1092 // Merge spikes.
672 // This crucially expects the merge routine to be stable with respect to spike locations, 1093 // This crucially expects the merge routine to be stable with respect to spike locations,
673 // and not to performing any pruning. That is be to done below simultaneously for γ. 1094 // and not to performing any pruning. That is be to done below simultaneously for γ.
674 if config.insertion.merge_now(&state) { 1095 if config.insertion.merge_now(&state) {
675 stats.merged += prox_penalty.merge_spikes( 1096 let m = prox_penalty.merge_spikes(
676 &mut μ, 1097 &mut μ,
677 &mut τv̆, 1098 &mut τv̆,
678 &μ̆, 1099 &μ̆,
679 τ, 1100 τ,
680 ε, 1101 ε,
681 &config.insertion, 1102 &config.insertion,
682 &reg, 1103 &reg,
683 Some(|μ̃: &RNDM<N, F>| f.apply(μ̃)), 1104 Some(|μ̃: &RNDM<N, F>| f.apply(μ̃)),
684 ); 1105 );
1106 if m > 0 {
1107 stats.merged += m;
1108 v = f.differential(&μ);
1109 }
685 } 1110 }
686 1111
687 γ.prune_compat(&mut μ, &mut stats); 1112 γ.prune_compat(&mut μ, &mut stats);
688 1113
689 let iter = state.iteration(); 1114 let iter = state.iteration();

mercurial