| 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 { |
| 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 } |