1#![allow(dead_code)]
2
3use std::cmp::Ordering as CmpOrdering;
4use std::collections::{BTreeMap, hash_map::DefaultHasher};
5use std::hash::{Hash, Hasher};
6use std::str::FromStr;
7use std::sync::{Arc, Mutex};
8use std::time::{Duration, Instant};
9
10#[cfg(feature = "complex")]
11use crate::algebra::bridge::BridgeScratch;
12#[allow(unused_imports)]
13use crate::algebra::prelude::*;
14use crate::config::kinds::{
15 AmgCoarseSolveKind, AmgCoarsenKind, AmgCycleKind, AmgInterpKind, AmgRelaxKind, AmgStrengthKind,
16};
17use crate::config::options::PcOptions;
18use crate::error::KError;
19#[cfg(not(feature = "complex"))]
20use crate::matrix::DistCsrOp;
21#[cfg(not(feature = "complex"))]
22use crate::matrix::convert::csr_from_linop;
23#[cfg(not(feature = "complex"))]
24use crate::matrix::dist::halo::HaloPlan;
25#[cfg(not(feature = "complex"))]
26use crate::matrix::dist::hierarchy::DistHierarchyMeta;
27#[cfg(all(feature = "complex", feature = "backend-faer"))]
28use crate::matrix::op::{CsrOp, DenseOp, GenericCsrOp};
29use crate::matrix::op::{LinOp, StructureId, ValuesId};
30#[cfg(feature = "complex")]
31use crate::matrix::parcsr::HaloPlan;
32use crate::matrix::{
33 dense_api::DenseMatRef,
34 sparse::CsrMatrix,
35 spmv::{csr_spmm_dense, spmv_scaled_f32_on_pattern},
36};
37#[cfg(feature = "simd")]
38use crate::matrix::{spmv::SpmvTuning, utils};
39#[cfg(feature = "complex")]
40use crate::ops::kpc::KPreconditioner;
41use crate::parallel::{Comm, UniverseComm};
42use crate::preconditioner::asm::Asm;
43#[cfg(not(feature = "complex"))]
44use crate::preconditioner::asm::{AsmCombine, AsmConfig, AsmLocalSolver};
45#[cfg(feature = "complex")]
46use crate::preconditioner::bridge::{
47 apply_pc_mut_s as bridge_apply_pc_mut_s, apply_pc_s as bridge_apply_pc_s,
48};
49use crate::preconditioner::chebyshev::{self, ChebBounds};
50use crate::preconditioner::deflation::{AmgCoarseSpace, DeflationOptions, ZSource};
51use crate::preconditioner::dist::{
52 DistCoarseRepartition, DistCoarseSolverRoute, DistCoarseStrategy,
53};
54use crate::preconditioner::ilu_csr::IluCsr;
55#[cfg(not(feature = "complex"))]
56use crate::preconditioner::ilu_csr::{IluCsrConfig, IluKind, PivotStrategy, ReorderingOptions};
57use crate::preconditioner::{PcCaps, PcSide, Preconditioner};
58#[cfg(all(not(feature = "complex"), feature = "superlu_dist"))]
59use crate::solver::superlu_dist;
60use crate::utils::conditioning::{ConditioningOptions, apply_csr_transforms};
61use faer::Mat;
62
63#[cfg(feature = "rayon")]
64use rayon::prelude::*;
65
66mod coarse_solver;
68pub use coarse_solver::CoarseSolve;
69pub mod coarsen;
70mod non_galerkin;
71pub(crate) mod prolong;
72pub use prolong::AdaptiveWeight;
73mod rap_ops;
74mod row_filter;
75mod scalar_core;
76pub mod strength;
77pub(crate) mod strength_nodal;
78pub(crate) mod util;
79
80use coarse_solver::{CoarseDenseLu, CoarseIlu, CoarseSolver};
81use coarsen::{
82 AggAlgo, AggOpts, build_aggregates, build_aggregates_nodal, lift_node_aggregates_to_dofs,
83};
84use non_galerkin::{NgRowFilter, non_galerkin_filter_coarse};
85use prolong::{
86 CFInfo, ClassicalParams, ClassicalVariant, Pcsr, TentativeNodal, TentativeP,
87 adaptive_fit_values_only, classical_pattern, classical_values_only, restrict_samples_to_coarse,
88 sample_low_modes, smooth_sa_values_only, smooth_sa_values_only_mf, smooth_sa_values_only_multi,
89 smooth_tentative_sa_mf, smooth_tentative_sa_multi,
90};
91use rap_ops::{CsrPattern, adjoint_csr_with_pos, rap_numeric, rap_symbolic};
92use row_filter::{
93 RowFilter, apply_filter_to_csr_values_in_place, compensate_nodal_diag, compensate_scalar_rows,
94 restrict_trials,
95};
96#[cfg(feature = "complex")]
97use scalar_core::AmgCore;
98use strength::Strength;
99use strength_nodal::strength_nodal_from_csr;
100use util::DofLayout;
101
102#[derive(Clone, Copy, Debug, PartialEq, Eq)]
106pub enum CoarsenType {
107 RS,
108 HMIS,
109 PMIS,
110 Falgout,
111}
112
113#[derive(Clone, Copy, Debug, PartialEq, Eq)]
115pub enum InterpType {
116 Classical,
117 Direct,
118 Multipass,
119 Extended,
120 Standard,
121 HE,
122}
123
124#[derive(Clone, Copy, Debug, PartialEq, Eq)]
126pub enum RankFallback {
127 RetryLooserInterp,
128 SwitchInterpKind,
129 Reaggregate,
130 Abort,
131}
132
133#[derive(Clone, Copy, Debug, PartialEq, Eq)]
135pub enum RelaxType {
136 Jacobi,
137 GaussSeidel,
138 GaussSeidelBackward,
139 SymmetricGaussSeidel,
140 HybridGaussSeidel,
141 L1Jacobi,
142 Chebyshev,
143 ChebyshevSafe,
144 SafeguardedGaussSeidel,
145 Ilu0,
146 Ras,
147 Fsai,
148}
149
150#[derive(Clone, Copy, Debug, PartialEq, Eq)]
152pub enum RelaxPhase {
153 Fine = 0,
154 Down = 1,
155 Up = 2,
156 Coarsest = 3,
157}
158
159enum RelaxWhere {
160 Pre,
161 Post,
162}
163
164impl RelaxPhase {
165 #[inline]
166 pub fn ix(self) -> usize {
167 self as usize
168 }
169 pub const ALL: [RelaxPhase; 4] = [
170 RelaxPhase::Fine,
171 RelaxPhase::Down,
172 RelaxPhase::Up,
173 RelaxPhase::Coarsest,
174 ];
175}
176
177#[cfg(test)]
178use std::cell::Cell;
179#[cfg(test)]
180thread_local! {
181 static RELAX_CALL_COUNTS: Cell<[usize; 4]> = Cell::new([0; 4]);
182}
183#[cfg(test)]
184thread_local! {
185 static BUILD_SYMBOLIC_COUNT: std::cell::Cell<usize> = std::cell::Cell::new(0);
186}
187
188#[cfg(test)]
189pub fn reset_relax_counts() {
190 RELAX_CALL_COUNTS.with(|counts| counts.set([0; 4]));
191}
192#[cfg(test)]
193pub fn get_relax_counts() -> [usize; 4] {
194 RELAX_CALL_COUNTS.with(|counts| counts.get())
195}
196
197#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
200pub enum CycleType {
201 #[default]
202 V,
203 W {
204 gamma: usize,
205 },
206}
207
208#[derive(Clone, Copy, Debug, PartialEq, Eq)]
209pub enum KrylovAlgo {
210 FCG,
211}
212
213#[derive(Clone, Debug, PartialEq, Eq)]
214pub struct KCycle {
215 pub levels: Vec<usize>,
217 pub iters: usize,
219 pub algo: KrylovAlgo,
220 pub place_post: bool,
222 pub place_pre: bool,
224}
225
226#[derive(Clone, Copy, Debug, PartialEq, Eq)]
227pub struct MixedPrecision {
228 pub smooth: bool,
229 pub residual: bool,
230}
231
232impl MixedPrecision {
233 #[inline]
234 fn smoothers_enabled(self) -> bool {
235 self.smooth
236 }
237
238 #[inline]
239 fn residual_enabled(self) -> bool {
240 self.residual
241 }
242}
243
244#[derive(Clone, Copy, Debug, PartialEq, Eq)]
245pub enum MixedStorage {
246 Cached,
247 Transient,
248}
249
250#[derive(Clone, Copy, Debug, PartialEq, Eq)]
251pub enum NgSymmetry {
252 None,
253 Symmetric,
254}
255
256#[derive(Clone, Copy, Debug, PartialEq, Eq)]
257pub enum NodalMode {
258 Off,
259 Nodal,
260}
261
262#[derive(Clone, Copy, Debug, PartialEq, Eq)]
263pub enum RowScaleMode {
264 ToNearNullspace,
267 SumToOne,
269 L2Unit,
271 DUnit,
273}
274
275#[derive(Clone, Copy, Debug, PartialEq)]
276pub enum PostInterpType {
277 None,
278 RowScaling(RowScaleMode),
280 LocalQR,
282 EnergyPolish {
284 sweeps: usize,
285 omega: f64,
286 },
287}
288
289#[derive(Clone, Debug)]
290pub struct NearNullspace {
291 pub basis: Vec<Vec<f64>>,
292}
293
294#[derive(Clone, Debug)]
295pub struct NonGalerkin {
296 pub enabled: bool,
297 pub start_level: usize,
298 pub drop_abs: f64,
299 pub drop_rel: f64,
300 pub cap_row: usize,
301 pub symmetry: NgSymmetry,
302 pub lump_diagonal: bool,
303 pub oc_target: Option<f64>,
304 pub oc_max_iter: usize,
305}
306
307#[derive(Clone, Debug)]
308pub struct AMGConfig {
309 pub max_levels: usize, pub strong_threshold: f64, pub coarse_threshold: usize, pub max_coarse_size: usize, pub min_coarse_size: usize, pub truncation_factor: f64, pub max_elements_per_row: usize, pub interpolation_truncation: f64,
317 pub rap_truncation_abs: f64,
318 pub rap_max_elements_per_row: usize,
319 pub keep_transpose: bool,
320 pub keep_pivot_in_rap: bool,
321 pub require_spd: bool,
322 pub spd_diag_floor: f64,
323 pub forbid_non_galerkin_in_spd: bool,
324 pub grid_relax_type: [RelaxType; 4], pub num_grid_sweeps: [usize; 4], pub pre_sweeps: usize, pub post_sweeps: usize, pub coarsen_type: CoarsenType, pub interp_type: InterpType, pub relax_type: RelaxType, pub logging_level: usize,
333 pub print_level: usize,
334 pub tolerance: f64, pub max_iterations: usize,
336 pub min_iterations: usize,
337 pub ieee_checks: bool,
338 pub optimize_workspace: bool,
339 pub jacobi_omega: f64,
340 pub adaptive_interp: bool,
341 pub adaptive_samples: usize,
342 pub adaptive_smooth_steps: usize,
343 pub adaptive_smooth_omega: f64,
344 pub adaptive_lambda: f64,
345 pub adaptive_enforce_sum1: bool,
346 pub adaptive_weight_mode: AdaptiveWeight,
347 pub chebyshev_recompute_esteig: bool,
348 pub chebyshev_safety: f64,
349 pub chebyshev_lower_ratio: f64,
350 pub chebyshev_power_steps: usize,
351 pub chebyshev_degree: usize,
352 pub use_level_scheduling: bool,
353 pub drop_tol: f64, pub stats_eps: f64, pub verify_p_rank: bool,
357 pub rank_sketch_cols: usize,
358 pub rank_cond_threshold: f64,
359 pub rank_min_col_norm: f64,
360 pub verify_galerkin: bool,
362 pub galerkin_samples: usize,
363 pub galerkin_rel_tol: f64,
364 pub on_rank_failure: RankFallback,
366 pub fsai_dist: usize,
368 pub fsai_use_strength: bool,
369 pub fsai_max_per_row: usize,
370 pub fsai_lambda: f64,
371 pub fsai_adaptive_passes: usize,
372 pub fsai_drop_tol: f64,
373 pub fsai_damping: f64,
374 pub normalize_strength: bool,
376 pub coarse_solve: CoarseSolve,
377 pub ilu_drop_tol: f64,
378 pub ilu_fill_per_row: usize,
379 pub max_operator_complexity: Option<f64>,
380 pub agg_num_levels: usize,
381 pub aggressive_mis_k: usize,
382 pub max_strong_per_row: Option<usize>,
383 pub cycle_type: CycleType,
384 pub kcycle: Option<KCycle>,
385 pub fmg_nu_pre: usize,
386 pub fmg_nu_post: usize,
387 pub fmg_gamma: usize,
388 pub fmg_levels_use: Option<usize>,
389 pub non_galerkin: NonGalerkin,
390 pub nodal: NodalMode,
391 pub block_size: usize,
392 pub num_functions: usize,
393 pub near_nullspace: Option<NearNullspace>,
394 pub filter_trial_vectors: Option<Vec<Vec<f64>>>,
395 pub filter_omega: f64,
396 pub filter_after_non_galerkin: bool,
397 pub post_interp: PostInterpType,
398 pub flexible_level: Option<usize>,
399 pub flexible_iters: usize,
400 pub flexible_rtol: f64,
401 pub flexible_pc_sweeps: usize,
402 pub mixed_precision: Option<MixedPrecision>,
403 pub mixed_storage: MixedStorage,
404 pub conditioning: ConditioningOptions,
405 pub dist_coarse_strategy: DistCoarseStrategy,
406 pub dist_coarse_repartition: DistCoarseRepartition,
407 pub dist_coarse_solver_route: DistCoarseSolverRoute,
408 pub dist_apply_instrumentation: bool,
409 pub dist_coarse_ghost_scale: f64,
410 pub require_native_complex_hierarchy: bool,
413 pub level_relax_overrides: BTreeMap<usize, RelaxType>,
414 pub level_sweep_overrides: BTreeMap<usize, (usize, usize)>,
415 pub level_coarse_overrides: BTreeMap<usize, CoarseSolve>,
416}
417
418impl Default for AMGConfig {
419 fn default() -> Self {
420 let mut cfg = Self {
421 max_levels: 25,
422 strong_threshold: 0.25,
423 coarse_threshold: 9,
424 max_coarse_size: 9,
425 min_coarse_size: 1,
426 truncation_factor: 0.0,
427 max_elements_per_row: 0,
428 interpolation_truncation: 0.0,
429 rap_truncation_abs: 0.0,
430 rap_max_elements_per_row: 0,
431 keep_transpose: true,
432 keep_pivot_in_rap: true,
433 require_spd: true,
434 spd_diag_floor: 0.0,
435 forbid_non_galerkin_in_spd: true,
436 grid_relax_type: [RelaxType::GaussSeidel; 4],
437 num_grid_sweeps: [1; 4],
438 pre_sweeps: 1,
439 post_sweeps: 1,
440 coarsen_type: CoarsenType::HMIS,
441 interp_type: InterpType::Extended,
442 relax_type: RelaxType::SymmetricGaussSeidel,
443 logging_level: 0,
444 print_level: 0,
445 tolerance: 1e-6,
446 max_iterations: 20,
447 min_iterations: 0,
448 ieee_checks: true,
449 optimize_workspace: true,
450 jacobi_omega: 2.0 / 3.0,
451 adaptive_interp: false,
452 adaptive_samples: 4,
453 adaptive_smooth_steps: 3,
454 adaptive_smooth_omega: 2.0 / 3.0,
455 adaptive_lambda: 1e-10,
456 adaptive_enforce_sum1: true,
457 adaptive_weight_mode: AdaptiveWeight::Diag,
458 chebyshev_recompute_esteig: true,
459 chebyshev_safety: 1.10,
460 chebyshev_lower_ratio: 0.10,
461 chebyshev_power_steps: 4,
462 chebyshev_degree: 2,
463 use_level_scheduling: false,
464 drop_tol: 1e-12,
465 stats_eps: 1e-12,
466 verify_p_rank: true,
467 rank_sketch_cols: 8,
468 rank_cond_threshold: 1e8,
469 rank_min_col_norm: 1e-12,
470 verify_galerkin: true,
471 galerkin_samples: 2,
472 galerkin_rel_tol: 1e-10,
473 on_rank_failure: RankFallback::RetryLooserInterp,
474 fsai_dist: 1,
475 fsai_use_strength: true,
476 fsai_max_per_row: 12,
477 fsai_lambda: 1e-12,
478 fsai_adaptive_passes: 0,
479 fsai_drop_tol: 1e-12,
480 fsai_damping: 0.3,
481 normalize_strength: true,
482 coarse_solve: CoarseSolve::CG,
483 ilu_drop_tol: 1e-2,
484 ilu_fill_per_row: 0,
485 max_operator_complexity: None,
486 agg_num_levels: 1,
487 aggressive_mis_k: 2,
488 max_strong_per_row: None,
489 cycle_type: CycleType::V,
490 kcycle: None,
491 fmg_nu_pre: 1,
492 fmg_nu_post: 1,
493 fmg_gamma: 1,
494 fmg_levels_use: None,
495 non_galerkin: NonGalerkin {
496 enabled: false,
497 start_level: 1,
498 drop_abs: 0.0,
499 drop_rel: 0.0,
500 cap_row: 0,
501 symmetry: NgSymmetry::Symmetric,
502 lump_diagonal: true,
503 oc_target: None,
504 oc_max_iter: 4,
505 },
506 nodal: NodalMode::Off,
507 block_size: 1,
508 num_functions: 1,
509 near_nullspace: None,
510 filter_trial_vectors: None,
511 filter_omega: 1.0,
512 filter_after_non_galerkin: true,
513 post_interp: PostInterpType::None,
514 flexible_level: None,
515 flexible_iters: 0,
516 flexible_rtol: 0.0,
517 flexible_pc_sweeps: 1,
518 mixed_precision: None,
519 mixed_storage: MixedStorage::Cached,
520 conditioning: ConditioningOptions::default(),
521 dist_coarse_strategy: DistCoarseStrategy::RootGather,
522 dist_coarse_repartition: DistCoarseRepartition::Keep,
523 dist_coarse_solver_route: DistCoarseSolverRoute::Auto,
524 dist_apply_instrumentation: false,
525 dist_coarse_ghost_scale: 0.0,
526 require_native_complex_hierarchy: false,
527 level_relax_overrides: BTreeMap::new(),
528 level_sweep_overrides: BTreeMap::new(),
529 level_coarse_overrides: BTreeMap::new(),
530 };
531 cfg.grid_relax_type = [
532 cfg.relax_type,
533 cfg.relax_type,
534 cfg.relax_type,
535 RelaxType::GaussSeidel,
536 ];
537 cfg.num_grid_sweeps = [cfg.pre_sweeps, cfg.pre_sweeps, cfg.post_sweeps, 1];
538 cfg.stats_eps = cfg.drop_tol;
539 cfg
540 }
541}
542
543impl AMGConfig {
544 fn validate(&self) -> Result<(), KError> {
545 validate_relax_policy(self, self.coarse_solve)?;
546 validate_truncation_and_caps(self)?;
547 if self.max_levels == 0 {
548 return Err(KError::InvalidInput("max_levels must be at least 1".into()));
549 }
550 if self.min_coarse_size == 0 {
551 return Err(KError::InvalidInput(
552 "min_coarse_size must be at least 1".into(),
553 ));
554 }
555 if self.max_coarse_size > 0 && self.max_coarse_size < self.min_coarse_size {
556 return Err(KError::InvalidInput(
557 "max_coarse_size must be ≥ min_coarse_size".into(),
558 ));
559 }
560 if self.max_iterations < self.min_iterations {
561 return Err(KError::InvalidInput(
562 "max_iterations must be ≥ min_iterations".into(),
563 ));
564 }
565 if self.jacobi_omega <= 0.0 {
566 return Err(KError::InvalidInput("jacobi_omega must be positive".into()));
567 }
568 if self.chebyshev_power_steps == 0 {
569 return Err(KError::InvalidInput(
570 "chebyshev_power_steps must be ≥ 1".into(),
571 ));
572 }
573 if !(0.0 < self.chebyshev_lower_ratio && self.chebyshev_lower_ratio < 1.0) {
574 return Err(KError::InvalidInput(
575 "chebyshev_lower_ratio must be in (0, 1)".into(),
576 ));
577 }
578 if self.chebyshev_safety <= 0.0 {
579 return Err(KError::InvalidInput(
580 "chebyshev_safety must be positive".into(),
581 ));
582 }
583 if self.drop_tol < 0.0 {
584 return Err(KError::InvalidInput("drop_tol must be ≥ 0".into()));
585 }
586 if !self.dist_coarse_ghost_scale.is_finite() || self.dist_coarse_ghost_scale < 0.0 {
587 return Err(KError::InvalidInput(
588 "dist_coarse_ghost_scale must be finite and ≥ 0".into(),
589 ));
590 }
591 if self.non_galerkin.enabled && self.non_galerkin.start_level >= self.max_levels {
592 return Err(KError::InvalidInput(
593 "non_galerkin.start_level must be less than max_levels".into(),
594 ));
595 }
596 if self.verify_galerkin && self.galerkin_samples == 0 {
597 return Err(KError::InvalidInput(
598 "galerkin_samples must be > 0 when verify_galerkin is enabled".into(),
599 ));
600 }
601 Ok(())
602 }
603
604 fn set_smoothing_sweeps(&mut self, pre: usize, post: usize) {
605 self.pre_sweeps = pre;
606 self.post_sweeps = post;
607 self.num_grid_sweeps[RelaxPhase::Fine.ix()] = pre;
608 self.num_grid_sweeps[RelaxPhase::Down.ix()] = pre;
609 self.num_grid_sweeps[RelaxPhase::Up.ix()] = post;
610 }
611
612 fn apply_relax_type(&mut self, value: &str) -> Result<(), KError> {
613 let relax = map_relax(AmgRelaxKind::from_str(value)?);
614 self.relax_type = relax;
615 for phase in RelaxPhase::ALL {
616 self.grid_relax_type[phase.ix()] = relax;
617 }
618 self.grid_relax_type[RelaxPhase::Coarsest.ix()] = RelaxType::GaussSeidel;
619 Ok(())
620 }
621
622 pub fn try_from_opts(opts: &PcOptions) -> Result<Self, KError> {
623 let mut cfg = Self::default();
624 let mut per_phase_relax = false;
625 let mut pre_override: Option<usize> = None;
626 let mut post_override: Option<usize> = None;
627 let coarse_sweeps_explicit = opts.amg_sweeps_coarse.is_some();
628
629 if let Some(levels) = opts.amg_levels {
630 cfg.max_levels = levels;
631 }
632 if let Some(threshold) = opts.amg_strength_threshold {
633 let threshold = ensure_finite("amg_strength_threshold", threshold)?;
634 if !(threshold > 0.0 && threshold <= 1.0) {
635 return Err(KError::InvalidInput(
636 "amg_strength_threshold must be in (0, 1]".into(),
637 ));
638 }
639 cfg.strong_threshold = threshold;
640 }
641 if let Some(ref strength) = opts.amg_strength_type {
642 cfg.normalize_strength = map_strength_kind(AmgStrengthKind::from_str(strength)?);
643 }
644 if let Some(pre) = opts.amg_nu_pre {
645 cfg.set_smoothing_sweeps(pre, cfg.post_sweeps);
646 }
647 if let Some(post) = opts.amg_nu_post {
648 cfg.set_smoothing_sweeps(cfg.pre_sweeps, post);
649 }
650 if let Some(ref cycle) = opts.amg_cycle_type {
651 cfg.cycle_type = map_cycle_type(AmgCycleKind::from_str(cycle)?, opts.amg_cycle_w_gamma);
652 } else if let Some(gamma) = opts.amg_cycle_w_gamma {
653 cfg.cycle_type = CycleType::W {
654 gamma: gamma.max(2),
655 };
656 }
657 if let Some(threshold) = opts.amg_coarse_threshold {
658 cfg.coarse_threshold = threshold;
659 }
660 if let Some(max) = opts.amg_max_coarse_size {
661 cfg.max_coarse_size = max;
662 }
663 if let Some(min) = opts.amg_min_coarse_size {
664 cfg.min_coarse_size = min;
665 }
666 if let Some(trunc) = opts.amg_truncation_factor {
667 cfg.truncation_factor = trunc;
668 }
669 if let Some(cap) = opts.amg_max_elements_per_row {
670 cfg.max_elements_per_row = cap;
671 }
672 if let Some(interop) = opts.amg_interpolation_truncation {
673 cfg.interpolation_truncation = interop;
674 }
675 if let Some(abs) = opts.amg_rap_truncation_abs {
676 cfg.rap_truncation_abs = abs;
677 }
678 if let Some(cap) = opts.amg_rap_max_elements_per_row {
679 cfg.rap_max_elements_per_row = cap;
680 }
681 if let Some(ref coarsen) = opts.amg_coarsen_type {
682 cfg.coarsen_type = map_coarsen(AmgCoarsenKind::from_str(coarsen)?);
683 }
684 if let Some(ref interp) = opts.amg_interp_type {
685 cfg.interp_type = map_interp(AmgInterpKind::from_str(interp)?);
686 } else if let Some(ref interp) = opts.amg_interp_variant {
687 cfg.interp_type = map_interp(AmgInterpKind::from_str(interp)?);
688 }
689 if let Some(ref smoother) = opts.amg_smoother {
690 cfg.apply_relax_type(smoother)?;
691 } else if let Some(ref relax) = opts.amg_relax_type {
692 cfg.apply_relax_type(relax)?;
693 }
694 if let Some(ref smoother) = opts.amg_smoother_fine {
695 cfg.grid_relax_type[RelaxPhase::Fine.ix()] =
696 map_relax(AmgRelaxKind::from_str(smoother)?);
697 per_phase_relax = true;
698 }
699 if let Some(ref smoother) = opts.amg_smoother_down {
700 cfg.grid_relax_type[RelaxPhase::Down.ix()] =
701 map_relax(AmgRelaxKind::from_str(smoother)?);
702 per_phase_relax = true;
703 }
704 if let Some(ref smoother) = opts.amg_smoother_up {
705 cfg.grid_relax_type[RelaxPhase::Up.ix()] = map_relax(AmgRelaxKind::from_str(smoother)?);
706 per_phase_relax = true;
707 }
708 if let Some(ref smoother) = opts.amg_smoother_coarse {
709 cfg.grid_relax_type[RelaxPhase::Coarsest.ix()] =
710 map_relax(AmgRelaxKind::from_str(smoother)?);
711 per_phase_relax = true;
712 }
713 if let Some(steps) = opts.amg_smoother_steps {
714 cfg.set_smoothing_sweeps(steps, steps);
715 }
716 if let Some(fine) = opts.amg_sweeps_fine {
717 cfg.num_grid_sweeps[RelaxPhase::Fine.ix()] = fine;
718 if opts.amg_sweeps_down.is_none() {
719 pre_override = Some(fine);
720 }
721 }
722 if let Some(down) = opts.amg_sweeps_down {
723 cfg.num_grid_sweeps[RelaxPhase::Down.ix()] = down;
724 pre_override = Some(down);
725 }
726 if let Some(up) = opts.amg_sweeps_up {
727 cfg.num_grid_sweeps[RelaxPhase::Up.ix()] = up;
728 post_override = Some(up);
729 }
730 if let Some(coarse) = opts.amg_sweeps_coarse {
731 cfg.num_grid_sweeps[RelaxPhase::Coarsest.ix()] = coarse;
732 }
733 if let Some(pre) = pre_override {
734 cfg.pre_sweeps = pre;
735 }
736 if let Some(post) = post_override {
737 cfg.post_sweeps = post;
738 }
739 if let Some(omega) = opts.amg_smoother_omega {
740 let omega = ensure_finite("amg_smoother_omega", omega)?;
741 if omega <= 0.0 {
742 return Err(KError::InvalidInput(
743 "amg_smoother_omega must be > 0".into(),
744 ));
745 }
746 cfg.jacobi_omega = omega;
747 cfg.adaptive_smooth_omega = omega;
748 }
749 if let Some(val) = opts.amg_logging_level {
750 cfg.logging_level = val;
751 }
752 if let Some(val) = opts.amg_print_level {
753 cfg.print_level = val;
754 }
755 if let Some(val) = opts.amg_tolerance {
756 cfg.tolerance = val;
757 }
758 if let Some(val) = opts.amg_max_iterations {
759 cfg.max_iterations = val;
760 }
761 if let Some(val) = opts.amg_min_iterations {
762 cfg.min_iterations = val;
763 }
764 if let Some(flag) = opts.amg_ieee_checks {
765 cfg.ieee_checks = flag;
766 }
767 if let Some(flag) = opts.amg_optimize_workspace {
768 cfg.optimize_workspace = flag;
769 }
770 if let Some(ref coarse) = opts.amg_coarse_solver {
771 cfg.coarse_solve = map_coarse_solve(AmgCoarseSolveKind::from_str(coarse)?);
772 if matches!(cfg.coarse_solve, CoarseSolve::DirectDense) && !coarse_sweeps_explicit {
773 cfg.num_grid_sweeps[RelaxPhase::Coarsest.ix()] = 0;
774 }
775 }
776 if let Some(flag) = opts.amg_keep_transpose {
777 cfg.keep_transpose = flag;
778 }
779 if let Some(flag) = opts.amg_keep_pivot_in_rap {
780 cfg.keep_pivot_in_rap = flag;
781 }
782 if let Some(flag) = opts.amg_require_spd {
783 cfg.require_spd = flag;
784 }
785 if let Some(flag) = opts.amg_require_native_complex_hierarchy {
786 cfg.require_native_complex_hierarchy = flag;
787 }
788 if cfg.require_spd && !cfg.keep_transpose {
789 return Err(KError::InvalidInput(
790 "SPD mode requires pc_amg_keep_transpose true".into(),
791 ));
792 }
793 if let Some(flag) = opts.amg_print_setup {
794 if flag {
795 cfg.print_level = cfg.print_level.max(1);
796 cfg.logging_level = cfg.logging_level.max(2);
797 }
798 }
799 if let Some(ref mode) = opts.amg_dist_apply_mode {
800 cfg.dist_coarse_strategy = parse_dist_apply_mode(mode)?;
801 }
802 if let Some(ref policy) = opts.amg_dist_coarse_policy {
803 cfg.dist_coarse_strategy = parse_dist_apply_mode(policy)?;
804 }
805 if let Some(ref repartition) = opts.amg_dist_coarse_repartition {
806 cfg.dist_coarse_repartition = parse_dist_coarse_repartition(repartition)?;
807 }
808 if let Some(ref route) = opts.amg_dist_coarse_solver_route {
809 cfg.dist_coarse_solver_route = parse_dist_coarse_solver_route(route)?;
810 }
811 if let Some(flag) = opts.amg_dist_instrumentation {
812 cfg.dist_apply_instrumentation = flag;
813 }
814 if let Some(scale) = opts.amg_dist_coarse_ghost_scale {
815 cfg.dist_coarse_ghost_scale = ensure_finite("amg_dist_coarse_ghost_scale", scale)?;
816 }
817 for (level, scoped) in &opts.pc_amg_level_scoped_options {
818 if let Some(relax) = scoped
819 .amg_smoother
820 .as_deref()
821 .or(scoped.amg_relax_type.as_deref())
822 {
823 cfg.level_relax_overrides
824 .insert(*level, map_relax(AmgRelaxKind::from_str(relax)?));
825 }
826 if scoped.amg_smoother_steps.is_some()
827 || scoped.amg_sweeps_down.is_some()
828 || scoped.amg_sweeps_up.is_some()
829 {
830 let sweeps = scoped.amg_smoother_steps.unwrap_or(1);
831 let pre = scoped.amg_sweeps_down.unwrap_or(sweeps);
832 let post = scoped.amg_sweeps_up.unwrap_or(sweeps);
833 cfg.level_sweep_overrides.insert(*level, (pre, post));
834 }
835 if let Some(coarse) = scoped.amg_coarse_solver.as_deref() {
836 cfg.level_coarse_overrides.insert(
837 *level,
838 map_coarse_solve(AmgCoarseSolveKind::from_str(coarse)?),
839 );
840 }
841 }
842 if per_phase_relax {
843 cfg.relax_type = cfg.grid_relax_type[RelaxPhase::Fine.ix()];
844 }
845 cfg.conditioning = opts.conditioning_options()?;
846 cfg.validate()?;
847 Ok(cfg)
848 }
849}
850
851fn ensure_finite(name: &str, value: f64) -> Result<f64, KError> {
852 if value.is_finite() {
853 Ok(value)
854 } else {
855 Err(KError::InvalidInput(format!("{name} must be finite")))
856 }
857}
858
859fn map_relax(kind: AmgRelaxKind) -> RelaxType {
860 match kind {
861 AmgRelaxKind::Jacobi => RelaxType::Jacobi,
862 AmgRelaxKind::Gs => RelaxType::GaussSeidel,
863 AmgRelaxKind::Gsr => RelaxType::GaussSeidelBackward,
864 AmgRelaxKind::Sgs => RelaxType::SymmetricGaussSeidel,
865 AmgRelaxKind::Hgs => RelaxType::HybridGaussSeidel,
866 AmgRelaxKind::L1Jacobi => RelaxType::L1Jacobi,
867 AmgRelaxKind::Chebyshev => RelaxType::Chebyshev,
868 AmgRelaxKind::ChebyshevSafe => RelaxType::ChebyshevSafe,
869 AmgRelaxKind::SafeguardedGs => RelaxType::SafeguardedGaussSeidel,
870 AmgRelaxKind::Ilu0 => RelaxType::Ilu0,
871 AmgRelaxKind::Ras => RelaxType::Ras,
872 }
873}
874
875fn map_coarsen(kind: AmgCoarsenKind) -> CoarsenType {
876 match kind {
877 AmgCoarsenKind::Rs => CoarsenType::RS,
878 AmgCoarsenKind::Hmis => CoarsenType::HMIS,
879 AmgCoarsenKind::Pmis => CoarsenType::PMIS,
880 AmgCoarsenKind::Falgout => CoarsenType::Falgout,
881 }
882}
883
884fn map_cycle_type(kind: AmgCycleKind, gamma: Option<usize>) -> CycleType {
885 match kind {
886 AmgCycleKind::V => CycleType::V,
887 AmgCycleKind::W => CycleType::W {
888 gamma: gamma.unwrap_or(2).max(2),
889 },
890 }
891}
892
893fn dist_route_fallback_order(
894 selected: DistCoarseSolverRoute,
895 strategy: DistCoarseStrategy,
896) -> Vec<DistCoarseSolverRoute> {
897 let mut routes = Vec::new();
898 let mut push_unique = |route: DistCoarseSolverRoute| {
899 if !routes.contains(&route) {
900 routes.push(route);
901 }
902 };
903 if selected != DistCoarseSolverRoute::Auto {
904 push_unique(selected);
905 }
906 match strategy {
907 DistCoarseStrategy::DistributedCsr => {
908 push_unique(DistCoarseSolverRoute::Local);
909 push_unique(DistCoarseSolverRoute::Root);
910 push_unique(DistCoarseSolverRoute::SuperLuDist);
911 }
912 DistCoarseStrategy::RootGather => {
913 push_unique(DistCoarseSolverRoute::Root);
914 push_unique(DistCoarseSolverRoute::Local);
915 push_unique(DistCoarseSolverRoute::SuperLuDist);
916 }
917 DistCoarseStrategy::LocalPrototype => {
918 push_unique(DistCoarseSolverRoute::Local);
919 push_unique(DistCoarseSolverRoute::Root);
920 push_unique(DistCoarseSolverRoute::SuperLuDist);
921 }
922 DistCoarseStrategy::SuperLuDist => {
923 push_unique(DistCoarseSolverRoute::SuperLuDist);
924 push_unique(DistCoarseSolverRoute::Root);
925 push_unique(DistCoarseSolverRoute::Local);
926 }
927 DistCoarseStrategy::None => {
928 push_unique(DistCoarseSolverRoute::Local);
929 push_unique(DistCoarseSolverRoute::Root);
930 }
931 }
932 routes
933}
934
935fn dist_route_label(route: DistCoarseSolverRoute, strategy: DistCoarseStrategy) -> &'static str {
936 match route {
937 DistCoarseSolverRoute::Root => "root_gather",
938 DistCoarseSolverRoute::Local => "local_prototype",
939 DistCoarseSolverRoute::SuperLuDist => "superlu_dist",
940 DistCoarseSolverRoute::Auto => match strategy {
941 DistCoarseStrategy::DistributedCsr => "distributed_csr",
942 DistCoarseStrategy::RootGather => "root_gather",
943 DistCoarseStrategy::LocalPrototype => "local_prototype",
944 DistCoarseStrategy::SuperLuDist => "superlu_dist",
945 DistCoarseStrategy::None => "none",
946 },
947 }
948}
949
950fn dist_strategy_label(strategy: DistCoarseStrategy) -> &'static str {
951 match strategy {
952 DistCoarseStrategy::DistributedCsr => "distributed_csr",
953 DistCoarseStrategy::RootGather => "root_gather",
954 DistCoarseStrategy::LocalPrototype => "local_prototype",
955 DistCoarseStrategy::SuperLuDist => "superlu_dist",
956 DistCoarseStrategy::None => "none",
957 }
958}
959
960fn dist_route_fallback_labels(
961 selected: DistCoarseSolverRoute,
962 strategy: DistCoarseStrategy,
963) -> Vec<String> {
964 dist_route_fallback_order(selected, strategy)
965 .into_iter()
966 .map(|route| dist_route_label(route, strategy).to_string())
967 .collect()
968}
969
970fn amg_dist_route_reports_distributed(
971 route: DistCoarseSolverRoute,
972 strategy: DistCoarseStrategy,
973) -> bool {
974 match route {
975 DistCoarseSolverRoute::SuperLuDist => true,
976 DistCoarseSolverRoute::Auto => matches!(strategy, DistCoarseStrategy::SuperLuDist),
977 DistCoarseSolverRoute::Root | DistCoarseSolverRoute::Local => false,
978 }
979}
980
981fn parse_dist_apply_mode(value: &str) -> Result<DistCoarseStrategy, KError> {
982 DistCoarseStrategy::from_str(value)
983 .map_err(|_| KError::InvalidInput(format!("invalid amg_dist_apply_mode: {value}")))
984}
985
986fn parse_dist_coarse_repartition(value: &str) -> Result<DistCoarseRepartition, KError> {
987 DistCoarseRepartition::from_str(value)
988 .map_err(|_| KError::InvalidInput(format!("invalid amg_dist_coarse_repartition: {value}")))
989}
990
991fn parse_dist_coarse_solver_route(value: &str) -> Result<DistCoarseSolverRoute, KError> {
992 DistCoarseSolverRoute::from_str(value)
993 .map_err(|_| KError::InvalidInput(format!("invalid amg_dist_coarse_solver_route: {value}")))
994}
995
996fn map_interp(kind: AmgInterpKind) -> InterpType {
997 match kind {
998 AmgInterpKind::Classical => InterpType::Classical,
999 AmgInterpKind::Direct => InterpType::Direct,
1000 AmgInterpKind::Multipass => InterpType::Multipass,
1001 AmgInterpKind::Extended => InterpType::Extended,
1002 AmgInterpKind::Standard => InterpType::Standard,
1003 AmgInterpKind::He => InterpType::HE,
1004 }
1005}
1006
1007fn map_coarse_solve(kind: AmgCoarseSolveKind) -> CoarseSolve {
1008 match kind {
1009 AmgCoarseSolveKind::Cg => CoarseSolve::CG,
1010 AmgCoarseSolveKind::Direct => CoarseSolve::DirectDense,
1011 AmgCoarseSolveKind::Ilu => CoarseSolve::ILU,
1012 AmgCoarseSolveKind::Smoother => CoarseSolve::Smoother,
1013 }
1014}
1015
1016fn map_strength_kind(kind: AmgStrengthKind) -> bool {
1017 match kind {
1018 AmgStrengthKind::Classical => false,
1019 AmgStrengthKind::Symmetric | AmgStrengthKind::Normalized => true,
1020 }
1021}
1022
1023#[cfg(test)]
1024mod config_mapping_tests {
1025 use super::*;
1026
1027 fn opts_from(args: &[&str]) -> PcOptions {
1028 PcOptions::from_args(args).expect("valid AMG args")
1029 }
1030
1031 #[test]
1032 fn amg_config_applies_cli_overrides() {
1033 let opts = opts_from(&[
1034 "-pc_type",
1035 "amg",
1036 "-pc_amg_levels",
1037 "6",
1038 "-pc_amg_strength_threshold",
1039 "0.25",
1040 "-pc_amg_coarsen",
1041 "hmis",
1042 "-pc_amg_interp",
1043 "extended",
1044 "-pc_amg_smoother",
1045 "chebyshev",
1046 "-pc_amg_smoother_steps",
1047 "2",
1048 "-pc_amg_smoother_omega",
1049 "0.8",
1050 "-pc_amg_truncation_factor",
1051 "0.2",
1052 "-pc_amg_interp_maxnnz",
1053 "8",
1054 "-pc_amg_rap_truncation_abs",
1055 "0.0",
1056 "-pc_amg_rap_maxnnz",
1057 "16",
1058 "-pc_amg_keep_transpose",
1059 "true",
1060 "-pc_amg_keep_pivot_in_rap",
1061 "true",
1062 "-pc_amg_require_spd",
1063 "true",
1064 "-pc_amg_print_setup",
1065 "true",
1066 ]);
1067 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1068 assert_eq!(cfg.max_levels, 6);
1069 assert!((cfg.strong_threshold - 0.25).abs() < 1e-12);
1070 assert_eq!(cfg.coarsen_type, CoarsenType::HMIS);
1071 assert_eq!(cfg.interp_type, InterpType::Extended);
1072 assert_eq!(cfg.relax_type, RelaxType::Chebyshev);
1073 assert_eq!(cfg.pre_sweeps, 2);
1074 assert_eq!(cfg.post_sweeps, 2);
1075 assert!((cfg.jacobi_omega - 0.8).abs() < 1e-12);
1076 assert!((cfg.adaptive_smooth_omega - 0.8).abs() < 1e-12);
1077 assert_eq!(cfg.truncation_factor, 0.2);
1078 assert_eq!(cfg.max_elements_per_row, 8);
1079 assert_eq!(cfg.rap_truncation_abs, 0.0);
1080 assert_eq!(cfg.rap_max_elements_per_row, 16);
1081 assert!(cfg.keep_transpose);
1082 assert!(cfg.keep_pivot_in_rap);
1083 assert!(cfg.require_spd);
1084 assert!(cfg.logging_level >= 2);
1085 assert!(cfg.print_level >= 1);
1086 }
1087
1088 #[test]
1089 fn amg_config_denies_keep_transpose_when_spd() {
1090 let opts = opts_from(&["-pc_type", "amg", "-pc_amg_keep_transpose", "false"]);
1091 assert!(AMGConfig::try_from_opts(&opts).is_err());
1092 }
1093
1094 #[test]
1095 fn amg_config_maps_native_complex_hierarchy_requirement() {
1096 let opts = opts_from(&[
1097 "-pc_type",
1098 "amg",
1099 "-pc_amg_require_native_complex_hierarchy",
1100 "true",
1101 ]);
1102 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1103 assert!(cfg.require_native_complex_hierarchy);
1104 }
1105
1106 #[test]
1107 fn amg_config_allows_keep_transpose_when_spd_off() {
1108 let opts = opts_from(&[
1109 "-pc_type",
1110 "amg",
1111 "-pc_amg_require_spd",
1112 "false",
1113 "-pc_amg_keep_transpose",
1114 "false",
1115 ]);
1116 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1117 assert!(!cfg.require_spd);
1118 assert!(!cfg.keep_transpose);
1119 }
1120
1121 #[test]
1122 fn amg_config_cycle_and_coarse_solver_overrides() {
1123 let opts = opts_from(&[
1124 "-pc_type",
1125 "amg",
1126 "-pc_amg_cycle_type",
1127 "w",
1128 "-pc_amg_cycle_w_gamma",
1129 "3",
1130 "-pc_amg_coarse_solver",
1131 "direct",
1132 ]);
1133 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1134 assert_eq!(cfg.cycle_type, CycleType::W { gamma: 3 });
1135 assert_eq!(cfg.coarse_solve, CoarseSolve::DirectDense);
1136 assert_eq!(cfg.num_grid_sweeps[RelaxPhase::Coarsest.ix()], 0);
1137 }
1138
1139 #[test]
1140 fn amg_config_phase_smoothers_and_sweeps_override() {
1141 let opts = opts_from(&[
1142 "-pc_type",
1143 "amg",
1144 "-pc_amg_require_spd",
1145 "false",
1146 "-pc_amg_smoother_fine",
1147 "jacobi",
1148 "-pc_amg_smoother_down",
1149 "gs",
1150 "-pc_amg_smoother_up",
1151 "sgs",
1152 "-pc_amg_smoother_coarse",
1153 "chebyshev",
1154 "-pc_amg_sweeps_fine",
1155 "2",
1156 "-pc_amg_sweeps_down",
1157 "3",
1158 "-pc_amg_sweeps_up",
1159 "4",
1160 "-pc_amg_sweeps_coarse",
1161 "0",
1162 ]);
1163 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1164 assert_eq!(
1165 cfg.grid_relax_type[RelaxPhase::Fine.ix()],
1166 RelaxType::Jacobi
1167 );
1168 assert_eq!(
1169 cfg.grid_relax_type[RelaxPhase::Down.ix()],
1170 RelaxType::GaussSeidel
1171 );
1172 assert_eq!(
1173 cfg.grid_relax_type[RelaxPhase::Up.ix()],
1174 RelaxType::SymmetricGaussSeidel
1175 );
1176 assert_eq!(
1177 cfg.grid_relax_type[RelaxPhase::Coarsest.ix()],
1178 RelaxType::Chebyshev
1179 );
1180 assert_eq!(cfg.num_grid_sweeps[RelaxPhase::Fine.ix()], 2);
1181 assert_eq!(cfg.num_grid_sweeps[RelaxPhase::Down.ix()], 3);
1182 assert_eq!(cfg.num_grid_sweeps[RelaxPhase::Up.ix()], 4);
1183 assert_eq!(cfg.num_grid_sweeps[RelaxPhase::Coarsest.ix()], 0);
1184 assert_eq!(cfg.pre_sweeps, 3);
1185 assert_eq!(cfg.post_sweeps, 4);
1186 assert_eq!(cfg.relax_type, RelaxType::Jacobi);
1187 }
1188
1189 #[test]
1190 fn amg_config_strength_and_interp_variants() {
1191 let opts = opts_from(&[
1192 "-pc_type",
1193 "amg",
1194 "-pc_amg_strength_type",
1195 "classical",
1196 "-pc_amg_interp_variant",
1197 "direct",
1198 ]);
1199 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1200 assert!(!cfg.normalize_strength);
1201 assert_eq!(cfg.interp_type, InterpType::Direct);
1202 }
1203
1204 #[test]
1205 fn amg_dist_coarse_policy_overrides_apply_mode() {
1206 let opts = opts_from(&[
1207 "-pc_type",
1208 "amg",
1209 "-pc_amg_dist_apply_mode",
1210 "root",
1211 "-pc_mg_dist_coarse_policy",
1212 "local",
1213 ]);
1214 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1215 assert_eq!(cfg.dist_coarse_strategy, DistCoarseStrategy::LocalPrototype);
1216 }
1217 #[test]
1218 fn amg_levels_zero_is_invalid() {
1219 let opts = opts_from(&["-pc_type", "amg", "-pc_amg_levels", "0"]);
1220 assert!(AMGConfig::try_from_opts(&opts).is_err());
1221 }
1222
1223 #[test]
1224 fn amg_strength_threshold_negative_is_invalid() {
1225 let opts = opts_from(&["-pc_type", "amg", "-pc_amg_strength_threshold", "-0.1"]);
1226 assert!(AMGConfig::try_from_opts(&opts).is_err());
1227 }
1228
1229 #[test]
1230 fn pc_amg_alias_sets_pc_type() {
1231 let opts = opts_from(&["-pc_amg"]);
1232 assert_eq!(opts.pc_type.as_deref(), Some("amg"));
1233 let cfg = AMGConfig::try_from_opts(&opts).unwrap();
1234 assert_eq!(cfg.max_levels, AMGConfig::default().max_levels);
1235 }
1236
1237 #[cfg(not(feature = "complex"))]
1238 #[test]
1239 fn dist_route_resolution_preserves_forced_routes() {
1240 let mut cfg = AMGConfig::default();
1241 cfg.dist_coarse_strategy = DistCoarseStrategy::RootGather;
1242 cfg.dist_coarse_solver_route = DistCoarseSolverRoute::Local;
1243 let amg = AMG::with_config(cfg);
1244 let comm = UniverseComm::NoComm(crate::parallel::NoComm);
1245 let (strategy, route) = amg.resolve_dist_coarse_strategy(&comm).unwrap();
1246 assert_eq!(strategy, DistCoarseStrategy::LocalPrototype);
1247 assert_eq!(route, DistCoarseSolverRoute::Local);
1248 }
1249
1250 #[cfg(all(not(feature = "complex"), not(feature = "superlu_dist")))]
1251 #[test]
1252 fn dist_route_resolution_errors_when_forced_superlu_missing() {
1253 let mut cfg = AMGConfig::default();
1254 cfg.dist_coarse_solver_route = DistCoarseSolverRoute::SuperLuDist;
1255 let amg = AMG::with_config(cfg);
1256 let comm = UniverseComm::NoComm(crate::parallel::NoComm);
1257 let err = amg
1258 .resolve_dist_coarse_strategy(&comm)
1259 .expect_err("missing feature");
1260 assert!(err.to_string().contains("explicitly requested"));
1261 }
1262
1263 #[test]
1264 fn dist_route_fallback_order_chain_is_deterministic() {
1265 let chain = dist_route_fallback_order(
1266 DistCoarseSolverRoute::SuperLuDist,
1267 DistCoarseStrategy::RootGather,
1268 );
1269 assert_eq!(
1270 chain,
1271 vec![
1272 DistCoarseSolverRoute::SuperLuDist,
1273 DistCoarseSolverRoute::Root,
1274 DistCoarseSolverRoute::Local
1275 ]
1276 );
1277 }
1278
1279 #[test]
1280 fn dist_route_fallback_order_omits_auto_placeholder() {
1281 let chain =
1282 dist_route_fallback_order(DistCoarseSolverRoute::Auto, DistCoarseStrategy::RootGather);
1283 assert_eq!(
1284 chain,
1285 vec![
1286 DistCoarseSolverRoute::Root,
1287 DistCoarseSolverRoute::Local,
1288 DistCoarseSolverRoute::SuperLuDist
1289 ]
1290 );
1291 assert_eq!(
1292 dist_route_fallback_labels(DistCoarseSolverRoute::Auto, DistCoarseStrategy::RootGather),
1293 vec![
1294 "root_gather".to_string(),
1295 "local_prototype".to_string(),
1296 "superlu_dist".to_string()
1297 ]
1298 );
1299 }
1300
1301 #[test]
1302 fn dist_route_support_reporting_excludes_fallback_prototypes() {
1303 assert!(!amg_dist_route_reports_distributed(
1304 DistCoarseSolverRoute::Root,
1305 DistCoarseStrategy::RootGather
1306 ));
1307 assert!(!amg_dist_route_reports_distributed(
1308 DistCoarseSolverRoute::Local,
1309 DistCoarseStrategy::LocalPrototype
1310 ));
1311 assert!(!amg_dist_route_reports_distributed(
1312 DistCoarseSolverRoute::Auto,
1313 DistCoarseStrategy::RootGather
1314 ));
1315 assert!(amg_dist_route_reports_distributed(
1316 DistCoarseSolverRoute::SuperLuDist,
1317 DistCoarseStrategy::RootGather
1318 ));
1319 assert!(amg_dist_route_reports_distributed(
1320 DistCoarseSolverRoute::Auto,
1321 DistCoarseStrategy::SuperLuDist
1322 ));
1323 }
1324}
1325
1326pub struct AMGBuilder {
1328 cfg: AMGConfig,
1329}
1330
1331impl AMGBuilder {
1332 pub fn new() -> Self {
1333 Self {
1334 cfg: AMGConfig::default(),
1335 }
1336 }
1337 pub fn cycle_v(mut self) -> Self {
1338 self.cfg.cycle_type = CycleType::V;
1339 self
1340 }
1341 pub fn cycle_w(mut self, gamma: usize) -> Self {
1342 let g = gamma.max(2);
1343 self.cfg.cycle_type = CycleType::W { gamma: g };
1344 self
1345 }
1346 pub fn cycle(mut self, c: CycleType) -> Self {
1347 self.cfg.cycle_type = c;
1348 self
1349 }
1350 pub fn kcycle(mut self, kc: KCycle) -> Self {
1351 self.cfg.kcycle = Some(kc);
1352 self
1353 }
1354 pub fn disable_kcycle(mut self) -> Self {
1355 self.cfg.kcycle = None;
1356 self
1357 }
1358 pub fn max_levels(mut self, v: usize) -> Self {
1359 self.cfg.max_levels = v;
1360 self
1361 }
1362 pub fn strong_threshold(mut self, v: f64) -> Self {
1363 self.cfg.strong_threshold = v;
1364 self
1365 }
1366 pub fn coarse_threshold(mut self, v: usize) -> Self {
1367 self.cfg.coarse_threshold = v;
1368 self
1369 }
1370 pub fn max_coarse_size(mut self, v: usize) -> Self {
1371 self.cfg.max_coarse_size = v;
1372 self
1373 }
1374 pub fn min_coarse_size(mut self, v: usize) -> Self {
1375 self.cfg.min_coarse_size = v;
1376 self
1377 }
1378 pub fn truncation_factor(mut self, v: f64) -> Self {
1379 self.cfg.truncation_factor = v;
1380 self
1381 }
1382 pub fn interpolation_drop_abs(mut self, v: f64) -> Self {
1383 self.cfg.interpolation_truncation = v;
1384 self
1385 }
1386 pub fn interpolation_cap(mut self, k: usize) -> Self {
1387 self.cfg.max_elements_per_row = k;
1388 self
1389 }
1390 pub fn rap_drop_abs(mut self, v: f64) -> Self {
1391 self.cfg.rap_truncation_abs = v;
1392 self
1393 }
1394 pub fn rap_cap(mut self, k: usize) -> Self {
1395 self.cfg.rap_max_elements_per_row = k;
1396 self
1397 }
1398 pub fn keep_transpose(mut self, on: bool) -> Self {
1399 self.cfg.keep_transpose = on;
1400 self
1401 }
1402 pub fn keep_pivot_in_rap(mut self, yes: bool) -> Self {
1403 self.cfg.keep_pivot_in_rap = yes;
1404 self
1405 }
1406 pub fn require_spd(mut self, on: bool) -> Self {
1407 self.cfg.require_spd = on;
1408 self
1409 }
1410 pub fn verify_p_rank(mut self, on: bool) -> Self {
1411 self.cfg.verify_p_rank = on;
1412 self
1413 }
1414 pub fn rank_cond_threshold(mut self, v: f64) -> Self {
1415 self.cfg.rank_cond_threshold = v;
1416 self
1417 }
1418 pub fn rank_min_col_norm(mut self, v: f64) -> Self {
1419 self.cfg.rank_min_col_norm = v;
1420 self
1421 }
1422 pub fn verify_galerkin(mut self, on: bool) -> Self {
1423 self.cfg.verify_galerkin = on;
1424 self
1425 }
1426 pub fn galerkin_rel_tol(mut self, v: f64) -> Self {
1427 self.cfg.galerkin_rel_tol = v;
1428 self
1429 }
1430 pub fn spd_diag_floor(mut self, eps: f64) -> Self {
1431 self.cfg.spd_diag_floor = eps.max(0.0);
1432 self
1433 }
1434 pub fn forbid_non_galerkin_in_spd(mut self, on: bool) -> Self {
1435 self.cfg.forbid_non_galerkin_in_spd = on;
1436 self
1437 }
1438 pub fn interpolation_truncation(self, v: f64) -> Self {
1439 self.interpolation_drop_abs(v)
1440 }
1441 pub fn smoothing_sweeps(mut self, pre: usize, post: usize) -> Self {
1442 self.cfg.pre_sweeps = pre;
1443 self.cfg.post_sweeps = post;
1444 self.cfg.num_grid_sweeps[RelaxPhase::Fine.ix()] = pre;
1445 self.cfg.num_grid_sweeps[RelaxPhase::Down.ix()] = pre;
1446 self.cfg.num_grid_sweeps[RelaxPhase::Up.ix()] = post;
1447 self
1449 }
1450 pub fn coarsening_type(mut self, v: CoarsenType) -> Self {
1451 self.cfg.coarsen_type = v;
1452 self
1453 }
1454 pub fn agg_num_levels(mut self, v: usize) -> Self {
1455 self.cfg.agg_num_levels = v;
1456 self
1457 }
1458 pub fn aggressive_mis_k(mut self, v: usize) -> Self {
1459 self.cfg.aggressive_mis_k = v;
1460 self
1461 }
1462 pub fn max_strong_per_row(mut self, k: usize) -> Self {
1463 self.cfg.max_strong_per_row = Some(k);
1464 self
1465 }
1466 pub fn interpolation_type(mut self, v: InterpType) -> Self {
1467 self.cfg.interp_type = v;
1468 self
1469 }
1470 pub fn relaxation_type(mut self, v: RelaxType) -> Self {
1471 self.cfg.relax_type = v;
1472 for ph in RelaxPhase::ALL {
1473 self.cfg.grid_relax_type[ph.ix()] = v;
1474 }
1475 self.cfg.grid_relax_type[RelaxPhase::Coarsest.ix()] = RelaxType::GaussSeidel;
1476 self
1477 }
1478 pub fn grid_relax_type(mut self, phase: RelaxPhase, t: RelaxType) -> Self {
1479 self.cfg.grid_relax_type[phase.ix()] = t;
1480 self
1481 }
1482 pub fn num_grid_sweeps(mut self, phase: RelaxPhase, k: usize) -> Self {
1483 self.cfg.num_grid_sweeps[phase.ix()] = k;
1484 self
1485 }
1486 pub fn grid_relax_type_all(mut self, t: RelaxType) -> Self {
1487 for ph in RelaxPhase::ALL {
1488 self.cfg.grid_relax_type[ph.ix()] = t;
1489 }
1490 self
1491 }
1492 pub fn num_grid_sweeps_all(mut self, k: usize) -> Self {
1493 for ph in RelaxPhase::ALL {
1494 self.cfg.num_grid_sweeps[ph.ix()] = k;
1495 }
1496 self
1497 }
1498 pub fn enable_logging(mut self) -> Self {
1499 self.cfg.logging_level = 1;
1500 self
1501 }
1502 pub fn logging_level(mut self, lvl: usize) -> Self {
1503 self.cfg.logging_level = lvl;
1504 self
1505 }
1506 pub fn enable_printing(mut self) -> Self {
1507 self.cfg.print_level = 1;
1508 self
1509 }
1510 pub fn print_level(mut self, lvl: usize) -> Self {
1511 self.cfg.print_level = lvl;
1512 self
1513 }
1514 pub fn jacobi_omega(mut self, w: f64) -> Self {
1515 self.cfg.jacobi_omega = w;
1516 self.cfg.adaptive_smooth_omega = w;
1517 self
1518 }
1519 pub fn adaptive_interp(mut self, on: bool) -> Self {
1520 self.cfg.adaptive_interp = on;
1521 self
1522 }
1523 pub fn adaptive_samples(mut self, r: usize) -> Self {
1524 self.cfg.adaptive_samples = r;
1525 self
1526 }
1527 pub fn adaptive_smooth_steps(mut self, nu: usize) -> Self {
1528 self.cfg.adaptive_smooth_steps = nu;
1529 self
1530 }
1531 pub fn adaptive_smooth_omega(mut self, w: f64) -> Self {
1532 self.cfg.adaptive_smooth_omega = w;
1533 self
1534 }
1535 pub fn adaptive_lambda(mut self, lam: f64) -> Self {
1536 self.cfg.adaptive_lambda = lam;
1537 self
1538 }
1539 pub fn adaptive_enforce_sum1(mut self, on: bool) -> Self {
1540 self.cfg.adaptive_enforce_sum1 = on;
1541 self
1542 }
1543 pub fn adaptive_weight_mode(mut self, mode: AdaptiveWeight) -> Self {
1544 self.cfg.adaptive_weight_mode = mode;
1545 self
1546 }
1547 pub fn chebyshev_recompute_esteig(mut self, on: bool) -> Self {
1548 self.cfg.chebyshev_recompute_esteig = on;
1549 self
1550 }
1551 pub fn chebyshev_safety(mut self, s: f64) -> Self {
1552 self.cfg.chebyshev_safety = s;
1553 self
1554 }
1555 pub fn chebyshev_lower_ratio(mut self, r: f64) -> Self {
1556 self.cfg.chebyshev_lower_ratio = r;
1557 self
1558 }
1559 pub fn chebyshev_power_steps(mut self, steps: usize) -> Self {
1560 self.cfg.chebyshev_power_steps = steps;
1561 self
1562 }
1563 pub fn chebyshev_degree(mut self, d: usize) -> Self {
1564 self.cfg.chebyshev_degree = d;
1565 self
1566 }
1567 pub fn filter_trials(mut self, trials: Vec<Vec<f64>>) -> Self {
1568 self.cfg.filter_trial_vectors = Some(trials);
1569 self
1570 }
1571 pub fn filter_omega(mut self, omega: f64) -> Self {
1572 self.cfg.filter_omega = omega;
1573 self
1574 }
1575 pub fn filter_after_non_galerkin(mut self, on: bool) -> Self {
1576 self.cfg.filter_after_non_galerkin = on;
1577 self
1578 }
1579 pub fn use_level_scheduling(mut self, v: bool) -> Self {
1580 self.cfg.use_level_scheduling = v;
1581 self
1582 }
1583 pub fn drop_tolerance(mut self, t: f64) -> Self {
1584 self.cfg.drop_tol = t;
1585 self.cfg.stats_eps = t;
1586 self
1587 }
1588 pub fn stats_eps(mut self, t: f64) -> Self {
1589 self.cfg.stats_eps = t;
1590 self
1591 }
1592 pub fn fsai_dist(mut self, dist: usize) -> Self {
1593 self.cfg.fsai_dist = dist.max(1);
1594 self
1595 }
1596 pub fn fsai_use_strength(mut self, use_strength: bool) -> Self {
1597 self.cfg.fsai_use_strength = use_strength;
1598 self
1599 }
1600 pub fn fsai_max_per_row(mut self, cap: usize) -> Self {
1601 self.cfg.fsai_max_per_row = cap.max(1);
1602 self
1603 }
1604 pub fn fsai_lambda(mut self, lambda: f64) -> Self {
1605 self.cfg.fsai_lambda = if lambda >= 0.0 { lambda } else { 0.0 };
1606 self
1607 }
1608 pub fn fsai_adaptive_passes(mut self, passes: usize) -> Self {
1609 self.cfg.fsai_adaptive_passes = passes;
1610 self
1611 }
1612 pub fn fsai_drop_tol(mut self, drop: f64) -> Self {
1613 self.cfg.fsai_drop_tol = if drop >= 0.0 { drop } else { 0.0 };
1614 self
1615 }
1616 pub fn fsai_damping(mut self, tau: f64) -> Self {
1617 self.cfg.fsai_damping = tau;
1618 self
1619 }
1620
1621 pub fn coarse_solve(mut self, v: CoarseSolve) -> Self {
1622 self.cfg.coarse_solve = v;
1623 self
1624 }
1625
1626 pub fn ilu_params(mut self, drop_tol: f64, fill_per_row: usize) -> Self {
1627 self.cfg.ilu_drop_tol = drop_tol;
1628 self.cfg.ilu_fill_per_row = fill_per_row;
1629 self
1630 }
1631
1632 pub fn non_galerkin(
1633 mut self,
1634 enabled: bool,
1635 start_level: usize,
1636 drop_abs: f64,
1637 drop_rel: f64,
1638 cap_row: usize,
1639 ) -> Self {
1640 self.cfg.non_galerkin.enabled = enabled;
1641 self.cfg.non_galerkin.start_level = start_level;
1642 self.cfg.non_galerkin.drop_abs = drop_abs;
1643 self.cfg.non_galerkin.drop_rel = drop_rel;
1644 self.cfg.non_galerkin.cap_row = cap_row;
1645 self
1646 }
1647
1648 pub fn non_galerkin_symmetry(mut self, sym: NgSymmetry, lump_diag: bool) -> Self {
1649 self.cfg.non_galerkin.symmetry = sym;
1650 self.cfg.non_galerkin.lump_diagonal = lump_diag;
1651 self
1652 }
1653
1654 pub fn non_galerkin_oc_target(mut self, target: Option<f64>, iters: usize) -> Self {
1655 self.cfg.non_galerkin.oc_target = target;
1656 self.cfg.non_galerkin.oc_max_iter = iters;
1657 self
1658 }
1659
1660 pub fn nodal(mut self, on: bool, block_size: usize) -> Self {
1661 self.cfg.nodal = if on { NodalMode::Nodal } else { NodalMode::Off };
1662 self.cfg.block_size = block_size.max(1);
1663 self
1664 }
1665
1666 pub fn num_functions(mut self, r: usize) -> Self {
1667 self.cfg.num_functions = r.max(1);
1668 self
1669 }
1670
1671 pub fn near_nullspace(mut self, basis: Vec<Vec<f64>>) -> Self {
1672 let count = basis.len().max(1);
1673 self.cfg.near_nullspace = Some(NearNullspace { basis });
1674 self.cfg.num_functions = count;
1675 self
1676 }
1677
1678 pub fn post_interp(mut self, t: PostInterpType) -> Self {
1679 self.cfg.post_interp = t;
1680 self
1681 }
1682
1683 pub fn flexible_level(mut self, level: usize) -> Self {
1684 self.cfg.flexible_level = Some(level);
1685 self
1686 }
1687
1688 pub fn flexible_iters(mut self, iters: usize) -> Self {
1689 self.cfg.flexible_iters = iters;
1690 self
1691 }
1692
1693 pub fn flexible_rtol(mut self, tol: f64) -> Self {
1694 self.cfg.flexible_rtol = tol;
1695 self
1696 }
1697
1698 pub fn flexible_pc_sweeps(mut self, sweeps: usize) -> Self {
1699 self.cfg.flexible_pc_sweeps = sweeps;
1700 self
1701 }
1702
1703 pub fn mixed_precision(mut self, mp: Option<MixedPrecision>) -> Self {
1704 self.cfg.mixed_precision = mp;
1705 self
1706 }
1707
1708 pub fn mixed_storage(mut self, storage: MixedStorage) -> Self {
1709 self.cfg.mixed_storage = storage;
1710 self
1711 }
1712
1713 pub fn require_native_complex_hierarchy(mut self, on: bool) -> Self {
1714 self.cfg.require_native_complex_hierarchy = on;
1715 self
1716 }
1717
1718 pub fn dist_coarse_strategy(mut self, strategy: DistCoarseStrategy) -> Self {
1719 self.cfg.dist_coarse_strategy = strategy;
1720 self
1721 }
1722
1723 pub fn dist_apply_instrumentation(mut self, on: bool) -> Self {
1724 self.cfg.dist_apply_instrumentation = on;
1725 self
1726 }
1727
1728 pub fn build(self, _matrix: &Mat<f64>) -> Result<AMG, KError> {
1729 Ok(AMG::with_config(self.cfg))
1730 }
1731}
1732
1733impl Default for AMGBuilder {
1734 fn default() -> Self {
1735 Self::new()
1736 }
1737}
1738
1739fn validate_relax_policy(cfg: &AMGConfig, coarse_solver: CoarseSolve) -> Result<(), KError> {
1742 let is_symmetric_relax = |rt: RelaxType| {
1743 matches!(
1744 rt,
1745 RelaxType::Jacobi
1746 | RelaxType::SymmetricGaussSeidel
1747 | RelaxType::L1Jacobi
1748 | RelaxType::Chebyshev
1749 | RelaxType::Fsai
1750 )
1751 };
1752
1753 if matches!(coarse_solver, CoarseSolve::DirectDense)
1754 && cfg.num_grid_sweeps[RelaxPhase::Coarsest.ix()] != 0
1755 {
1756 return Err(KError::InvalidInput(
1757 "num_grid_sweeps[Coarsest] must be 0 when coarse_solve is DirectDense".into(),
1758 ));
1759 }
1760
1761 for (i, &rt) in cfg.grid_relax_type.iter().enumerate() {
1762 match rt {
1763 RelaxType::Jacobi
1764 | RelaxType::GaussSeidel
1765 | RelaxType::GaussSeidelBackward
1766 | RelaxType::SymmetricGaussSeidel
1767 | RelaxType::L1Jacobi
1768 | RelaxType::Chebyshev
1769 | RelaxType::ChebyshevSafe
1770 | RelaxType::SafeguardedGaussSeidel
1771 | RelaxType::Ilu0
1772 | RelaxType::Ras
1773 | RelaxType::Fsai => {}
1774 _ => {
1775 return Err(KError::InvalidInput(format!(
1776 "RelaxType {rt:?} not yet supported (phase index {i})"
1777 )));
1778 }
1779 }
1780 }
1781
1782 if let Some(mp) = cfg.mixed_precision
1783 && mp.smoothers_enabled()
1784 {
1785 for (i, &rt) in cfg.grid_relax_type.iter().enumerate() {
1786 if i == RelaxPhase::Coarsest.ix() {
1787 continue;
1788 }
1789 match rt {
1790 RelaxType::Jacobi
1791 | RelaxType::L1Jacobi
1792 | RelaxType::Chebyshev
1793 | RelaxType::Fsai => {}
1794 other => {
1795 return Err(KError::InvalidInput(format!(
1796 "Mixed-precision smoothing only supports Jacobi, L1Jacobi, Chebyshev, or Fsai; got {other:?}"
1797 )));
1798 }
1799 }
1800 }
1801 }
1802
1803 for (i, &k) in cfg.num_grid_sweeps.iter().enumerate() {
1804 if i != RelaxPhase::Coarsest.ix() && k == 0 {
1805 return Err(KError::InvalidInput(format!(
1806 "num_grid_sweeps for phase {i} must be >= 1"
1807 )));
1808 }
1809 }
1810
1811 if cfg.flexible_iters > 0 {
1812 let level = cfg.flexible_level.ok_or_else(|| {
1813 KError::InvalidInput("flexible_iters > 0 requires flexible_level to be set".into())
1814 })?;
1815 if level >= cfg.max_levels {
1816 return Err(KError::InvalidInput(
1817 "flexible_level must be less than max_levels".into(),
1818 ));
1819 }
1820 if cfg.flexible_rtol < 0.0 {
1821 return Err(KError::InvalidInput(
1822 "flexible_rtol must be non-negative".into(),
1823 ));
1824 }
1825 }
1826
1827 if cfg.require_spd {
1828 if matches!(coarse_solver, CoarseSolve::ILU) {
1829 return Err(KError::InvalidInput(
1830 "SPD mode requires DirectDense or CG as the coarse solver; ILU is not SPD-safe"
1831 .into(),
1832 ));
1833 }
1834
1835 if cfg.non_galerkin.enabled && cfg.forbid_non_galerkin_in_spd {
1836 return Err(KError::InvalidInput(
1837 concat!(
1838 "Non-Galerkin filtering is disabled when require_spd is true ",
1839 "(set forbid_non_galerkin_in_spd = false to override)."
1840 )
1841 .into(),
1842 ));
1843 }
1844
1845 if !is_symmetric_relax(cfg.relax_type) {
1846 return Err(KError::InvalidInput(
1847 "SPD mode requires symmetric or Jacobi-type smoothers".into(),
1848 ));
1849 }
1850
1851 let down_sweeps = cfg.num_grid_sweeps[RelaxPhase::Down.ix()];
1852 let up_sweeps = cfg.num_grid_sweeps[RelaxPhase::Up.ix()];
1853 if down_sweeps != up_sweeps {
1854 return Err(KError::InvalidInput(
1855 "SPD mode requires symmetric pre/post smoothing counts".into(),
1856 ));
1857 }
1858 if cfg.pre_sweeps != cfg.post_sweeps {
1859 return Err(KError::InvalidInput(
1860 "SPD mode requires pre_sweeps == post_sweeps".into(),
1861 ));
1862 }
1863
1864 let down_type = cfg.grid_relax_type[RelaxPhase::Down.ix()];
1865 let up_type = cfg.grid_relax_type[RelaxPhase::Up.ix()];
1866 if down_type != up_type {
1867 return Err(KError::InvalidInput(
1868 "SPD mode requires matching relax types for Down and Up phases".into(),
1869 ));
1870 }
1871
1872 for phase in [RelaxPhase::Fine, RelaxPhase::Down, RelaxPhase::Up] {
1873 match cfg.grid_relax_type[phase.ix()] {
1874 RelaxType::GaussSeidelBackward
1875 | RelaxType::HybridGaussSeidel
1876 | RelaxType::SafeguardedGaussSeidel
1877 | RelaxType::Ilu0
1878 | RelaxType::Ras
1879 | RelaxType::ChebyshevSafe => {
1880 return Err(KError::InvalidInput(
1881 "SPD mode does not support asymmetric Gauss-Seidel variants".into(),
1882 ));
1883 }
1884 RelaxType::Jacobi
1885 | RelaxType::GaussSeidel
1886 | RelaxType::SymmetricGaussSeidel
1887 | RelaxType::L1Jacobi
1888 | RelaxType::Chebyshev
1889 | RelaxType::Fsai => {}
1890 }
1891 }
1892
1893 if cfg.flexible_iters > 0 {
1894 let level = cfg.flexible_level.expect("validated above");
1895 let phase = if level == 0 {
1896 RelaxPhase::Fine
1897 } else {
1898 RelaxPhase::Down
1899 };
1900 match cfg.grid_relax_type[phase.ix()] {
1901 RelaxType::Jacobi | RelaxType::L1Jacobi | RelaxType::SymmetricGaussSeidel => {}
1902 other => {
1903 return Err(KError::InvalidInput(format!(
1904 "SPD mode requires a symmetric positive definite smoother for flexible presmoothing; got {other:?}"
1905 )));
1906 }
1907 }
1908 }
1909 }
1910
1911 if cfg.non_galerkin.enabled && cfg.non_galerkin.symmetry == NgSymmetry::Symmetric {
1912 let down = cfg.grid_relax_type[RelaxPhase::Down.ix()];
1913 let up = cfg.grid_relax_type[RelaxPhase::Up.ix()];
1914 if down != up {
1915 return Err(KError::InvalidInput(
1916 "Symmetric non-Galerkin mode requires matching Down/Up relaxers".into(),
1917 ));
1918 }
1919 if cfg.num_grid_sweeps[RelaxPhase::Down.ix()] != cfg.num_grid_sweeps[RelaxPhase::Up.ix()] {
1920 return Err(KError::InvalidInput(
1921 "Symmetric non-Galerkin mode requires matching Down/Up sweep counts".into(),
1922 ));
1923 }
1924 if !is_symmetric_relax(down) {
1925 return Err(KError::InvalidInput(format!(
1926 "Symmetric non-Galerkin mode requires a symmetric smoother; got {down:?}"
1927 )));
1928 }
1929 }
1930
1931 Ok(())
1932}
1933
1934fn validate_truncation_and_caps(cfg: &AMGConfig) -> Result<(), KError> {
1935 if !(0.0..1.0).contains(&cfg.truncation_factor) {
1936 return Err(KError::InvalidInput(
1937 "truncation_factor must satisfy 0 ≤ τ_rel < 1".into(),
1938 ));
1939 }
1940 if cfg.interpolation_truncation < 0.0 || cfg.rap_truncation_abs < 0.0 {
1941 return Err(KError::InvalidInput(
1942 "absolute drop tolerances must be ≥ 0".into(),
1943 ));
1944 }
1945 Ok(())
1946}
1947
1948#[derive(Debug)]
1949struct AMGWorkspace {
1950 temp: Vec<R>,
1951 work: Vec<R>,
1952 residual: Vec<R>,
1953 coarse_rhs: Vec<R>,
1954 coarse_sol: Vec<Vec<R>>,
1955 fine_corr: Vec<R>,
1956 k_zeta: Vec<R>,
1957 k_p: Vec<R>,
1958 k_ap: Vec<R>,
1959 k_temp: Vec<R>,
1960 k_work: Vec<R>,
1961 k_residual: Vec<R>,
1962 mp: Option<MixedWs>,
1963}
1964
1965impl AMGWorkspace {
1966 fn new(cap: usize) -> Self {
1967 Self {
1968 temp: vec![R::zero(); cap],
1969 work: vec![R::zero(); cap],
1970 residual: vec![R::zero(); cap],
1971 coarse_rhs: vec![R::zero(); cap],
1972 coarse_sol: Vec::new(),
1973 fine_corr: vec![R::zero(); cap],
1974 k_zeta: vec![R::zero(); cap],
1975 k_p: vec![R::zero(); cap],
1976 k_ap: vec![R::zero(); cap],
1977 k_temp: vec![R::zero(); cap],
1978 k_work: vec![R::zero(); cap],
1979 k_residual: vec![R::zero(); cap],
1980 mp: None,
1981 }
1982 }
1983 fn ensure(&mut self, n: usize) {
1984 let grow = |v: &mut Vec<R>, n: usize| {
1985 if v.len() < n {
1986 v.resize(n, R::zero())
1987 }
1988 };
1989 grow(&mut self.temp, n);
1990 grow(&mut self.work, n);
1991 grow(&mut self.residual, n);
1992 grow(&mut self.coarse_rhs, n);
1993 grow(&mut self.fine_corr, n);
1994 grow(&mut self.k_zeta, n);
1995 grow(&mut self.k_p, n);
1996 grow(&mut self.k_ap, n);
1997 grow(&mut self.k_temp, n);
1998 grow(&mut self.k_work, n);
1999 grow(&mut self.k_residual, n);
2000 }
2001
2002 fn take_coarse_sol(&mut self, level: usize, n: usize) -> Vec<R> {
2003 if self.coarse_sol.len() <= level {
2004 self.coarse_sol.resize_with(level + 1, Vec::new);
2005 }
2006 let mut sol = std::mem::take(&mut self.coarse_sol[level]);
2007 sol.resize(n, R::zero());
2008 sol.fill(R::zero());
2009 sol
2010 }
2011
2012 fn put_coarse_sol(&mut self, level: usize, sol: Vec<R>) {
2013 debug_assert!(level < self.coarse_sol.len());
2014 self.coarse_sol[level] = sol;
2015 }
2016
2017 fn ensure_mixed(&mut self, n: usize) {
2018 if self.mp.is_none() {
2019 self.mp = Some(MixedWs::with_capacity(n));
2020 }
2021 if let Some(ref mut mp) = self.mp {
2022 mp.ensure_vectors(n);
2023 }
2024 }
2025}
2026
2027#[derive(Debug)]
2028struct MixedWs {
2029 temp32: Vec<f32>,
2030 work32: Vec<f32>,
2031 residual32: Vec<f32>,
2032 coarse_rhs32: Vec<f32>,
2033 fine_corr32: Vec<f32>,
2034 a_vals32: Vec<f32>,
2035 diag_inv32: Vec<f32>,
2036 l1_inv32: Vec<f32>,
2037 d_sqrt_inv32: Vec<f32>,
2038 fsai_g_vals32: Vec<f32>,
2039 fsai_gt_vals32: Vec<f32>,
2040}
2041
2042impl MixedWs {
2043 fn with_capacity(n: usize) -> Self {
2044 Self {
2045 temp32: vec![0.0; n],
2046 work32: vec![0.0; n],
2047 residual32: vec![0.0; n],
2048 coarse_rhs32: vec![0.0; n],
2049 fine_corr32: vec![0.0; n],
2050 a_vals32: Vec::new(),
2051 diag_inv32: Vec::new(),
2052 l1_inv32: Vec::new(),
2053 d_sqrt_inv32: Vec::new(),
2054 fsai_g_vals32: Vec::new(),
2055 fsai_gt_vals32: Vec::new(),
2056 }
2057 }
2058
2059 fn ensure_vectors(&mut self, n: usize) {
2060 let grow = |v: &mut Vec<f32>, n: usize| {
2061 if v.len() < n {
2062 v.resize(n, 0.0);
2063 }
2064 };
2065 grow(&mut self.temp32, n);
2066 grow(&mut self.work32, n);
2067 grow(&mut self.residual32, n);
2068 grow(&mut self.coarse_rhs32, n);
2069 grow(&mut self.fine_corr32, n);
2070 }
2071
2072 fn ensure_vals(&mut self, nnz: usize) -> &mut Vec<f32> {
2073 if self.a_vals32.len() < nnz {
2074 self.a_vals32.resize(nnz, 0.0);
2075 }
2076 &mut self.a_vals32
2077 }
2078
2079 fn ensure_diag(&mut self, n: usize) -> &mut Vec<f32> {
2080 if self.diag_inv32.len() < n {
2081 self.diag_inv32.resize(n, 0.0);
2082 }
2083 &mut self.diag_inv32
2084 }
2085
2086 fn ensure_l1(&mut self, n: usize) -> &mut Vec<f32> {
2087 if self.l1_inv32.len() < n {
2088 self.l1_inv32.resize(n, 0.0);
2089 }
2090 &mut self.l1_inv32
2091 }
2092
2093 fn ensure_d_sqrt(&mut self, n: usize) -> &mut Vec<f32> {
2094 if self.d_sqrt_inv32.len() < n {
2095 self.d_sqrt_inv32.resize(n, 0.0);
2096 }
2097 &mut self.d_sqrt_inv32
2098 }
2099
2100 fn ensure_fsai_g(&mut self, nnz: usize) -> &mut Vec<f32> {
2101 if self.fsai_g_vals32.len() < nnz {
2102 self.fsai_g_vals32.resize(nnz, 0.0);
2103 }
2104 &mut self.fsai_g_vals32
2105 }
2106
2107 fn ensure_fsai_gt(&mut self, nnz: usize) -> &mut Vec<f32> {
2108 if self.fsai_gt_vals32.len() < nnz {
2109 self.fsai_gt_vals32.resize(nnz, 0.0);
2110 }
2111 &mut self.fsai_gt_vals32
2112 }
2113}
2114
2115#[derive(Clone)]
2116struct FsaiData {
2117 g: CsrMatrix<f64>,
2118 gt: CsrMatrix<f64>,
2119 g2gt_pos: Vec<usize>,
2120}
2121
2122struct AMGLevel {
2123 a: CsrMatrix<f64>,
2125 p: CsrMatrix<f64>,
2127 r: CsrMatrix<f64>,
2129 diag_inv: Vec<f64>,
2131 d_sqrt_inv: Option<Vec<f64>>,
2133 l1_inv: Option<Vec<f64>>,
2135 diag_inv_safe: Option<Vec<f64>>,
2137 d_sqrt_inv_safe: Option<Vec<f64>>,
2139 cheb: Option<ChebData>,
2141 cheb_safe: Option<ChebData>,
2143 agg_of: Vec<usize>,
2145 is_c: Vec<bool>,
2147 cf: Option<CFInfo>,
2149 p2r_pos: Vec<usize>,
2151 num_functions: usize,
2153 row_basis: Option<Vec<f64>>,
2155 layout: Option<DofLayout>,
2157 nns: Option<Vec<Vec<f64>>>,
2159 a_next_pat: Option<CsrPattern>,
2161 a_next_pat_ng: Option<CsrPattern>,
2163 rap_full2ng_pos: Option<Vec<Option<usize>>>,
2165 r_row_ptr: Option<Vec<usize>>,
2167 r_col_idx: Option<Vec<usize>>,
2168 r_vals_scratch: Option<Vec<f64>>,
2169 #[allow(clippy::redundant_allocation)]
2170 coarse_solver: Option<Mutex<Box<dyn CoarseSolver + Send>>>,
2171 ilu0: Option<Mutex<IluCsr>>,
2172 ras: Option<Mutex<Asm>>,
2173 fsai: Option<FsaiData>,
2174 a_vals_f32: Option<Vec<f32>>,
2175 diag_inv_f32: Option<Vec<f32>>,
2176 d_sqrt_inv_f32: Option<Vec<f32>>,
2177 l1_inv_f32: Option<Vec<f32>>,
2178 fsai_g_vals_f32: Option<Vec<f32>>,
2179 fsai_gt_vals_f32: Option<Vec<f32>>,
2180}
2181
2182#[cfg(feature = "simd")]
2183fn build_level_spmv_plans(level: &mut AMGLevel, tuning: &SpmvTuning) {
2184 level.a.build_spmv_plan(tuning);
2185 level.p.build_spmv_plan(tuning);
2186 level.r.build_spmv_plan(tuning);
2187}
2188
2189#[derive(Clone)]
2190struct ChebData {
2191 lambda_max: f64,
2192 lambda_min: f64,
2193}
2194
2195#[derive(Clone)]
2196struct RelaxPolicy {
2197 kind: [RelaxType; 4],
2198 sweeps: [usize; 4],
2199 omega: f64,
2200}
2201
2202trait CyclePolicy: Send + Sync {
2203 fn gamma_visits(&self, level: usize) -> usize;
2204 fn k_presmooth(&self, _level: usize) -> Option<(KrylovAlgo, usize)> {
2205 None
2206 }
2207 fn k_postsmooth(&self, _level: usize) -> Option<(KrylovAlgo, usize)> {
2208 None
2209 }
2210}
2211
2212struct VPolicy;
2213
2214impl CyclePolicy for VPolicy {
2215 fn gamma_visits(&self, _level: usize) -> usize {
2216 1
2217 }
2218}
2219
2220struct WPolicy {
2221 gamma: usize,
2222}
2223
2224impl CyclePolicy for WPolicy {
2225 fn gamma_visits(&self, _level: usize) -> usize {
2226 self.gamma.max(1)
2227 }
2228}
2229
2230struct KPolicy {
2231 base: CycleType,
2232 cfg: KCycle,
2233}
2234
2235impl KPolicy {
2236 fn new(base: CycleType, mut cfg: KCycle) -> Self {
2237 cfg.levels.sort_unstable();
2238 cfg.levels.dedup();
2239 cfg.iters = cfg.iters.max(1);
2240 Self { base, cfg }
2241 }
2242
2243 fn contains(&self, level: usize) -> bool {
2244 self.cfg.levels.binary_search(&level).is_ok()
2245 }
2246}
2247
2248impl CyclePolicy for KPolicy {
2249 fn gamma_visits(&self, _level: usize) -> usize {
2250 match self.base {
2251 CycleType::V => 1,
2252 CycleType::W { gamma } => gamma.max(2),
2253 }
2254 .max(1)
2255 }
2256
2257 fn k_presmooth(&self, level: usize) -> Option<(KrylovAlgo, usize)> {
2258 if self.cfg.place_pre && self.contains(level) {
2259 Some((self.cfg.algo, self.cfg.iters))
2260 } else {
2261 None
2262 }
2263 }
2264
2265 fn k_postsmooth(&self, level: usize) -> Option<(KrylovAlgo, usize)> {
2266 if self.cfg.place_post && self.contains(level) {
2267 Some((self.cfg.algo, self.cfg.iters))
2268 } else {
2269 None
2270 }
2271 }
2272}
2273
2274struct AmgHierarchy {
2275 levels: Vec<AMGLevel>, policy: RelaxPolicy,
2277 coarse_solve: CoarseSolve,
2278}
2279
2280impl AmgHierarchy {
2281 fn finest(&self) -> &AMGLevel {
2282 &self.levels[0]
2283 }
2284 fn coarsest_ix(&self) -> usize {
2285 self.levels.len() - 1
2286 }
2287}
2288
2289fn build_r_from_p(lvl: &mut AMGLevel) -> CsrMatrix<f64> {
2290 let rr = lvl.r_row_ptr.as_ref().expect("missing R pattern");
2291 let rc = lvl.r_col_idx.as_ref().expect("missing R pattern");
2292 let p2r = &lvl.p2r_pos;
2293 let pvals = lvl.p.values();
2294 let rvals = lvl.r_vals_scratch.as_mut().expect("missing R scratch");
2295 sync_adjoint_values_from_forward(pvals, p2r, rvals);
2296 CsrMatrix::from_csr(
2297 lvl.p.ncols(),
2298 lvl.p.nrows(),
2299 rr.clone(),
2300 rc.clone(),
2301 rvals.clone(),
2302 )
2303}
2304
2305fn rap_numeric_with_pt(lvl: &mut AMGLevel, out_vals: &mut [f64]) -> Result<(), KError> {
2306 let r_tmp = build_r_from_p(lvl);
2307 let pat = lvl.a_next_pat.as_ref().expect("missing pattern").clone();
2308 rap_numeric(&pat, &r_tmp, &lvl.a, &lvl.p, out_vals);
2309 Ok(())
2310}
2311
2312fn sync_adjoint_values_from_forward<T: KrystScalar>(
2313 forward_values: &[T],
2314 forward_to_adjoint_pos: &[usize],
2315 adjoint_values: &mut [T],
2316) {
2317 debug_assert_eq!(
2318 forward_values.len(),
2319 forward_to_adjoint_pos.len(),
2320 "forward/adjoin position length mismatch"
2321 );
2322 for (pi, &ri) in forward_to_adjoint_pos.iter().enumerate() {
2323 debug_assert!(ri < adjoint_values.len(), "adjoint position out of bounds");
2324 adjoint_values[ri] = forward_values[pi].conj();
2325 }
2326}
2327
2328fn make_trial_matrix(cfg: &AMGConfig, n: usize) -> Result<Option<Mat<f64>>, KError> {
2329 if cfg.filter_omega <= 0.0 {
2330 return Ok(None);
2331 }
2332 if cfg.require_spd && cfg.filter_trial_vectors.is_none() {
2333 return Ok(None);
2334 }
2335 let mat = if let Some(basis) = cfg.filter_trial_vectors.as_ref() {
2336 if basis.is_empty() {
2337 return Ok(None);
2338 }
2339 let r = basis.len();
2340 let mut m = Mat::<f64>::zeros(n, r);
2341 for (col, vec) in basis.iter().enumerate() {
2342 if vec.len() != n {
2343 return Err(KError::InvalidInput(
2344 "filter_trial_vectors entry has mismatched length".into(),
2345 ));
2346 }
2347 for i in 0..n {
2348 m[(i, col)] = vec[i];
2349 }
2350 }
2351 m
2352 } else {
2353 let mut m = Mat::<f64>::zeros(n, 1);
2354 for i in 0..n {
2355 m[(i, 0)] = 1.0;
2356 }
2357 m
2358 };
2359 Ok(Some(mat))
2360}
2361
2362fn apply_trial_compensation(
2363 cfg: &AMGConfig,
2364 a: &mut CsrMatrix<f64>,
2365 trials: Option<&Mat<f64>>,
2366 block_size: usize,
2367) -> Result<(), KError> {
2368 if cfg.filter_omega <= 0.0 {
2369 return Ok(());
2370 }
2371 let Some(trials_mat) = trials else {
2372 return Ok(());
2373 };
2374 let min_diag = if cfg.require_spd {
2375 Some(cfg.spd_diag_floor.max(1e-12))
2376 } else {
2377 None
2378 };
2379 if block_size > 1 && matches!(cfg.nodal, NodalMode::Nodal) && a.nrows() % block_size == 0 {
2380 compensate_nodal_diag(
2381 a,
2382 trials_mat.as_ref(),
2383 block_size,
2384 cfg.filter_omega,
2385 min_diag,
2386 )
2387 } else {
2388 compensate_scalar_rows(a, trials_mat.as_ref(), cfg.filter_omega, min_diag)
2389 }
2390}
2391
2392struct LevelPostContext<'a> {
2393 r: usize,
2394 agg_of: &'a [usize],
2395 nns: Option<Vec<&'a [f64]>>,
2396 a: Option<&'a CsrMatrix<f64>>,
2397 d_inv: Option<&'a [f64]>,
2398}
2399
2400pub(crate) fn row_scaling<T: KrystScalar<Real = f64>>(
2401 mode: RowScaleMode,
2402 r: usize,
2403 nns: Option<&[&[f64]]>,
2404 agg_of: &[usize],
2405 d_inv: Option<&[f64]>,
2406 p_row_ptr: &[usize],
2407 p_col_idx: &[usize],
2408 p_vals: &mut [T],
2409) -> Result<(), KError> {
2410 let n = p_row_ptr.len() - 1;
2411 let eps = 1e-30;
2412 for i in 0..n {
2413 let rs = p_row_ptr[i];
2414 let re = p_row_ptr[i + 1];
2415 match mode {
2416 RowScaleMode::SumToOne => {
2417 let sum = p_vals[rs..re]
2418 .iter()
2419 .copied()
2420 .fold(T::zero(), |acc, v| acc + v);
2421 if sum.abs() > eps {
2422 let s = T::one() / sum;
2423 for k in rs..re {
2424 p_vals[k] = p_vals[k] * s;
2425 }
2426 }
2427 }
2428 RowScaleMode::L2Unit => {
2429 for alpha in 0..r {
2430 let mut n2 = 0.0;
2431 for k in rs..re {
2432 if p_col_idx[k] % r == alpha {
2433 n2 += p_vals[k].abs2();
2434 }
2435 }
2436 if n2 > eps {
2437 let s = T::from_real(1.0 / n2.sqrt());
2438 for k in rs..re {
2439 if p_col_idx[k] % r == alpha {
2440 p_vals[k] = p_vals[k] * s;
2441 }
2442 }
2443 }
2444 }
2445 }
2446 RowScaleMode::DUnit => {
2447 let d = d_inv.expect("DUnit requires diag_inv");
2448 for alpha in 0..r {
2449 let mut n2 = 0.0;
2450 for k in rs..re {
2451 if p_col_idx[k] % r == alpha {
2452 n2 += p_vals[k].abs2();
2453 }
2454 }
2455 let w = d[i].abs().recip().sqrt().max(1e-15);
2456 if n2 > eps {
2457 let s = T::from_real(1.0 / (w * n2.sqrt()));
2458 for k in rs..re {
2459 if p_col_idx[k] % r == alpha {
2460 p_vals[k] = p_vals[k] * s;
2461 }
2462 }
2463 }
2464 }
2465 }
2466 RowScaleMode::ToNearNullspace => {
2467 let t = nns.expect("ToNearNullspace requires NNS basis");
2468 for alpha in 0..r {
2469 let target = T::from_real(t[alpha][i]);
2470 let mut sum = T::zero();
2471 for k in rs..re {
2472 if p_col_idx[k] % r == alpha {
2473 sum = sum + p_vals[k];
2474 }
2475 }
2476 if sum.abs() > eps {
2477 let s = target / sum;
2478 for k in rs..re {
2479 if p_col_idx[k] % r == alpha {
2480 p_vals[k] = p_vals[k] * s;
2481 }
2482 }
2483 } else {
2484 let own_c = agg_of[i] * r + alpha;
2485 for k in rs..re {
2486 if p_col_idx[k] == own_c {
2487 p_vals[k] = target;
2488 break;
2489 }
2490 }
2491 }
2492 }
2493 }
2494 }
2495 }
2496 Ok(())
2497}
2498
2499fn local_qr<T: KrystScalar<Real = f64>>(
2500 r: usize,
2501 agg_of: &[usize],
2502 p_row_ptr: &[usize],
2503 p_col_idx: &[usize],
2504 p_vals: &mut [T],
2505) -> Result<(), KError> {
2506 let n = agg_of.len();
2507 let n_aggs = 1 + agg_of.iter().copied().max().unwrap_or(0);
2508 let mut rows_in_agg: Vec<Vec<usize>> = vec![Vec::new(); n_aggs];
2509 for i in 0..n {
2510 rows_in_agg[agg_of[i]].push(i);
2511 }
2512 let mut pos_alpha: Vec<Vec<usize>> = vec![vec![usize::MAX; r]; n];
2513 for i in 0..n {
2514 let g = agg_of[i];
2515 let rs = p_row_ptr[i];
2516 let re = p_row_ptr[i + 1];
2517 for k in rs..re {
2518 let col = p_col_idx[k];
2519 if col / r == g {
2520 pos_alpha[i][col % r] = k;
2521 }
2522 }
2523 }
2524 for g in 0..n_aggs {
2525 let rows = &rows_in_agg[g];
2526 if rows.is_empty() {
2527 continue;
2528 }
2529 let m = rows.len();
2530 let mut q = vec![vec![T::zero(); r]; m];
2531 for (ii, &i) in rows.iter().enumerate() {
2532 for alpha in 0..r {
2533 let k = pos_alpha[i][alpha];
2534 if k != usize::MAX {
2535 q[ii][alpha] = p_vals[k];
2536 }
2537 }
2538 }
2539 for alpha in 0..r {
2540 for beta in 0..alpha {
2541 let mut dot = T::zero();
2542 for ii in 0..m {
2543 dot = dot + q[ii][beta].conj() * q[ii][alpha];
2544 }
2545 for ii in 0..m {
2546 q[ii][alpha] = q[ii][alpha] - q[ii][beta] * dot;
2547 }
2548 }
2549 let mut n2 = 0.0;
2550 for ii in 0..m {
2551 n2 += q[ii][alpha].abs2();
2552 }
2553 if n2 > 1e-30 {
2554 let inv = T::from_real(1.0 / n2.sqrt());
2555 for ii in 0..m {
2556 q[ii][alpha] = q[ii][alpha] * inv;
2557 }
2558 }
2559 }
2560 for (ii, &i) in rows.iter().enumerate() {
2561 for alpha in 0..r {
2562 let k = pos_alpha[i][alpha];
2563 if k != usize::MAX {
2564 p_vals[k] = q[ii][alpha];
2565 }
2566 }
2567 }
2568 }
2569 Ok(())
2570}
2571
2572fn energy_polish(
2573 a: &CsrMatrix<f64>,
2574 d_inv: &[f64],
2575 p_row_ptr: &[usize],
2576 p_col_idx: &[usize],
2577 p_vals: &mut [f64],
2578 sweeps: usize,
2579 omega: f64,
2580) -> Result<(), KError> {
2581 let n = a.nrows();
2582 for _ in 0..sweeps {
2583 let old = p_vals.to_vec();
2584 for i in 0..n {
2585 let di = d_inv[i];
2586 let rs = p_row_ptr[i];
2587 let re = p_row_ptr[i + 1];
2588 for k in rs..re {
2589 let c = p_col_idx[k];
2590 let mut sum = 0.0;
2591 let ars = a.row_ptr()[i];
2592 let are = a.row_ptr()[i + 1];
2593 for ap in ars..are {
2594 let j = a.col_idx()[ap];
2595 let aij = a.values()[ap];
2596 let prs = p_row_ptr[j];
2597 let pre = p_row_ptr[j + 1];
2598 for pk in prs..pre {
2599 if p_col_idx[pk] == c {
2600 sum += aij * old[pk];
2601 break;
2602 }
2603 }
2604 }
2605 p_vals[k] = old[k] - omega * di * sum;
2606 }
2607 }
2608 }
2609 Ok(())
2610}
2611
2612fn apply_post_interp(
2613 cfg: &AMGConfig,
2614 ctx: &LevelPostContext,
2615 p_row_ptr: &[usize],
2616 p_col_idx: &[usize],
2617 p_vals: &mut [f64],
2618) -> Result<(), KError> {
2619 match cfg.post_interp {
2620 PostInterpType::None => Ok(()),
2621 PostInterpType::RowScaling(mode) => {
2622 if matches!(mode, RowScaleMode::SumToOne) && ctx.r > 1 && ctx.nns.is_none() {
2623 return Ok(());
2624 }
2625 row_scaling(
2626 mode,
2627 ctx.r,
2628 ctx.nns.as_deref(),
2629 ctx.agg_of,
2630 ctx.d_inv,
2631 p_row_ptr,
2632 p_col_idx,
2633 p_vals,
2634 )
2635 }
2636 PostInterpType::LocalQR => local_qr(ctx.r, ctx.agg_of, p_row_ptr, p_col_idx, p_vals),
2637 PostInterpType::EnergyPolish { sweeps, omega } => {
2638 if sweeps == 0 {
2639 return Ok(());
2640 }
2641 energy_polish(
2642 ctx.a.expect("EnergyPolish requires A"),
2643 ctx.d_inv.expect("EnergyPolish requires diag_inv"),
2644 p_row_ptr,
2645 p_col_idx,
2646 p_vals,
2647 sweeps,
2648 omega,
2649 )
2650 }
2651 }
2652}
2653
2654#[derive(Clone, Debug, Default)]
2655struct RankDiagnostics {
2656 min_col_norm: f64,
2657 cond_estimate: f64,
2658 suspect: bool,
2659 degenerate_cols: Vec<usize>,
2660}
2661
2662#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2663enum RankFixOutcome {
2664 Fixed,
2665 Unfixed,
2666}
2667
2668fn p_column_norms2<T: KrystScalar<Real = f64>>(p: &CsrMatrix<T>) -> Vec<f64> {
2669 let mut n2 = vec![0.0; p.ncols()];
2670 let rp = p.row_ptr();
2671 let cj = p.col_idx();
2672 let vv = p.values();
2673 for i in 0..p.nrows() {
2674 let (rs, re) = (rp[i], rp[i + 1]);
2675 for k in rs..re {
2676 let j = cj[k];
2677 n2[j] += vv[k].abs2();
2678 }
2679 }
2680 n2
2681}
2682
2683fn symmetric_eigenvalues(mut m: Mat<f64>) -> Vec<f64> {
2684 let n = m.nrows();
2685 if n == 0 {
2686 return Vec::new();
2687 }
2688 let mut iter = 0usize;
2689 loop {
2690 let mut max_val = 0.0f64;
2691 let mut p = 0usize;
2692 let mut q = 1usize;
2693 for i in 0..n {
2694 for j in (i + 1)..n {
2695 let val = m[(i, j)].abs();
2696 if val > max_val {
2697 max_val = val;
2698 p = i;
2699 q = j;
2700 }
2701 }
2702 }
2703 if max_val < 1e-12 || iter > 64 * n * n {
2704 break;
2705 }
2706 iter += 1;
2707 let app = m[(p, p)];
2708 let aqq = m[(q, q)];
2709 let apq = m[(p, q)];
2710 if apq.abs() < 1e-30 {
2711 continue;
2712 }
2713 let tau = (aqq - app) / (2.0 * apq);
2714 let t = if tau >= 0.0_f64 {
2715 1.0_f64 / (tau + (1.0_f64 + tau * tau).sqrt())
2716 } else {
2717 -1.0_f64 / (-tau + (1.0_f64 + tau * tau).sqrt())
2718 };
2719 let c = 1.0_f64 / (1.0_f64 + t * t).sqrt();
2720 let s = t * c;
2721 for k in 0..n {
2722 if k != p && k != q {
2723 let mkp = m[(k, p)];
2724 let mkq = m[(k, q)];
2725 let new_kp = mkp * c - mkq * s;
2726 let new_kq = mkp * s + mkq * c;
2727 m[(k, p)] = new_kp;
2728 m[(p, k)] = new_kp;
2729 m[(k, q)] = new_kq;
2730 m[(q, k)] = new_kq;
2731 }
2732 }
2733 let app_new = app * c * c - 2.0 * apq * s * c + aqq * s * s;
2734 let aqq_new = app * s * s + 2.0 * apq * s * c + aqq * c * c;
2735 m[(p, p)] = app_new;
2736 m[(q, q)] = aqq_new;
2737 m[(p, q)] = 0.0;
2738 m[(q, p)] = 0.0;
2739 }
2740 (0..n).map(|i| m[(i, i)]).collect()
2741}
2742
2743fn rank_condition_estimate(p: &CsrMatrix<f64>, s: usize, seed: u64) -> Result<(bool, f64), KError> {
2744 let nc = p.ncols();
2745 if nc == 0 {
2746 return Ok((true, 1.0));
2747 }
2748 let n = p.nrows();
2749 let s = s.max(1).min(nc);
2750 let mut w = Mat::<f64>::zeros(n, s);
2751 let mut x = vec![0.0f64; nc];
2752 let mut y = vec![0.0f64; n];
2753 let mut omega_cols: Vec<Vec<i8>> = Vec::with_capacity(s);
2754 for col in 0..s {
2755 let mut col_vals = vec![0i8; nc];
2756 for i in 0..nc {
2757 let mut h = seed
2758 ^ (i as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15)
2759 ^ (col as u64).wrapping_mul(0xD00D_F00D_F00D_F00D);
2760 h ^= h >> 30;
2761 h = h.wrapping_mul(0xBF58_476D_1CE4_E5B9);
2762 h ^= h >> 27;
2763 h = h.wrapping_mul(0x94D0_49BB_1331_11EB);
2764 h ^= h >> 31;
2765 col_vals[i] = if h & 1 == 0 { 1 } else { -1 };
2766 }
2767 if nc > 0 {
2768 'adjust: loop {
2769 let mut tweaked = false;
2770 for prev in 0..omega_cols.len() {
2771 let mut same = true;
2772 let mut neg = true;
2773 for i in 0..nc {
2774 let v = col_vals[i];
2775 let pv = omega_cols[prev][i];
2776 if v != pv {
2777 same = false;
2778 }
2779 if v != -pv {
2780 neg = false;
2781 }
2782 if !same && !neg {
2783 break;
2784 }
2785 }
2786 if same || neg {
2787 let idx = (col + prev) % nc;
2788 col_vals[idx] = -col_vals[idx];
2789 tweaked = true;
2790 break;
2791 }
2792 }
2793 if !tweaked {
2794 break 'adjust;
2795 }
2796 }
2797 }
2798 for i in 0..nc {
2799 x[i] = col_vals[i] as f64;
2800 }
2801 p.spmv_scaled(1.0, &x, 0.0, &mut y)?;
2802 for i in 0..n {
2803 w[(i, col)] = y[i];
2804 }
2805 omega_cols.push(col_vals);
2806 }
2807 let mut s_mat = Mat::<f64>::zeros(s, s);
2808 for i in 0..s {
2809 for j in i..s {
2810 let mut dot = 0.0f64;
2811 for k in 0..n {
2812 dot += w[(k, i)] * w[(k, j)];
2813 }
2814 s_mat[(i, j)] = dot;
2815 s_mat[(j, i)] = dot;
2816 }
2817 }
2818 let mut lam = symmetric_eigenvalues(s_mat);
2819 lam.sort_by(|a, b| a.partial_cmp(b).unwrap_or(CmpOrdering::Equal));
2820 let lam_min = lam.first().copied().unwrap_or(0.0).max(0.0);
2821 let lam_max = lam.last().copied().unwrap_or(0.0).max(0.0);
2822 let cond = if lam_min > 0.0 {
2823 lam_max / lam_min
2824 } else if lam_max == 0.0 {
2825 1.0
2826 } else {
2827 f64::INFINITY
2828 };
2829 let ok = lam_min.is_finite() && lam_max.is_finite() && cond.is_finite();
2830 Ok((ok, cond))
2831}
2832
2833fn check_p_rank_fast(p: &CsrMatrix<f64>, cfg: &AMGConfig) -> Result<RankDiagnostics, KError> {
2834 if p.ncols() == 0 {
2835 return Ok(RankDiagnostics::default());
2836 }
2837 let norms2 = p_column_norms2(p);
2838 let mut diag = RankDiagnostics::default();
2839 let mut min_norm = f64::MAX;
2840 for (j, &n2) in norms2.iter().enumerate() {
2841 let norm = n2.sqrt();
2842 if norm < cfg.rank_min_col_norm {
2843 diag.degenerate_cols.push(j);
2844 }
2845 if norm < min_norm {
2846 min_norm = norm;
2847 }
2848 }
2849 diag.min_col_norm = if min_norm.is_finite() { min_norm } else { 0.0 };
2850 let (ok, cond) = rank_condition_estimate(p, cfg.rank_sketch_cols, 0x00C0_FFEE_u64)?;
2851 diag.cond_estimate = cond;
2852 let cond_bad = !ok || !cond.is_finite() || cond > cfg.rank_cond_threshold;
2853 diag.suspect = !diag.degenerate_cols.is_empty() || cond_bad;
2854 Ok(diag)
2855}
2856
2857fn try_fix_rank(
2858 _level_idx: usize,
2859 a_l: &CsrMatrix<f64>,
2860 diag_inv_l: &[f64],
2861 tp: &TentativeP,
2862 ctx: &LevelPostContext,
2863 p_csr: &mut Pcsr,
2864 cfg: &mut AMGConfig,
2865) -> Result<RankFixOutcome, KError> {
2866 let old_drop = cfg.interpolation_truncation;
2867 let old_cap = cfg.max_elements_per_row;
2868 cfg.interpolation_truncation = (old_drop * 0.1).min(old_drop);
2869 if old_cap == 0 {
2870 cfg.max_elements_per_row = 16;
2871 } else {
2872 cfg.max_elements_per_row = old_cap.max(16);
2873 }
2874
2875 let mut new_vals = vec![0.0; p_csr.col_idx.len()];
2876 if tp.num_functions > 1 {
2877 smooth_sa_values_only_multi(
2878 a_l,
2879 diag_inv_l,
2880 tp,
2881 cfg.jacobi_omega,
2882 &p_csr.row_ptr,
2883 &p_csr.col_idx,
2884 &mut new_vals,
2885 )?;
2886 } else {
2887 smooth_sa_values_only(
2888 a_l,
2889 diag_inv_l,
2890 tp,
2891 cfg.jacobi_omega,
2892 &p_csr.row_ptr,
2893 &p_csr.col_idx,
2894 &mut new_vals,
2895 )?;
2896 }
2897 p_csr.vals.copy_from_slice(&new_vals);
2898 apply_post_interp(cfg, ctx, &p_csr.row_ptr, &p_csr.col_idx, &mut p_csr.vals)?;
2899 let p_tmp = CsrMatrix::from_csr(
2900 p_csr.m,
2901 p_csr.n,
2902 p_csr.row_ptr.clone(),
2903 p_csr.col_idx.clone(),
2904 p_csr.vals.clone(),
2905 );
2906 let diag = check_p_rank_fast(&p_tmp, cfg)?;
2907
2908 cfg.interpolation_truncation = old_drop;
2909 cfg.max_elements_per_row = old_cap;
2910
2911 if diag.suspect {
2912 Ok(RankFixOutcome::Unfixed)
2913 } else {
2914 Ok(RankFixOutcome::Fixed)
2915 }
2916}
2917
2918fn galerkin_sample_check(
2919 a_l: &CsrMatrix<f64>,
2920 p_l: &CsrMatrix<f64>,
2921 r_l: &CsrMatrix<f64>,
2922 a_c: &CsrMatrix<f64>,
2923 samples: usize,
2924 tol: f64,
2925 seed: u64,
2926) -> Result<(bool, f64), KError> {
2927 let n = a_l.nrows();
2928 let nc = a_c.nrows();
2929 if nc == 0 {
2930 return Ok((true, 0.0));
2931 }
2932 let s = samples.max(1);
2933 let mut worst = 0.0f64;
2934 let mut y = vec![0.0f64; nc];
2935 let mut x = vec![0.0f64; n];
2936 let mut ax = vec![0.0f64; n];
2937 let mut u = vec![0.0f64; nc];
2938 let mut v = vec![0.0f64; nc];
2939 for t in 0..s {
2940 for i in 0..nc {
2941 let h = seed ^ ((t as u64).wrapping_mul(0x9E37) ^ (i as u64).wrapping_mul(0xD00D));
2942 y[i] = if (h >> 1) & 1 == 0 { 1.0 } else { -1.0 };
2943 }
2944 p_l.spmv_scaled(1.0, &y, 0.0, &mut x)?;
2945 a_l.spmv_scaled(1.0, &x, 0.0, &mut ax)?;
2946 r_l.spmv_scaled(1.0, &ax, 0.0, &mut u)?;
2947 a_c.spmv_scaled(1.0, &y, 0.0, &mut v)?;
2948 let mut num = 0.0f64;
2949 let mut den = 0.0f64;
2950 for i in 0..nc {
2951 let d = u[i] - v[i];
2952 num += d * d;
2953 den += v[i] * v[i];
2954 }
2955 let rel = num.sqrt() / den.sqrt().max(1e-300);
2956 if rel > worst {
2957 worst = rel;
2958 }
2959 }
2960 Ok((worst <= tol, worst))
2961}
2962
2963fn csr_pattern_hash<T: KrystScalar>(a: &CsrMatrix<T>) -> u64 {
2964 let mut hasher = DefaultHasher::new();
2965 a.row_ptr().hash(&mut hasher);
2966 a.col_idx().hash(&mut hasher);
2967 hasher.finish()
2968}
2969
2970#[cfg(feature = "complex")]
2971fn csr_from_dense_complex(mat: &Mat<S>, drop_tol: R) -> CsrMatrix<S> {
2972 let m = mat.nrows();
2973 let n = mat.ncols();
2974 let mut row_ptr = Vec::with_capacity(m + 1);
2975 let mut col_idx = Vec::new();
2976 let mut vals = Vec::new();
2977 row_ptr.push(0);
2978 for i in 0..m {
2979 for j in 0..n {
2980 let v = mat[(i, j)];
2981 if v.abs() > drop_tol {
2982 col_idx.push(j);
2983 vals.push(v);
2984 }
2985 }
2986 row_ptr.push(col_idx.len());
2987 }
2988 CsrMatrix::from_csr(m, n, row_ptr, col_idx, vals)
2989}
2990
2991#[cfg(feature = "complex")]
2992fn csr_from_linop_complex(op: &dyn LinOp<S = S>, drop_tol: R) -> Result<Arc<CsrMatrix<S>>, KError> {
2993 if let Some(csr) = op.as_any().downcast_ref::<CsrMatrix<S>>() {
2994 return Ok(Arc::new(csr.clone()));
2995 }
2996 #[cfg(feature = "backend-faer")]
2997 if let Some(csr_op) = op.as_any().downcast_ref::<CsrOp<S>>() {
2998 return Ok(Arc::new(csr_op.inner().clone()));
2999 }
3000 #[cfg(feature = "backend-faer")]
3001 if let Some(generic) = op.as_any().downcast_ref::<GenericCsrOp<S>>() {
3002 let mat = generic.matrix();
3003 return Ok(Arc::new(CsrMatrix::from_csr(
3004 mat.nrows,
3005 mat.ncols,
3006 mat.rowptr.clone(),
3007 mat.colind.clone(),
3008 mat.values.clone(),
3009 )));
3010 }
3011 if let Some(mat) = op.as_any().downcast_ref::<Mat<S>>() {
3012 return Ok(Arc::new(csr_from_dense_complex(mat, drop_tol)));
3013 }
3014 #[cfg(feature = "backend-faer")]
3015 if let Some(dense_op) = op.as_any().downcast_ref::<DenseOp<S>>() {
3016 return Ok(Arc::new(csr_from_dense_complex(dense_op.inner(), drop_tol)));
3017 }
3018 Err(KError::InvalidInput(
3019 "AMG: unsupported complex LinOp; provide CSR/Dense or GenericCsrOp".into(),
3020 ))
3021}
3022
3023#[cfg(feature = "complex")]
3024fn csr_real_metadata_from_complex(csr: &CsrMatrix<S>) -> CsrMatrix<f64> {
3028 let values = csr.values().iter().map(|v| v.real()).collect();
3029 CsrMatrix::from_csr(
3030 csr.nrows(),
3031 csr.ncols(),
3032 csr.row_ptr().to_vec(),
3033 csr.col_idx().to_vec(),
3034 values,
3035 )
3036}
3037
3038#[cfg(feature = "complex")]
3039fn complex_single_level_stats(csr: &CsrMatrix<S>, cfg: &AMGConfig) -> AmgStats {
3040 AmgStats::from_scalar_core(std::iter::once((csr, 0, 0)), cfg)
3041}
3042
3043#[cfg(feature = "complex")]
3044fn complex_diagonal_inverse(csr: &CsrMatrix<S>, drop_tol: R) -> Result<Option<Vec<S>>, KError> {
3045 if csr.nrows() != csr.ncols() {
3046 return Ok(None);
3047 }
3048 let n = csr.nrows();
3049 let mut diag = vec![S::zero(); n];
3050 let mut seen = vec![false; n];
3051 for i in 0..n {
3052 for p in csr.row_ptr()[i]..csr.row_ptr()[i + 1] {
3053 let j = csr.col_idx()[p];
3054 let v = csr.values()[p];
3055 if j == i {
3056 diag[i] = v;
3057 seen[i] = true;
3058 } else if v.abs() > drop_tol {
3059 return Ok(None);
3060 }
3061 }
3062 }
3063 let mut inv = Vec::with_capacity(n);
3064 for i in 0..n {
3065 if !seen[i] {
3066 return Ok(None);
3067 }
3068 if diag[i].abs() <= drop_tol {
3069 return Err(KError::InvalidInput(format!(
3070 "AMG: zero diagonal in complex diagonal fast path at row {i}"
3071 )));
3072 }
3073 inv.push(diag[i].inv());
3074 }
3075 Ok(Some(inv))
3076}
3077
3078#[cfg(not(feature = "complex"))]
3079fn pack_message_u64(message: &str) -> (u64, Vec<u64>) {
3080 let bytes = message.as_bytes();
3081 let len = bytes.len();
3082 if len == 0 {
3083 return (0, Vec::new());
3084 }
3085 let words = (len + 7) / 8;
3086 let mut data = vec![0u64; words];
3087 for (idx, &b) in bytes.iter().enumerate() {
3088 let word = idx / 8;
3089 let shift = (idx % 8) * 8;
3090 data[word] |= (b as u64) << shift;
3091 }
3092 (len as u64, data)
3093}
3094
3095#[cfg(not(feature = "complex"))]
3096fn unpack_message_u64(words: &[u64], len: usize) -> String {
3097 if len == 0 {
3098 return String::new();
3099 }
3100 let mut bytes = Vec::with_capacity(len);
3101 for (i, &word) in words.iter().enumerate() {
3102 let base = i * 8;
3103 for j in 0..8 {
3104 let idx = base + j;
3105 if idx >= len {
3106 break;
3107 }
3108 bytes.push(((word >> (j * 8)) & 0xFF) as u8);
3109 }
3110 }
3111 String::from_utf8_lossy(&bytes).to_string()
3112}
3113
3114#[cfg(not(feature = "complex"))]
3115fn broadcast_message<C: Comm>(comm: &C, root: usize, message: Option<String>) -> String {
3116 let rank = comm.rank();
3117 let size = comm.size();
3118 if size <= 1 {
3119 return message.unwrap_or_default();
3120 }
3121 if rank == root {
3122 let msg = message.unwrap_or_default();
3123 let (len, data) = pack_message_u64(&msg);
3124 let len_buf = [len];
3125 let mut reqs = Vec::new();
3126 for r in 0..size {
3127 if r == root {
3128 continue;
3129 }
3130 reqs.push(comm.isend_to_u64(&len_buf, r as i32));
3131 }
3132 comm.wait_all(&mut reqs);
3133 if len > 0 {
3134 let mut data_reqs = Vec::new();
3135 for r in 0..size {
3136 if r == root {
3137 continue;
3138 }
3139 data_reqs.push(comm.isend_to_u64(&data, r as i32));
3140 }
3141 comm.wait_all(&mut data_reqs);
3142 }
3143 msg
3144 } else {
3145 let mut len_buf = [0u64];
3146 {
3147 let mut reqs = vec![comm.irecv_from_u64(&mut len_buf, root as i32)];
3148 comm.wait_all(&mut reqs);
3149 }
3150 let len = len_buf[0] as usize;
3151 if len == 0 {
3152 return String::new();
3153 }
3154 let words = (len + 7) / 8;
3155 let mut data = vec![0u64; words];
3156 {
3157 let mut data_reqs = vec![comm.irecv_from_u64(&mut data, root as i32)];
3158 comm.wait_all(&mut data_reqs);
3159 }
3160 unpack_message_u64(&data, len)
3161 }
3162}
3163
3164#[cfg(not(feature = "complex"))]
3165fn collect_error_message<C: Comm>(comm: &C, root: usize, local_message: Option<String>) -> String {
3166 let rank = comm.rank();
3167 let size = comm.size();
3168 let local = local_message.unwrap_or_default();
3169 if size <= 1 {
3170 return local;
3171 }
3172 let (len, data) = pack_message_u64(&local);
3173 if rank != root {
3174 let len_buf = [len];
3175 let mut reqs = vec![comm.isend_to_u64(&len_buf, root as i32)];
3176 comm.wait_all(&mut reqs);
3177 if len > 0 {
3178 let mut data_reqs = vec![comm.isend_to_u64(&data, root as i32)];
3179 comm.wait_all(&mut data_reqs);
3180 }
3181 return broadcast_message(comm, root, None);
3182 }
3183
3184 let mut len_bufs = vec![[0u64; 1]; size];
3185 let mut len_reqs = Vec::new();
3186 for r in 0..size {
3187 if r == root {
3188 continue;
3189 }
3190 let buf = unsafe { &mut *len_bufs.as_mut_ptr().add(r) };
3191 len_reqs.push(comm.irecv_from_u64(buf, r as i32));
3192 }
3193 comm.wait_all(&mut len_reqs);
3194
3195 let mut data_bufs: Vec<Vec<u64>> = vec![Vec::new(); size];
3196 {
3197 let mut data_reqs = Vec::new();
3198 for r in 0..size {
3199 if r == root {
3200 continue;
3201 }
3202 let msg_len = len_bufs[r][0] as usize;
3203 if msg_len == 0 {
3204 continue;
3205 }
3206 let words = (msg_len + 7) / 8;
3207 let buf = unsafe { &mut *data_bufs.as_mut_ptr().add(r) };
3208 *buf = vec![0u64; words];
3209 data_reqs.push(comm.irecv_from_u64(buf.as_mut_slice(), r as i32));
3210 }
3211 comm.wait_all(&mut data_reqs);
3212 }
3213
3214 let mut messages = vec![String::new(); size];
3215 messages[root] = local;
3216 for r in 0..size {
3217 if r == root {
3218 continue;
3219 }
3220 let msg_len = len_bufs[r][0] as usize;
3221 if msg_len == 0 {
3222 continue;
3223 }
3224 messages[r] = unpack_message_u64(&data_bufs[r], msg_len);
3225 }
3226 let chosen = messages
3227 .iter()
3228 .enumerate()
3229 .find(|(_, msg)| !msg.is_empty())
3230 .map(|(_, msg)| msg.clone())
3231 .unwrap_or_default();
3232 broadcast_message(comm, root, Some(chosen))
3233}
3234
3235#[cfg(not(feature = "complex"))]
3236fn gather_dist_csr(dist: &DistCsrOp, root: usize) -> Result<Option<CsrMatrix<f64>>, KError> {
3237 let comm = dist.comm();
3238 let rank = comm.rank();
3239 let size = comm.size();
3240 let row_part = dist.row_partition();
3241 if row_part.len() != size + 1 {
3242 return Err(KError::InvalidInput(
3243 "distributed row partition is inconsistent with communicator size".into(),
3244 ));
3245 }
3246 let local = dist.local_matrix();
3247 let local_row_ptr = local.row_ptr().to_vec();
3248 let local_col_idx = local.col_idx().to_vec();
3249 let local_vals = local.values().to_vec();
3250 let local_nnz = local_col_idx.len() as u64;
3251 let mut nnz_counts = Vec::new();
3252 comm.gather(&[local_nnz], &mut nnz_counts, root);
3253
3254 if rank != root {
3255 let row_ptr_u64: Vec<u64> = local_row_ptr.iter().map(|&v| v as u64).collect();
3256 let col_idx_u64: Vec<u64> = local_col_idx.iter().map(|&v| v as u64).collect();
3257 let mut reqs = Vec::new();
3258 reqs.push(comm.isend_to_u64(&row_ptr_u64, root as i32));
3259 reqs.push(comm.isend_to_u64(&col_idx_u64, root as i32));
3260 reqs.push(comm.isend_to(&local_vals, root as i32));
3261 comm.wait_all(&mut reqs);
3262 return Ok(None);
3263 }
3264
3265 let mut recv_row_ptr_u64: Vec<Vec<u64>> = vec![Vec::new(); size];
3266 let mut recv_col_idx_u64: Vec<Vec<u64>> = vec![Vec::new(); size];
3267 let mut recv_vals: Vec<Vec<f64>> = vec![Vec::new(); size];
3268 for r in 0..size {
3269 if r == root {
3270 continue;
3271 }
3272 let n_local = row_part[r + 1] - row_part[r];
3273 let nnz = *nnz_counts.get(r).unwrap_or(&0) as usize;
3274 recv_row_ptr_u64[r] = vec![0u64; n_local + 1];
3275 recv_col_idx_u64[r] = vec![0u64; nnz];
3276 recv_vals[r] = vec![0.0; nnz];
3277 let mut reqs = Vec::with_capacity(3);
3278 reqs.push(comm.irecv_from_u64(recv_row_ptr_u64[r].as_mut_slice(), r as i32));
3279 reqs.push(comm.irecv_from_u64(recv_col_idx_u64[r].as_mut_slice(), r as i32));
3280 reqs.push(comm.irecv_from(recv_vals[r].as_mut_slice(), r as i32));
3281 comm.wait_all(&mut reqs);
3282 }
3283
3284 let mut row_ptrs: Vec<Vec<usize>> = vec![Vec::new(); size];
3285 let mut col_idxs: Vec<Vec<usize>> = vec![Vec::new(); size];
3286 let mut vals: Vec<Vec<f64>> = vec![Vec::new(); size];
3287 row_ptrs[root] = local_row_ptr;
3288 col_idxs[root] = local_col_idx;
3289 vals[root] = local_vals;
3290 for r in 0..size {
3291 if r == root {
3292 continue;
3293 }
3294 row_ptrs[r] = recv_row_ptr_u64[r].iter().map(|&v| v as usize).collect();
3295 col_idxs[r] = recv_col_idx_u64[r].iter().map(|&v| v as usize).collect();
3296 vals[r] = recv_vals[r].clone();
3297 }
3298
3299 let n_global = dist.n_global;
3300 let mut row_nnz = vec![0usize; n_global];
3301 for r in 0..size {
3302 let row_start = row_part[r];
3303 let n_local = row_part[r + 1] - row_part[r];
3304 if n_local + 1 != row_ptrs[r].len() {
3305 return Err(KError::InvalidInput(format!(
3306 "rank {r} row_ptr length mismatch: expected {}, got {}",
3307 n_local + 1,
3308 row_ptrs[r].len()
3309 )));
3310 }
3311 for i in 0..n_local {
3312 let nnz = row_ptrs[r][i + 1] - row_ptrs[r][i];
3313 row_nnz[row_start + i] = nnz;
3314 }
3315 }
3316 let mut row_ptr_global = vec![0usize; n_global + 1];
3317 for i in 0..n_global {
3318 row_ptr_global[i + 1] = row_ptr_global[i] + row_nnz[i];
3319 }
3320 let total_nnz = row_ptr_global[n_global];
3321 let mut col_idx_global = vec![0usize; total_nnz];
3322 let mut vals_global = vec![0.0f64; total_nnz];
3323 let mut next_pos = row_ptr_global.clone();
3324 for r in 0..size {
3325 let row_start = row_part[r];
3326 let n_local = row_part[r + 1] - row_part[r];
3327 for i in 0..n_local {
3328 let global_row = row_start + i;
3329 let mut slot = next_pos[global_row];
3330 let rs = row_ptrs[r][i];
3331 let re = row_ptrs[r][i + 1];
3332 for p in rs..re {
3333 col_idx_global[slot] = col_idxs[r][p];
3334 vals_global[slot] = vals[r][p];
3335 slot += 1;
3336 }
3337 next_pos[global_row] = slot;
3338 }
3339 }
3340
3341 Ok(Some(CsrMatrix::from_csr(
3342 n_global,
3343 n_global,
3344 row_ptr_global,
3345 col_idx_global,
3346 vals_global,
3347 )))
3348}
3349
3350#[cfg(not(feature = "complex"))]
3351fn gather_vector(
3352 comm: &UniverseComm,
3353 row_part: &[usize],
3354 root: usize,
3355 local: &[f64],
3356) -> Result<Option<Vec<f64>>, KError> {
3357 let rank = comm.rank();
3360 let size = comm.size();
3361 if row_part.len() != size + 1 {
3362 return Err(KError::InvalidInput(
3363 "distributed row partition is inconsistent with communicator size".into(),
3364 ));
3365 }
3366 let (start, end) = (row_part[rank], row_part[rank + 1]);
3367 if local.len() != end.saturating_sub(start) {
3368 return Err(KError::InvalidInput(
3369 "distributed vector length does not match local row partition".into(),
3370 ));
3371 }
3372 if let Some(chunk) = uniform_positive_partition_len(row_part) {
3373 debug_assert_eq!(chunk, local.len());
3374 let mut gathered = Vec::new();
3375 comm.gather(local, &mut gathered, root);
3376 if rank == root {
3377 return Ok(Some(gathered));
3378 }
3379 return Ok(None);
3380 }
3381 if rank != root {
3382 let mut reqs = vec![comm.isend_to(local, root as i32)];
3383 comm.wait_all(&mut reqs);
3384 return Ok(None);
3385 }
3386 let n_global = *row_part.last().unwrap_or(&0);
3387 let mut global = vec![0.0f64; n_global];
3388 global[start..end].copy_from_slice(local);
3389 for r in 0..size {
3390 if r == root {
3391 continue;
3392 }
3393 let rs = row_part[r];
3394 let re = row_part[r + 1];
3395 let mut reqs = Vec::with_capacity(1);
3396 reqs.push(comm.irecv_from(&mut global[rs..re], r as i32));
3397 comm.wait_all(&mut reqs);
3398 }
3399 Ok(Some(global))
3400}
3401
3402#[cfg(not(feature = "complex"))]
3403fn scatter_vector(
3404 comm: &UniverseComm,
3405 row_part: &[usize],
3406 root: usize,
3407 global: Option<&[f64]>,
3408 local_out: &mut [f64],
3409) -> Result<(), KError> {
3410 let rank = comm.rank();
3413 let size = comm.size();
3414 if row_part.len() != size + 1 {
3415 return Err(KError::InvalidInput(
3416 "distributed row partition is inconsistent with communicator size".into(),
3417 ));
3418 }
3419 let (start, end) = (row_part[rank], row_part[rank + 1]);
3420 if local_out.len() != end.saturating_sub(start) {
3421 return Err(KError::InvalidInput(
3422 "distributed vector length does not match local row partition".into(),
3423 ));
3424 }
3425 if uniform_positive_partition_len(row_part).is_some() {
3426 if rank == root {
3427 let global = global.ok_or_else(|| {
3428 KError::InvalidInput("root rank missing global vector for scatter".into())
3429 })?;
3430 let n_global = *row_part.last().unwrap_or(&0);
3431 if global.len() < n_global {
3432 return Err(KError::InvalidInput(
3433 "global vector length does not match distributed partition".into(),
3434 ));
3435 }
3436 comm.scatter(&global[..n_global], local_out, root);
3437 } else {
3438 comm.scatter(&[], local_out, root);
3439 }
3440 return Ok(());
3441 }
3442 if rank == root {
3443 let global = global.ok_or_else(|| {
3444 KError::InvalidInput("root rank missing global vector for scatter".into())
3445 })?;
3446 if global.len() < *row_part.last().unwrap_or(&0) {
3447 return Err(KError::InvalidInput(
3448 "global vector length does not match distributed partition".into(),
3449 ));
3450 }
3451 local_out.copy_from_slice(&global[start..end]);
3452 let mut reqs = Vec::new();
3453 for r in 0..size {
3454 if r == root {
3455 continue;
3456 }
3457 let rs = row_part[r];
3458 let re = row_part[r + 1];
3459 reqs.push(comm.isend_to(&global[rs..re], r as i32));
3460 }
3461 comm.wait_all(&mut reqs);
3462 return Ok(());
3463 }
3464 let mut reqs = vec![comm.irecv_from(local_out, root as i32)];
3465 comm.wait_all(&mut reqs);
3466 Ok(())
3467}
3468
3469#[cfg(not(feature = "complex"))]
3470fn uniform_positive_partition_len(row_part: &[usize]) -> Option<usize> {
3471 let first = row_part.get(1)?.checked_sub(*row_part.first()?)?;
3472 if first == 0 {
3473 return None;
3474 }
3475 row_part
3476 .windows(2)
3477 .all(|w| w[1].checked_sub(w[0]) == Some(first))
3478 .then_some(first)
3479}
3480
3481#[cfg(not(feature = "complex"))]
3482fn owner_of_row(row_part: &[usize], gcol: usize) -> usize {
3483 let mut lo = 0usize;
3484 let mut hi = row_part.len().saturating_sub(2);
3485 while lo <= hi {
3486 let mid = (lo + hi) / 2;
3487 if gcol < row_part[mid + 1] {
3488 if gcol >= row_part[mid] {
3489 return mid;
3490 }
3491 if mid == 0 {
3492 break;
3493 }
3494 hi = mid - 1;
3495 } else {
3496 lo = mid + 1;
3497 }
3498 }
3499 lo.min(row_part.len().saturating_sub(2))
3500}
3501
3502#[cfg(not(feature = "complex"))]
3503fn build_amg_halo_plan(
3504 comm: UniverseComm,
3505 row_part: Arc<Vec<usize>>,
3506 row_start: usize,
3507 row_end: usize,
3508 local: &CsrMatrix<f64>,
3509) -> Result<HaloPlan, KError> {
3510 let rank = comm.rank();
3511 let mut recv_map: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
3512 let row_ptr = local.row_ptr();
3513 let col_idx = local.col_idx();
3514 for row in 0..local.nrows() {
3515 for idx in row_ptr[row]..row_ptr[row + 1] {
3516 let gcol = col_idx[idx];
3517 let owner = owner_of_row(&row_part, gcol);
3518 if owner != rank {
3519 recv_map.entry(owner).or_default().push(gcol);
3520 }
3521 }
3522 }
3523 HaloPlan::new(comm, row_part, row_start, row_end, recv_map)
3524}
3525
3526#[cfg(debug_assertions)]
3527fn debug_check_csr<T: KrystScalar>(a: &CsrMatrix<T>, name: &str) {
3528 let row_ptr = a.row_ptr();
3529 let nrows = a.nrows();
3530 let nnz = a.nnz();
3531 let col_idx = a.col_idx();
3532 let vals = a.values();
3533 debug_assert_eq!(
3534 row_ptr.len(),
3535 nrows + 1,
3536 "{name}: row_ptr.len() mismatch ({} vs {})",
3537 row_ptr.len(),
3538 nrows + 1
3539 );
3540 debug_assert_eq!(
3541 col_idx.len(),
3542 nnz,
3543 "{name}: col_idx.len() ({}) != nnz ({})",
3544 col_idx.len(),
3545 nnz
3546 );
3547 debug_assert_eq!(
3548 vals.len(),
3549 nnz,
3550 "{name}: vals.len() ({}) != nnz ({})",
3551 vals.len(),
3552 nnz
3553 );
3554 debug_assert_eq!(
3555 row_ptr[nrows], nnz,
3556 "{name}: row_ptr[nrows] ({}) != nnz ({})",
3557 row_ptr[nrows], nnz
3558 );
3559 let ncols = a.ncols();
3560 for row in 0..nrows {
3561 let start = row_ptr[row];
3562 let end = row_ptr[row + 1];
3563 debug_assert!(
3564 start <= end,
3565 "{name}: row {row} pointers out-of-order ({start}..{end})"
3566 );
3567 let mut last_col = None;
3568 for idx in start..end {
3569 let col = col_idx[idx];
3570 debug_assert!(
3571 col < ncols,
3572 "{name}: column index {} (row {}) out of bounds (ncols={ncols})",
3573 col,
3574 row
3575 );
3576 if let Some(prev) = last_col {
3577 debug_assert!(
3578 col >= prev,
3579 "{name}: column index decreased at row {}: {} < {}",
3580 row,
3581 col,
3582 prev
3583 );
3584 }
3585 last_col = Some(col);
3586 let val = vals[idx];
3587 debug_assert!(
3588 val.is_finite(),
3589 "{name}: non-finite value at row {} (idx {})",
3590 row,
3591 idx
3592 );
3593 }
3594 }
3595}
3596
3597enum AmgState {
3598 Uninitialized,
3599 SymbolicOnly {
3600 hierarchy: Box<AmgHierarchy>,
3601 last_structure_id: StructureId,
3602 pattern_hash: u64,
3603 },
3604 Ready {
3605 hierarchy: Box<AmgHierarchy>,
3606 last_structure_id: StructureId,
3607 last_values_id: ValuesId,
3608 pattern_hash: u64,
3609 },
3610}
3611
3612impl AmgState {
3613 fn as_ref(&self) -> Option<&AmgHierarchy> {
3614 match self {
3615 AmgState::SymbolicOnly { hierarchy, .. } => Some(hierarchy),
3616 AmgState::Ready { hierarchy, .. } => Some(hierarchy),
3617 _ => None,
3618 }
3619 }
3620}
3621
3622struct DistAmgInfo {
3623 comm: UniverseComm,
3624 root: usize,
3625 row_part: Arc<Vec<usize>>,
3626 n_global: usize,
3627 local_amg: Option<Box<AMG>>,
3628 local_matrix: Option<Arc<CsrMatrix<f64>>>,
3629 #[cfg(not(feature = "complex"))]
3630 distributed_matrix: Option<Arc<DistCsrOp>>,
3631 halo_plan: Option<HaloPlan>,
3632 #[cfg(not(feature = "complex"))]
3633 distributed_hierarchy: Option<DistAmgHierarchy>,
3634}
3635
3636impl DistAmgInfo {
3637 fn local_range(&self) -> (usize, usize) {
3638 let rank = self.comm.rank();
3639 let start = self.row_part.get(rank).copied().unwrap_or_default();
3640 let end = self.row_part.get(rank + 1).copied().unwrap_or(start);
3641 (start, end)
3642 }
3643
3644 fn local_nrows(&self) -> usize {
3645 let (start, end) = self.local_range();
3646 end.saturating_sub(start)
3647 }
3648}
3649
3650#[cfg(not(feature = "complex"))]
3651struct DistAmgLevel {
3652 a: Arc<DistCsrOp>,
3653 row_part: Arc<Vec<usize>>,
3654 halo: Arc<HaloPlan>,
3655 global_row_start: usize,
3656 global_row_end: usize,
3657}
3658
3659#[cfg(not(feature = "complex"))]
3660struct DistAmgHierarchy {
3661 levels: Vec<DistAmgLevel>,
3662 meta: DistHierarchyMeta,
3663}
3664
3665#[cfg(not(feature = "complex"))]
3666#[derive(Clone, Copy, Debug, Default)]
3667struct DistAmgHierarchyShape {
3668 levels: usize,
3669 global_rows: usize,
3670 local_nnz: usize,
3671}
3672
3673#[cfg(not(feature = "complex"))]
3674#[derive(Default)]
3675struct DistCsrCorrectionSummary {
3676 local_apply: Duration,
3677 halo_exchange: Duration,
3678 comm_bytes: usize,
3679}
3680
3681#[cfg(not(feature = "complex"))]
3682impl DistAmgHierarchy {
3683 fn from_fine(comm: UniverseComm, row_part: Arc<Vec<usize>>, fine: Arc<DistCsrOp>) -> Self {
3684 let halo = Arc::new(HaloPlan::from_shared_index(fine.halo_index()));
3685 let level = DistAmgLevel {
3686 a: fine,
3687 row_part: row_part.clone(),
3688 halo: halo.clone(),
3689 global_row_start: row_part.get(comm.rank()).copied().unwrap_or_default(),
3690 global_row_end: row_part
3691 .get(comm.rank() + 1)
3692 .copied()
3693 .unwrap_or_else(|| row_part.get(comm.rank()).copied().unwrap_or_default()),
3694 };
3695 let meta = DistHierarchyMeta {
3696 comm,
3697 row_part: row_part.clone(),
3698 level_row_parts: vec![row_part],
3699 level_halos: vec![halo],
3700 coarse_owners: Vec::new(),
3701 coarse_offsets: Vec::new(),
3702 coarse_global_size: 0,
3703 };
3704 let hierarchy = Self {
3705 levels: vec![level],
3706 meta,
3707 };
3708 debug_assert!(hierarchy.validate_metadata().is_ok());
3709 hierarchy
3710 }
3711
3712 fn validate_metadata(&self) -> Result<(), KError> {
3713 if self.levels.len() != self.meta.level_row_parts.len()
3714 || self.levels.len() != self.meta.level_halos.len()
3715 {
3716 return Err(KError::InvalidInput(
3717 "AMG distributed CSR hierarchy metadata level count mismatch".into(),
3718 ));
3719 }
3720 for (ix, level) in self.levels.iter().enumerate() {
3721 if level.row_part.as_ref() != self.meta.level_row_parts[ix].as_ref() {
3722 return Err(KError::InvalidInput(format!(
3723 "AMG distributed CSR hierarchy row partition mismatch at level {ix}"
3724 )));
3725 }
3726 if level.global_row_start
3727 != level
3728 .row_part
3729 .get(self.meta.comm.rank())
3730 .copied()
3731 .unwrap_or_default()
3732 || level.global_row_end
3733 != level
3734 .row_part
3735 .get(self.meta.comm.rank() + 1)
3736 .copied()
3737 .unwrap_or(level.global_row_start)
3738 {
3739 return Err(KError::InvalidInput(format!(
3740 "AMG distributed CSR hierarchy local row range mismatch at level {ix}"
3741 )));
3742 }
3743 }
3744 Ok(())
3745 }
3746
3747 fn shape(&self) -> DistAmgHierarchyShape {
3748 DistAmgHierarchyShape {
3749 levels: self.meta.level_row_parts.len(),
3750 global_rows: self.meta.row_part.last().copied().unwrap_or_default(),
3751 local_nnz: self
3752 .finest()
3753 .map(|level| level.a.local_matrix().nnz())
3754 .unwrap_or_default(),
3755 }
3756 }
3757
3758 fn finest(&self) -> Option<&DistAmgLevel> {
3759 self.levels.first()
3760 }
3761
3762 fn num_levels(&self) -> usize {
3763 self.levels.len()
3764 }
3765
3766 fn finest_local_nrows(&self) -> Option<usize> {
3767 self.finest()
3768 .map(|level| level.global_row_end.saturating_sub(level.global_row_start))
3769 }
3770
3771 fn finest_comm_bytes(&self) -> usize {
3772 self.finest()
3773 .map(|level| {
3774 level
3775 .halo
3776 .index
3777 .recv_map
3778 .values()
3779 .chain(level.halo.index.send_map.values())
3780 .map(|cols| cols.len() * std::mem::size_of::<f64>())
3781 .sum()
3782 })
3783 .unwrap_or(0)
3784 }
3785
3786 fn finest_residual(&self, x: &[f64], rhs: &[f64], residual: &mut [f64]) -> Result<(), KError> {
3787 let level = self
3788 .finest()
3789 .ok_or_else(|| KError::InvalidInput("AMG distributed CSR hierarchy is empty".into()))?;
3790 let n = level.global_row_end.saturating_sub(level.global_row_start);
3791 if x.len() != n || rhs.len() != n || residual.len() < n {
3792 return Err(KError::InvalidInput(
3793 "AMG distributed CSR finest residual length mismatch".into(),
3794 ));
3795 }
3796 level.a.try_matvec(x, &mut residual[..n])?;
3797 #[cfg(feature = "rayon")]
3798 residual[..n]
3799 .par_iter_mut()
3800 .enumerate()
3801 .for_each(|(i, value)| *value = rhs[i] - *value);
3802 #[cfg(not(feature = "rayon"))]
3803 for i in 0..n {
3804 residual[i] = rhs[i] - residual[i];
3805 }
3806 Ok(())
3807 }
3808
3809 fn apply_finest_residual_correction(
3810 &self,
3811 local_amg: &AMG,
3812 side: PcSide,
3813 rhs: &[f64],
3814 out: &mut [f64],
3815 work: &mut DistCsrApplyWorkspace,
3816 ) -> Result<DistCsrCorrectionSummary, KError> {
3817 let n = self
3818 .finest_local_nrows()
3819 .ok_or_else(|| KError::InvalidInput("AMG distributed CSR hierarchy is empty".into()))?;
3820 if rhs.len() != n || out.len() != n {
3821 return Err(KError::InvalidInput(
3822 "AMG distributed CSR apply length mismatch".into(),
3823 ));
3824 }
3825 work.ensure(n);
3826
3827 let mut summary = DistCsrCorrectionSummary::default();
3828 let t_local = tic();
3829 local_amg.apply_local(side, rhs, out)?;
3830 summary.local_apply += toc(t_local);
3831
3832 let t_halo = tic();
3833 self.finest_residual(out, rhs, &mut work.residual[..n])?;
3834 summary.halo_exchange += toc(t_halo);
3835 summary.comm_bytes = summary.comm_bytes.saturating_add(self.finest_comm_bytes());
3836
3837 let t_local = tic();
3838 local_amg.apply_local(side, &work.residual[..n], &mut work.correction[..n])?;
3839 summary.local_apply += toc(t_local);
3840
3841 #[cfg(feature = "rayon")]
3842 out.par_iter_mut()
3843 .zip(work.correction[..n].par_iter())
3844 .for_each(|(zi, correction)| *zi += *correction);
3845 #[cfg(not(feature = "rayon"))]
3846 for i in 0..n {
3847 out[i] += work.correction[i];
3848 }
3849
3850 Ok(summary)
3851 }
3852}
3853
3854#[derive(Default)]
3855struct DistCsrApplyWorkspace {
3856 residual: Vec<f64>,
3857 correction: Vec<f64>,
3858}
3859
3860impl DistCsrApplyWorkspace {
3861 fn ensure(&mut self, n: usize) {
3862 self.residual.resize(n, 0.0);
3863 self.correction.resize(n, 0.0);
3864 }
3865}
3866
3867#[cfg(feature = "complex")]
3870#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3871pub enum AmgComplexSetupMode {
3872 Unset,
3873 NativeDiagonal,
3874 NativeCoarse,
3875 NativeHierarchy,
3876 ProjectedRealHierarchy,
3877}
3878
3879#[cfg(feature = "complex")]
3880impl AmgComplexSetupMode {
3881 pub fn as_str(self) -> &'static str {
3882 match self {
3883 Self::Unset => "unset",
3884 Self::NativeDiagonal => "native_diagonal",
3885 Self::NativeCoarse => "native_coarse",
3886 Self::NativeHierarchy => "native_hierarchy",
3887 Self::ProjectedRealHierarchy => "projected_real_hierarchy",
3888 }
3889 }
3890}
3891
3892pub struct AMG {
3893 csr: Option<Arc<CsrMatrix<f64>>>,
3894 #[cfg(feature = "complex")]
3895 complex_diag_inv: Option<Vec<S>>,
3896 #[cfg(feature = "complex")]
3897 complex_coarse_solver: Option<Mutex<CoarseDenseLu<S>>>,
3898 #[cfg(feature = "complex")]
3899 complex_core: Option<Mutex<AmgCore<S>>>,
3900 #[cfg(feature = "complex")]
3901 complex_setup_mode: AmgComplexSetupMode,
3902 #[cfg(feature = "complex")]
3903 complex_setup_fallback_reason: Option<String>,
3904 state: AmgState,
3905 cycle_policy: Box<dyn CyclePolicy + Send + Sync>,
3906 cfg: AMGConfig,
3907 stats: Option<AmgStats>,
3908 runtime: Mutex<AmgRuntime>,
3909 workspace_pool: Mutex<Vec<AMGWorkspace>>,
3910 dist_csr_workspace_pool: Mutex<Vec<DistCsrApplyWorkspace>>,
3911 dist: Option<DistAmgInfo>,
3912 transfer_overrides: BTreeMap<usize, AmgTransferOperators>,
3913 coarse_level_overrides: BTreeMap<usize, CoarseSolve>,
3914 relax_level_overrides: BTreeMap<usize, RelaxType>,
3915 sweep_level_overrides: BTreeMap<usize, (usize, usize)>,
3916}
3917
3918#[derive(Clone, Debug)]
3919pub struct AmgTransferOperators {
3920 pub prolongation: CsrMatrix<S>,
3921 pub restriction: CsrMatrix<S>,
3922}
3923
3924impl AmgTransferOperators {
3925 pub fn from_prolongation_adjoint(prolongation: CsrMatrix<S>) -> Self {
3931 let (restriction, _) = adjoint_csr_with_pos(&prolongation);
3932 Self {
3933 prolongation,
3934 restriction,
3935 }
3936 }
3937}
3938
3939impl Default for AMG {
3940 fn default() -> Self {
3941 let cfg = AMGConfig::default();
3942 Self {
3943 csr: None,
3944 #[cfg(feature = "complex")]
3945 complex_diag_inv: None,
3946 #[cfg(feature = "complex")]
3947 complex_coarse_solver: None,
3948 #[cfg(feature = "complex")]
3949 complex_core: None,
3950 #[cfg(feature = "complex")]
3951 complex_setup_mode: AmgComplexSetupMode::Unset,
3952 #[cfg(feature = "complex")]
3953 complex_setup_fallback_reason: None,
3954 state: AmgState::Uninitialized,
3955 cycle_policy: Self::make_cycle_policy(&cfg),
3956 cfg,
3957 stats: None,
3958 runtime: Mutex::new(AmgRuntime::default()),
3959 workspace_pool: Mutex::new(Vec::new()),
3960 dist_csr_workspace_pool: Mutex::new(Vec::new()),
3961 dist: None,
3962 transfer_overrides: BTreeMap::new(),
3963 coarse_level_overrides: BTreeMap::new(),
3964 relax_level_overrides: BTreeMap::new(),
3965 sweep_level_overrides: BTreeMap::new(),
3966 }
3967 }
3968}
3969
3970impl AMG {
3971 pub fn new(_matrix: &Mat<f64>, _max_levels: usize, _coarsening_threshold: f64) -> Self {
3972 AMG::default()
3973 }
3974 pub fn builder() -> AMGBuilder {
3975 AMGBuilder::new()
3976 }
3977 pub fn with_config(mut cfg: AMGConfig) -> Self {
3978 if cfg.spd_diag_floor < 0.0 {
3979 cfg.spd_diag_floor = 0.0;
3980 }
3981 let coarse_level_overrides = cfg.level_coarse_overrides.clone();
3982 let relax_level_overrides = cfg.level_relax_overrides.clone();
3983 let sweep_level_overrides = cfg.level_sweep_overrides.clone();
3984 Self {
3985 cycle_policy: Self::make_cycle_policy(&cfg),
3986 coarse_level_overrides,
3987 relax_level_overrides,
3988 sweep_level_overrides,
3989 cfg,
3990 state: AmgState::Uninitialized,
3991 ..Default::default()
3992 }
3993 }
3994
3995 pub fn set_level_transfer_operators(&mut self, level: usize, operators: AmgTransferOperators) {
3996 self.transfer_overrides.insert(level, operators);
3997 }
3998
3999 pub fn set_level_coarse_solver(&mut self, level: usize, solve: CoarseSolve) {
4000 self.coarse_level_overrides.insert(level, solve);
4001 }
4002
4003 pub fn clear_hierarchy_overrides(&mut self) {
4004 self.transfer_overrides.clear();
4005 self.coarse_level_overrides.clear();
4006 self.relax_level_overrides.clear();
4007 self.sweep_level_overrides.clear();
4008 }
4009
4010 pub fn set_level_relax_type(&mut self, level: usize, relax: RelaxType) {
4011 self.relax_level_overrides.insert(level, relax);
4012 }
4013
4014 pub fn set_level_sweeps(&mut self, level: usize, pre: usize, post: usize) {
4015 self.sweep_level_overrides.insert(level, (pre, post));
4016 }
4017
4018 fn make_cycle_policy(cfg: &AMGConfig) -> Box<dyn CyclePolicy + Send + Sync> {
4019 match (&cfg.cycle_type, cfg.kcycle.as_ref()) {
4020 (CycleType::V, None) => Box::new(VPolicy),
4021 (CycleType::W { gamma }, None) => Box::new(WPolicy {
4022 gamma: (*gamma).max(2),
4023 }),
4024 (base, Some(kc)) => Box::new(KPolicy::new(*base, kc.clone())),
4025 }
4026 }
4027
4028 pub fn extract_coarse_space(&self, opts: &DeflationOptions) -> Result<AmgCoarseSpace, KError> {
4029 let state = match &self.state {
4030 AmgState::Ready { hierarchy, .. } => hierarchy,
4031 _ => return Err(KError::InvalidInput("AMG not set up".into())),
4032 };
4033 match opts.z_source {
4034 ZSource::CoarsestIdentity { cap_k } => {
4035 let coarse_ix = state.coarsest_ix();
4036 let n_coarse = state.levels[coarse_ix].a.nrows();
4037 let k_full = n_coarse;
4038 let cap = cap_k.unwrap_or(k_full).min(k_full);
4039 let mut z = Mat::<f64>::zeros(n_coarse, cap);
4040 for i in 0..cap {
4041 z[(i, i)] = 1.0;
4042 }
4043 let mut current = z;
4044 for lvl in (0..coarse_ix).rev() {
4045 let p = &state.levels[lvl].p;
4046 let mut next = Mat::<f64>::zeros(p.nrows(), cap);
4047 csr_spmm_dense(p, current.as_ref(), next.as_mut())?;
4048 current = next;
4049 }
4050 Ok(AmgCoarseSpace {
4051 z: current,
4052 local_range: None,
4053 })
4054 }
4055 ZSource::NearNullspace => {
4056 let finest = &state.levels[0];
4057 let basis = finest
4058 .nns
4059 .as_ref()
4060 .ok_or_else(|| KError::InvalidInput("near-nullspace unavailable".into()))?;
4061 let k = basis.len();
4062 if k == 0 {
4063 return Err(KError::InvalidInput("near-nullspace empty".into()));
4064 }
4065 let n = finest.a.nrows();
4066 let mut z = Mat::<f64>::zeros(n, k);
4067 for (j, vec) in basis.iter().enumerate() {
4068 if vec.len() != n {
4069 return Err(KError::InvalidInput(
4070 "near-nullspace vector has wrong length".into(),
4071 ));
4072 }
4073 for i in 0..n {
4074 z[(i, j)] = vec[i];
4075 }
4076 }
4077 Ok(AmgCoarseSpace {
4078 z,
4079 local_range: None,
4080 })
4081 }
4082 ZSource::External => Err(KError::InvalidInput(
4083 "ZSource::External requires user-provided coarse space".into(),
4084 )),
4085 }
4086 }
4087
4088 #[cfg(not(feature = "complex"))]
4091 fn resolve_dist_coarse_strategy(
4092 &self,
4093 comm: &UniverseComm,
4094 ) -> Result<(DistCoarseStrategy, DistCoarseSolverRoute), KError> {
4095 let is_route_available = |route: DistCoarseSolverRoute| -> bool {
4096 match route {
4097 DistCoarseSolverRoute::Auto => true,
4098 DistCoarseSolverRoute::Root | DistCoarseSolverRoute::Local => true,
4099 DistCoarseSolverRoute::SuperLuDist => {
4100 #[cfg(feature = "superlu_dist")]
4101 {
4102 true
4103 }
4104 #[cfg(not(feature = "superlu_dist"))]
4105 {
4106 false
4107 }
4108 }
4109 }
4110 };
4111
4112 let route_to_strategy = |route: DistCoarseSolverRoute| -> DistCoarseStrategy {
4113 match route {
4114 DistCoarseSolverRoute::Auto => self.cfg.dist_coarse_strategy,
4115 DistCoarseSolverRoute::Root => DistCoarseStrategy::RootGather,
4116 DistCoarseSolverRoute::Local => DistCoarseStrategy::LocalPrototype,
4117 DistCoarseSolverRoute::SuperLuDist => DistCoarseStrategy::SuperLuDist,
4118 }
4119 };
4120
4121 let strategy_to_route = |strategy: DistCoarseStrategy| -> DistCoarseSolverRoute {
4122 match strategy {
4123 DistCoarseStrategy::DistributedCsr => DistCoarseSolverRoute::Auto,
4124 DistCoarseStrategy::RootGather => DistCoarseSolverRoute::Root,
4125 DistCoarseStrategy::LocalPrototype => DistCoarseSolverRoute::Local,
4126 DistCoarseStrategy::SuperLuDist => DistCoarseSolverRoute::SuperLuDist,
4127 DistCoarseStrategy::None => DistCoarseSolverRoute::Auto,
4128 }
4129 };
4130
4131 if self.cfg.dist_coarse_solver_route != DistCoarseSolverRoute::Auto {
4132 let forced_route = self.cfg.dist_coarse_solver_route;
4133 if !is_route_available(forced_route) {
4134 let _ = forced_route;
4135 return Err(KError::Unsupported(
4136 "AMG distributed coarse route was explicitly requested but is unavailable (missing feature support)",
4137 ));
4138 }
4139 return Ok((route_to_strategy(forced_route), forced_route));
4140 }
4141
4142 let strategy = self.cfg.dist_coarse_strategy;
4143 match strategy {
4144 DistCoarseStrategy::DistributedCsr => Ok((
4145 DistCoarseStrategy::DistributedCsr,
4146 DistCoarseSolverRoute::Auto,
4147 )),
4148 DistCoarseStrategy::None => {
4149 if comm.size() > 1 {
4150 log::warn!(
4151 "AMG distributed coarse strategy set to none; falling back to root_gather."
4152 );
4153 Ok((DistCoarseStrategy::RootGather, DistCoarseSolverRoute::Root))
4154 } else {
4155 Ok((DistCoarseStrategy::None, DistCoarseSolverRoute::Auto))
4156 }
4157 }
4158 DistCoarseStrategy::SuperLuDist => {
4159 if is_route_available(DistCoarseSolverRoute::SuperLuDist) {
4160 Ok((
4161 DistCoarseStrategy::SuperLuDist,
4162 DistCoarseSolverRoute::SuperLuDist,
4163 ))
4164 } else {
4165 log::warn!(
4166 "AMG distributed coarse strategy superlu_dist unavailable; falling back to root_gather."
4167 );
4168 Ok((DistCoarseStrategy::RootGather, DistCoarseSolverRoute::Root))
4169 }
4170 }
4171 _ => Ok((strategy, strategy_to_route(strategy))),
4172 }
4173 }
4174
4175 fn build_symbolic(&mut self, fine: &CsrMatrix<f64>) -> Result<Box<AmgHierarchy>, KError> {
4176 let mut cfg = self.cfg.clone();
4178 let primary = build_hierarchy(fine, &mut cfg, &self.transfer_overrides);
4179 let (hier, stats, used_cfg) = match primary {
4180 Ok((hier, stats)) => (hier, stats, cfg),
4181 Err(primary_err) => {
4182 let mut fallback_cfg = cfg.clone();
4183 fallback_cfg.coarsen_type = CoarsenType::RS;
4184 fallback_cfg.interp_type = InterpType::Classical;
4185 match build_hierarchy(fine, &mut fallback_cfg, &self.transfer_overrides) {
4186 Ok((hier, stats)) => (hier, stats, fallback_cfg),
4187 Err(fallback_err) => {
4188 let mut smoother_cfg = cfg.clone();
4189 smoother_cfg.coarse_solve = CoarseSolve::Smoother;
4190 match build_smoother_only_hierarchy(fine, &mut smoother_cfg) {
4191 Ok((hier, stats)) => (hier, stats, smoother_cfg),
4192 Err(_) => {
4193 let combined = format!(
4194 "AMG setup failed: {primary_err}; fallback failed: {fallback_err}"
4195 );
4196 return Err(match (&primary_err, &fallback_err) {
4197 (KError::SolveError(_), KError::SolveError(_)) => {
4198 KError::SolveError(combined)
4199 }
4200 _ => KError::InvalidInput(combined),
4201 });
4202 }
4203 }
4204 }
4205 }
4206 }
4207 };
4208 #[cfg(test)]
4209 BUILD_SYMBOLIC_COUNT.with(|c| c.set(c.get() + 1));
4210 self.cfg = used_cfg;
4211 if let Some((_, solve)) = self
4212 .coarse_level_overrides
4213 .iter()
4214 .filter(|(lvl, _)| **lvl < hier.levels.len())
4215 .max_by_key(|(lvl, _)| *lvl)
4216 {
4217 self.cfg.coarse_solve = *solve;
4219 }
4220 self.cycle_policy = Self::make_cycle_policy(&self.cfg);
4221 self.stats = stats;
4222 Ok(Box::new(hier))
4223 }
4224
4225 fn set_dist_route_stats_from_apply(&mut self, stats: &DistApplyStats) {
4226 self.stats = Some(AmgStats::from_dist_apply(stats));
4227 }
4228
4229 #[cfg(not(feature = "complex"))]
4230 fn set_dist_route_stats_from_apply_with_shape(
4231 &mut self,
4232 stats: &DistApplyStats,
4233 shape: Option<DistAmgHierarchyShape>,
4234 ) {
4235 self.set_dist_route_stats_from_apply(stats);
4236 if let (Some(amg_stats), Some(shape)) = (self.stats.as_mut(), shape) {
4237 amg_stats.num_levels = shape.levels;
4238 amg_stats.total_nnz = shape.local_nnz;
4239 amg_stats.grid_complexity = if shape.global_rows > 0 { 1.0 } else { 0.0 };
4240 amg_stats.operator_complexity = if shape.local_nnz > 0 { 1.0 } else { 0.0 };
4241 }
4242 }
4243
4244 #[cfg(not(feature = "complex"))]
4245 fn setup_dist(&mut self, dist: &DistCsrOp) -> Result<(), KError> {
4246 let setup_t0 = tic();
4247 let (strategy, selected_route) = self.resolve_dist_coarse_strategy(&dist.comm())?;
4248 if matches!(strategy, DistCoarseStrategy::DistributedCsr) {
4249 return self.setup_dist_local_mode(dist, strategy);
4250 }
4251 if matches!(strategy, DistCoarseStrategy::LocalPrototype) {
4252 return self.setup_dist_local_mode(dist, strategy);
4253 }
4254 if matches!(strategy, DistCoarseStrategy::SuperLuDist) {
4255 return self.setup_dist_superlu(dist, selected_route);
4256 }
4257 let comm = dist.comm();
4258 let rank = comm.rank();
4259 let row_part = self.dist_coarse_partition(dist.row_partition(), dist.n_global);
4260 let root = 0usize;
4261 self.dist = Some(DistAmgInfo {
4267 comm: comm.clone(),
4268 root,
4269 row_part: row_part.clone(),
4270 n_global: dist.n_global,
4271 local_amg: None,
4272 local_matrix: Some(Arc::new(dist.local_matrix())),
4273 distributed_matrix: None,
4274 halo_plan: None,
4275 distributed_hierarchy: None,
4276 });
4277 let prev_cfg = self.cfg.clone();
4278 let mut local_stage: Option<&'static str> = None;
4279 let mut local_detail: Option<String> = None;
4280 let record_error = |local_stage: &mut Option<&'static str>,
4281 local_detail: &mut Option<String>,
4282 stage: &'static str,
4283 err: KError| {
4284 if local_stage.is_none() {
4285 *local_stage = Some(stage);
4286 *local_detail = Some(err.to_string());
4287 }
4288 };
4289 let mut setup_state: Option<(
4290 CsrMatrix<f64>,
4291 Box<AmgHierarchy>,
4292 StructureId,
4293 ValuesId,
4294 u64,
4295 )> = None;
4296 let mut setup_profile: Option<&'static str> = None;
4297 let global = match gather_dist_csr(dist, root) {
4298 Ok(global) => global,
4299 Err(err) => {
4300 record_error(&mut local_stage, &mut local_detail, "gather_dist_csr", err);
4301 None
4302 }
4303 };
4304 if local_stage.is_none() && rank == root {
4305 match global {
4306 Some(mut fine) => {
4307 let cfg = self.cfg.clone();
4308 if cfg.conditioning.is_active() {
4309 if let Err(err) = apply_csr_transforms("AMG", &mut fine, &cfg.conditioning)
4310 {
4311 record_error(&mut local_stage, &mut local_detail, "conditioning", err);
4312 }
4313 }
4314 if local_stage.is_none() {
4315 let sid = dist.structure_id();
4316 let vid = dist.values_id();
4317 let pattern_hash = csr_pattern_hash(&fine);
4318 match self.build_symbolic(&fine) {
4319 Ok(hierarchy) => {
4320 setup_profile = Some("strict");
4321 setup_state = Some((fine, hierarchy, sid, vid, pattern_hash));
4322 }
4323 Err(primary_err) => {
4324 let mut permissive_cfg = prev_cfg.clone();
4325 permissive_cfg.require_spd = false;
4326 permissive_cfg.verify_galerkin = false;
4327 permissive_cfg.verify_p_rank = false;
4328 permissive_cfg.interp_type = InterpType::Classical;
4329 self.cfg = permissive_cfg;
4330 match self.build_symbolic(&fine) {
4331 Ok(hierarchy) => {
4332 setup_profile = Some("permissive");
4333 setup_state =
4334 Some((fine, hierarchy, sid, vid, pattern_hash));
4335 }
4336 Err(fallback_err) => {
4337 self.cfg = prev_cfg.clone();
4338 record_error(
4339 &mut local_stage,
4340 &mut local_detail,
4341 "build_symbolic",
4342 KError::InvalidInput(format!(
4343 "strict setup failed: {primary_err}; permissive fallback failed: {fallback_err}"
4344 )),
4345 );
4346 }
4347 }
4348 }
4349 }
4350 }
4351 }
4352 None => {
4353 record_error(
4354 &mut local_stage,
4355 &mut local_detail,
4356 "gather_dist_csr",
4357 KError::InvalidInput("root rank missing assembled CSR matrix".into()),
4358 );
4359 }
4360 }
4361 }
4362 let local_failure = if local_stage.is_none() { 0.0 } else { 1.0 };
4363 let failure_sum = comm.all_reduce_f64(local_failure);
4364 if failure_sum > 0.0 {
4365 let local_message = local_stage.map(|stage| {
4366 let detail = local_detail
4367 .clone()
4368 .unwrap_or_else(|| "unknown error".to_string());
4369 format!("stage={stage}: {detail}")
4370 });
4371 let error_message = collect_error_message(&comm, root, local_message);
4372 self.cfg = prev_cfg;
4373 self.cycle_policy = Self::make_cycle_policy(&self.cfg);
4374 self.state = AmgState::Uninitialized;
4375 self.stats = None;
4376 self.csr = None;
4377 return Err(KError::InvalidInput(format!(
4378 "AMG distributed setup failed: {}",
4379 if error_message.is_empty() {
4380 "unknown error".to_string()
4381 } else {
4382 error_message
4383 }
4384 )));
4385 }
4386 if rank != root {
4387 self.state = AmgState::Uninitialized;
4388 self.csr = None;
4389 let mut ds = DistApplyStats::default();
4390 ds.mode = strategy;
4391 ds.coarse_repartition = self.cfg.dist_coarse_repartition;
4392 ds.coarse_solver_route = selected_route;
4393 ds.setup_total = toc(setup_t0);
4394 ds.reductions = 1;
4395 ds.setup_gathered_fine_matrix = true;
4396 self.set_dist_route_stats_from_apply(&ds);
4397 if let Ok(mut rt) = self.runtime.lock() {
4398 rt.last_dist_apply = Some(ds);
4399 }
4400 return Ok(());
4401 }
4402 let (fine, hierarchy, sid, vid, pattern_hash) = setup_state.ok_or_else(|| {
4403 KError::InvalidInput("AMG distributed setup missing hierarchy state".into())
4404 })?;
4405 self.state = AmgState::Ready {
4406 hierarchy,
4407 last_structure_id: sid,
4408 last_values_id: vid,
4409 pattern_hash,
4410 };
4411 self.csr = Some(Arc::new(fine));
4412 let mut ds = DistApplyStats::default();
4413 ds.mode = strategy;
4414 ds.coarse_repartition = self.cfg.dist_coarse_repartition;
4415 ds.coarse_solver_route = selected_route;
4416 ds.setup_total = toc(setup_t0);
4417 ds.reductions = 1;
4418 ds.setup_gathered_fine_matrix = true;
4419 if let Some(stats) = self.stats.as_mut() {
4420 stats.selected_dist_coarse_route =
4421 Some(dist_route_label(selected_route, strategy).to_string());
4422 stats.dist_route_fallback = dist_route_fallback_labels(selected_route, strategy);
4423 } else {
4424 self.set_dist_route_stats_from_apply(&ds);
4425 }
4426 if let Some(profile) = setup_profile
4427 && profile == "permissive"
4428 && self.cfg.print_level >= 1
4429 {
4430 println!("AMG distributed setup succeeded using permissive configuration.");
4431 }
4432 if self.cfg.logging_level >= 2
4433 && self.cfg.print_level >= 1
4434 && let Some(s) = self.stats.as_ref()
4435 {
4436 print_setup_tables(s);
4437 }
4438 if let Ok(mut rt) = self.runtime.lock() {
4439 rt.last_dist_apply = Some(ds);
4440 }
4441 Ok(())
4442 }
4443
4444 #[cfg(not(feature = "complex"))]
4445 fn setup_dist_local_mode(
4446 &mut self,
4447 dist: &DistCsrOp,
4448 strategy: DistCoarseStrategy,
4449 ) -> Result<(), KError> {
4450 let setup_t0 = tic();
4451 let comm = dist.comm();
4452 let rank = comm.rank();
4453 let row_part = self.dist_coarse_partition(dist.row_partition(), dist.n_global);
4454 let root = 0usize;
4455
4456 let mut local_stage: Option<&'static str> = None;
4457 let mut local_detail: Option<String> = None;
4458 let record_error = |local_stage: &mut Option<&'static str>,
4459 local_detail: &mut Option<String>,
4460 stage: &'static str,
4461 err: KError| {
4462 if local_stage.is_none() {
4463 *local_stage = Some(stage);
4464 *local_detail = Some(err.to_string());
4465 }
4466 };
4467
4468 let local_matrix = dist.local_matrix();
4469 let local_block = dist.local_block_csr();
4470 let mut local_amg: Option<Box<AMG>> = None;
4471 let mut halo_plan: Option<HaloPlan> = None;
4472
4473 if local_stage.is_none() {
4474 let mut candidate = Box::new(AMG::with_config(self.cfg.clone()));
4475 match candidate.setup(&local_block) {
4476 Ok(()) => {
4477 local_amg = Some(candidate);
4478 }
4479 Err(err) => {
4480 record_error(
4481 &mut local_stage,
4482 &mut local_detail,
4483 "build_symbolic(local)",
4484 err,
4485 );
4486 }
4487 }
4488 }
4489
4490 if local_stage.is_none() && matches!(strategy, DistCoarseStrategy::LocalPrototype) {
4491 match build_amg_halo_plan(
4492 comm.clone(),
4493 row_part.clone(),
4494 dist.row_start,
4495 dist.row_end,
4496 &local_matrix,
4497 ) {
4498 Ok(plan) => halo_plan = Some(plan),
4499 Err(err) => {
4500 record_error(&mut local_stage, &mut local_detail, "build_halo_plan", err);
4501 }
4502 }
4503 }
4504
4505 let distributed_matrix = if matches!(strategy, DistCoarseStrategy::DistributedCsr) {
4506 match DistCsrOp::from_local_rows(
4507 dist.n_global,
4508 dist.row_start,
4509 &local_matrix,
4510 dist.row_partition().as_ref(),
4511 comm.clone(),
4512 ) {
4513 Ok(op) => Some(Arc::new(op)),
4514 Err(err) => {
4515 record_error(&mut local_stage, &mut local_detail, "build_dist_csr", err);
4516 None
4517 }
4518 }
4519 } else {
4520 None
4521 };
4522 let distributed_hierarchy = distributed_matrix
4523 .as_ref()
4524 .map(|op| DistAmgHierarchy::from_fine(comm.clone(), row_part.clone(), op.clone()));
4525 let distributed_shape = distributed_hierarchy.as_ref().map(DistAmgHierarchy::shape);
4526
4527 let local_failure = if local_stage.is_none() { 0.0 } else { 1.0 };
4528 let failure_sum = comm.all_reduce_f64(local_failure);
4529 if failure_sum > 0.0 {
4530 let local_message = local_stage.map(|stage| {
4531 let detail = local_detail
4532 .clone()
4533 .unwrap_or_else(|| "unknown error".to_string());
4534 format!("stage={stage}: {detail}")
4535 });
4536 let error_message = collect_error_message(&comm, root, local_message);
4537 self.state = AmgState::Uninitialized;
4538 self.stats = None;
4539 self.csr = None;
4540 return Err(KError::InvalidInput(format!(
4541 "AMG distributed local setup failed: {}",
4542 if error_message.is_empty() {
4543 "unknown error".to_string()
4544 } else {
4545 error_message
4546 }
4547 )));
4548 }
4549
4550 self.dist = Some(DistAmgInfo {
4551 comm: comm.clone(),
4552 root,
4553 row_part: row_part.clone(),
4554 n_global: dist.n_global,
4555 local_amg,
4556 local_matrix: Some(Arc::new(local_matrix)),
4557 distributed_matrix,
4558 halo_plan,
4559 distributed_hierarchy,
4560 });
4561 self.state = AmgState::Uninitialized;
4562 self.csr = None;
4563 let mut ds = DistApplyStats::default();
4564 ds.mode = strategy;
4565 ds.coarse_repartition = self.cfg.dist_coarse_repartition;
4566 ds.coarse_solver_route = if matches!(strategy, DistCoarseStrategy::DistributedCsr) {
4567 DistCoarseSolverRoute::Auto
4568 } else {
4569 DistCoarseSolverRoute::Local
4570 };
4571 ds.setup_total = toc(setup_t0);
4572 ds.reductions = 1;
4573 self.set_dist_route_stats_from_apply_with_shape(&ds, distributed_shape);
4574 if let Ok(mut rt) = self.runtime.lock() {
4575 rt.last_dist_apply = Some(ds);
4576 }
4577
4578 if self.cfg.print_level >= 1 && self.cfg.logging_level >= 1 {
4579 log::info!(
4580 "AMG {} setup complete: rank={} local_rows={}",
4581 dist_strategy_label(strategy),
4582 rank,
4583 dist.local_nrows()
4584 );
4585 }
4586 Ok(())
4587 }
4588
4589 #[cfg(not(feature = "complex"))]
4590 fn try_update_dist_local_numeric(
4591 &mut self,
4592 dist: &DistCsrOp,
4593 strategy: DistCoarseStrategy,
4594 ) -> Result<bool, KError> {
4595 let setup_t0 = tic();
4596 if !matches!(
4597 strategy,
4598 DistCoarseStrategy::DistributedCsr | DistCoarseStrategy::LocalPrototype
4599 ) {
4600 return Ok(false);
4601 }
4602
4603 let comm = dist.comm();
4604 let row_part = dist.row_partition();
4605 let local_matrix = dist.local_matrix();
4606 let local_block = dist.local_block_csr();
4607 let compatible = self.dist.as_ref().is_some_and(|state| {
4608 state.n_global == dist.n_global
4609 && state.row_part.as_ref() == row_part.as_ref()
4610 && state.local_range() == (dist.row_start, dist.row_end)
4611 && state.local_amg.is_some()
4612 && state.local_matrix.as_ref().is_some_and(|stored| {
4613 stored.row_ptr() == local_matrix.row_ptr()
4614 && stored.col_idx() == local_matrix.col_idx()
4615 })
4616 && match strategy {
4617 DistCoarseStrategy::DistributedCsr => {
4618 state.distributed_matrix.as_ref().is_some_and(|stored| {
4619 stored.local_matrix().row_ptr() == local_matrix.row_ptr()
4620 && stored.local_matrix().col_idx() == local_matrix.col_idx()
4621 })
4622 }
4623 DistCoarseStrategy::LocalPrototype => state.distributed_matrix.is_none(),
4624 _ => false,
4625 }
4626 });
4627 if comm.all_reduce_f64(if compatible { 0.0 } else { 1.0 }) > 0.0 {
4628 return Ok(false);
4629 }
4630
4631 let mut local_error = None;
4632 if let Some(state) = self.dist.as_mut() {
4633 if let Some(local_amg) = state.local_amg.as_mut()
4634 && let Err(err) = local_amg.update_numeric(&local_block)
4635 {
4636 local_error = Some(err);
4637 }
4638 if local_error.is_none() && matches!(strategy, DistCoarseStrategy::DistributedCsr) {
4639 match DistCsrOp::from_local_rows(
4640 dist.n_global,
4641 dist.row_start,
4642 &local_matrix,
4643 dist.row_partition().as_ref(),
4644 comm.clone(),
4645 ) {
4646 Ok(op) => {
4647 let op = Arc::new(op);
4648 state.distributed_hierarchy = Some(DistAmgHierarchy::from_fine(
4649 comm.clone(),
4650 state.row_part.clone(),
4651 op.clone(),
4652 ));
4653 state.distributed_matrix = Some(op);
4654 }
4655 Err(err) => {
4656 local_error = Some(err);
4657 }
4658 }
4659 } else if local_error.is_none() {
4660 state.distributed_hierarchy = None;
4661 }
4662 if local_error.is_none() {
4663 state.local_matrix = Some(Arc::new(local_matrix));
4664 }
4665 }
4666
4667 if comm.all_reduce_f64(if local_error.is_some() { 1.0 } else { 0.0 }) > 0.0 {
4668 return Ok(false);
4669 }
4670 let mut ds = DistApplyStats::default();
4671 ds.mode = strategy;
4672 ds.coarse_repartition = self.cfg.dist_coarse_repartition;
4673 ds.coarse_solver_route = if matches!(strategy, DistCoarseStrategy::DistributedCsr) {
4674 DistCoarseSolverRoute::Auto
4675 } else {
4676 DistCoarseSolverRoute::Local
4677 };
4678 ds.setup_total = toc(setup_t0);
4679 ds.reductions = 2;
4680 let distributed_shape = self
4681 .dist
4682 .as_ref()
4683 .and_then(|state| state.distributed_hierarchy.as_ref())
4684 .map(DistAmgHierarchy::shape);
4685 self.set_dist_route_stats_from_apply_with_shape(&ds, distributed_shape);
4686 if let Ok(mut rt) = self.runtime.lock() {
4687 rt.last_dist_apply = Some(ds);
4688 }
4689 Ok(true)
4690 }
4691
4692 #[cfg(not(feature = "complex"))]
4693 fn setup_dist_superlu(
4694 &mut self,
4695 dist: &DistCsrOp,
4696 selected_route: DistCoarseSolverRoute,
4697 ) -> Result<(), KError> {
4698 let setup_t0 = tic();
4699 let comm = dist.comm();
4700 let root = 0usize;
4701 let row_part = self.dist_coarse_partition(dist.row_partition(), dist.n_global);
4702 self.dist = Some(DistAmgInfo {
4703 comm: comm.clone(),
4704 root,
4705 row_part: row_part.clone(),
4706 n_global: dist.n_global,
4707 local_amg: None,
4708 local_matrix: Some(Arc::new(dist.local_matrix())),
4709 distributed_matrix: None,
4710 halo_plan: None,
4711 distributed_hierarchy: None,
4712 });
4713 self.state = AmgState::Uninitialized;
4714 self.csr = None;
4715 let mut ds = DistApplyStats::default();
4716 ds.mode = DistCoarseStrategy::SuperLuDist;
4717 ds.coarse_repartition = self.cfg.dist_coarse_repartition;
4718 ds.coarse_solver_route = selected_route;
4719 ds.setup_total = toc(setup_t0);
4720 self.set_dist_route_stats_from_apply(&ds);
4721 if let Ok(mut rt) = self.runtime.lock() {
4722 rt.last_dist_apply = Some(ds);
4723 }
4724 Ok(())
4725 }
4726
4727 fn apply_local(&self, side: PcSide, r: &[f64], z: &mut [f64]) -> Result<(), KError> {
4728 if r.len() != z.len() {
4729 return Err(KError::InvalidInput(format!(
4730 "AMG.apply: r/z size mismatch: {} vs {}",
4731 r.len(),
4732 z.len()
4733 )));
4734 }
4735 if self.cfg.require_spd && side != PcSide::Left {
4736 return Err(KError::InvalidInput(
4737 "AMG in SPD mode supports only Left preconditioning for CG-safe use".into(),
4738 ));
4739 }
4740 let h = match &self.state {
4741 AmgState::Ready { hierarchy, .. } => hierarchy,
4742 _ => {
4743 return Err(KError::InvalidInput("AMG not set up".into()));
4744 }
4745 };
4746 self.apply_with_hierarchy(side, r, z, h)
4747 }
4748
4749 fn apply_with_hierarchy(
4750 &self,
4751 #[allow(unused_variables)] side: PcSide,
4752 r: &[f64],
4753 z: &mut [f64],
4754 h: &AmgHierarchy,
4755 ) -> Result<(), KError> {
4756 let n = h.finest().a.nrows();
4757 let mut ws = self
4758 .workspace_pool
4759 .lock()
4760 .unwrap_or_else(std::sync::PoisonError::into_inner)
4761 .pop()
4762 .unwrap_or_else(|| AMGWorkspace::new(n));
4763 ws.ensure(n);
4764 let do_prof = self.cfg.logging_level >= 2;
4765 let result = if do_prof {
4766 let mut cyc = CycleTimings::default();
4767 let t_all = tic();
4768 z.fill(R::default());
4769 let cycle_result = self.cycle_profiled(0, r, z, &mut ws, Some(&mut cyc));
4770 if cycle_result.is_ok() {
4771 cyc.total_cycle = toc(t_all);
4772 cyc.cycle_type = self.cfg.cycle_type;
4773 cyc.kcycle = self.cfg.kcycle.clone();
4774 if let Ok(mut rt) = self.runtime.lock() {
4775 rt.last_cycle = Some(cyc.clone());
4776 }
4777 if self.cfg.print_level >= 2 {
4778 print_cycle_table(&cyc);
4779 }
4780 }
4781 cycle_result
4782 } else {
4783 z.fill(R::default());
4784 self.cycle(0, r, z, &mut ws)
4785 };
4786 self.workspace_pool
4787 .lock()
4788 .unwrap_or_else(std::sync::PoisonError::into_inner)
4789 .push(ws);
4790 result
4791 }
4792
4793 #[cfg(not(feature = "complex"))]
4794 fn apply_dist(
4795 &self,
4796 side: PcSide,
4797 r: &[f64],
4798 z: &mut [f64],
4799 dist: &DistAmgInfo,
4800 ) -> Result<(), KError> {
4801 let do_prof = self.cfg.dist_apply_instrumentation;
4802 let (strategy, selected_route) = self.resolve_dist_coarse_strategy(&dist.comm)?;
4803 let mut stats = if do_prof {
4804 let mut stats = DistApplyStats::default();
4805 stats.mode = strategy;
4806 stats.coarse_solver_route = selected_route;
4807 stats.coarse_repartition = self.cfg.dist_coarse_repartition;
4808 if let Ok(rt) = self.runtime.lock()
4809 && let Some(setup_stats) = rt.last_dist_apply.as_ref()
4810 {
4811 stats.setup_total = setup_stats.setup_total;
4812 stats.reductions = setup_stats.reductions;
4813 stats.setup_gathered_fine_matrix = setup_stats.setup_gathered_fine_matrix;
4814 }
4815 Some(stats)
4816 } else {
4817 None
4818 };
4819
4820 if let Some(stats_ref) = stats.as_mut()
4821 && stats_ref.per_level_comm_bytes.is_empty()
4822 {
4823 let levels = dist
4824 .distributed_hierarchy
4825 .as_ref()
4826 .map(DistAmgHierarchy::num_levels)
4827 .unwrap_or(1)
4828 .max(1);
4829 stats_ref.per_level_comm_bytes = vec![0; levels];
4830 }
4831
4832 let result = match strategy {
4833 DistCoarseStrategy::DistributedCsr => {
4834 self.apply_dist_csr(side, r, z, dist, stats.as_mut())
4835 }
4836 DistCoarseStrategy::RootGather => {
4837 self.apply_dist_root(side, r, z, dist, stats.as_mut())
4838 }
4839 DistCoarseStrategy::LocalPrototype => {
4840 self.apply_dist_local(side, r, z, dist, stats.as_mut())
4841 }
4842 DistCoarseStrategy::None => Err(KError::Unsupported(
4843 "AMG distributed apply requires a coarse strategy (distributed_csr, root_gather, local_prototype, or superlu_dist)"
4844 .into(),
4845 )),
4846 DistCoarseStrategy::SuperLuDist => {
4847 self.apply_dist_superlu(side, r, z, dist, stats.as_mut())
4848 }
4849 };
4850
4851 if do_prof {
4852 if let Some(stats_ref) = stats.as_mut()
4853 && let Ok(rt) = self.runtime.lock()
4854 && let Some(last_cycle) = rt.last_cycle.as_ref()
4855 {
4856 stats_ref.per_level_apply = last_cycle
4857 .per_level
4858 .iter()
4859 .map(|lvl| {
4860 lvl.pre_smooth
4861 + lvl.matvec
4862 + lvl.residual_axpy
4863 + lvl.restrict
4864 + lvl.coarse_solve
4865 + lvl.prolong
4866 + lvl.post_smooth
4867 })
4868 .collect();
4869 }
4870 if let Ok(mut rt) = self.runtime.lock() {
4871 rt.last_dist_apply = stats;
4872 }
4873 }
4874
4875 result
4876 }
4877
4878 #[cfg(not(feature = "complex"))]
4879 fn apply_dist_csr(
4880 &self,
4881 side: PcSide,
4882 r: &[f64],
4883 z: &mut [f64],
4884 dist: &DistAmgInfo,
4885 mut stats: Option<&mut DistApplyStats>,
4886 ) -> Result<(), KError> {
4887 let local_amg = dist.local_amg.as_ref().ok_or_else(|| {
4888 KError::InvalidInput("AMG distributed CSR local solver not initialized".into())
4889 })?;
4890 let hierarchy = dist.distributed_hierarchy.as_ref().ok_or_else(|| {
4891 KError::InvalidInput("AMG distributed CSR hierarchy state not initialized".into())
4892 })?;
4893 let n = hierarchy
4894 .finest_local_nrows()
4895 .ok_or_else(|| KError::InvalidInput("AMG distributed CSR hierarchy is empty".into()))?;
4896 if r.len() != n || z.len() != n {
4897 return Err(KError::InvalidInput(
4898 "AMG distributed CSR apply length mismatch".into(),
4899 ));
4900 }
4901
4902 let mut ws = self
4903 .dist_csr_workspace_pool
4904 .lock()
4905 .unwrap_or_else(std::sync::PoisonError::into_inner)
4906 .pop()
4907 .unwrap_or_default();
4908 ws.ensure(n);
4909
4910 let result = hierarchy.apply_finest_residual_correction(local_amg, side, r, z, &mut ws);
4911 if let (Some(stats), Ok(summary)) = (stats.as_mut(), result.as_ref()) {
4912 stats.local_apply += summary.local_apply;
4913 stats.halo_exchange += summary.halo_exchange;
4914 stats.comm_bytes = stats.comm_bytes.saturating_add(summary.comm_bytes);
4915 if !stats.per_level_comm_bytes.is_empty() {
4916 stats.per_level_comm_bytes[0] =
4917 stats.per_level_comm_bytes[0].saturating_add(summary.comm_bytes);
4918 }
4919 }
4920 self.dist_csr_workspace_pool
4921 .lock()
4922 .unwrap_or_else(std::sync::PoisonError::into_inner)
4923 .push(ws);
4924 result.map(|_| ())
4925 }
4926
4927 #[cfg(not(feature = "complex"))]
4928 fn apply_dist_root(
4929 &self,
4930 side: PcSide,
4931 r: &[f64],
4932 z: &mut [f64],
4933 dist: &DistAmgInfo,
4934 mut stats: Option<&mut DistApplyStats>,
4935 ) -> Result<(), KError> {
4936 let comm = &dist.comm;
4937 let root = dist.root;
4938 let row_part = dist.row_part.as_ref();
4939 let rank = comm.rank();
4940 let t_gather = stats.as_ref().map(|_| tic());
4941 let global_r = gather_vector(comm, row_part, root, r)?;
4942 if let (Some(stats), Some(t0)) = (stats.as_mut(), t_gather) {
4943 stats.gather = toc(t0);
4944 let n = row_part.last().copied().unwrap_or_default();
4945 let bytes = n * std::mem::size_of::<f64>();
4946 stats.comm_bytes = stats.comm_bytes.saturating_add(bytes);
4947 if !stats.per_level_comm_bytes.is_empty() {
4948 stats.per_level_comm_bytes[0] = stats.per_level_comm_bytes[0].saturating_add(bytes);
4949 }
4950 }
4951 let mut global_z = if rank == root {
4952 vec![0.0f64; dist.n_global]
4953 } else {
4954 Vec::new()
4955 };
4956 if rank == root {
4957 let rhs = global_r
4958 .as_ref()
4959 .ok_or_else(|| KError::InvalidInput("root rank missing gathered RHS".into()))?;
4960 let t_local = stats.as_ref().map(|_| tic());
4961 self.apply_local(side, rhs, &mut global_z)?;
4962 if let (Some(stats), Some(t0)) = (stats.as_mut(), t_local) {
4963 stats.local_apply = toc(t0);
4964 }
4965 }
4966 let global_ref = if rank == root {
4967 Some(global_z.as_slice())
4968 } else {
4969 None
4970 };
4971 let t_scatter = stats.as_ref().map(|_| tic());
4972 scatter_vector(comm, row_part, root, global_ref, z)?;
4973 if let (Some(stats), Some(t0)) = (stats.as_mut(), t_scatter) {
4974 stats.scatter = toc(t0);
4975 let n = row_part.last().copied().unwrap_or_default();
4976 let bytes = n * std::mem::size_of::<f64>();
4977 stats.comm_bytes = stats.comm_bytes.saturating_add(bytes);
4978 if !stats.per_level_comm_bytes.is_empty() {
4979 stats.per_level_comm_bytes[0] = stats.per_level_comm_bytes[0].saturating_add(bytes);
4980 }
4981 }
4982 Ok(())
4983 }
4984
4985 #[cfg(not(feature = "complex"))]
4986 fn apply_dist_local(
4987 &self,
4988 side: PcSide,
4989 r: &[f64],
4990 z: &mut [f64],
4991 dist: &DistAmgInfo,
4992 mut stats: Option<&mut DistApplyStats>,
4993 ) -> Result<(), KError> {
4994 let local_amg = dist.local_amg.as_ref().ok_or_else(|| {
4995 KError::InvalidInput("AMG local prototype solver not initialized".into())
4996 })?;
4997 if r.len() != z.len() || r.len() != local_amg.dims().0 {
4998 return Err(KError::InvalidInput(
4999 "AMG local prototype apply length mismatch".into(),
5000 ));
5001 }
5002 let t_local = stats.as_ref().map(|_| tic());
5003 local_amg.apply_local(side, r, z)?;
5004 if let (Some(stats), Some(t0)) = (stats.as_mut(), t_local) {
5005 stats.local_apply = toc(t0);
5006 }
5007
5008 if let (Some(halo), Some(local_matrix)) =
5009 (dist.halo_plan.as_ref(), dist.local_matrix.as_ref())
5010 {
5011 let t_halo = stats.as_ref().map(|_| tic());
5012 let req = halo.post_halo(r);
5013 halo.complete_halo(req);
5014 if let (Some(stats), Some(t0)) = (stats.as_mut(), t_halo) {
5015 stats.halo_exchange = toc(t0);
5016 let halo_bytes: usize = halo
5017 .index
5018 .recv_map
5019 .values()
5020 .map(|cols| cols.len() * std::mem::size_of::<f64>())
5021 .sum();
5022 stats.comm_bytes = stats.comm_bytes.saturating_add(halo_bytes);
5023 if !stats.per_level_comm_bytes.is_empty() {
5024 stats.per_level_comm_bytes[0] =
5025 stats.per_level_comm_bytes[0].saturating_add(halo_bytes);
5026 }
5027 }
5028 if self.cfg.dist_coarse_ghost_scale > 0.0 {
5029 let ghost = halo.ghost_slice_ref();
5030 let row_ptr = local_matrix.row_ptr();
5031 let col_idx = local_matrix.col_idx();
5032 for row in 0..local_matrix.nrows() {
5033 let mut acc = 0.0f64;
5034 let mut count = 0usize;
5035 for idx in row_ptr[row]..row_ptr[row + 1] {
5036 let gcol = col_idx[idx];
5037 if let Some(&slot) = halo.index.ghost_index_of.get(&gcol) {
5038 acc += ghost[slot];
5039 count += 1;
5040 }
5041 }
5042 if count > 0 {
5043 z[row] += self.cfg.dist_coarse_ghost_scale * acc / (count as f64);
5044 }
5045 }
5046 }
5047 }
5048 Ok(())
5049 }
5050
5051 #[cfg(not(feature = "complex"))]
5052 fn apply_dist_superlu(
5053 &self,
5054 _side: PcSide,
5055 r: &[f64],
5056 z: &mut [f64],
5057 dist: &DistAmgInfo,
5058 stats: Option<&mut DistApplyStats>,
5059 ) -> Result<(), KError> {
5060 let local_matrix = dist.local_matrix.as_ref().ok_or_else(|| {
5061 KError::InvalidInput("AMG superlu_dist matrix state not initialized".into())
5062 })?;
5063 if r.len() != z.len() || r.len() != dist.local_nrows() {
5064 return Err(KError::InvalidInput(
5065 "AMG superlu_dist apply length mismatch".into(),
5066 ));
5067 }
5068 #[allow(unused_variables)]
5069 let t_local = stats.as_ref().map(|_| tic());
5070 #[cfg(feature = "superlu_dist")]
5071 {
5072 superlu_dist::solve(local_matrix.as_ref(), r, z, &dist.comm)?;
5073 }
5074 #[cfg(not(feature = "superlu_dist"))]
5075 {
5076 let _ = (local_matrix, r, z);
5077 return Err(KError::Unsupported(
5078 "AMG superlu_dist route requires feature superlu_dist".into(),
5079 ));
5080 }
5081 #[allow(unreachable_code)]
5082 if let (Some(stats), Some(t0)) = (stats.as_mut(), t_local) {
5083 stats.local_apply = toc(t0);
5084 }
5085 Ok(())
5086 }
5087
5088 fn dist_coarse_partition(
5089 &self,
5090 fine_row_part: Arc<Vec<usize>>,
5091 n_global: usize,
5092 ) -> Arc<Vec<usize>> {
5093 match self.cfg.dist_coarse_repartition {
5094 DistCoarseRepartition::Keep => fine_row_part,
5095 DistCoarseRepartition::Uniform => {
5096 if fine_row_part.is_empty() {
5097 return fine_row_part;
5098 }
5099 let size = fine_row_part.len().saturating_sub(1);
5100 let mut part = Vec::with_capacity(size + 1);
5101 part.push(0);
5102 for r in 0..size {
5103 let end = n_global * (r + 1) / size;
5104 part.push(end);
5105 }
5106 Arc::new(part)
5107 }
5108 DistCoarseRepartition::Root => {
5109 let size = fine_row_part.len().saturating_sub(1);
5110 let mut part = vec![0usize; size + 1];
5111 if let Some(last) = part.last_mut() {
5112 *last = n_global;
5113 }
5114 Arc::new(part)
5115 }
5116 }
5117 }
5118
5119 fn ensure_symbolic_structure(
5120 &mut self,
5121 fine: &CsrMatrix<f64>,
5122 sid: StructureId,
5123 pattern_hash: u64,
5124 ) -> Result<(), KError> {
5125 let needs_rebuild = match &self.state {
5126 AmgState::Uninitialized => true,
5127 AmgState::SymbolicOnly {
5128 last_structure_id,
5129 pattern_hash: hash,
5130 ..
5131 } => *last_structure_id != sid || *hash != pattern_hash,
5132 AmgState::Ready {
5133 last_structure_id,
5134 pattern_hash: hash,
5135 ..
5136 } => *last_structure_id != sid || *hash != pattern_hash,
5137 };
5138 if needs_rebuild {
5139 let hierarchy = self.build_symbolic(fine)?;
5140 self.state = AmgState::SymbolicOnly {
5141 hierarchy,
5142 last_structure_id: sid,
5143 pattern_hash,
5144 };
5145 }
5146 Ok(())
5147 }
5148
5149 fn refresh_numeric_ready(
5150 &mut self,
5151 fine: &CsrMatrix<f64>,
5152 sid: StructureId,
5153 vid: ValuesId,
5154 pattern_hash: u64,
5155 ) -> Result<(), KError> {
5156 let mut hierarchy = match std::mem::replace(&mut self.state, AmgState::Uninitialized) {
5157 AmgState::SymbolicOnly { hierarchy, .. } => hierarchy,
5158 AmgState::Ready { hierarchy, .. } => hierarchy,
5159 AmgState::Uninitialized => {
5160 return Err(KError::InvalidInput(
5161 "AMG internal state inconsistent".into(),
5162 ));
5163 }
5164 };
5165 self.refresh_numeric(fine, &mut hierarchy)?;
5166 self.state = AmgState::Ready {
5167 hierarchy,
5168 last_structure_id: sid,
5169 last_values_id: vid,
5170 pattern_hash,
5171 };
5172 Ok(())
5173 }
5174
5175 fn refresh_numeric(
5176 &mut self,
5177 fine: &CsrMatrix<f64>,
5178 h: &mut AmgHierarchy,
5179 ) -> Result<(), KError> {
5180 if h.levels.is_empty() {
5181 return Err(KError::InvalidInput("AMG hierarchy empty".into()));
5182 }
5183
5184 #[cfg(feature = "simd")]
5185 let spmv_tuning = utils::default_spmv_tuning();
5186
5187 h.levels[0].a = fine.clone();
5189 let need_l1 = self.cfg.grid_relax_type.contains(&RelaxType::L1Jacobi);
5190 let need_cheb = self.cfg.grid_relax_type.contains(&RelaxType::Chebyshev);
5191 let need_cheb_safe = self.cfg.grid_relax_type.contains(&RelaxType::ChebyshevSafe);
5192 let need_safe_diag = self
5193 .cfg
5194 .grid_relax_type
5195 .contains(&RelaxType::SafeguardedGaussSeidel)
5196 || need_cheb_safe;
5197 let need_ilu0 = self.cfg.grid_relax_type.contains(&RelaxType::Ilu0);
5198 let need_ras = self.cfg.grid_relax_type.contains(&RelaxType::Ras);
5199 let allow_safeguard = need_safe_diag || need_ilu0 || need_ras;
5200 let need_fsai = self.cfg.grid_relax_type.contains(&RelaxType::Fsai);
5201
5202 h.levels[0].diag_inv =
5203 diag_inv_from_csr_cfg_fallback(&h.levels[0].a, &self.cfg, allow_safeguard)?;
5204 let mut trials_current = make_trial_matrix(&self.cfg, h.levels[0].a.nrows())?;
5205 let mut diag_stats: Vec<AmgLevelStats> = Vec::new();
5206 diag_stats.push(AmgLevelStats {
5207 p_min_col_norm: 0.0,
5208 p_cond_sketched: 0.0,
5209 galerkin_worst_rel: 0.0,
5210 });
5211
5212 for l in 0..h.coarsest_ix() {
5214 let pr = h.levels[l].p.row_ptr().to_vec();
5216 let pc = h.levels[l].p.col_idx().to_vec();
5217 let mut p_new_vals = vec![0.0f64; pc.len()];
5218 let mut tp_opt: Option<TentativeP> = None;
5219 if let Some(ref cf) = h.levels[l].cf {
5220 let s = Strength::from_csr(
5221 &h.levels[l].a,
5222 self.cfg.strong_threshold,
5223 self.cfg.normalize_strength,
5224 );
5225 let s_sym = s.symmetrize();
5226 let params = ClassicalParams {
5227 variant: match self.cfg.interp_type {
5228 InterpType::Direct => ClassicalVariant::Direct,
5229 InterpType::HE => ClassicalVariant::HE,
5230 InterpType::Standard | InterpType::Classical | InterpType::Extended => {
5231 ClassicalVariant::Standard
5232 }
5233 _ => ClassicalVariant::Standard,
5234 },
5235 extended: matches!(self.cfg.interp_type, InterpType::Extended),
5236 drop_abs: self.cfg.interpolation_truncation,
5237 trunc_rel: self.cfg.truncation_factor,
5238 cap_row: self.cfg.max_elements_per_row,
5239 keep_at_least_one: true,
5240 };
5241 classical_values_only(
5242 &h.levels[l].a,
5243 &s_sym,
5244 cf,
5245 ¶ms,
5246 &pr,
5247 &pc,
5248 &mut p_new_vals,
5249 )?;
5250 } else {
5251 let tp = TentativeP {
5252 agg_of: h.levels[l].agg_of.clone(),
5253 n_coarse: h.levels[l + 1].a.nrows(),
5254 num_functions: h.levels[l].num_functions,
5255 nns: h.levels[l].nns.clone(),
5256 comp_of: h.levels[l].layout.as_ref().map(|lay| lay.comp_of.clone()),
5257 };
5258 if let Some(ref rb) = h.levels[l].row_basis {
5259 let n_agg = 1 + h.levels[l].agg_of.iter().copied().max().unwrap_or(0);
5260 let tn = TentativeNodal {
5261 agg_of: tp.agg_of.clone(),
5262 n_agg,
5263 mfun: tp.num_functions,
5264 row_basis: rb.clone(),
5265 };
5266 smooth_sa_values_only_mf(
5267 &h.levels[l].a,
5268 &h.levels[l].diag_inv,
5269 &tn,
5270 self.cfg.jacobi_omega,
5271 &pr,
5272 &pc,
5273 &mut p_new_vals,
5274 )?;
5275 } else {
5276 smooth_sa_values_only_multi(
5277 &h.levels[l].a,
5278 &h.levels[l].diag_inv,
5279 &tp,
5280 self.cfg.jacobi_omega,
5281 &pr,
5282 &pc,
5283 &mut p_new_vals,
5284 )?;
5285 }
5286 tp_opt = Some(tp);
5287 }
5288 let ctx = LevelPostContext {
5289 r: h.levels[l].num_functions,
5290 agg_of: &h.levels[l].agg_of,
5291 nns: h.levels[l]
5292 .nns
5293 .as_ref()
5294 .map(|v| v.iter().map(|b| b.as_slice()).collect()),
5295 a: Some(&h.levels[l].a),
5296 d_inv: Some(&h.levels[l].diag_inv),
5297 };
5298 apply_post_interp(&self.cfg, &ctx, &pr, &pc, &mut p_new_vals)?;
5299 if self.cfg.adaptive_interp
5300 && self.cfg.adaptive_samples > 0
5301 && h.levels[l].cf.is_none()
5302 && h.levels[l + 1].a.nrows() > self.cfg.max_coarse_size
5303 && tp_opt.as_ref().is_some_and(|tp| tp.num_functions == 1)
5304 {
5305 let omega = if self.cfg.adaptive_smooth_omega == 0.0 {
5306 self.cfg.jacobi_omega
5307 } else {
5308 self.cfg.adaptive_smooth_omega
5309 };
5310 let samples = sample_low_modes(
5311 &h.levels[l].a,
5312 &h.levels[l].diag_inv,
5313 self.cfg.adaptive_samples,
5314 self.cfg.adaptive_smooth_steps,
5315 omega,
5316 0xC0FFEE,
5317 )?;
5318 if let Some(ref tp) = tp_opt {
5319 let coarse_samples = restrict_samples_to_coarse(
5320 &h.levels[l].a,
5321 tp,
5322 &samples,
5323 self.cfg.adaptive_weight_mode,
5324 );
5325 adaptive_fit_values_only(
5326 &pr,
5327 &pc,
5328 &mut p_new_vals,
5329 tp,
5330 &samples,
5331 &coarse_samples,
5332 self.cfg.adaptive_lambda,
5333 self.cfg.adaptive_enforce_sum1,
5334 self.cfg.interpolation_truncation,
5335 )?;
5336 }
5337 }
5338 let mut p_tmp = Pcsr {
5339 m: h.levels[l].p.nrows(),
5340 n: h.levels[l].p.ncols(),
5341 row_ptr: pr,
5342 col_idx: pc,
5343 vals: p_new_vals,
5344 };
5345 let mut rank_diag = RankDiagnostics::default();
5346 let check_rank = self.cfg.verify_p_rank && tp_opt.is_some();
5347 if check_rank {
5348 let p_view = CsrMatrix::from_csr(
5349 p_tmp.m,
5350 p_tmp.n,
5351 p_tmp.row_ptr.clone(),
5352 p_tmp.col_idx.clone(),
5353 p_tmp.vals.clone(),
5354 );
5355 rank_diag = check_p_rank_fast(&p_view, &self.cfg)?;
5356 if rank_diag.suspect {
5357 let mut cond_report = rank_diag.cond_estimate;
5358 match self.cfg.on_rank_failure {
5359 RankFallback::RetryLooserInterp => {
5360 let tp = tp_opt.as_ref().ok_or_else(|| {
5361 KError::InvalidInput(
5362 "AMG: RetryLooserInterp requires SA interpolation".into(),
5363 )
5364 })?;
5365 match try_fix_rank(
5366 l,
5367 &h.levels[l].a,
5368 &h.levels[l].diag_inv,
5369 tp,
5370 &ctx,
5371 &mut p_tmp,
5372 &mut self.cfg,
5373 )? {
5374 RankFixOutcome::Fixed => {
5375 let p_view = CsrMatrix::from_csr(
5376 p_tmp.m,
5377 p_tmp.n,
5378 p_tmp.row_ptr.clone(),
5379 p_tmp.col_idx.clone(),
5380 p_tmp.vals.clone(),
5381 );
5382 rank_diag = check_p_rank_fast(&p_view, &self.cfg)?;
5383 cond_report = rank_diag.cond_estimate;
5384 if rank_diag.suspect {
5385 return Err(KError::InvalidInput(format!(
5386 "AMG: P rank suspect at level {l}, cond≈{cond_report:.3e}"
5387 )));
5388 }
5389 }
5390 RankFixOutcome::Unfixed => {
5391 return Err(KError::InvalidInput(format!(
5392 "AMG: P rank suspect at level {l}, cond≈{cond_report:.3e}"
5393 )));
5394 }
5395 }
5396 }
5397 RankFallback::Abort => {
5398 return Err(KError::InvalidInput(format!(
5399 "AMG: P rank suspect at level {l}, cond≈{cond_report:.3e}"
5400 )));
5401 }
5402 other => {
5403 return Err(KError::InvalidInput(format!(
5404 "AMG: rank fallback {other:?} not implemented at level {l}"
5405 )));
5406 }
5407 }
5408 }
5409 }
5410 h.levels[l].p.values_mut().copy_from_slice(&p_tmp.vals);
5411 if self.cfg.keep_transpose {
5413 let pvals = h.levels[l].p.values().to_vec();
5414 let p2r = h.levels[l].p2r_pos.clone();
5415 let rvalsm = h.levels[l].r.values_mut();
5416 sync_adjoint_values_from_forward(&pvals, &p2r, rvalsm);
5417 if cfg!(debug_assertions) {
5418 let step = (pvals.len() / 7).max(1);
5419 for s in (0..pvals.len()).step_by(step) {
5420 let ri = p2r[s];
5421 let dv = (pvals[s].conj() - rvalsm[ri]).abs();
5422 debug_assert!(dv <= 1e-12, "R != P^H at sample {s}");
5423 }
5424 }
5425 }
5426 let mut galerkin_worst = 0.0;
5427 let has_ng = h.levels[l].a_next_pat_ng.is_some();
5428 let allow_galerkin = self.cfg.verify_galerkin
5429 && self.cfg.filter_omega <= 0.0
5430 && !has_ng
5431 && self.cfg.galerkin_samples > 0;
5432 if let Some(pat_ref) = h.levels[l].a_next_pat.as_ref() {
5434 let pat = pat_ref.clone();
5435 let nnz = pat.col_idx.len();
5436 let r_tmp_storage = if self.cfg.keep_transpose {
5437 None
5438 } else {
5439 Some(build_r_from_p(&mut h.levels[l]))
5440 };
5441 let r_for_ops = r_tmp_storage.as_ref().unwrap_or(&h.levels[l].r);
5442 let trials_next = if let Some(ref trials) = trials_current {
5443 let mut next = Mat::<f64>::zeros(r_for_ops.nrows(), trials.ncols());
5444 restrict_trials(r_for_ops, trials.as_ref(), next.as_mut())?;
5445 Some(next)
5446 } else {
5447 None
5448 };
5449 let mut vals = vec![0.0; nnz];
5450 rap_numeric(&pat, r_for_ops, &h.levels[l].a, &h.levels[l].p, &mut vals);
5451 {
5452 let mut rf = |row: usize| RowFilter {
5453 tau_abs: self.cfg.rap_truncation_abs,
5454 tau_rel: self.cfg.truncation_factor,
5455 k_max: self.cfg.rap_max_elements_per_row,
5456 must_keep: if self.cfg.keep_pivot_in_rap {
5457 Some(row)
5458 } else {
5459 None
5460 },
5461 };
5462 apply_filter_to_csr_values_in_place(
5463 pat.nrows,
5464 &pat.row_ptr,
5465 &pat.col_idx,
5466 &mut vals,
5467 &mut rf,
5468 );
5469 }
5470 let block_size_next = h.levels[l].num_functions.max(1);
5471 let mut a_full = CsrMatrix::from_csr(
5472 pat.nrows,
5473 pat.ncols,
5474 pat.row_ptr.clone(),
5475 pat.col_idx.clone(),
5476 vals,
5477 );
5478 let use_ng = has_ng;
5479 if !use_ng || !self.cfg.filter_after_non_galerkin {
5480 apply_trial_compensation(
5481 &self.cfg,
5482 &mut a_full,
5483 trials_next.as_ref(),
5484 block_size_next,
5485 )?;
5486 }
5487 let mut a_coarse = if let (Some(ng_pat), Some(map)) =
5488 (&h.levels[l].a_next_pat_ng, &h.levels[l].rap_full2ng_pos)
5489 {
5490 let mut vals_ng = vec![0.0; ng_pat.col_idx.len()];
5491 let full_vals = a_full.values();
5492 for (k_full, &maybe) in map.iter().enumerate() {
5493 if let Some(k_ng) = maybe {
5494 vals_ng[k_ng] += full_vals[k_full];
5495 }
5496 }
5497 let mut a_ng = CsrMatrix::from_csr(
5498 ng_pat.nrows,
5499 ng_pat.ncols,
5500 ng_pat.row_ptr.clone(),
5501 ng_pat.col_idx.clone(),
5502 vals_ng,
5503 );
5504 if self.cfg.filter_after_non_galerkin {
5505 apply_trial_compensation(
5506 &self.cfg,
5507 &mut a_ng,
5508 trials_next.as_ref(),
5509 block_size_next,
5510 )?;
5511 }
5512 a_ng
5513 } else {
5514 a_full
5515 };
5516 if allow_galerkin {
5517 let (ok, worst) = galerkin_sample_check(
5518 &h.levels[l].a,
5519 &h.levels[l].p,
5520 r_for_ops,
5521 &a_coarse,
5522 self.cfg.galerkin_samples,
5523 self.cfg.galerkin_rel_tol,
5524 0xBEEF,
5525 )?;
5526 galerkin_worst = worst;
5527 if !ok {
5528 let a_fix = rap(r_for_ops, &h.levels[l].a, &h.levels[l].p)?;
5529 let (ok2, worst2) = galerkin_sample_check(
5530 &h.levels[l].a,
5531 &h.levels[l].p,
5532 r_for_ops,
5533 &a_fix,
5534 self.cfg.galerkin_samples,
5535 self.cfg.galerkin_rel_tol,
5536 0xBEEF,
5537 )?;
5538 if ok2 {
5539 galerkin_worst = worst2;
5540 a_coarse = a_fix;
5541 } else {
5542 return Err(KError::InvalidInput(format!(
5543 "AMG: Galerkin identity failed at level {l}: worst rel={worst:.3e} (retry={worst2:.3e})"
5544 )));
5545 }
5546 }
5547 }
5548 h.levels[l + 1].a = a_coarse;
5549 h.levels[l + 1].diag_inv =
5550 diag_inv_from_csr_cfg_fallback(&h.levels[l + 1].a, &self.cfg, allow_safeguard)?;
5551 #[cfg(debug_assertions)]
5552 debug_check_csr(&h.levels[l + 1].a, "coarse A");
5553 trials_current = trials_next;
5554 } else {
5555 let mut _r_tmp_owned: Option<CsrMatrix<f64>> = None;
5557 let r_used = if self.cfg.keep_transpose {
5558 &h.levels[l].r
5559 } else {
5560 _r_tmp_owned = Some(build_r_from_p(&mut h.levels[l]));
5561 _r_tmp_owned.as_ref().unwrap()
5562 };
5563 let mut a_coarse = rap(r_used, &h.levels[l].a, &h.levels[l].p)?;
5564 if allow_galerkin {
5565 let (ok, worst) = galerkin_sample_check(
5566 &h.levels[l].a,
5567 &h.levels[l].p,
5568 r_used,
5569 &a_coarse,
5570 self.cfg.galerkin_samples,
5571 self.cfg.galerkin_rel_tol,
5572 0xBEEF,
5573 )?;
5574 galerkin_worst = worst;
5575 if !ok {
5576 let a_fix = rap(r_used, &h.levels[l].a, &h.levels[l].p)?;
5577 let (ok2, worst2) = galerkin_sample_check(
5578 &h.levels[l].a,
5579 &h.levels[l].p,
5580 r_used,
5581 &a_fix,
5582 self.cfg.galerkin_samples,
5583 self.cfg.galerkin_rel_tol,
5584 0xBEEF,
5585 )?;
5586 if ok2 {
5587 galerkin_worst = worst2;
5588 a_coarse = a_fix;
5589 } else {
5590 return Err(KError::InvalidInput(format!(
5591 "AMG: Galerkin identity failed at level {l}: worst rel={worst:.3e} (retry={worst2:.3e})"
5592 )));
5593 }
5594 }
5595 }
5596 h.levels[l + 1].diag_inv =
5597 diag_inv_from_csr_cfg_fallback(&a_coarse, &self.cfg, allow_safeguard)?;
5598 h.levels[l + 1].a = a_coarse;
5599 trials_current = None;
5600 }
5601 diag_stats.push(AmgLevelStats {
5602 p_min_col_norm: rank_diag.min_col_norm,
5603 p_cond_sketched: rank_diag.cond_estimate,
5604 galerkin_worst_rel: galerkin_worst,
5605 });
5606 if l + 1 == h.coarsest_ix() {
5607 let levelc = &mut h.levels[l + 1];
5608 let prefer_dense = matches!(self.cfg.coarse_solve, CoarseSolve::DirectDense)
5609 || levelc.a.nrows() <= self.cfg.max_coarse_size;
5610 if prefer_dense {
5611 if let Some(m) = &levelc.coarse_solver {
5612 m.lock().unwrap().setup(&levelc.a)?;
5613 } else {
5614 let mut solver = CoarseDenseLu::new();
5615 solver.setup(&levelc.a)?;
5616 levelc.coarse_solver = Some(Mutex::new(Box::new(solver)));
5617 }
5618 } else if matches!(self.cfg.coarse_solve, CoarseSolve::ILU) {
5619 if let Some(m) = &levelc.coarse_solver {
5620 m.lock().unwrap().setup(&levelc.a)?;
5621 } else {
5622 let mut solver = CoarseIlu::new(
5623 self.cfg.tolerance,
5624 levelc.a.nrows().min(self.cfg.max_iterations.max(50)),
5625 self.cfg.ilu_drop_tol,
5626 self.cfg.ilu_fill_per_row,
5627 );
5628 solver.setup(&levelc.a)?;
5629 levelc.coarse_solver = Some(Mutex::new(Box::new(solver)));
5630 }
5631 } else {
5632 levelc.coarse_solver = None;
5633 }
5634 }
5635 }
5636
5637 diag_stats.push(AmgLevelStats {
5638 p_min_col_norm: 0.0,
5639 p_cond_sketched: 0.0,
5640 galerkin_worst_rel: 0.0,
5641 });
5642 for lvl in 0..=h.coarsest_ix() {
5643 let recompute = self.cfg.chebyshev_recompute_esteig || h.levels[lvl].cheb.is_none();
5644 update_level_caches(
5645 &self.cfg,
5646 &mut h.levels[lvl],
5647 need_l1,
5648 need_cheb,
5649 need_safe_diag,
5650 need_cheb_safe,
5651 need_ilu0,
5652 need_ras,
5653 recompute,
5654 )?;
5655 if need_fsai {
5656 let level = &mut h.levels[lvl];
5657 if level.fsai.is_none() {
5658 let strength_opt = if self.cfg.fsai_use_strength {
5659 Some(Strength::from_csr(
5660 &level.a,
5661 self.cfg.strong_threshold,
5662 self.cfg.normalize_strength,
5663 ))
5664 } else {
5665 None
5666 };
5667 level.fsai = Some(fsai_build_for_level(
5668 &self.cfg,
5669 &level.a,
5670 strength_opt.as_ref(),
5671 )?);
5672 }
5673 if let Some(mut data) = level.fsai.take() {
5674 fsai_refresh_numeric(&level.a, &mut data, self.cfg.fsai_lambda)?;
5675 level.fsai = Some(data);
5676 }
5677 refresh_mixed_precision_shadows(&self.cfg, level);
5678 } else {
5679 h.levels[lvl].fsai = None;
5680 refresh_mixed_precision_shadows(&self.cfg, &mut h.levels[lvl]);
5681 }
5682 }
5683
5684 #[cfg(feature = "simd")]
5685 {
5686 for level in &mut h.levels {
5687 build_level_spmv_plans(level, &spmv_tuning);
5688 }
5689 }
5690
5691 if self.cfg.logging_level > 0 {
5692 let mut st = AmgStats::from_hierarchy(h);
5693 st.levels = collect_level_stats(
5694 h,
5695 &self.cfg,
5696 Some(&self.relax_level_overrides),
5697 Some(&self.sweep_level_overrides),
5698 );
5699 st.total_smoothing_work = st
5700 .levels
5701 .iter()
5702 .map(|l| l.pre_work_estimate + l.post_work_estimate)
5703 .sum();
5704 st.selected_dist_coarse_route = Some(
5705 dist_route_label(
5706 self.cfg.dist_coarse_solver_route,
5707 self.cfg.dist_coarse_strategy,
5708 )
5709 .to_string(),
5710 );
5711 st.dist_route_fallback = dist_route_fallback_labels(
5712 self.cfg.dist_coarse_solver_route,
5713 self.cfg.dist_coarse_strategy,
5714 );
5715 st.diagnostics = diag_stats;
5716 self.stats = Some(st);
5717 }
5718 Ok(())
5719 }
5720
5721 fn jacobi_smooth_sparse(
5724 omega: f64,
5725 a: &CsrMatrix<f64>,
5726 diag_inv: &[f64],
5727 r: &[f64],
5728 z: &mut [f64],
5729 iters: usize,
5730 ws: &mut AMGWorkspace,
5731 ) -> Result<(), KError> {
5732 if iters == 0 {
5733 return Ok(());
5734 }
5735 let n = a.nrows();
5736 if diag_inv.len() != n || r.len() != n || z.len() != n {
5737 return Err(KError::InvalidInput("Jacobi: dimension mismatch".into()));
5738 }
5739 ws.ensure(n);
5740 ws.temp[..n].copy_from_slice(z);
5741
5742 for _ in 0..iters {
5743 a.spmv_scaled(1.0, &ws.temp[..n], 0.0, &mut ws.work[..n])?;
5745 #[cfg(feature = "rayon")]
5747 ws.temp[..n].par_iter_mut().enumerate().for_each(|(i, zi)| {
5748 *zi += omega * diag_inv[i] * (r[i] - ws.work[i]);
5749 });
5750 #[cfg(not(feature = "rayon"))]
5751 for i in 0..n {
5752 ws.temp[i] += omega * diag_inv[i] * (r[i] - ws.work[i]);
5753 }
5754 }
5755 z.copy_from_slice(&ws.temp[..n]);
5756 Ok(())
5757 }
5758
5759 fn jacobi_smooth_sparse_mp(
5760 omega: f32,
5761 level: &AMGLevel,
5762 rhs: &[f64],
5763 z: &mut [f64],
5764 iters: usize,
5765 ws: &mut AMGWorkspace,
5766 cfg: &AMGConfig,
5767 ) -> Result<(), KError> {
5768 if iters == 0 {
5769 return Ok(());
5770 }
5771 let n = level.a.nrows();
5772 if rhs.len() != n || z.len() != n {
5773 return Err(KError::InvalidInput("Jacobi: dimension mismatch".into()));
5774 }
5775 ws.ensure(n);
5776 ws.ensure_mixed(n);
5777 let mp_ws = ws
5778 .mp
5779 .as_mut()
5780 .expect("mixed workspace missing after ensure_mixed");
5781 mp_ws.temp32[..n]
5782 .iter_mut()
5783 .zip(z.iter())
5784 .for_each(|(dst, &src)| *dst = src as f32);
5785 mp_ws.residual32[..n]
5786 .iter_mut()
5787 .zip(rhs.iter())
5788 .for_each(|(dst, &src)| *dst = src as f32);
5789 let mut diag_owned = Vec::new();
5790 let diag_slice: &[f32] = match cfg.mixed_storage {
5791 MixedStorage::Cached => level
5792 .diag_inv_f32
5793 .as_ref()
5794 .ok_or_else(|| KError::InvalidInput("Jacobi mixed cache missing".into()))?,
5795 MixedStorage::Transient => {
5796 let buf = mp_ws.ensure_diag(n);
5797 buf.iter_mut()
5798 .zip(level.diag_inv.iter())
5799 .for_each(|(d, &s)| *d = s as f32);
5800 diag_owned.extend_from_slice(&buf[..n]);
5801 &diag_owned
5802 }
5803 };
5804 let mut vals_owned = Vec::new();
5805 let vals32: &[f32] = match cfg.mixed_storage {
5806 MixedStorage::Cached => level
5807 .a_vals_f32
5808 .as_ref()
5809 .ok_or_else(|| KError::InvalidInput("Jacobi mixed matrix cache missing".into()))?,
5810 MixedStorage::Transient => {
5811 let buf = mp_ws.ensure_vals(level.a.nnz());
5812 buf.iter_mut()
5813 .zip(level.a.values().iter())
5814 .for_each(|(d, &s)| *d = s as f32);
5815 vals_owned.extend_from_slice(buf.as_slice());
5816 &vals_owned
5817 }
5818 };
5819 let row_ptr = level.a.row_ptr();
5820 let col_idx = level.a.col_idx();
5821 for _ in 0..iters {
5822 spmv_scaled_f32_on_pattern(
5823 n,
5824 row_ptr,
5825 col_idx,
5826 vals32,
5827 1.0,
5828 &mp_ws.temp32[..n],
5829 0.0,
5830 &mut mp_ws.work32[..n],
5831 );
5832 for i in 0..n {
5833 mp_ws.temp32[i] += omega * diag_slice[i] * (mp_ws.residual32[i] - mp_ws.work32[i]);
5834 }
5835 }
5836 z.iter_mut()
5837 .zip(mp_ws.temp32[..n].iter())
5838 .for_each(|(dst, &src)| *dst = src as f64);
5839 Ok(())
5840 }
5841
5842 fn l1_jacobi(
5843 omega: f64,
5844 a: &CsrMatrix<f64>,
5845 l1_inv: &[f64],
5846 r: &[f64],
5847 z: &mut [f64],
5848 iters: usize,
5849 ws: &mut AMGWorkspace,
5850 ) -> Result<(), KError> {
5851 if iters == 0 {
5852 return Ok(());
5853 }
5854 let n = a.nrows();
5855 if l1_inv.len() != n || r.len() != n || z.len() != n {
5856 return Err(KError::InvalidInput("L1-Jacobi: dimension mismatch".into()));
5857 }
5858 ws.ensure(n);
5859 ws.temp[..n].copy_from_slice(z);
5860 for _ in 0..iters {
5861 a.spmv_scaled(1.0, &ws.temp[..n], 0.0, &mut ws.work[..n])?;
5862 for i in 0..n {
5863 ws.temp[i] += omega * l1_inv[i] * (r[i] - ws.work[i]);
5864 }
5865 }
5866 z.copy_from_slice(&ws.temp[..n]);
5867 Ok(())
5868 }
5869
5870 fn gs_forward(
5871 omega: f64,
5872 a: &CsrMatrix<f64>,
5873 diag_inv: &[f64],
5874 r: &[f64],
5875 z: &mut [f64],
5876 sweeps: usize,
5877 ) -> Result<(), KError> {
5878 let n = a.nrows();
5879 if diag_inv.len() != n || r.len() != n || z.len() != n {
5880 return Err(KError::InvalidInput("GS: dimension mismatch".into()));
5881 }
5882 for _ in 0..sweeps {
5883 for i in 0..n {
5884 let mut s = 0.0;
5885 let rs = a.row_ptr()[i];
5886 let re = a.row_ptr()[i + 1];
5887 for p in rs..re {
5888 s += a.values()[p] * z[a.col_idx()[p]];
5889 }
5890 z[i] += omega * diag_inv[i] * (r[i] - s);
5891 }
5892 }
5893 Ok(())
5894 }
5895
5896 fn l1_jacobi_mp(
5897 omega: f32,
5898 level: &AMGLevel,
5899 rhs: &[f64],
5900 z: &mut [f64],
5901 iters: usize,
5902 ws: &mut AMGWorkspace,
5903 cfg: &AMGConfig,
5904 ) -> Result<(), KError> {
5905 if iters == 0 {
5906 return Ok(());
5907 }
5908 let n = level.a.nrows();
5909 if rhs.len() != n || z.len() != n {
5910 return Err(KError::InvalidInput("L1Jacobi: dimension mismatch".into()));
5911 }
5912 ws.ensure(n);
5913 ws.ensure_mixed(n);
5914 let mp_ws = ws.mp.as_mut().expect("mixed workspace missing");
5915 mp_ws.temp32[..n]
5916 .iter_mut()
5917 .zip(z.iter())
5918 .for_each(|(dst, &src)| *dst = src as f32);
5919 mp_ws.residual32[..n]
5920 .iter_mut()
5921 .zip(rhs.iter())
5922 .for_each(|(dst, &src)| *dst = src as f32);
5923 let mut l1_owned = Vec::new();
5924 let l1_slice: &[f32] = match cfg.mixed_storage {
5925 MixedStorage::Cached => level
5926 .l1_inv_f32
5927 .as_ref()
5928 .ok_or_else(|| KError::InvalidInput("L1Jacobi mixed cache missing".into()))?,
5929 MixedStorage::Transient => {
5930 let buf = mp_ws.ensure_l1(n);
5931 buf.iter_mut()
5932 .zip(
5933 level
5934 .l1_inv
5935 .as_ref()
5936 .ok_or_else(|| KError::InvalidInput("L1Jacobi cache missing".into()))?
5937 .iter(),
5938 )
5939 .for_each(|(d, &s)| *d = s as f32);
5940 l1_owned.extend_from_slice(&buf[..n]);
5941 &l1_owned
5942 }
5943 };
5944 let mut vals_owned = Vec::new();
5945 let vals32: &[f32] = match cfg.mixed_storage {
5946 MixedStorage::Cached => level.a_vals_f32.as_ref().ok_or_else(|| {
5947 KError::InvalidInput("L1Jacobi mixed matrix cache missing".into())
5948 })?,
5949 MixedStorage::Transient => {
5950 let buf = mp_ws.ensure_vals(level.a.nnz());
5951 buf.iter_mut()
5952 .zip(level.a.values().iter())
5953 .for_each(|(d, &s)| *d = s as f32);
5954 vals_owned.extend_from_slice(buf.as_slice());
5955 &vals_owned
5956 }
5957 };
5958 let row_ptr = level.a.row_ptr();
5959 let col_idx = level.a.col_idx();
5960 for _ in 0..iters {
5961 spmv_scaled_f32_on_pattern(
5962 n,
5963 row_ptr,
5964 col_idx,
5965 vals32,
5966 1.0,
5967 &mp_ws.temp32[..n],
5968 0.0,
5969 &mut mp_ws.work32[..n],
5970 );
5971 for i in 0..n {
5972 mp_ws.temp32[i] += omega * l1_slice[i] * (mp_ws.residual32[i] - mp_ws.work32[i]);
5973 }
5974 }
5975 z.iter_mut()
5976 .zip(mp_ws.temp32[..n].iter())
5977 .for_each(|(dst, &src)| *dst = src as f64);
5978 Ok(())
5979 }
5980
5981 fn gs_backward(
5982 omega: f64,
5983 a: &CsrMatrix<f64>,
5984 diag_inv: &[f64],
5985 r: &[f64],
5986 z: &mut [f64],
5987 sweeps: usize,
5988 ) -> Result<(), KError> {
5989 let n = a.nrows();
5990 if diag_inv.len() != n || r.len() != n || z.len() != n {
5991 return Err(KError::InvalidInput("GS: dimension mismatch".into()));
5992 }
5993 for _ in 0..sweeps {
5994 for i in (0..n).rev() {
5995 let mut s = 0.0;
5996 let rs = a.row_ptr()[i];
5997 let re = a.row_ptr()[i + 1];
5998 for p in rs..re {
5999 s += a.values()[p] * z[a.col_idx()[p]];
6000 }
6001 z[i] += omega * diag_inv[i] * (r[i] - s);
6002 }
6003 }
6004 Ok(())
6005 }
6006
6007 fn sym_gs(
6008 omega: f64,
6009 a: &CsrMatrix<f64>,
6010 diag_inv: &[f64],
6011 r: &[f64],
6012 z: &mut [f64],
6013 sweeps: usize,
6014 ) -> Result<(), KError> {
6015 for _ in 0..sweeps {
6016 Self::gs_forward(omega, a, diag_inv, r, z, 1)?;
6017 Self::gs_backward(omega, a, diag_inv, r, z, 1)?;
6018 }
6019 Ok(())
6020 }
6021
6022 fn ilu0_smooth(
6023 omega: f64,
6024 level: &AMGLevel,
6025 r: &[f64],
6026 z: &mut [f64],
6027 sweeps: usize,
6028 ws: &mut AMGWorkspace,
6029 ) -> Result<(), KError> {
6030 #[cfg(feature = "complex")]
6031 {
6032 let _ = (omega, level, r, z, sweeps, ws);
6033 return Err(KError::Unsupported(
6034 "RelaxType::Ilu0 is not supported in complex AMG mode".into(),
6035 ));
6036 }
6037 #[cfg(not(feature = "complex"))]
6038 {
6039 if sweeps == 0 {
6040 return Ok(());
6041 }
6042 let ilu = level
6043 .ilu0
6044 .as_ref()
6045 .ok_or_else(|| KError::InvalidInput("ILU0 cache missing".into()))?;
6046 let n = level.a.nrows();
6047 ws.ensure(n);
6048 for _ in 0..sweeps {
6049 level.a.spmv_scaled(1.0, z, 0.0, &mut ws.work[..n])?;
6050 for i in 0..n {
6051 ws.residual[i] = r[i] - ws.work[i];
6052 }
6053 ws.temp[..n].fill(R::default());
6054 ilu.lock().expect("ILU0 mutex poisoned").apply(
6055 PcSide::Left,
6056 &ws.residual[..n],
6057 &mut ws.temp[..n],
6058 )?;
6059 for i in 0..n {
6060 z[i] += omega * ws.temp[i];
6061 }
6062 }
6063 Ok(())
6064 }
6065 }
6066
6067 fn ras_smooth(
6068 omega: f64,
6069 level: &AMGLevel,
6070 r: &[f64],
6071 z: &mut [f64],
6072 sweeps: usize,
6073 ws: &mut AMGWorkspace,
6074 ) -> Result<(), KError> {
6075 #[cfg(feature = "complex")]
6076 {
6077 let _ = (omega, level, r, z, sweeps, ws);
6078 return Err(KError::Unsupported(
6079 "RelaxType::Ras is not supported in complex AMG mode".into(),
6080 ));
6081 }
6082 #[cfg(not(feature = "complex"))]
6083 {
6084 if sweeps == 0 {
6085 return Ok(());
6086 }
6087 let ras = level
6088 .ras
6089 .as_ref()
6090 .ok_or_else(|| KError::InvalidInput("RAS cache missing".into()))?;
6091 let n = level.a.nrows();
6092 ws.ensure(n);
6093 for _ in 0..sweeps {
6094 level.a.spmv_scaled(1.0, z, 0.0, &mut ws.work[..n])?;
6095 for i in 0..n {
6096 ws.residual[i] = r[i] - ws.work[i];
6097 }
6098 ws.temp[..n].fill(R::default());
6099 ras.lock().expect("RAS mutex poisoned").apply(
6100 PcSide::Left,
6101 &ws.residual[..n],
6102 &mut ws.temp[..n],
6103 )?;
6104 for i in 0..n {
6105 z[i] += omega * ws.temp[i];
6106 }
6107 }
6108 Ok(())
6109 }
6110 }
6111
6112 fn apply_chebyshev(
6113 a: &CsrMatrix<f64>,
6114 d_inv: &[f64],
6115 rhs: &[f64],
6116 z: &mut [f64],
6117 degree: usize,
6118 data: &ChebData,
6119 ws: &mut AMGWorkspace,
6120 ) -> Result<(), KError> {
6121 if degree == 0 {
6122 return Ok(());
6123 }
6124 let n = a.nrows();
6125 ws.ensure(n);
6126 let bounds = ChebBounds {
6127 lam_max: data.lambda_max,
6128 lam_min: data.lambda_min,
6129 };
6130 chebyshev::chebyshev_smooth_csr(
6131 a,
6132 d_inv,
6133 rhs,
6134 z,
6135 degree,
6136 &bounds,
6137 &mut ws.residual[..n],
6138 &mut ws.temp[..n],
6139 &mut ws.work[..n],
6140 )
6141 }
6142
6143 fn fsai_smooth_core(
6144 g: &CsrMatrix<f64>,
6145 gt: &CsrMatrix<f64>,
6146 a: &CsrMatrix<f64>,
6147 rhs: &[f64],
6148 z: &mut [f64],
6149 tau: f64,
6150 residual: &mut [f64],
6151 work: &mut [f64],
6152 tmp: &mut [f64],
6153 ) -> Result<(), KError> {
6154 let n = a.nrows();
6155 a.spmv_scaled(1.0, z, 0.0, work)?;
6156 for i in 0..n {
6157 residual[i] = rhs[i] - work[i];
6158 }
6159 gt.spmv_scaled(1.0, residual, 0.0, tmp)?;
6160 g.spmv_scaled(1.0, tmp, 0.0, work)?;
6161 for i in 0..n {
6162 z[i] += tau * work[i];
6163 }
6164 Ok(())
6165 }
6166
6167 fn chebyshev_smooth_csr_mp(
6168 level: &AMGLevel,
6169 rhs: &[f64],
6170 z: &mut [f64],
6171 degree: usize,
6172 data: &ChebData,
6173 ws: &mut AMGWorkspace,
6174 cfg: &AMGConfig,
6175 ) -> Result<(), KError> {
6176 if degree == 0 {
6177 return Ok(());
6178 }
6179 let n = level.a.nrows();
6180 if rhs.len() != n || z.len() != n {
6181 return Err(KError::InvalidInput("Chebyshev: dimension mismatch".into()));
6182 }
6183 ws.ensure(n);
6184 ws.ensure_mixed(n);
6185 let mp_ws = ws.mp.as_mut().expect("mixed workspace missing");
6186 mp_ws.temp32[..n]
6187 .iter_mut()
6188 .zip(z.iter())
6189 .for_each(|(dst, &src)| *dst = src as f32);
6190 mp_ws.residual32[..n]
6191 .iter_mut()
6192 .zip(rhs.iter())
6193 .for_each(|(dst, &src)| *dst = src as f32);
6194
6195 let mut vals_owned = Vec::new();
6196 let vals32: &[f32] = match cfg.mixed_storage {
6197 MixedStorage::Cached => level.a_vals_f32.as_ref().ok_or_else(|| {
6198 KError::InvalidInput("Chebyshev mixed matrix cache missing".into())
6199 })?,
6200 MixedStorage::Transient => {
6201 let buf = mp_ws.ensure_vals(level.a.nnz());
6202 buf.iter_mut()
6203 .zip(level.a.values().iter())
6204 .for_each(|(d, &s)| *d = s as f32);
6205 vals_owned.extend_from_slice(buf.as_slice());
6206 &vals_owned
6207 }
6208 };
6209 let mut diag_owned = Vec::new();
6210 let diag32: &[f32] = match cfg.mixed_storage {
6211 MixedStorage::Cached => level
6212 .diag_inv_f32
6213 .as_ref()
6214 .ok_or_else(|| KError::InvalidInput("Chebyshev mixed diag cache missing".into()))?,
6215 MixedStorage::Transient => {
6216 let buf = mp_ws.ensure_diag(n);
6217 buf.iter_mut()
6218 .zip(level.diag_inv.iter())
6219 .for_each(|(d, &s)| *d = s as f32);
6220 diag_owned.extend_from_slice(&buf[..n]);
6221 &diag_owned
6222 }
6223 };
6224
6225 let row_ptr = level.a.row_ptr();
6226 let col_idx = level.a.col_idx();
6227 spmv_scaled_f32_on_pattern(
6228 n,
6229 row_ptr,
6230 col_idx,
6231 vals32,
6232 1.0,
6233 &mp_ws.temp32[..n],
6234 0.0,
6235 &mut mp_ws.work32[..n],
6236 );
6237 for i in 0..n {
6238 mp_ws.residual32[i] -= mp_ws.work32[i];
6239 }
6240
6241 let theta = (0.5 * (data.lambda_max + data.lambda_min)).max(1e-12) as f32;
6242 let delta = (0.5 * (data.lambda_max - data.lambda_min)) as f32;
6243 let mut alpha = 1.0f32 / theta;
6244
6245 for i in 0..n {
6246 mp_ws.fine_corr32[i] = diag32[i] * mp_ws.residual32[i];
6247 }
6248 for i in 0..n {
6249 mp_ws.temp32[i] += alpha * mp_ws.fine_corr32[i];
6250 }
6251 spmv_scaled_f32_on_pattern(
6252 n,
6253 row_ptr,
6254 col_idx,
6255 vals32,
6256 1.0,
6257 &mp_ws.fine_corr32[..n],
6258 0.0,
6259 &mut mp_ws.work32[..n],
6260 );
6261 for i in 0..n {
6262 mp_ws.residual32[i] -= alpha * mp_ws.work32[i];
6263 }
6264
6265 for _ in 1..degree {
6266 for i in 0..n {
6267 mp_ws.fine_corr32[i] = diag32[i] * mp_ws.residual32[i];
6268 }
6269 let beta = 0.25f32 * delta * delta * alpha;
6270 alpha = 1.0f32 / (theta - beta);
6271 for i in 0..n {
6272 mp_ws.temp32[i] += alpha * mp_ws.fine_corr32[i];
6273 }
6274 spmv_scaled_f32_on_pattern(
6275 n,
6276 row_ptr,
6277 col_idx,
6278 vals32,
6279 1.0,
6280 &mp_ws.fine_corr32[..n],
6281 0.0,
6282 &mut mp_ws.work32[..n],
6283 );
6284 for i in 0..n {
6285 mp_ws.residual32[i] -= alpha * mp_ws.work32[i];
6286 }
6287 }
6288
6289 z.iter_mut()
6290 .zip(mp_ws.temp32[..n].iter())
6291 .for_each(|(dst, &src)| *dst = src as f64);
6292 Ok(())
6293 }
6294
6295 fn fsai_smooth(
6296 g: &CsrMatrix<f64>,
6297 gt: &CsrMatrix<f64>,
6298 a: &CsrMatrix<f64>,
6299 rhs: &[f64],
6300 z: &mut [f64],
6301 tau: f64,
6302 ws: &mut AMGWorkspace,
6303 ) -> Result<(), KError> {
6304 let n = a.nrows();
6305 ws.ensure(n);
6306 let residual = &mut ws.residual[..n];
6307 let work = &mut ws.work[..n];
6308 let tmp = &mut ws.fine_corr[..n];
6309 Self::fsai_smooth_core(g, gt, a, rhs, z, tau, residual, work, tmp)
6310 }
6311
6312 fn fsai_smooth_mp(
6313 level: &AMGLevel,
6314 data: &FsaiData,
6315 rhs: &[f64],
6316 z: &mut [f64],
6317 tau: f32,
6318 ws: &mut AMGWorkspace,
6319 cfg: &AMGConfig,
6320 ) -> Result<(), KError> {
6321 let n = level.a.nrows();
6322 if rhs.len() != n || z.len() != n {
6323 return Err(KError::InvalidInput("FSAI: dimension mismatch".into()));
6324 }
6325 ws.ensure(n);
6326 ws.ensure_mixed(n);
6327 let mp_ws = ws.mp.as_mut().expect("mixed workspace missing");
6328 mp_ws.temp32[..n]
6329 .iter_mut()
6330 .zip(z.iter())
6331 .for_each(|(dst, &src)| *dst = src as f32);
6332 mp_ws.residual32[..n]
6333 .iter_mut()
6334 .zip(rhs.iter())
6335 .for_each(|(dst, &src)| *dst = src as f32);
6336
6337 let mut a_vals_owned = Vec::new();
6338 let a_vals32: &[f32] = match cfg.mixed_storage {
6339 MixedStorage::Cached => level
6340 .a_vals_f32
6341 .as_ref()
6342 .ok_or_else(|| KError::InvalidInput("FSAI mixed matrix cache missing".into()))?,
6343 MixedStorage::Transient => {
6344 let buf = mp_ws.ensure_vals(level.a.nnz());
6345 buf.iter_mut()
6346 .zip(level.a.values().iter())
6347 .for_each(|(d, &s)| *d = s as f32);
6348 a_vals_owned.extend_from_slice(buf.as_slice());
6349 &a_vals_owned
6350 }
6351 };
6352 let row_ptr = level.a.row_ptr();
6353 let col_idx = level.a.col_idx();
6354 spmv_scaled_f32_on_pattern(
6355 n,
6356 row_ptr,
6357 col_idx,
6358 a_vals32,
6359 1.0,
6360 &mp_ws.temp32[..n],
6361 0.0,
6362 &mut mp_ws.work32[..n],
6363 );
6364 for i in 0..n {
6365 mp_ws.residual32[i] -= mp_ws.work32[i];
6366 }
6367
6368 let mut g_vals_owned = Vec::new();
6369 let g_vals32: &[f32] = match cfg.mixed_storage {
6370 MixedStorage::Cached => level
6371 .fsai_g_vals_f32
6372 .as_ref()
6373 .ok_or_else(|| KError::InvalidInput("FSAI mixed G cache missing".into()))?,
6374 MixedStorage::Transient => {
6375 let buf = mp_ws.ensure_fsai_g(data.g.nnz());
6376 buf.iter_mut()
6377 .zip(data.g.values().iter())
6378 .for_each(|(d, &s)| *d = s as f32);
6379 g_vals_owned.extend_from_slice(buf.as_slice());
6380 &g_vals_owned
6381 }
6382 };
6383 spmv_scaled_f32_on_pattern(
6384 data.g.nrows(),
6385 data.g.row_ptr(),
6386 data.g.col_idx(),
6387 g_vals32,
6388 1.0,
6389 &mp_ws.residual32[..data.g.ncols()],
6390 0.0,
6391 &mut mp_ws.work32[..data.g.nrows()],
6392 );
6393
6394 let mut gt_vals_owned = Vec::new();
6395 let gt_vals32: &[f32] = match cfg.mixed_storage {
6396 MixedStorage::Cached => level
6397 .fsai_gt_vals_f32
6398 .as_ref()
6399 .ok_or_else(|| KError::InvalidInput("FSAI mixed Gt cache missing".into()))?,
6400 MixedStorage::Transient => {
6401 let buf = mp_ws.ensure_fsai_gt(data.gt.nnz());
6402 buf.iter_mut()
6403 .zip(data.gt.values().iter())
6404 .for_each(|(d, &s)| *d = s as f32);
6405 gt_vals_owned.extend_from_slice(buf.as_slice());
6406 >_vals_owned
6407 }
6408 };
6409 spmv_scaled_f32_on_pattern(
6410 data.gt.nrows(),
6411 data.gt.row_ptr(),
6412 data.gt.col_idx(),
6413 gt_vals32,
6414 1.0,
6415 &mp_ws.work32[..data.gt.ncols()],
6416 0.0,
6417 &mut mp_ws.coarse_rhs32[..data.gt.nrows()],
6418 );
6419 for i in 0..n {
6420 mp_ws.temp32[i] += tau * mp_ws.coarse_rhs32[i];
6421 }
6422 z.iter_mut()
6423 .zip(mp_ws.temp32[..n].iter())
6424 .for_each(|(dst, &src)| *dst = src as f64);
6425 Ok(())
6426 }
6427
6428 fn mixed_spmv(
6429 level: &AMGLevel,
6430 cfg: &AMGConfig,
6431 x: &[f64],
6432 ws: &mut AMGWorkspace,
6433 ) -> Result<(), KError> {
6434 let n = level.a.nrows();
6435 if x.len() != n {
6436 return Err(KError::InvalidInput(
6437 "mixed SpMV: dimension mismatch".into(),
6438 ));
6439 }
6440 ws.ensure_mixed(n);
6441 let mp_ws = ws.mp.as_mut().expect("mixed workspace missing");
6442 mp_ws.temp32[..n]
6443 .iter_mut()
6444 .zip(x.iter())
6445 .for_each(|(dst, &src)| *dst = src as f32);
6446 let mut vals_owned = Vec::new();
6447 let vals32: &[f32] = match cfg.mixed_storage {
6448 MixedStorage::Cached => level
6449 .a_vals_f32
6450 .as_ref()
6451 .ok_or_else(|| KError::InvalidInput("mixed matrix cache missing".into()))?,
6452 MixedStorage::Transient => {
6453 let buf = mp_ws.ensure_vals(level.a.nnz());
6454 buf.iter_mut()
6455 .zip(level.a.values().iter())
6456 .for_each(|(d, &s)| *d = s as f32);
6457 vals_owned.extend_from_slice(buf.as_slice());
6458 &vals_owned
6459 }
6460 };
6461 spmv_scaled_f32_on_pattern(
6462 n,
6463 level.a.row_ptr(),
6464 level.a.col_idx(),
6465 vals32,
6466 1.0,
6467 &mp_ws.temp32[..n],
6468 0.0,
6469 &mut mp_ws.work32[..n],
6470 );
6471 Ok(())
6472 }
6473
6474 fn flexible_relax_supported(relax: RelaxType) -> bool {
6475 matches!(
6476 relax,
6477 RelaxType::Jacobi
6478 | RelaxType::L1Jacobi
6479 | RelaxType::SymmetricGaussSeidel
6480 | RelaxType::Fsai
6481 )
6482 }
6483
6484 fn apply_smoother_as_pc(
6485 &self,
6486 relax: RelaxType,
6487 lvl: &AMGLevel,
6488 sweeps: usize,
6489 omega: f64,
6490 ws: &mut AMGWorkspace,
6491 ) -> Result<(), KError> {
6492 let n = lvl.a.nrows();
6493 ws.ensure(n);
6494 let r = &ws.residual[..n];
6495 let z = &mut ws.fine_corr[..n];
6496 let work = &mut ws.work[..n];
6497 if sweeps == 0 {
6498 z.copy_from_slice(r);
6499 return Ok(());
6500 }
6501 match relax {
6502 RelaxType::Jacobi => {
6503 z.fill(R::zero());
6504 for _ in 0..sweeps {
6505 lvl.a.spmv_scaled(1.0, &z[..n], 0.0, work)?;
6506 for i in 0..n {
6507 z[i] += omega * lvl.diag_inv[i] * (r[i] - work[i]);
6508 }
6509 }
6510 Ok(())
6511 }
6512 RelaxType::L1Jacobi => {
6513 let l1 = lvl
6514 .l1_inv
6515 .as_ref()
6516 .ok_or_else(|| KError::InvalidInput("L1Jacobi cache missing".into()))?;
6517 if l1.len() != n {
6518 return Err(KError::InvalidInput(
6519 "L1Jacobi cache has incorrect length".into(),
6520 ));
6521 }
6522 z.fill(R::zero());
6523 for _ in 0..sweeps {
6524 lvl.a.spmv_scaled(1.0, &z[..n], 0.0, work)?;
6525 for i in 0..n {
6526 z[i] += omega * l1[i] * (r[i] - work[i]);
6527 }
6528 }
6529 Ok(())
6530 }
6531 RelaxType::SymmetricGaussSeidel => {
6532 z.fill(R::zero());
6533 Self::sym_gs(omega, &lvl.a, &lvl.diag_inv, r, z, sweeps)?;
6534 Ok(())
6535 }
6536 RelaxType::Fsai => {
6537 let data = lvl
6538 .fsai
6539 .as_ref()
6540 .ok_or_else(|| KError::InvalidInput("FSAI cache missing".into()))?;
6541 ws.fine_corr[..n].fill(R::zero());
6542 for _ in 0..sweeps {
6543 Self::fsai_smooth_core(
6544 &data.g,
6545 &data.gt,
6546 &lvl.a,
6547 r,
6548 &mut ws.fine_corr[..n],
6549 self.cfg.fsai_damping,
6550 &mut ws.work[..n],
6551 &mut ws.coarse_rhs[..n],
6552 &mut ws.temp[..n],
6553 )?;
6554 }
6555 Ok(())
6556 }
6557 other => Err(KError::InvalidInput(format!(
6558 "RelaxType {other:?} cannot be used as a flexible preconditioner"
6559 ))),
6560 }
6561 }
6562
6563 fn fcg_presmooth(
6564 &self,
6565 lvl: &AMGLevel,
6566 rhs: &[f64],
6567 sol: &mut [f64],
6568 iters: usize,
6569 rtol: f64,
6570 pc_sweeps: usize,
6571 relax: RelaxType,
6572 omega: f64,
6573 ws: &mut AMGWorkspace,
6574 ) -> Result<(), KError> {
6575 if iters == 0 {
6576 return Ok(());
6577 }
6578 let n = lvl.a.nrows();
6579 if rhs.len() != n || sol.len() != n {
6580 return Err(KError::InvalidInput(
6581 "fcg_presmooth: dimension mismatch".into(),
6582 ));
6583 }
6584 ws.ensure(n);
6585
6586 lvl.a.spmv_scaled(1.0, sol, 0.0, &mut ws.work[..n])?;
6587 for i in 0..n {
6588 ws.residual[i] = rhs[i] - ws.work[i];
6589 }
6590
6591 self.apply_smoother_as_pc(relax, lvl, pc_sweeps, omega, ws)?;
6592 ws.temp[..n].copy_from_slice(&ws.fine_corr[..n]);
6593 let mut rho = dot(&ws.residual[..n], &ws.fine_corr[..n]);
6594 if !rho.is_finite() {
6595 return Err(KError::InvalidInput(
6596 "fcg_presmooth: non-finite rho encountered".into(),
6597 ));
6598 }
6599 let mut rnorm_sq = dot(&ws.residual[..n], &ws.residual[..n]);
6600 let mut tol = -1.0f64;
6601 if rtol > 0.0 {
6602 let base = rnorm_sq.sqrt().max(1e-30);
6603 tol = base * rtol.max(1e-15);
6604 if base <= tol {
6605 return Ok(());
6606 }
6607 }
6608
6609 for _ in 0..iters {
6610 lvl.a
6611 .spmv_scaled(1.0, &ws.temp[..n], 0.0, &mut ws.work[..n])?;
6612 let p_ap = dot(&ws.temp[..n], &ws.work[..n]);
6613 if !p_ap.is_finite() || p_ap.abs() < 1e-300 {
6614 break;
6615 }
6616
6617 let alpha = rho / p_ap;
6618 for i in 0..n {
6619 sol[i] += alpha * ws.temp[i];
6620 ws.residual[i] -= alpha * ws.work[i];
6621 }
6622
6623 if tol > 0.0 {
6624 rnorm_sq = dot(&ws.residual[..n], &ws.residual[..n]);
6625 if rnorm_sq.sqrt() <= tol {
6626 break;
6627 }
6628 }
6629
6630 ws.coarse_rhs[..n].copy_from_slice(&ws.fine_corr[..n]);
6631 self.apply_smoother_as_pc(relax, lvl, pc_sweeps, omega, ws)?;
6632 let z = &ws.fine_corr[..n];
6633 let mut rz_diff = R::default();
6634 for i in 0..n {
6635 rz_diff += ws.residual[i] * (z[i] - ws.coarse_rhs[i]);
6636 }
6637 let beta = if rho.abs() > R::default() {
6638 rz_diff / rho
6639 } else {
6640 R::default()
6641 };
6642 for i in 0..n {
6643 ws.temp[i] = z[i] + beta * ws.temp[i];
6644 }
6645 let rho_new = dot(&ws.residual[..n], z);
6646 if !rho_new.is_finite() {
6647 break;
6648 }
6649 rho = rho_new;
6650 }
6651 Ok(())
6652 }
6653
6654 fn apply_relax(
6656 pol: &RelaxPolicy,
6657 phase: RelaxPhase,
6658 where_: RelaxWhere,
6659 lvl: &AMGLevel,
6660 rhs: &[f64],
6661 sol: &mut [f64],
6662 ws: &mut AMGWorkspace,
6663 cfg: &AMGConfig,
6664 ) -> Result<(), KError> {
6665 let k = pol.sweeps[phase.ix()];
6666 if k == 0 {
6667 return Ok(());
6668 }
6669 let use_mp_smooth = cfg
6670 .mixed_precision
6671 .map(|mp| mp.smoothers_enabled())
6672 .unwrap_or(false);
6673 #[cfg(test)]
6674 {
6675 RELAX_CALL_COUNTS.with(|counts| {
6676 let mut data = counts.get();
6677 data[phase.ix()] += 1;
6678 counts.set(data);
6679 });
6680 }
6681 let a = &lvl.a;
6682 match pol.kind[phase.ix()] {
6683 RelaxType::Jacobi => {
6684 if use_mp_smooth {
6685 Self::jacobi_smooth_sparse_mp(pol.omega as f32, lvl, rhs, sol, k, ws, cfg)
6686 } else {
6687 Self::jacobi_smooth_sparse(pol.omega, a, &lvl.diag_inv, rhs, sol, k, ws)
6688 }
6689 }
6690 RelaxType::GaussSeidel => {
6691 if matches!(where_, RelaxWhere::Pre) {
6692 Self::gs_forward(1.0, a, &lvl.diag_inv, rhs, sol, k)
6693 } else {
6694 Self::gs_backward(1.0, a, &lvl.diag_inv, rhs, sol, k)
6695 }
6696 }
6697 RelaxType::SafeguardedGaussSeidel => {
6698 let diag = lvl
6699 .diag_inv_safe
6700 .as_ref()
6701 .ok_or_else(|| KError::InvalidInput("Safeguarded GS cache missing".into()))?;
6702 if matches!(where_, RelaxWhere::Pre) {
6703 Self::gs_forward(1.0, a, diag, rhs, sol, k)
6704 } else {
6705 Self::gs_backward(1.0, a, diag, rhs, sol, k)
6706 }
6707 }
6708 RelaxType::GaussSeidelBackward => Self::gs_backward(1.0, a, &lvl.diag_inv, rhs, sol, k),
6709 RelaxType::SymmetricGaussSeidel => Self::sym_gs(1.0, a, &lvl.diag_inv, rhs, sol, k),
6710 RelaxType::L1Jacobi => {
6711 if let Some(ref l1) = lvl.l1_inv {
6712 if use_mp_smooth {
6713 Self::l1_jacobi_mp(pol.omega as f32, lvl, rhs, sol, k, ws, cfg)
6714 } else {
6715 Self::l1_jacobi(pol.omega, a, l1, rhs, sol, k, ws)
6716 }
6717 } else {
6718 Err(KError::InvalidInput("L1Jacobi cache missing".into()))
6719 }
6720 }
6721 RelaxType::Chebyshev => {
6722 let cheb = lvl
6723 .cheb
6724 .as_ref()
6725 .ok_or_else(|| KError::InvalidInput("Chebyshev cache missing".into()))?;
6726 let degree = cfg.chebyshev_degree.max(1);
6727 for _ in 0..k {
6728 if use_mp_smooth {
6729 Self::chebyshev_smooth_csr_mp(lvl, rhs, sol, degree, cheb, ws, cfg)?;
6730 } else {
6731 Self::apply_chebyshev(a, &lvl.diag_inv, rhs, sol, degree, cheb, ws)?;
6732 }
6733 }
6734 Ok(())
6735 }
6736 RelaxType::ChebyshevSafe => {
6737 if use_mp_smooth {
6738 return Err(KError::InvalidInput(
6739 "ChebyshevSafe does not support mixed-precision smoothing".into(),
6740 ));
6741 }
6742 let cheb = lvl
6743 .cheb_safe
6744 .as_ref()
6745 .ok_or_else(|| KError::InvalidInput("ChebyshevSafe cache missing".into()))?;
6746 let diag = lvl.diag_inv_safe.as_ref().ok_or_else(|| {
6747 KError::InvalidInput("ChebyshevSafe diag cache missing".into())
6748 })?;
6749 let degree = cfg.chebyshev_degree.max(1);
6750 for _ in 0..k {
6751 Self::apply_chebyshev(a, diag, rhs, sol, degree, cheb, ws)?;
6752 }
6753 Ok(())
6754 }
6755 RelaxType::Ilu0 => Self::ilu0_smooth(pol.omega, lvl, rhs, sol, k, ws),
6756 RelaxType::Ras => Self::ras_smooth(pol.omega, lvl, rhs, sol, k, ws),
6757 RelaxType::Fsai => {
6758 let data = lvl
6759 .fsai
6760 .as_ref()
6761 .ok_or_else(|| KError::InvalidInput("FSAI cache missing".into()))?;
6762 for _ in 0..k {
6763 if use_mp_smooth {
6764 Self::fsai_smooth_mp(
6765 lvl,
6766 data,
6767 rhs,
6768 sol,
6769 cfg.fsai_damping as f32,
6770 ws,
6771 cfg,
6772 )?;
6773 } else {
6774 Self::fsai_smooth(&data.g, &data.gt, a, rhs, sol, cfg.fsai_damping, ws)?;
6775 }
6776 }
6777 Ok(())
6778 }
6779 other => Err(KError::InvalidInput(format!(
6780 "RelaxType {other:?} not yet supported"
6781 ))),
6782 }
6783 }
6784
6785 fn smooth_dispatch(
6788 relax: RelaxType,
6789 sweeps: usize,
6790 lvl: &AMGLevel,
6791 rhs: &[f64],
6792 sol: &mut [f64],
6793 ws: &mut AMGWorkspace,
6794 cfg: &AMGConfig,
6795 ) -> Result<(), KError> {
6796 if sweeps == 0 {
6797 return Ok(());
6798 }
6799 match relax {
6800 RelaxType::Jacobi => Self::jacobi_smooth_sparse(
6801 cfg.jacobi_omega,
6802 &lvl.a,
6803 &lvl.diag_inv,
6804 rhs,
6805 sol,
6806 sweeps,
6807 ws,
6808 ),
6809 RelaxType::GaussSeidel => {
6810 Self::gs_forward(1.0, &lvl.a, &lvl.diag_inv, rhs, sol, sweeps)
6811 }
6812 RelaxType::SafeguardedGaussSeidel => {
6813 let diag = lvl
6814 .diag_inv_safe
6815 .as_ref()
6816 .ok_or_else(|| KError::InvalidInput("Safeguarded GS cache missing".into()))?;
6817 Self::gs_forward(1.0, &lvl.a, diag, rhs, sol, sweeps)
6818 }
6819 RelaxType::GaussSeidelBackward => {
6820 Self::gs_backward(1.0, &lvl.a, &lvl.diag_inv, rhs, sol, sweeps)
6821 }
6822 RelaxType::SymmetricGaussSeidel => {
6823 Self::sym_gs(1.0, &lvl.a, &lvl.diag_inv, rhs, sol, sweeps)
6824 }
6825 RelaxType::L1Jacobi => {
6826 let l1 = lvl
6827 .l1_inv
6828 .as_ref()
6829 .ok_or_else(|| KError::InvalidInput("L1Jacobi cache missing".into()))?;
6830 Self::l1_jacobi(cfg.jacobi_omega, &lvl.a, l1, rhs, sol, sweeps, ws)
6831 }
6832 RelaxType::Chebyshev => {
6833 let cheb = lvl
6834 .cheb
6835 .as_ref()
6836 .ok_or_else(|| KError::InvalidInput("Chebyshev cache missing".into()))?;
6837 let degree = cfg.chebyshev_degree.max(1);
6838 for _ in 0..sweeps {
6839 Self::apply_chebyshev(&lvl.a, &lvl.diag_inv, rhs, sol, degree, cheb, ws)?;
6840 }
6841 Ok(())
6842 }
6843 RelaxType::ChebyshevSafe => {
6844 let cheb = lvl
6845 .cheb_safe
6846 .as_ref()
6847 .ok_or_else(|| KError::InvalidInput("ChebyshevSafe cache missing".into()))?;
6848 let diag = lvl.diag_inv_safe.as_ref().ok_or_else(|| {
6849 KError::InvalidInput("ChebyshevSafe diag cache missing".into())
6850 })?;
6851 let degree = cfg.chebyshev_degree.max(1);
6852 for _ in 0..sweeps {
6853 Self::apply_chebyshev(&lvl.a, diag, rhs, sol, degree, cheb, ws)?;
6854 }
6855 Ok(())
6856 }
6857 RelaxType::Ilu0 => Self::ilu0_smooth(cfg.jacobi_omega, lvl, rhs, sol, sweeps, ws),
6858 RelaxType::Ras => Self::ras_smooth(cfg.jacobi_omega, lvl, rhs, sol, sweeps, ws),
6859 RelaxType::Fsai => {
6860 let data = lvl
6861 .fsai
6862 .as_ref()
6863 .ok_or_else(|| KError::InvalidInput("FSAI cache missing".into()))?;
6864 for _ in 0..sweeps {
6865 Self::fsai_smooth(&data.g, &data.gt, &lvl.a, rhs, sol, cfg.fsai_damping, ws)?;
6866 }
6867 Ok(())
6868 }
6869 other => Err(KError::InvalidInput(format!(
6870 "RelaxType {other:?} not yet supported",
6871 ))),
6872 }
6873 }
6874
6875 fn apply_precond_one_sweep(
6876 &self,
6877 level: usize,
6878 r: &[f64],
6879 out: &mut [f64],
6880 work: &mut [f64],
6881 temp: &mut [f64],
6882 residual: &mut [f64],
6883 ) -> Result<(), KError> {
6884 let h = match &self.state {
6885 AmgState::Ready { hierarchy, .. } => hierarchy,
6886 _ => return Err(KError::InvalidInput("AMG not set up".into())),
6887 };
6888 let lvl = h
6889 .levels
6890 .get(level)
6891 .ok_or_else(|| KError::InvalidInput("level out of range".into()))?;
6892 let n = lvl.a.nrows();
6893 if r.len() != n
6894 || out.len() != n
6895 || work.len() != n
6896 || temp.len() != n
6897 || residual.len() != n
6898 {
6899 return Err(KError::InvalidInput(
6900 "krylov preconditioner: level size mismatch".into(),
6901 ));
6902 }
6903 out.fill(R::default());
6904
6905 match self.cfg.relax_type {
6906 RelaxType::Jacobi => {
6907 for i in 0..n {
6908 out[i] = self.cfg.jacobi_omega * lvl.diag_inv[i] * r[i];
6909 }
6910 Ok(())
6911 }
6912 RelaxType::GaussSeidel => Self::gs_forward(1.0, &lvl.a, &lvl.diag_inv, r, out, 1),
6913 RelaxType::SafeguardedGaussSeidel => {
6914 let diag = lvl
6915 .diag_inv_safe
6916 .as_ref()
6917 .ok_or_else(|| KError::InvalidInput("Safeguarded GS cache missing".into()))?;
6918 Self::gs_forward(1.0, &lvl.a, diag, r, out, 1)
6919 }
6920 RelaxType::GaussSeidelBackward => {
6921 Self::gs_backward(1.0, &lvl.a, &lvl.diag_inv, r, out, 1)
6922 }
6923 RelaxType::SymmetricGaussSeidel => Self::sym_gs(1.0, &lvl.a, &lvl.diag_inv, r, out, 1),
6924 RelaxType::L1Jacobi => {
6925 let l1 = lvl
6926 .l1_inv
6927 .as_ref()
6928 .ok_or_else(|| KError::InvalidInput("L1Jacobi cache missing".into()))?;
6929 work.fill(R::default());
6930 lvl.a.spmv_scaled(1.0, out, 0.0, work)?;
6931 for i in 0..n {
6932 out[i] += self.cfg.jacobi_omega * l1[i] * (r[i] - work[i]);
6933 }
6934 Ok(())
6935 }
6936 RelaxType::Chebyshev => {
6937 let cheb = lvl
6938 .cheb
6939 .as_ref()
6940 .ok_or_else(|| KError::InvalidInput("Chebyshev cache missing".into()))?;
6941 let bounds = ChebBounds {
6942 lam_max: cheb.lambda_max,
6943 lam_min: cheb.lambda_min,
6944 };
6945 chebyshev::chebyshev_smooth_csr(
6946 &lvl.a,
6947 &lvl.diag_inv,
6948 r,
6949 out,
6950 self.cfg.chebyshev_degree.max(1),
6951 &bounds,
6952 residual,
6953 temp,
6954 work,
6955 )
6956 }
6957 RelaxType::ChebyshevSafe => {
6958 let cheb = lvl
6959 .cheb_safe
6960 .as_ref()
6961 .ok_or_else(|| KError::InvalidInput("ChebyshevSafe cache missing".into()))?;
6962 let diag = lvl.diag_inv_safe.as_ref().ok_or_else(|| {
6963 KError::InvalidInput("ChebyshevSafe diag cache missing".into())
6964 })?;
6965 let bounds = ChebBounds {
6966 lam_max: cheb.lambda_max,
6967 lam_min: cheb.lambda_min,
6968 };
6969 chebyshev::chebyshev_smooth_csr(
6970 &lvl.a,
6971 diag,
6972 r,
6973 out,
6974 self.cfg.chebyshev_degree.max(1),
6975 &bounds,
6976 residual,
6977 temp,
6978 work,
6979 )
6980 }
6981 #[cfg(not(feature = "complex"))]
6982 RelaxType::Ilu0 => {
6983 let ilu = lvl
6984 .ilu0
6985 .as_ref()
6986 .ok_or_else(|| KError::InvalidInput("ILU0 cache missing".into()))?;
6987 ilu.lock()
6988 .expect("ILU0 mutex poisoned")
6989 .apply(PcSide::Left, r, out)
6990 }
6991 #[cfg(feature = "complex")]
6992 RelaxType::Ilu0 => Err(KError::Unsupported(
6993 "RelaxType::Ilu0 is not supported in complex AMG mode".into(),
6994 )),
6995 #[cfg(not(feature = "complex"))]
6996 RelaxType::Ras => {
6997 let ras = lvl
6998 .ras
6999 .as_ref()
7000 .ok_or_else(|| KError::InvalidInput("RAS cache missing".into()))?;
7001 ras.lock()
7002 .expect("RAS mutex poisoned")
7003 .apply(PcSide::Left, r, out)
7004 }
7005 #[cfg(feature = "complex")]
7006 RelaxType::Ras => Err(KError::Unsupported(
7007 "RelaxType::Ras is not supported in complex AMG mode".into(),
7008 )),
7009 RelaxType::Fsai => {
7010 let data = lvl
7011 .fsai
7012 .as_ref()
7013 .ok_or_else(|| KError::InvalidInput("FSAI cache missing".into()))?;
7014 residual.fill(R::default());
7015 work.fill(R::default());
7016 temp.fill(R::default());
7017 Self::fsai_smooth_core(
7018 &data.g,
7019 &data.gt,
7020 &lvl.a,
7021 r,
7022 out,
7023 self.cfg.fsai_damping,
7024 residual,
7025 work,
7026 temp,
7027 )
7028 }
7029 other => Err(KError::InvalidInput(format!(
7030 "RelaxType {other:?} not yet supported",
7031 ))),
7032 }
7033 }
7034 fn krylov_smooth(
7035 &self,
7036 algo: KrylovAlgo,
7037 iters: usize,
7038 level: usize,
7039 rhs: &[f64],
7040 sol: &mut [f64],
7041 ws: &mut AMGWorkspace,
7042 ) -> Result<(), KError> {
7043 if iters == 0 {
7044 return Ok(());
7045 }
7046 match algo {
7047 KrylovAlgo::FCG => self.krylov_smooth_pcg(iters, level, rhs, sol, ws),
7048 }
7049 }
7050
7051 fn krylov_smooth_pcg(
7052 &self,
7053 iters: usize,
7054 level: usize,
7055 rhs: &[f64],
7056 sol: &mut [f64],
7057 ws: &mut AMGWorkspace,
7058 ) -> Result<(), KError> {
7059 if iters == 0 {
7060 return Ok(());
7061 }
7062 let h = match &self.state {
7063 AmgState::Ready { hierarchy, .. } => hierarchy,
7064 _ => return Err(KError::InvalidInput("AMG not set up".into())),
7065 };
7066 let lvl = h
7067 .levels
7068 .get(level)
7069 .ok_or_else(|| KError::InvalidInput("level out of range".into()))?;
7070 let a = &lvl.a;
7071 let n = a.nrows();
7072 if rhs.len() != n || sol.len() != n {
7073 return Err(KError::InvalidInput(
7074 "krylov smoother: dimension mismatch".into(),
7075 ));
7076 }
7077 ws.ensure(n);
7078
7079 a.spmv_scaled(1.0, sol, 0.0, &mut ws.work[..n])?;
7080 for i in 0..n {
7081 ws.residual[i] = rhs[i] - ws.work[i];
7082 }
7083
7084 ws.k_residual[..n].copy_from_slice(&ws.residual[..n]);
7085 self.apply_precond_one_sweep(
7086 level,
7087 &ws.k_residual[..n],
7088 &mut ws.k_zeta[..n],
7089 &mut ws.k_work[..n],
7090 &mut ws.k_temp[..n],
7091 &mut ws.k_ap[..n],
7092 )?;
7093 ws.k_p[..n].copy_from_slice(&ws.k_zeta[..n]);
7094
7095 let mut rz_old = dot(&ws.residual[..n], &ws.k_zeta[..n]);
7096 if !rz_old.is_finite() || rz_old.abs() < 1e-300 {
7097 return Ok(());
7098 }
7099
7100 for _ in 0..iters {
7101 a.spmv_scaled(1.0, &ws.k_p[..n], 0.0, &mut ws.k_ap[..n])?;
7102 let denom = dot(&ws.k_p[..n], &ws.k_ap[..n]);
7103 if !denom.is_finite() || denom.abs() < 1e-300 {
7104 break;
7105 }
7106 let alpha = rz_old / denom;
7107 for i in 0..n {
7108 sol[i] += alpha * ws.k_p[i];
7109 ws.residual[i] -= alpha * ws.k_ap[i];
7110 }
7111
7112 ws.k_residual[..n].copy_from_slice(&ws.residual[..n]);
7113 self.apply_precond_one_sweep(
7114 level,
7115 &ws.k_residual[..n],
7116 &mut ws.k_zeta[..n],
7117 &mut ws.k_work[..n],
7118 &mut ws.k_temp[..n],
7119 &mut ws.k_ap[..n],
7120 )?;
7121 let rz_new = dot(&ws.residual[..n], &ws.k_zeta[..n]);
7122 if !rz_new.is_finite() {
7123 break;
7124 }
7125 if rz_new.abs() < 1e-300 {
7126 break;
7127 }
7128 let beta = rz_new / rz_old;
7129 for i in 0..n {
7130 ws.k_p[i] = ws.k_zeta[i] + beta * ws.k_p[i];
7131 }
7132 rz_old = rz_new;
7133 }
7134 Ok(())
7135 }
7136
7137 fn effective_relax_type(&self, level: usize, phase: RelaxPhase, lc: usize) -> RelaxType {
7138 if level == lc {
7139 self.relax_level_overrides
7140 .get(&level)
7141 .copied()
7142 .unwrap_or(self.cfg.grid_relax_type[RelaxPhase::Coarsest.ix()])
7143 } else if let Some(relax) = self.relax_level_overrides.get(&level).copied() {
7144 relax
7145 } else {
7146 self.cfg.grid_relax_type[phase.ix()]
7147 }
7148 }
7149
7150 fn effective_relax_sweeps(&self, level: usize, phase: RelaxPhase, lc: usize) -> usize {
7151 if level == lc {
7152 return self
7153 .sweep_level_overrides
7154 .get(&level)
7155 .map(|(pre, _)| *pre)
7156 .unwrap_or(self.cfg.num_grid_sweeps[RelaxPhase::Coarsest.ix()]);
7157 }
7158 if let Some((pre, post)) = self.sweep_level_overrides.get(&level) {
7159 return if matches!(phase, RelaxPhase::Up) {
7160 *post
7161 } else {
7162 *pre
7163 };
7164 }
7165 self.cfg.num_grid_sweeps[phase.ix()]
7166 }
7167
7168 fn apply_relax_effective(
7169 &self,
7170 phase: RelaxPhase,
7171 where_: RelaxWhere,
7172 level_ix: usize,
7173 lvl: &AMGLevel,
7174 rhs: &[f64],
7175 sol: &mut [f64],
7176 ws: &mut AMGWorkspace,
7177 ) -> Result<(), KError> {
7178 let lc = match &self.state {
7179 AmgState::Ready { hierarchy, .. } => hierarchy.coarsest_ix(),
7180 _ => return Err(KError::InvalidInput("AMG not set up".into())),
7181 };
7182 let mut eff = RelaxPolicy {
7183 kind: self.cfg.grid_relax_type,
7184 sweeps: self.cfg.num_grid_sweeps,
7185 omega: self.cfg.jacobi_omega,
7186 };
7187 let relax = self.effective_relax_type(level_ix, phase, lc);
7188 let sweeps = self.effective_relax_sweeps(level_ix, phase, lc);
7189 eff.kind[phase.ix()] = relax;
7190 eff.sweeps[phase.ix()] = sweeps;
7191 Self::apply_relax(&eff, phase, where_, lvl, rhs, sol, ws, &self.cfg)
7192 }
7193
7194 fn restrict_apply(
7195 lvl: &AMGLevel,
7196 fine_res: &[f64],
7197 coarse_rhs: &mut [f64],
7198 ) -> Result<(), KError> {
7199 if lvl.r_row_ptr.is_some() {
7200 lvl.p.spmv_transpose_scaled(1.0, fine_res, 0.0, coarse_rhs)
7201 } else {
7202 lvl.r.spmv_scaled(1.0, fine_res, 0.0, coarse_rhs)
7203 }
7204 }
7205
7206 fn solve_coarse_level(
7207 &self,
7208 h: &AmgHierarchy,
7209 level: usize,
7210 rhs: &[f64],
7211 sol: &mut [f64],
7212 ws: &mut AMGWorkspace,
7213 ) -> Result<(), KError> {
7214 let levelc = &h.levels[level];
7215 let a = &levelc.a;
7216 let n = a.nrows();
7217 if matches!(h.coarse_solve, CoarseSolve::Smoother) {
7218 ws.ensure(n);
7219 return self.apply_relax_effective(
7220 RelaxPhase::Coarsest,
7221 RelaxWhere::Pre,
7222 level,
7223 levelc,
7224 rhs,
7225 sol,
7226 ws,
7227 );
7228 }
7229
7230 let prefer_dense =
7231 matches!(h.coarse_solve, CoarseSolve::DirectDense) || n <= self.cfg.max_coarse_size;
7232 if prefer_dense {
7233 if let Some(m) = &levelc.coarse_solver {
7234 return m
7235 .lock()
7236 .expect("coarse solver mutex poisoned")
7237 .solve(rhs, sol);
7238 }
7239 let mut solver = CoarseDenseLu::new();
7240 solver.setup(a)?;
7241 return solver.solve(rhs, sol);
7242 }
7243
7244 match h.coarse_solve {
7245 CoarseSolve::CG => cg_sparse(
7246 a,
7247 rhs,
7248 sol,
7249 self.cfg.tolerance,
7250 n.min(self.cfg.max_iterations.max(50)),
7251 ),
7252 CoarseSolve::ILU => {
7253 if let Some(m) = &levelc.coarse_solver {
7254 let mut guard = m.lock().expect("coarse solver mutex poisoned");
7255 guard.solve(rhs, sol)
7256 } else {
7257 let mut solver = CoarseIlu::new(
7258 self.cfg.tolerance,
7259 n.min(self.cfg.max_iterations.max(50)),
7260 self.cfg.ilu_drop_tol,
7261 self.cfg.ilu_fill_per_row,
7262 );
7263 solver.setup(a)?;
7264 solver.solve(rhs, sol)
7265 }
7266 }
7267 CoarseSolve::DirectDense => unreachable!(),
7268 CoarseSolve::Smoother => unreachable!(),
7269 }
7270 }
7271
7272 fn cycle_profiled(
7273 &self,
7274 level: usize,
7275 rhs: &[f64],
7276 sol: &mut [f64],
7277 ws: &mut AMGWorkspace,
7278 mut cyc: Option<&mut CycleTimings>,
7279 ) -> Result<(), KError> {
7280 let h = match &self.state {
7281 AmgState::Ready { hierarchy, .. } => hierarchy,
7282 _ => return Err(KError::InvalidInput("AMG not set up".into())),
7283 };
7284 let lc = h.coarsest_ix();
7285
7286 let a = &h.levels[level].a;
7287 let pol = &h.policy;
7288 let cycle_pol = &*self.cycle_policy;
7289 let mut lv = CycleLevelTiming {
7290 level,
7291 ..Default::default()
7292 };
7293 let prof = cyc.is_some();
7294
7295 if level == lc {
7296 with_timing(prof, &mut lv.coarse_solve, || {
7297 self.solve_coarse_level(h, level, rhs, sol, ws)
7298 })?;
7299 if let Some(c) = cyc {
7300 c.per_level.push(lv);
7301 }
7302 return Ok(());
7303 }
7304
7305 let n = a.nrows();
7306 ws.ensure(n);
7307 let use_mp_residual = self
7308 .cfg
7309 .mixed_precision
7310 .map(|mp| mp.residual_enabled())
7311 .unwrap_or(false);
7312
7313 let phase_pre = if level == 0 {
7315 RelaxPhase::Fine
7316 } else {
7317 RelaxPhase::Down
7318 };
7319 with_timing(prof, &mut lv.pre_smooth, || {
7320 let use_flexible =
7321 self.cfg.flexible_level == Some(level) && self.cfg.flexible_iters > 0;
7322 if use_flexible {
7323 let relax = pol.kind[phase_pre.ix()];
7324 if !Self::flexible_relax_supported(relax) {
7325 return Err(KError::InvalidInput(format!(
7326 "RelaxType {relax:?} cannot be used for flexible presmoothing",
7327 )));
7328 }
7329 self.fcg_presmooth(
7330 &h.levels[level],
7331 rhs,
7332 sol,
7333 self.cfg.flexible_iters,
7334 self.cfg.flexible_rtol,
7335 self.cfg.flexible_pc_sweeps,
7336 relax,
7337 pol.omega,
7338 ws,
7339 )
7340 } else {
7341 self.apply_relax_effective(
7342 phase_pre,
7343 RelaxWhere::Pre,
7344 level,
7345 &h.levels[level],
7346 rhs,
7347 sol,
7348 ws,
7349 )
7350 }
7351 })?;
7352
7353 if use_mp_residual {
7355 with_timing(prof, &mut lv.matvec, || {
7356 Self::mixed_spmv(&h.levels[level], &self.cfg, sol, ws)
7357 })?;
7358 with_timing(prof, &mut lv.residual_axpy, || {
7359 let mp_ws = ws
7360 .mp
7361 .as_ref()
7362 .expect("mixed workspace missing after mixed_spmv");
7363 #[cfg(feature = "rayon")]
7364 {
7365 for i in 0..n {
7366 ws.work[i] = mp_ws.work32[i] as f64;
7367 }
7368 ws.residual[..n]
7369 .par_iter_mut()
7370 .enumerate()
7371 .for_each(|(i, ri)| {
7372 *ri = rhs[i] - ws.work[i];
7373 });
7374 }
7375 #[cfg(not(feature = "rayon"))]
7376 for i in 0..n {
7377 let az = mp_ws.work32[i] as f64;
7378 ws.work[i] = az;
7379 ws.residual[i] = rhs[i] - az;
7380 }
7381 });
7382 } else {
7383 with_timing(prof, &mut lv.matvec, || {
7384 a.spmv_scaled(1.0, sol, 0.0, &mut ws.work[..n])
7385 })?;
7386 with_timing(prof, &mut lv.residual_axpy, || {
7387 #[cfg(feature = "rayon")]
7388 ws.residual[..n]
7389 .par_iter_mut()
7390 .enumerate()
7391 .for_each(|(i, ri)| {
7392 *ri = rhs[i] - ws.work[i];
7393 });
7394 #[cfg(not(feature = "rayon"))]
7395 for i in 0..n {
7396 ws.residual[i] = rhs[i] - ws.work[i];
7397 }
7398 });
7399 }
7400
7401 let p = &h.levels[level].p;
7403 let nc = h.levels[level + 1].a.nrows();
7404
7405 let mut local_coarse = std::mem::take(&mut ws.coarse_rhs);
7406 local_coarse.resize(nc, R::zero());
7407 with_timing(prof, &mut lv.restrict, || {
7408 Self::restrict_apply(&h.levels[level], &ws.residual[..n], &mut local_coarse[..nc])
7409 })?;
7410
7411 let gamma = cycle_pol.gamma_visits(level).max(1);
7412 for t in 0..gamma {
7413 let mut zc = ws.take_coarse_sol(level, nc);
7414 let coarse_result = if level + 1 == lc {
7415 with_timing(prof, &mut lv.coarse_solve, || {
7416 self.solve_coarse_level(h, level + 1, &local_coarse[..nc], &mut zc, ws)
7417 })
7418 } else {
7419 self.cycle_profiled(
7420 level + 1,
7421 &local_coarse[..nc],
7422 &mut zc,
7423 ws,
7424 cyc.as_deref_mut(),
7425 )
7426 };
7427 if let Err(err) = coarse_result {
7428 ws.put_coarse_sol(level, zc);
7429 return Err(err);
7430 }
7431 let prolong_result = with_timing(prof, &mut lv.prolong, || {
7432 ws.fine_corr[..n].fill(R::zero());
7433 p.spmv_scaled(1.0, &zc, 0.0, &mut ws.fine_corr[..n])
7434 });
7435 ws.put_coarse_sol(level, zc);
7436 prolong_result?;
7437 for i in 0..n {
7438 sol[i] += ws.fine_corr[i];
7439 }
7440
7441 if t + 1 < gamma {
7442 with_timing(prof, &mut lv.matvec, || {
7443 a.spmv_scaled(1.0, sol, 0.0, &mut ws.work[..n])
7444 })?;
7445 with_timing(prof, &mut lv.residual_axpy, || {
7446 #[cfg(feature = "rayon")]
7447 ws.residual[..n]
7448 .par_iter_mut()
7449 .enumerate()
7450 .for_each(|(i, ri)| {
7451 *ri = rhs[i] - ws.work[i];
7452 });
7453 #[cfg(not(feature = "rayon"))]
7454 for i in 0..n {
7455 ws.residual[i] = rhs[i] - ws.work[i];
7456 }
7457 });
7458 with_timing(prof, &mut lv.restrict, || {
7459 Self::restrict_apply(
7460 &h.levels[level],
7461 &ws.residual[..n],
7462 &mut local_coarse[..nc],
7463 )
7464 })?;
7465 }
7466 }
7467 ws.coarse_rhs = local_coarse;
7468
7469 if let Some((algo, iters)) = cycle_pol.k_presmooth(level) {
7470 with_timing(prof, &mut lv.post_smooth, || {
7471 self.krylov_smooth(algo, iters, level, rhs, sol, ws)
7472 })?;
7473 }
7474
7475 let phase_post = if level == 0 {
7477 RelaxPhase::Fine
7478 } else {
7479 RelaxPhase::Up
7480 };
7481 with_timing(prof, &mut lv.post_smooth, || {
7482 self.apply_relax_effective(
7483 phase_post,
7484 RelaxWhere::Post,
7485 level,
7486 &h.levels[level],
7487 rhs,
7488 sol,
7489 ws,
7490 )
7491 })?;
7492
7493 if let Some((algo, iters)) = cycle_pol.k_postsmooth(level) {
7494 with_timing(prof, &mut lv.post_smooth, || {
7495 self.krylov_smooth(algo, iters, level, rhs, sol, ws)
7496 })?;
7497 }
7498
7499 if let Some(c) = cyc {
7500 c.per_level.push(lv);
7501 }
7502 Ok(())
7503 }
7504
7505 fn cycle(
7506 &self,
7507 level: usize,
7508 rhs: &[f64],
7509 sol: &mut [f64],
7510 ws: &mut AMGWorkspace,
7511 ) -> Result<(), KError> {
7512 self.cycle_profiled(level, rhs, sol, ws, None)
7513 }
7514
7515 #[inline]
7516 fn v_cycle(
7517 &self,
7518 level: usize,
7519 rhs: &[f64],
7520 sol: &mut [f64],
7521 ws: &mut AMGWorkspace,
7522 ) -> Result<(), KError> {
7523 self.cycle(level, rhs, sol, ws)
7524 }
7525
7526 pub fn fmg_solve(&self, _b: &[f64], _x: &mut [f64]) -> Result<(), KError> {
7527 Err(KError::NotImplemented(
7528 "FMG solve not yet implemented".into(),
7529 ))
7530 }
7531
7532 pub fn cascade_solve(&self, _b: &[f64], _x: &mut [f64]) -> Result<(), KError> {
7533 Err(KError::NotImplemented(
7534 "Cascade solve not yet implemented".into(),
7535 ))
7536 }
7537
7538 #[cfg(not(feature = "complex"))]
7540 pub fn apply(&self, side: PcSide, x: &[f64], y: &mut [f64]) -> Result<(), KError> {
7541 Preconditioner::apply(self, side, x, y)
7542 }
7543 pub fn stats(&self) -> Option<AmgStats> {
7544 let mut out = self.stats.clone();
7545 if let (Some(s), Ok(rt)) = (out.as_mut(), self.runtime.lock()) {
7546 s.last_cycle = rt.last_cycle.clone();
7547 }
7548 #[cfg(feature = "complex")]
7549 if let Some(s) = out.as_mut() {
7550 s.complex_setup_mode = self.complex_setup_mode;
7551 s.complex_setup_fallback_reason = self.complex_setup_fallback_reason.clone();
7552 }
7553 out
7554 }
7555
7556 #[cfg(feature = "complex")]
7557 pub fn complex_setup_mode(&self) -> AmgComplexSetupMode {
7558 self.complex_setup_mode
7559 }
7560
7561 #[cfg(feature = "complex")]
7562 pub fn complex_setup_mode_label(&self) -> &'static str {
7563 self.complex_setup_mode.as_str()
7564 }
7565
7566 #[cfg(feature = "complex")]
7567 pub fn complex_setup_fallback_reason(&self) -> Option<&str> {
7568 self.complex_setup_fallback_reason.as_deref()
7569 }
7570
7571 pub fn dist_apply_stats(&self) -> Option<DistApplyStats> {
7572 if let Ok(rt) = self.runtime.lock() {
7573 rt.last_dist_apply.clone()
7574 } else {
7575 None
7576 }
7577 }
7578
7579 pub fn dist_setup_uses_fine_matrix_gather(&self) -> bool {
7582 self.dist_apply_stats()
7583 .is_some_and(|stats| stats.setup_uses_fine_matrix_gather())
7584 }
7585
7586 #[cfg(test)]
7587 pub(crate) fn debug_levels_r_equals_pt(&self) -> bool {
7588 let h = match &self.state {
7589 AmgState::Ready { hierarchy, .. } => hierarchy,
7590 _ => return true,
7591 };
7592 if !self.cfg.keep_transpose {
7593 return true;
7594 }
7595 for lvl in 0..h.coarsest_ix() {
7596 let pvals = h.levels[lvl].p.values();
7597 let rvals = h.levels[lvl].r.values();
7598 let map = &h.levels[lvl].p2r_pos;
7599 if pvals.len() != map.len() || rvals.len() < map.len() {
7600 return false;
7601 }
7602 for (pi, &ri) in map.iter().enumerate() {
7603 if ri >= rvals.len() {
7604 return false;
7605 }
7606 if (pvals[pi] - rvals[ri]).abs() > 1e-12 {
7607 return false;
7608 }
7609 }
7610 }
7611 true
7612 }
7613
7614 #[cfg(all(debug_assertions, not(feature = "complex")))]
7615 fn spd_probe(&self) -> Result<(), KError> {
7616 if !self.cfg.require_spd {
7617 return Ok(());
7618 }
7619 let h = match &self.state {
7620 AmgState::Ready { hierarchy, .. } => hierarchy,
7621 _ => return Err(KError::InvalidInput("AMG not set up".into())),
7622 };
7623 let n = h.finest().a.nrows();
7624 if n == 0 {
7625 return Ok(());
7626 }
7627 let mut x = vec![R::default(); n];
7628 let mut y = vec![R::default(); n];
7629 for t in 0..3 {
7630 for i in 0..n {
7631 x[i] = ((i + 7919 * t) % 127) as f64 - 63.0;
7632 }
7633 y.fill(R::default());
7634 self.apply(PcSide::Left, &x, &mut y)?;
7635 let qf = x.iter().zip(&y).map(|(a, b)| a * b).sum::<f64>();
7636 let x_norm2 = x.iter().map(|v| v * v).sum::<f64>();
7637 let tol = 1e-12_f64.max(1e-10 * x_norm2.abs());
7638 debug_assert!(
7639 qf.is_finite() && qf > -tol,
7640 "Preconditioned operator is not SPD: qf={qf}, tol={tol}"
7641 );
7642 }
7643 Ok(())
7644 }
7645}
7646
7647#[cfg(not(feature = "complex"))]
7650impl Preconditioner for AMG {
7651 fn dims(&self) -> (usize, usize) {
7652 if let Some(dist) = &self.dist {
7653 let n = dist.local_nrows();
7654 return (n, n);
7655 }
7656 if let AmgState::Ready { hierarchy, .. } = &self.state {
7657 let n = hierarchy.finest().a.nrows();
7658 (n, n)
7659 } else if let Some(csr) = self.csr.as_ref() {
7660 (csr.nrows(), csr.ncols())
7661 } else {
7662 (0, 0)
7663 }
7664 }
7665
7666 fn required_format(&self) -> crate::matrix::format::OpFormat {
7667 crate::matrix::format::OpFormat::Csr
7668 }
7669
7670 fn setup(&mut self, op: &dyn LinOp<S = f64>) -> Result<(), KError> {
7671 self.cfg.validate()?;
7672 if let Some(dist) = op.as_any().downcast_ref::<DistCsrOp>() {
7673 if op.comm().size() > 1
7674 || matches!(
7675 self.cfg.dist_coarse_strategy,
7676 DistCoarseStrategy::DistributedCsr
7677 )
7678 {
7679 let (strategy, _) = self.resolve_dist_coarse_strategy(&dist.comm())?;
7680 if self.try_update_dist_local_numeric(dist, strategy)? {
7681 return Ok(());
7682 }
7683 return self.setup_dist(dist);
7684 }
7685 }
7686 self.dist = None;
7687 let sid = op.structure_id();
7688 let vid = op.values_id();
7689 let csr = csr_from_linop(op, self.cfg.drop_tol)?;
7690 self.setup_from_csr_real(csr, sid, vid, false)
7691 }
7692
7693 fn apply(&self, side: PcSide, r: &[f64], z: &mut [f64]) -> Result<(), KError> {
7694 if let Some(dist) = &self.dist {
7695 return self.apply_dist(side, r, z, dist);
7696 }
7697 self.apply_local(side, r, z)
7698 }
7699
7700 fn capabilities(&self) -> PcCaps {
7701 let mut caps = PcCaps::default();
7702 if self.cfg.require_spd {
7703 caps.is_spd = true;
7704 caps.side_restriction = Some(PcSide::Left);
7705 }
7706 caps
7707 }
7708
7709 fn supports_numeric_update(&self) -> bool {
7710 true
7711 }
7712
7713 fn update_numeric(&mut self, op: &dyn LinOp<S = f64>) -> Result<(), KError> {
7714 self.cfg.validate()?;
7715 if let Some(dist) = op.as_any().downcast_ref::<DistCsrOp>() {
7716 if op.comm().size() > 1
7717 || matches!(
7718 self.cfg.dist_coarse_strategy,
7719 DistCoarseStrategy::DistributedCsr
7720 )
7721 {
7722 let (strategy, _) = self.resolve_dist_coarse_strategy(&dist.comm())?;
7723 if self.try_update_dist_local_numeric(dist, strategy)? {
7724 return Ok(());
7725 }
7726 return self.setup_dist(dist);
7727 }
7728 }
7729 self.dist = None;
7730 let csr = csr_from_linop(op, self.cfg.drop_tol)?;
7731 let sid = op.structure_id();
7732 let vid = op.values_id();
7733 let pattern_hash = csr_pattern_hash(&csr);
7734 self.ensure_symbolic_structure(csr.as_ref(), sid, pattern_hash)?;
7735 self.refresh_numeric_ready(csr.as_ref(), sid, vid, pattern_hash)?;
7736 self.csr = Some(csr.clone());
7737 if self.cfg.logging_level >= 2
7738 && self.cfg.print_level >= 1
7739 && let Some(s) = self.stats.as_ref()
7740 {
7741 print_setup_tables(s);
7742 }
7743 Ok(())
7744 }
7745
7746 fn update_symbolic(&mut self, op: &dyn LinOp<S = f64>) -> Result<(), KError> {
7747 self.cfg.validate()?;
7748 if let Some(dist) = op.as_any().downcast_ref::<DistCsrOp>() {
7749 if op.comm().size() > 1
7750 || matches!(
7751 self.cfg.dist_coarse_strategy,
7752 DistCoarseStrategy::DistributedCsr
7753 )
7754 {
7755 return self.setup_dist(dist);
7756 }
7757 }
7758 self.dist = None;
7759 let csr = csr_from_linop(op, self.cfg.drop_tol)?;
7760 let sid = op.structure_id();
7761 let vid = op.values_id();
7762 let pattern_hash = csr_pattern_hash(&csr);
7763 let hierarchy = self.build_symbolic(csr.as_ref())?;
7764 self.state = AmgState::SymbolicOnly {
7765 hierarchy,
7766 last_structure_id: sid,
7767 pattern_hash,
7768 };
7769 self.refresh_numeric_ready(csr.as_ref(), sid, vid, pattern_hash)?;
7770 self.csr = Some(csr.clone());
7771 if self.cfg.logging_level >= 2
7772 && self.cfg.print_level >= 1
7773 && let Some(s) = self.stats.as_ref()
7774 {
7775 print_setup_tables(s);
7776 }
7777 #[cfg(all(debug_assertions, not(feature = "complex")))]
7778 if self.cfg.require_spd {
7779 self.spd_probe()?;
7780 }
7781 Ok(())
7782 }
7783
7784 fn distributed_support(&self) -> crate::preconditioner::PcDistributedSupport {
7785 if self.dist.is_some() {
7786 if let Ok(rt) = self.runtime.lock()
7787 && let Some(last_dist) = rt.last_dist_apply.as_ref()
7788 {
7789 return if last_dist.reports_distributed_support() {
7790 crate::preconditioner::PcDistributedSupport::Distributed
7791 } else {
7792 crate::preconditioner::PcDistributedSupport::LocalOnly
7793 };
7794 }
7795 if amg_dist_route_reports_distributed(
7796 self.cfg.dist_coarse_solver_route,
7797 self.cfg.dist_coarse_strategy,
7798 ) {
7799 return crate::preconditioner::PcDistributedSupport::Distributed;
7800 }
7801 }
7802 crate::preconditioner::PcDistributedSupport::LocalOnly
7803 }
7804}
7805
7806impl AMG {
7807 fn setup_from_csr_real(
7808 &mut self,
7809 csr: Arc<CsrMatrix<f64>>,
7810 sid: StructureId,
7811 vid: ValuesId,
7812 force_numeric: bool,
7813 ) -> Result<(), KError> {
7814 let csr = if self.cfg.conditioning.is_active() {
7815 let mut local = (*csr).clone();
7816 apply_csr_transforms("AMG", &mut local, &self.cfg.conditioning)?;
7817 Arc::new(local)
7818 } else {
7819 csr
7820 };
7821 let csr_ref = csr.as_ref();
7822 #[cfg(debug_assertions)]
7823 debug_check_csr(csr_ref, "setup csr");
7824 let pattern_hash = csr_pattern_hash(csr_ref);
7825
7826 self.ensure_symbolic_structure(csr_ref, sid, pattern_hash)?;
7827
7828 let need_numeric = force_numeric
7829 || match &self.state {
7830 AmgState::Ready { last_values_id, .. } => *last_values_id != vid,
7831 AmgState::SymbolicOnly { .. } => true,
7832 AmgState::Uninitialized => true,
7833 };
7834
7835 if need_numeric {
7836 self.refresh_numeric_ready(csr_ref, sid, vid, pattern_hash)?;
7837 }
7838
7839 self.csr = Some(csr.clone());
7840 if self.cfg.logging_level >= 2
7841 && self.cfg.print_level >= 1
7842 && let Some(s) = self.stats.as_ref()
7843 {
7844 print_setup_tables(s);
7845 }
7846 #[cfg(all(debug_assertions, not(feature = "complex")))]
7847 if self.cfg.require_spd {
7848 self.spd_probe()?;
7849 }
7850 Ok(())
7851 }
7852
7853 #[cfg(feature = "complex")]
7854 fn setup_complex(
7855 &mut self,
7856 op: &dyn LinOp<S = S>,
7857 _force_numeric: bool,
7858 require_same_pattern: bool,
7859 ) -> Result<(), KError> {
7860 match self.cfg.coarse_solve {
7861 CoarseSolve::ILU => {
7862 return Err(KError::InvalidInput(
7863 "AMG complex setup does not support coarse_solve=ILU yet; use CG for HPD problems or DirectDense for nonsymmetric complex problems"
7864 .into(),
7865 ));
7866 }
7867 CoarseSolve::CG | CoarseSolve::DirectDense | CoarseSolve::Smoother => {}
7868 }
7869 self.cfg.validate()?;
7870 let csr_complex = csr_from_linop_complex(op, self.cfg.drop_tol)?;
7871 if require_same_pattern {
7872 let current = self
7873 .csr
7874 .as_ref()
7875 .ok_or_else(|| KError::InvalidInput("AMG not set up".into()))?;
7876 if current.nrows() != csr_complex.nrows()
7877 || current.ncols() != csr_complex.ncols()
7878 || csr_pattern_hash(current.as_ref()) != csr_pattern_hash(csr_complex.as_ref())
7879 {
7880 return Err(KError::InvalidInput(
7881 "AMG complex numeric update requires unchanged sparsity; call update_symbolic instead"
7882 .into(),
7883 ));
7884 }
7885 }
7886
7887 if require_same_pattern
7888 && self.complex_setup_mode == AmgComplexSetupMode::NativeHierarchy
7889 && let Some(core) = self.complex_core.as_ref()
7890 {
7891 let mut core = core
7892 .lock()
7893 .unwrap_or_else(std::sync::PoisonError::into_inner);
7894 core.update_numeric(csr_complex.as_ref())?;
7895 self.stats = Some(core.stats());
7896 self.complex_setup_fallback_reason = None;
7897 self.csr = Some(Arc::new(csr_real_metadata_from_complex(
7898 csr_complex.as_ref(),
7899 )));
7900 self.state = AmgState::Uninitialized;
7901 return Ok(());
7902 }
7903
7904 self.dist = None;
7905 self.complex_setup_mode = AmgComplexSetupMode::Unset;
7906 self.complex_setup_fallback_reason = None;
7907 self.complex_diag_inv = complex_diagonal_inverse(csr_complex.as_ref(), self.cfg.drop_tol)?;
7908 self.complex_coarse_solver = None;
7909 self.complex_core = None;
7910 if self.complex_diag_inv.is_some() {
7911 self.complex_setup_mode = AmgComplexSetupMode::NativeDiagonal;
7912 self.csr = Some(Arc::new(csr_real_metadata_from_complex(
7913 csr_complex.as_ref(),
7914 )));
7915 self.state = AmgState::Uninitialized;
7916 self.stats = Some(complex_single_level_stats(csr_complex.as_ref(), &self.cfg));
7917 return Ok(());
7918 }
7919 if csr_complex.nrows() <= self.cfg.max_coarse_size && self.transfer_overrides.is_empty() {
7920 let mut solver = CoarseDenseLu::<S>::new();
7921 solver.setup(csr_complex.as_ref())?;
7922 self.complex_coarse_solver = Some(Mutex::new(solver));
7923 self.complex_setup_mode = AmgComplexSetupMode::NativeCoarse;
7924 self.csr = Some(Arc::new(csr_real_metadata_from_complex(
7925 csr_complex.as_ref(),
7926 )));
7927 self.state = AmgState::Uninitialized;
7928 self.stats = Some(complex_single_level_stats(csr_complex.as_ref(), &self.cfg));
7929 return Ok(());
7930 }
7931
7932 let transfer_overrides = self
7933 .transfer_overrides
7934 .iter()
7935 .map(|(&level, ops)| (level, ops.clone()))
7936 .collect::<Vec<_>>();
7937 let core = AmgCore::<S>::setup_with_transfer_overrides(
7938 csr_complex.as_ref(),
7939 &self.cfg,
7940 &transfer_overrides,
7941 )?;
7942 self.stats = Some(core.stats());
7943 self.complex_core = Some(Mutex::new(core));
7944 self.complex_setup_mode = AmgComplexSetupMode::NativeHierarchy;
7945 self.complex_setup_fallback_reason = None;
7946 self.csr = Some(Arc::new(csr_real_metadata_from_complex(
7947 csr_complex.as_ref(),
7948 )));
7949 self.state = AmgState::Uninitialized;
7950 Ok(())
7951 }
7952}
7953
7954#[cfg(feature = "complex")]
7955impl Preconditioner for AMG {
7956 fn dims(&self) -> (usize, usize) {
7957 if let Some(dist) = &self.dist {
7958 let n = dist.local_nrows();
7959 return (n, n);
7960 }
7961 if let AmgState::Ready { hierarchy, .. } = &self.state {
7962 let n = hierarchy.finest().a.nrows();
7963 (n, n)
7964 } else if let Some(core) = self.complex_core.as_ref() {
7965 let n = core
7966 .lock()
7967 .unwrap_or_else(std::sync::PoisonError::into_inner)
7968 .stats()
7969 .levels
7970 .first()
7971 .map(|level| level.n)
7972 .unwrap_or(0);
7973 (n, n)
7974 } else if let Some(csr) = self.csr.as_ref() {
7975 (csr.nrows(), csr.ncols())
7976 } else {
7977 (0, 0)
7978 }
7979 }
7980
7981 fn required_format(&self) -> crate::matrix::format::OpFormat {
7982 crate::matrix::format::OpFormat::Csr
7983 }
7984
7985 fn setup(&mut self, op: &dyn LinOp<S = S>) -> Result<(), KError> {
7986 self.setup_complex(op, false, false)
7987 }
7988
7989 fn apply(&self, side: PcSide, r: &[S], z: &mut [S]) -> Result<(), KError> {
7990 if r.len() != z.len() {
7991 return Err(KError::InvalidInput(format!(
7992 "AMG.apply: r/z size mismatch: {} vs {}",
7993 r.len(),
7994 z.len()
7995 )));
7996 }
7997 if self.cfg.require_spd && side != PcSide::Left {
7998 return Err(KError::InvalidInput(
7999 "AMG in SPD mode supports only Left preconditioning for CG-safe use".into(),
8000 ));
8001 }
8002 let n = r.len();
8003 if let Some(diag_inv) = self.complex_diag_inv.as_ref() {
8004 if diag_inv.len() != n {
8005 return Err(KError::InvalidInput(
8006 "AMG.apply: complex diagonal cache size mismatch".into(),
8007 ));
8008 }
8009 for i in 0..n {
8010 z[i] = diag_inv[i] * r[i];
8011 }
8012 return Ok(());
8013 }
8014 if let Some(solver) = self.complex_coarse_solver.as_ref() {
8015 return solver
8016 .lock()
8017 .unwrap_or_else(std::sync::PoisonError::into_inner)
8018 .solve(r, z);
8019 }
8020 if let Some(core) = self.complex_core.as_ref() {
8021 return core
8022 .lock()
8023 .unwrap_or_else(std::sync::PoisonError::into_inner)
8024 .apply(r, z);
8025 }
8026 Err(KError::InvalidInput(
8027 "AMG complex apply requires native complex setup state; call setup/update_symbolic before apply"
8028 .into(),
8029 ))
8030 }
8031
8032 fn capabilities(&self) -> PcCaps {
8033 let mut caps = PcCaps::default();
8034 if self.cfg.require_spd {
8035 caps.is_spd = true;
8036 caps.side_restriction = Some(PcSide::Left);
8037 }
8038 caps
8039 }
8040
8041 fn supports_numeric_update(&self) -> bool {
8042 true
8043 }
8044
8045 fn update_numeric(&mut self, op: &dyn LinOp<S = S>) -> Result<(), KError> {
8046 self.setup_complex(op, true, true)
8047 }
8048
8049 fn update_symbolic(&mut self, op: &dyn LinOp<S = S>) -> Result<(), KError> {
8050 self.setup(op)
8051 }
8052
8053 fn distributed_support(&self) -> crate::preconditioner::PcDistributedSupport {
8054 crate::preconditioner::PcDistributedSupport::LocalOnly
8055 }
8056}
8057
8058#[cfg(feature = "complex")]
8059impl KPreconditioner for AMG {
8060 type Scalar = S;
8061
8062 #[inline]
8063 fn dims(&self) -> (usize, usize) {
8064 <Self as Preconditioner>::dims(self)
8065 }
8066
8067 fn apply_s(
8068 &self,
8069 side: PcSide,
8070 x: &[S],
8071 y: &mut [S],
8072 scratch: &mut BridgeScratch,
8073 ) -> Result<(), KError> {
8074 bridge_apply_pc_s(self, side, x, y, scratch)
8075 }
8076
8077 fn apply_mut_s(
8078 &mut self,
8079 side: PcSide,
8080 x: &[S],
8081 y: &mut [S],
8082 scratch: &mut BridgeScratch,
8083 ) -> Result<(), KError> {
8084 bridge_apply_pc_mut_s(self, side, x, y, scratch)
8085 }
8086}
8087
8088#[cfg(not(feature = "complex"))]
8091impl crate::preconditioner::legacy::Preconditioner<Mat<f64>, Vec<f64>> for AMG {
8092 fn setup(&mut self, a: &Mat<f64>) -> Result<(), KError> {
8093 Preconditioner::setup(self, a)
8094 }
8095 fn apply(&self, side: PcSide, r: &Vec<f64>, z: &mut Vec<f64>) -> Result<(), KError> {
8096 Preconditioner::apply(self, side, r.as_slice(), z.as_mut_slice())
8097 }
8098}
8099
8100fn build_hierarchy(
8103 fine: &CsrMatrix<f64>,
8104 cfg: &mut AMGConfig,
8105 transfer_overrides: &BTreeMap<usize, AmgTransferOperators>,
8106) -> Result<(AmgHierarchy, Option<AmgStats>), KError> {
8107 let mut levels: Vec<AMGLevel> = Vec::with_capacity(cfg.max_levels);
8108 let mut a_cur = fine.clone();
8109 let do_stats = cfg.logging_level > 0;
8110 let mut level_stats: Vec<LevelStats> = Vec::new();
8111 let mut timings: Vec<LevelSetupTiming> = Vec::new();
8112 let mut diag_stats: Vec<AmgLevelStats> = Vec::new();
8113 let t_setup_all = if do_stats { Some(tic()) } else { None };
8114 #[cfg(feature = "simd")]
8115 let spmv_tuning: SpmvTuning = utils::default_spmv_tuning();
8116
8117 let need_l1 = cfg.grid_relax_type.contains(&RelaxType::L1Jacobi);
8118 let need_cheb = cfg.grid_relax_type.contains(&RelaxType::Chebyshev);
8119 let need_cheb_safe = cfg.grid_relax_type.contains(&RelaxType::ChebyshevSafe);
8120 let need_safe_diag = cfg
8121 .grid_relax_type
8122 .contains(&RelaxType::SafeguardedGaussSeidel)
8123 || need_cheb_safe;
8124 let need_ilu0 = cfg.grid_relax_type.contains(&RelaxType::Ilu0);
8125 let need_ras = cfg.grid_relax_type.contains(&RelaxType::Ras);
8126 let allow_safeguard = need_safe_diag || need_ilu0 || need_ras;
8127 let need_fsai = cfg.grid_relax_type.contains(&RelaxType::Fsai);
8128
8129 let mut lt0 = LevelSetupTiming::default();
8131 let t = tic();
8132 let diag0 = diag_inv_from_csr_cfg_fallback(&a_cur, cfg, allow_safeguard)?;
8133 if do_stats {
8134 lt0.diag = toc(t);
8135 lt0.total = lt0.diag;
8136 }
8137 let layout0 = if cfg.nodal == NodalMode::Nodal {
8138 Some(DofLayout::new(a_cur.nrows(), cfg.block_size))
8139 } else {
8140 None
8141 };
8142 let l0_num_functions = cfg
8143 .near_nullspace
8144 .as_ref()
8145 .map(|nns| nns.basis.len().max(1))
8146 .or_else(|| layout0.as_ref().map(|layout| layout.block_size.max(1)))
8147 .unwrap_or_else(|| cfg.num_functions.max(1));
8148 let l0 = AMGLevel {
8149 a: a_cur.clone(),
8150 p: CsrMatrix::identity(a_cur.nrows()),
8151 r: CsrMatrix::identity(a_cur.nrows()),
8152 diag_inv: diag0,
8153 d_sqrt_inv: None,
8154 l1_inv: None,
8155 diag_inv_safe: None,
8156 d_sqrt_inv_safe: None,
8157 cheb: None,
8158 cheb_safe: None,
8159 agg_of: (0..a_cur.nrows()).collect(),
8160 is_c: Vec::new(),
8161 cf: None,
8162 p2r_pos: Vec::new(),
8163 num_functions: l0_num_functions,
8164 row_basis: None,
8165 layout: layout0.clone(),
8166 nns: cfg.near_nullspace.as_ref().map(|nns| nns.basis.clone()),
8167 a_next_pat: None,
8168 a_next_pat_ng: None,
8169 rap_full2ng_pos: None,
8170 r_row_ptr: None,
8171 r_col_idx: None,
8172 r_vals_scratch: None,
8173 coarse_solver: None,
8174 ilu0: None,
8175 ras: None,
8176 fsai: None,
8177 a_vals_f32: None,
8178 diag_inv_f32: None,
8179 d_sqrt_inv_f32: None,
8180 l1_inv_f32: None,
8181 fsai_g_vals_f32: None,
8182 fsai_gt_vals_f32: None,
8183 };
8184 let mut l0 = l0;
8185 update_level_caches(
8186 cfg,
8187 &mut l0,
8188 need_l1,
8189 need_cheb,
8190 need_safe_diag,
8191 need_cheb_safe,
8192 need_ilu0,
8193 need_ras,
8194 true,
8195 )?;
8196 if need_fsai {
8197 let strength0 = if cfg.fsai_use_strength {
8198 Some(Strength::from_csr(
8199 &l0.a,
8200 cfg.strong_threshold,
8201 cfg.normalize_strength,
8202 ))
8203 } else {
8204 None
8205 };
8206 l0.fsai = Some(fsai_build_for_level(cfg, &l0.a, strength0.as_ref())?);
8207 refresh_mixed_precision_shadows(cfg, &mut l0);
8208 }
8209 #[cfg(feature = "simd")]
8210 {
8211 let tuning = utils::default_spmv_tuning();
8212 build_level_spmv_plans(&mut l0, &tuning);
8213 }
8214 #[cfg(feature = "simd")]
8215 build_level_spmv_plans(&mut l0, &spmv_tuning);
8216 levels.push(l0);
8217 if do_stats {
8218 level_stats.push(LevelStats {
8219 level: 0,
8220 n: a_cur.nrows(),
8221 nnz_a: a_cur.nnz(),
8222 nnz_p: 0,
8223 nnz_r: 0,
8224 max_row_sum_a: max_row_sum_abs(&a_cur),
8225 eff_nnz_a: Some(eff_nnz(&a_cur, cfg.stats_eps)),
8226 pre_sweeps: cfg.num_grid_sweeps[RelaxPhase::Down.ix()],
8227 post_sweeps: cfg.num_grid_sweeps[RelaxPhase::Up.ix()],
8228 pre_work_estimate: (cfg.num_grid_sweeps[RelaxPhase::Down.ix()] as f64)
8229 * a_cur.nnz() as f64,
8230 post_work_estimate: (cfg.num_grid_sweeps[RelaxPhase::Up.ix()] as f64)
8231 * a_cur.nnz() as f64,
8232 selected_relax_pre: format!("{:?}", cfg.grid_relax_type[RelaxPhase::Down.ix()]),
8233 selected_relax_post: format!("{:?}", cfg.grid_relax_type[RelaxPhase::Up.ix()]),
8234 coarse_solver: None,
8235 });
8236 timings.push(lt0);
8237 }
8238 diag_stats.push(AmgLevelStats {
8239 p_min_col_norm: 0.0,
8240 p_cond_sketched: 0.0,
8241 galerkin_worst_rel: 0.0,
8242 });
8243
8244 let mut block_size_cur = if cfg.nodal == NodalMode::Nodal {
8245 cfg.block_size
8246 } else {
8247 1
8248 };
8249 let mut trials_current = make_trial_matrix(cfg, a_cur.nrows())?;
8250 for level in 0..cfg.max_levels {
8252 let n = a_cur.nrows();
8253 if n <= cfg.coarse_threshold || n <= cfg.min_coarse_size {
8254 break;
8255 }
8256
8257 let mut lt = LevelSetupTiming::default();
8258
8259 let layout = if cfg.nodal == NodalMode::Nodal {
8260 Some(DofLayout::new(n, block_size_cur))
8261 } else {
8262 None
8263 };
8264 let (s, nodal_strength_opt) = with_timing(do_stats, &mut lt.strength, || {
8266 if let Some(ref lay) = layout {
8267 let nodal = strength_nodal_from_csr(
8268 &a_cur,
8269 lay.block_size,
8270 cfg.strong_threshold,
8271 cfg.normalize_strength,
8272 );
8273 let strength = Strength {
8274 row_ptr: nodal.row_ptr.clone(),
8275 col_idx: nodal.col_idx.clone(),
8276 };
8277 (strength, Some(nodal))
8278 } else {
8279 (
8280 Strength::from_csr(&a_cur, cfg.strong_threshold, cfg.normalize_strength),
8281 None,
8282 )
8283 }
8284 });
8285 let mis_k = if level < cfg.agg_num_levels {
8287 cfg.aggressive_mis_k.max(2)
8288 } else {
8289 1
8290 };
8291 let agg_algo = match cfg.coarsen_type {
8292 CoarsenType::RS => AggAlgo::RSGreedy,
8293 CoarsenType::HMIS => AggAlgo::HMIS,
8294 CoarsenType::PMIS => AggAlgo::PMIS,
8295 CoarsenType::Falgout => AggAlgo::Falgout,
8296 };
8297 let (agg_node, is_c_node) = with_timing(do_stats, &mut lt.aggregate, || {
8298 match (layout.as_ref(), nodal_strength_opt.as_ref()) {
8299 (Some(_), Some(nodal)) => build_aggregates_nodal(
8300 nodal,
8301 agg_algo,
8302 &AggOpts {
8303 mis_k,
8304 cap_per_row: cfg.max_strong_per_row,
8305 },
8306 ),
8307 _ => build_aggregates(
8308 &s,
8309 agg_algo,
8310 &AggOpts {
8311 mis_k,
8312 cap_per_row: cfg.max_strong_per_row,
8313 },
8314 ),
8315 }
8316 });
8317 let (agg, is_c) = if let Some(ref lay) = layout {
8318 lift_node_aggregates_to_dofs(&agg_node, &is_c_node, lay)
8319 } else {
8320 (agg_node, is_c_node)
8321 };
8322 let mut nns_basis: Vec<Vec<f64>> = if level == 0 {
8323 if let Some(ref nns) = cfg.near_nullspace {
8324 for (k, vec) in nns.basis.iter().enumerate() {
8325 if vec.len() != n {
8326 return Err(KError::InvalidInput(format!(
8327 "AMG: near-nullspace vector {k} has length {}, expected {n}",
8328 vec.len()
8329 )));
8330 }
8331 }
8332 nns.basis.clone()
8333 } else {
8334 Vec::new()
8335 }
8336 } else {
8337 Vec::new()
8338 };
8339 let user_supplied_nns = !nns_basis.is_empty();
8340 if nns_basis.is_empty() {
8341 if let Some(ref lay) = layout {
8342 let mut basis = vec![vec![0.0; n]; lay.block_size.max(1)];
8343 for i in 0..n {
8344 let comp = lay.comp_of[i];
8345 basis[comp][i] = 1.0;
8346 }
8347 nns_basis = basis;
8348 } else {
8349 nns_basis.push(vec![1.0; n]);
8350 }
8351 }
8352 let mut target_functions = if user_supplied_nns {
8353 nns_basis.len().max(1)
8354 } else if level == 0 {
8355 cfg.num_functions.max(1)
8356 } else if let Some(ref lay) = layout {
8357 lay.block_size.max(1)
8358 } else {
8359 block_size_cur.max(1)
8360 };
8361 if !user_supplied_nns {
8362 if let Some(ref lay) = layout {
8363 target_functions = target_functions.max(lay.block_size.max(1));
8364 }
8365 target_functions = target_functions.max(nns_basis.len().max(1));
8366 while nns_basis.len() < target_functions {
8367 nns_basis.push(vec![0.0; n]);
8368 }
8369 }
8370 let num_functions = target_functions;
8371 let nns_opt = Some(nns_basis.clone());
8372 let comp_opt = if user_supplied_nns {
8373 None
8374 } else {
8375 layout.as_ref().map(|lay| lay.comp_of.clone())
8376 };
8377 let tp = TentativeP {
8378 n_coarse: 1 + agg.iter().copied().max().unwrap_or(0),
8379 agg_of: agg.clone(),
8380 num_functions,
8381 nns: nns_opt.clone(),
8382 comp_of: comp_opt.clone(),
8383 };
8384 let block_size_next = tp.num_functions.max(1);
8385 let d = diag_inv_from_csr_cfg_fallback(&a_cur, cfg, allow_safeguard)?;
8386 let diag_weights: Vec<f64> = d
8387 .iter()
8388 .map(|&inv| {
8389 if inv.abs() > 0.0 {
8390 (1.0 / inv).abs().max(1e-30)
8391 } else {
8392 1.0
8393 }
8394 })
8395 .collect();
8396 let mut tn_opt: Option<TentativeNodal> = None;
8397 if layout.is_some() {
8398 let n_agg = tp.n_coarse;
8399 let mut rows_per_agg: Vec<Vec<usize>> = vec![Vec::new(); n_agg];
8400 for (row, &g) in agg.iter().enumerate() {
8401 rows_per_agg[g].push(row);
8402 }
8403 let mut row_basis_vec = vec![0.0; n * num_functions];
8404 for rows in &rows_per_agg {
8405 orthonormalize_aggregate(
8406 rows,
8407 &nns_basis,
8408 &diag_weights,
8409 num_functions,
8410 &mut row_basis_vec,
8411 );
8412 }
8413 tn_opt = Some(TentativeNodal {
8414 agg_of: agg.clone(),
8415 n_agg,
8416 mfun: num_functions,
8417 row_basis: row_basis_vec,
8418 });
8419 }
8420 let s_sym = s.symmetrize();
8421 let tn_ref = tn_opt.as_ref();
8422 let (mut p_csr, cf_opt): (Pcsr, Option<CFInfo>) =
8423 with_timing(do_stats, &mut lt.prolong, || {
8424 if matches!(
8425 cfg.interp_type,
8426 InterpType::Direct
8427 | InterpType::Standard
8428 | InterpType::Extended
8429 | InterpType::Classical
8430 | InterpType::HE
8431 ) {
8432 let extended = matches!(cfg.interp_type, InterpType::Extended);
8433 let (pat, cf) = classical_pattern(&a_cur, &s_sym, &is_c, extended);
8434 let mut vals = vec![0.0; pat.col_idx.len()];
8435 let params = ClassicalParams {
8436 variant: match cfg.interp_type {
8437 InterpType::Direct => ClassicalVariant::Direct,
8438 InterpType::HE => ClassicalVariant::HE,
8439 InterpType::Standard | InterpType::Classical | InterpType::Extended => {
8440 ClassicalVariant::Standard
8441 }
8442 _ => ClassicalVariant::Standard,
8443 },
8444 extended,
8445 drop_abs: cfg.interpolation_truncation,
8446 trunc_rel: cfg.truncation_factor,
8447 cap_row: cfg.max_elements_per_row,
8448 keep_at_least_one: true,
8449 };
8450 classical_values_only(
8451 &a_cur,
8452 &s_sym,
8453 &cf,
8454 ¶ms,
8455 &pat.row_ptr,
8456 &pat.col_idx,
8457 &mut vals,
8458 )?;
8459 let mut p = pat.clone();
8460 p.vals = vals;
8461 Ok((p, Some(cf)))
8462 } else if let Some(tn) = tn_ref {
8463 Ok((
8464 smooth_tentative_sa_mf(
8465 &a_cur,
8466 &d,
8467 tn,
8468 cfg.jacobi_omega,
8469 cfg.interpolation_truncation,
8470 cfg.max_elements_per_row,
8471 ),
8472 None,
8473 ))
8474 } else {
8475 Ok((
8476 smooth_tentative_sa_multi(
8477 &a_cur,
8478 &d,
8479 &tp,
8480 cfg.jacobi_omega,
8481 cfg.interpolation_truncation,
8482 cfg.max_elements_per_row,
8483 cfg.truncation_factor,
8484 ),
8485 None,
8486 ))
8487 }
8488 })?;
8489 let row_basis_for_level = tn_opt.as_ref().map(|tn| tn.row_basis.clone());
8490 let ctx = LevelPostContext {
8491 r: num_functions,
8492 agg_of: &tp.agg_of,
8493 nns: nns_opt
8494 .as_ref()
8495 .map(|v| v.iter().map(|b| b.as_slice()).collect()),
8496 a: Some(&a_cur),
8497 d_inv: Some(&d),
8498 };
8499 apply_post_interp(cfg, &ctx, &p_csr.row_ptr, &p_csr.col_idx, &mut p_csr.vals)?;
8500 if cfg.adaptive_interp
8501 && cfg.adaptive_samples > 0
8502 && cf_opt.is_none()
8503 && tp.num_functions == 1
8504 && p_csr.n > cfg.max_coarse_size
8505 {
8506 let omega = if cfg.adaptive_smooth_omega == 0.0 {
8507 cfg.jacobi_omega
8508 } else {
8509 cfg.adaptive_smooth_omega
8510 };
8511 let samples = sample_low_modes(
8512 &a_cur,
8513 &d,
8514 cfg.adaptive_samples,
8515 cfg.adaptive_smooth_steps,
8516 omega,
8517 0xC0FFEE,
8518 )?;
8519 let coarse_samples =
8520 restrict_samples_to_coarse(&a_cur, &tp, &samples, cfg.adaptive_weight_mode);
8521 adaptive_fit_values_only(
8522 &p_csr.row_ptr,
8523 &p_csr.col_idx,
8524 &mut p_csr.vals,
8525 &tp,
8526 &samples,
8527 &coarse_samples,
8528 cfg.adaptive_lambda,
8529 cfg.adaptive_enforce_sum1,
8530 cfg.interpolation_truncation,
8531 )?;
8532 }
8533 let mut p = CsrMatrix::from_csr(
8534 p_csr.m,
8535 p_csr.n,
8536 p_csr.row_ptr.clone(),
8537 p_csr.col_idx.clone(),
8538 p_csr.vals.clone(),
8539 );
8540 let mut rank_diag = RankDiagnostics::default();
8541 let check_rank = cfg.verify_p_rank && cf_opt.is_none();
8542 if check_rank {
8543 rank_diag = check_p_rank_fast(&p, cfg)?;
8544 if rank_diag.suspect {
8545 let mut cond_report = rank_diag.cond_estimate;
8546 match cfg.on_rank_failure {
8547 RankFallback::RetryLooserInterp if cf_opt.is_none() => {
8548 match try_fix_rank(level, &a_cur, &d, &tp, &ctx, &mut p_csr, cfg)? {
8549 RankFixOutcome::Fixed => {
8550 p = CsrMatrix::from_csr(
8551 p_csr.m,
8552 p_csr.n,
8553 p_csr.row_ptr.clone(),
8554 p_csr.col_idx.clone(),
8555 p_csr.vals.clone(),
8556 );
8557 rank_diag = check_p_rank_fast(&p, cfg)?;
8558 cond_report = rank_diag.cond_estimate;
8559 if rank_diag.suspect {
8560 return Err(KError::InvalidInput(format!(
8561 "AMG: P rank suspect at level {level}, cond≈{cond_report:.3e}"
8562 )));
8563 }
8564 }
8565 RankFixOutcome::Unfixed => {
8566 return Err(KError::InvalidInput(format!(
8567 "AMG: P rank suspect at level {level}, cond≈{cond_report:.3e}"
8568 )));
8569 }
8570 }
8571 }
8572 RankFallback::Abort => {
8573 return Err(KError::InvalidInput(format!(
8574 "AMG: P rank suspect at level {level}, cond≈{cond_report:.3e}"
8575 )));
8576 }
8577 other => {
8578 return Err(KError::InvalidInput(format!(
8579 "AMG: rank fallback {other:?} not implemented at level {level}"
8580 )));
8581 }
8582 }
8583 }
8584 }
8585 if let Some(override_ops) = transfer_overrides.get(&level) {
8586 if override_ops.prolongation.nrows() != n {
8587 return Err(KError::InvalidInput(format!(
8588 "AMG transfer override level {level} has P rows {}, expected {n}",
8589 override_ops.prolongation.nrows()
8590 )));
8591 }
8592 if override_ops.prolongation.ncols() == 0 {
8593 return Err(KError::InvalidInput(format!(
8594 "AMG transfer override level {level} has zero coarse columns"
8595 )));
8596 }
8597 if override_ops.restriction.ncols() != n {
8598 return Err(KError::InvalidInput(format!(
8599 "AMG transfer override level {level} has R cols {}, expected {n}",
8600 override_ops.restriction.ncols()
8601 )));
8602 }
8603 if override_ops.prolongation.ncols() != override_ops.restriction.nrows() {
8604 return Err(KError::InvalidInput(format!(
8605 "AMG transfer override level {level} has inconsistent coarse dims P: {} cols, R: {} rows",
8606 override_ops.prolongation.ncols(),
8607 override_ops.restriction.nrows()
8608 )));
8609 }
8610 #[cfg(feature = "complex")]
8611 {
8612 p = csr_real_metadata_from_complex(&override_ops.prolongation);
8613 p_csr = Pcsr {
8614 m: p.nrows(),
8615 n: p.ncols(),
8616 row_ptr: p.row_ptr().to_vec(),
8617 col_idx: p.col_idx().to_vec(),
8618 vals: p.values().to_vec(),
8619 };
8620 }
8621 #[cfg(not(feature = "complex"))]
8622 {
8623 p = override_ops.prolongation.clone();
8624 p_csr = Pcsr {
8625 m: p.nrows(),
8626 n: p.ncols(),
8627 row_ptr: p.row_ptr().to_vec(),
8628 col_idx: p.col_idx().to_vec(),
8629 vals: p.values().to_vec(),
8630 };
8631 }
8632 }
8633 let (r_row_ptr, r_col_idx, r_vals, p2r_pos) =
8635 with_timing(do_stats, &mut lt.restrict, || {
8636 transpose_csr_with_pos(&p_csr)
8637 });
8638 let r = CsrMatrix::from_csr(
8639 p_csr.n,
8640 p_csr.m,
8641 r_row_ptr.clone(),
8642 r_col_idx.clone(),
8643 r_vals.clone(),
8644 );
8645 #[cfg(debug_assertions)]
8646 debug_check_csr(&r, "R pattern");
8647 debug_assert_eq!(p_csr.col_idx.len(), p2r_pos.len());
8648 debug_assert_eq!(r_vals.len(), p2r_pos.len());
8649 for (pi, &ri) in p2r_pos.iter().enumerate() {
8650 debug_assert!(ri < r_vals.len(), "p2r_pos out of range");
8651 debug_assert!(
8652 (p_csr.vals[pi] - r_vals[ri]).abs() <= 1e-12,
8653 "R != P^T at index {pi}"
8654 );
8655 }
8656
8657 let trials_next = if let Some(ref trials) = trials_current {
8658 let mut next = Mat::<f64>::zeros(r.nrows(), trials.ncols());
8659 restrict_trials(&r, trials.as_ref(), next.as_mut())?;
8660 Some(next)
8661 } else {
8662 None
8663 };
8664
8665 let pat = with_timing(do_stats, &mut lt.rap_symbolic, || {
8667 rap_symbolic(&r, &a_cur, &p)
8668 });
8669 let mut a_coarse_vals = vec![0.0; pat.col_idx.len()];
8670 with_timing(do_stats, &mut lt.rap_numeric, || {
8671 rap_numeric(&pat, &r, &a_cur, &p, &mut a_coarse_vals);
8672 });
8673 {
8674 let mut rf = |row: usize| RowFilter {
8675 tau_abs: cfg.rap_truncation_abs,
8676 tau_rel: cfg.truncation_factor,
8677 k_max: cfg.rap_max_elements_per_row,
8678 must_keep: if cfg.keep_pivot_in_rap {
8679 Some(row)
8680 } else {
8681 None
8682 },
8683 };
8684 apply_filter_to_csr_values_in_place(
8685 pat.nrows,
8686 &pat.row_ptr,
8687 &pat.col_idx,
8688 &mut a_coarse_vals,
8689 &mut rf,
8690 );
8691 }
8692
8693 let mut use_ng = cfg.non_galerkin.enabled && (level + 1) >= cfg.non_galerkin.start_level;
8694 if cfg.require_spd && cfg.forbid_non_galerkin_in_spd {
8695 use_ng = false;
8696 }
8697 let mut a_full = CsrMatrix::from_csr(
8698 pat.nrows,
8699 pat.ncols,
8700 pat.row_ptr.clone(),
8701 pat.col_idx.clone(),
8702 a_coarse_vals,
8703 );
8704 if !use_ng || !cfg.filter_after_non_galerkin {
8705 apply_trial_compensation(cfg, &mut a_full, trials_next.as_ref(), block_size_next)?;
8706 }
8707 let (mut a_coarse, mut ng_pat_opt, mut map_opt) = if use_ng {
8708 let (ng_pat, ng_vals, full2ng) = non_galerkin_filter_coarse(
8709 &pat,
8710 a_full.values(),
8711 cfg.non_galerkin.symmetry,
8712 NgRowFilter {
8713 tau_abs: cfg.non_galerkin.drop_abs,
8714 tau_rel: cfg.non_galerkin.drop_rel,
8715 k_max: cfg.non_galerkin.cap_row,
8716 lump_diag: cfg.non_galerkin.lump_diagonal,
8717 },
8718 );
8719 let mut a_ng = CsrMatrix::from_csr(
8720 ng_pat.nrows,
8721 ng_pat.ncols,
8722 ng_pat.row_ptr.clone(),
8723 ng_pat.col_idx.clone(),
8724 ng_vals,
8725 );
8726 if cfg.filter_after_non_galerkin {
8727 apply_trial_compensation(cfg, &mut a_ng, trials_next.as_ref(), block_size_next)?;
8728 }
8729 (a_ng, Some(ng_pat), Some(full2ng))
8730 } else {
8731 (a_full, None, None)
8732 };
8733 let mut galerkin_worst = 0.0;
8734 let allow_galerkin =
8735 cfg.verify_galerkin && cfg.filter_omega <= 0.0 && !use_ng && cfg.galerkin_samples > 0;
8736 if allow_galerkin {
8737 let (ok, worst) = galerkin_sample_check(
8738 &a_cur,
8739 &p,
8740 &r,
8741 &a_coarse,
8742 cfg.galerkin_samples,
8743 cfg.galerkin_rel_tol,
8744 0xBEEF,
8745 )?;
8746 galerkin_worst = worst;
8747 if !ok {
8748 let a_fix = rap(&r, &a_cur, &p)?;
8749 let (ok2, worst2) = galerkin_sample_check(
8750 &a_cur,
8751 &p,
8752 &r,
8753 &a_fix,
8754 cfg.galerkin_samples,
8755 cfg.galerkin_rel_tol,
8756 0xBEEF,
8757 )?;
8758 if ok2 {
8759 galerkin_worst = worst2;
8760 a_coarse = a_fix;
8761 ng_pat_opt = None;
8762 map_opt = None;
8763 } else {
8764 return Err(KError::InvalidInput(format!(
8765 "AMG: Galerkin identity failed at level {level}: worst rel={worst:.3e} (retry={worst2:.3e})"
8766 )));
8767 }
8768 }
8769 }
8770 let diag_inv_coarse = with_timing(do_stats, &mut lt.diag, || {
8771 diag_inv_from_csr_cfg_fallback(&a_coarse, cfg, allow_safeguard)
8772 })?;
8773 lt.total = lt.strength
8774 + lt.aggregate
8775 + lt.prolong
8776 + lt.restrict
8777 + lt.rap_symbolic
8778 + lt.rap_numeric
8779 + lt.diag;
8780 if do_stats {
8781 timings.push(lt);
8782 }
8783 diag_stats.push(AmgLevelStats {
8784 p_min_col_norm: rank_diag.min_col_norm,
8785 p_cond_sketched: rank_diag.cond_estimate,
8786 galerkin_worst_rel: galerkin_worst,
8787 });
8788
8789 let mut row_basis_owned = row_basis_for_level;
8791 if let Some(prev) = levels.last_mut() {
8792 prev.p = p.clone();
8793 prev.agg_of = tp.agg_of.clone();
8794 prev.is_c = is_c.clone();
8795 prev.cf = cf_opt.clone();
8796 prev.p2r_pos = p2r_pos;
8797 prev.num_functions = tp.num_functions;
8798 prev.row_basis = row_basis_owned.take();
8799 prev.nns = tp.nns.clone();
8800 prev.layout = layout.clone();
8801 prev.a_next_pat = Some(pat.clone());
8802 prev.a_next_pat_ng = ng_pat_opt.clone();
8803 prev.rap_full2ng_pos = map_opt;
8804 if cfg.keep_transpose {
8805 prev.r = r.clone();
8806 prev.r_row_ptr = None;
8807 prev.r_col_idx = None;
8808 prev.r_vals_scratch = None;
8809 } else {
8810 prev.r = CsrMatrix::identity(0);
8811 prev.r_row_ptr = Some(r_row_ptr);
8812 prev.r_col_idx = Some(r_col_idx);
8813 prev.r_vals_scratch = Some(vec![0.0; prev.r_col_idx.as_ref().unwrap().len()]);
8814 }
8815 #[cfg(feature = "simd")]
8816 build_level_spmv_plans(prev, &spmv_tuning);
8817 }
8818
8819 a_cur = a_coarse.clone();
8821 trials_current = trials_next;
8822 block_size_cur = tp.num_functions;
8823 let mut next_level = AMGLevel {
8824 a: a_coarse,
8825 p: CsrMatrix::identity(a_cur.nrows()),
8826 r: CsrMatrix::identity(a_cur.nrows()),
8827 diag_inv: diag_inv_coarse,
8828 d_sqrt_inv: None,
8829 l1_inv: None,
8830 diag_inv_safe: None,
8831 d_sqrt_inv_safe: None,
8832 cheb: None,
8833 cheb_safe: None,
8834 agg_of: (0..a_cur.nrows()).collect(),
8835 is_c: Vec::new(),
8836 cf: None,
8837 p2r_pos: Vec::new(),
8838 num_functions: 1,
8839 row_basis: None,
8840 layout: if cfg.nodal == NodalMode::Nodal {
8841 Some(DofLayout::new(a_cur.nrows(), block_size_cur))
8842 } else {
8843 None
8844 },
8845 nns: None,
8846 a_next_pat: None,
8847 a_next_pat_ng: None,
8848 rap_full2ng_pos: None,
8849 r_row_ptr: None,
8850 r_col_idx: None,
8851 r_vals_scratch: None,
8852 coarse_solver: None,
8853 ilu0: None,
8854 ras: None,
8855 fsai: None,
8856 a_vals_f32: None,
8857 diag_inv_f32: None,
8858 d_sqrt_inv_f32: None,
8859 l1_inv_f32: None,
8860 fsai_g_vals_f32: None,
8861 fsai_gt_vals_f32: None,
8862 };
8863 update_level_caches(
8864 cfg,
8865 &mut next_level,
8866 need_l1,
8867 need_cheb,
8868 need_safe_diag,
8869 need_cheb_safe,
8870 need_ilu0,
8871 need_ras,
8872 true,
8873 )?;
8874 if need_fsai {
8875 let strength_coarse = if cfg.fsai_use_strength {
8876 Some(Strength::from_csr(
8877 &next_level.a,
8878 cfg.strong_threshold,
8879 cfg.normalize_strength,
8880 ))
8881 } else {
8882 None
8883 };
8884 next_level.fsai = Some(fsai_build_for_level(
8885 cfg,
8886 &next_level.a,
8887 strength_coarse.as_ref(),
8888 )?);
8889 refresh_mixed_precision_shadows(cfg, &mut next_level);
8890 }
8891 #[cfg(feature = "simd")]
8892 build_level_spmv_plans(&mut next_level, &spmv_tuning);
8893 levels.push(next_level);
8894
8895 if do_stats {
8896 level_stats.push(LevelStats {
8897 level: levels.len() - 1,
8898 n: a_cur.nrows(),
8899 nnz_a: a_cur.nnz(),
8900 nnz_p: 0,
8901 nnz_r: 0,
8902 max_row_sum_a: max_row_sum_abs(&a_cur),
8903 eff_nnz_a: Some(eff_nnz(&a_cur, cfg.stats_eps)),
8904 pre_sweeps: cfg.num_grid_sweeps[RelaxPhase::Down.ix()],
8905 post_sweeps: cfg.num_grid_sweeps[RelaxPhase::Up.ix()],
8906 pre_work_estimate: (cfg.num_grid_sweeps[RelaxPhase::Down.ix()] as f64)
8907 * a_cur.nnz() as f64,
8908 post_work_estimate: (cfg.num_grid_sweeps[RelaxPhase::Up.ix()] as f64)
8909 * a_cur.nnz() as f64,
8910 selected_relax_pre: format!("{:?}", cfg.grid_relax_type[RelaxPhase::Down.ix()]),
8911 selected_relax_post: format!("{:?}", cfg.grid_relax_type[RelaxPhase::Up.ix()]),
8912 coarse_solver: None,
8913 });
8914 let ls_len = level_stats.len();
8915 if ls_len >= 2
8916 && let Some(prev) = level_stats.get_mut(ls_len - 2)
8917 {
8918 prev.nnz_p = p.nnz();
8919 prev.nnz_r = r.nnz();
8920 }
8921 }
8922
8923 if a_cur.nrows() >= n {
8924 break;
8925 } if a_cur.nrows() <= cfg.max_coarse_size {
8927 break;
8928 }
8929 if let Some(limit) = cfg.max_operator_complexity {
8930 let oc = operator_complexity_estimate(&levels);
8931 if oc > limit {
8932 break;
8933 }
8934 }
8935 }
8936
8937 diag_stats.push(AmgLevelStats {
8938 p_min_col_norm: 0.0,
8939 p_cond_sketched: 0.0,
8940 galerkin_worst_rel: 0.0,
8941 });
8942 if cfg.non_galerkin.enabled && cfg.non_galerkin.oc_target.is_some() {
8943 enforce_oc_target(&mut levels, cfg)?;
8944 }
8945
8946 #[allow(non_snake_case)]
8947 let L = levels.len() - 1;
8948 let prefer_dense = matches!(cfg.coarse_solve, CoarseSolve::DirectDense)
8949 || levels[L].a.nrows() <= cfg.max_coarse_size;
8950 if prefer_dense {
8951 let mut solver = CoarseDenseLu::new();
8952 solver.setup(&levels[L].a)?;
8953 levels[L].coarse_solver = Some(Mutex::new(Box::new(solver)));
8954 } else if matches!(cfg.coarse_solve, CoarseSolve::ILU) {
8955 let n = levels[L].a.nrows();
8956 let mut ilu = CoarseIlu::new(
8957 cfg.tolerance,
8958 n.min(cfg.max_iterations.max(50)),
8959 cfg.ilu_drop_tol,
8960 cfg.ilu_fill_per_row,
8961 );
8962 ilu.setup(&levels[L].a)?;
8963 levels[L].coarse_solver = Some(Mutex::new(Box::new(ilu)));
8964 }
8965
8966 let hier = AmgHierarchy {
8967 policy: RelaxPolicy {
8968 kind: cfg.grid_relax_type,
8969 sweeps: cfg.num_grid_sweeps,
8970 omega: cfg.jacobi_omega,
8971 },
8972 coarse_solve: cfg.coarse_solve,
8973 levels,
8974 };
8975
8976 let stats_opt = if do_stats {
8977 let mut stats = AmgStats::from_hierarchy(&hier);
8978 stats.levels = level_stats;
8979 stats.total_smoothing_work = stats
8980 .levels
8981 .iter()
8982 .map(|l| l.pre_work_estimate + l.post_work_estimate)
8983 .sum();
8984 stats.selected_dist_coarse_route = Some(
8985 dist_route_label(cfg.dist_coarse_solver_route, cfg.dist_coarse_strategy).to_string(),
8986 );
8987 stats.dist_route_fallback =
8988 dist_route_fallback_labels(cfg.dist_coarse_solver_route, cfg.dist_coarse_strategy);
8989 stats.diagnostics = diag_stats;
8990 let mut setup = SetupTimings::default();
8991 setup.per_level = timings;
8992 if let Some(t0) = t_setup_all {
8993 setup.total_setup = toc(t0);
8994 }
8995 for lt in &setup.per_level {
8996 setup.total_symbolic +=
8997 lt.strength + lt.aggregate + lt.prolong + lt.restrict + lt.rap_symbolic;
8998 setup.total_numeric += lt.rap_numeric + lt.diag;
8999 }
9000 stats.setup = setup;
9001 Some(stats)
9002 } else {
9003 None
9004 };
9005
9006 Ok((hier, stats_opt))
9007}
9008
9009fn build_smoother_only_hierarchy(
9010 fine: &CsrMatrix<f64>,
9011 cfg: &mut AMGConfig,
9012) -> Result<(AmgHierarchy, Option<AmgStats>), KError> {
9013 let allow_safeguard = cfg
9014 .grid_relax_type
9015 .contains(&RelaxType::SafeguardedGaussSeidel)
9016 || cfg.grid_relax_type.contains(&RelaxType::ChebyshevSafe)
9017 || cfg.grid_relax_type.contains(&RelaxType::Ilu0)
9018 || cfg.grid_relax_type.contains(&RelaxType::Ras);
9019 let need_l1 = cfg.grid_relax_type.contains(&RelaxType::L1Jacobi);
9020 let need_cheb = cfg.grid_relax_type.contains(&RelaxType::Chebyshev);
9021 let need_cheb_safe = cfg.grid_relax_type.contains(&RelaxType::ChebyshevSafe);
9022 let need_safe_diag = cfg
9023 .grid_relax_type
9024 .contains(&RelaxType::SafeguardedGaussSeidel)
9025 || need_cheb_safe;
9026 let need_ilu0 = cfg.grid_relax_type.contains(&RelaxType::Ilu0);
9027 let need_ras = cfg.grid_relax_type.contains(&RelaxType::Ras);
9028 let need_fsai = cfg.grid_relax_type.contains(&RelaxType::Fsai);
9029
9030 let diag0 = diag_inv_from_csr_cfg_fallback(fine, cfg, allow_safeguard)?;
9031 let layout0 = if cfg.nodal == NodalMode::Nodal {
9032 Some(DofLayout::new(fine.nrows(), cfg.block_size))
9033 } else {
9034 None
9035 };
9036 let l0_num_functions = cfg
9037 .near_nullspace
9038 .as_ref()
9039 .map(|nns| nns.basis.len().max(1))
9040 .or_else(|| layout0.as_ref().map(|layout| layout.block_size.max(1)))
9041 .unwrap_or_else(|| cfg.num_functions.max(1));
9042 let mut l0 = AMGLevel {
9043 a: fine.clone(),
9044 p: CsrMatrix::identity(fine.nrows()),
9045 r: CsrMatrix::identity(fine.nrows()),
9046 diag_inv: diag0,
9047 d_sqrt_inv: None,
9048 l1_inv: None,
9049 diag_inv_safe: None,
9050 d_sqrt_inv_safe: None,
9051 cheb: None,
9052 cheb_safe: None,
9053 agg_of: (0..fine.nrows()).collect(),
9054 is_c: Vec::new(),
9055 cf: None,
9056 p2r_pos: Vec::new(),
9057 num_functions: l0_num_functions,
9058 row_basis: None,
9059 layout: layout0,
9060 nns: cfg.near_nullspace.as_ref().map(|nns| nns.basis.clone()),
9061 a_next_pat: None,
9062 a_next_pat_ng: None,
9063 rap_full2ng_pos: None,
9064 r_row_ptr: None,
9065 r_col_idx: None,
9066 r_vals_scratch: None,
9067 coarse_solver: None,
9068 ilu0: None,
9069 ras: None,
9070 fsai: None,
9071 a_vals_f32: None,
9072 diag_inv_f32: None,
9073 d_sqrt_inv_f32: None,
9074 l1_inv_f32: None,
9075 fsai_g_vals_f32: None,
9076 fsai_gt_vals_f32: None,
9077 };
9078 update_level_caches(
9079 cfg,
9080 &mut l0,
9081 need_l1,
9082 need_cheb,
9083 need_safe_diag,
9084 need_cheb_safe,
9085 need_ilu0,
9086 need_ras,
9087 true,
9088 )?;
9089 if need_fsai {
9090 let strength0 = if cfg.fsai_use_strength {
9091 Some(Strength::from_csr(
9092 &l0.a,
9093 cfg.strong_threshold,
9094 cfg.normalize_strength,
9095 ))
9096 } else {
9097 None
9098 };
9099 l0.fsai = Some(fsai_build_for_level(cfg, &l0.a, strength0.as_ref())?);
9100 refresh_mixed_precision_shadows(cfg, &mut l0);
9101 }
9102
9103 let hier = AmgHierarchy {
9104 policy: RelaxPolicy {
9105 kind: cfg.grid_relax_type,
9106 sweeps: cfg.num_grid_sweeps,
9107 omega: cfg.jacobi_omega,
9108 },
9109 coarse_solve: CoarseSolve::Smoother,
9110 levels: vec![l0],
9111 };
9112
9113 let stats_opt = if cfg.logging_level > 0 {
9114 let mut stats = AmgStats::from_hierarchy(&hier);
9115 stats.levels = vec![LevelStats {
9116 level: 0,
9117 n: fine.nrows(),
9118 nnz_a: fine.nnz(),
9119 nnz_p: 0,
9120 nnz_r: 0,
9121 max_row_sum_a: max_row_sum_abs(fine),
9122 eff_nnz_a: Some(eff_nnz(fine, cfg.stats_eps)),
9123 pre_sweeps: cfg.num_grid_sweeps[RelaxPhase::Down.ix()],
9124 post_sweeps: cfg.num_grid_sweeps[RelaxPhase::Up.ix()],
9125 pre_work_estimate: (cfg.num_grid_sweeps[RelaxPhase::Down.ix()] as f64)
9126 * fine.nnz() as f64,
9127 post_work_estimate: (cfg.num_grid_sweeps[RelaxPhase::Up.ix()] as f64)
9128 * fine.nnz() as f64,
9129 selected_relax_pre: format!("{:?}", cfg.grid_relax_type[RelaxPhase::Down.ix()]),
9130 selected_relax_post: format!("{:?}", cfg.grid_relax_type[RelaxPhase::Up.ix()]),
9131 coarse_solver: Some(format!("{:?}", CoarseSolve::Smoother)),
9132 }];
9133 stats.total_smoothing_work = stats
9134 .levels
9135 .iter()
9136 .map(|l| l.pre_work_estimate + l.post_work_estimate)
9137 .sum();
9138 stats.selected_dist_coarse_route = Some(
9139 dist_route_label(cfg.dist_coarse_solver_route, cfg.dist_coarse_strategy).to_string(),
9140 );
9141 stats.dist_route_fallback =
9142 dist_route_fallback_labels(cfg.dist_coarse_solver_route, cfg.dist_coarse_strategy);
9143 stats.diagnostics = vec![AmgLevelStats {
9144 p_min_col_norm: 0.0,
9145 p_cond_sketched: 0.0,
9146 galerkin_worst_rel: 0.0,
9147 }];
9148 Some(stats)
9149 } else {
9150 None
9151 };
9152
9153 Ok((hier, stats_opt))
9154}
9155
9156fn enforce_oc_target(levels: &mut Vec<AMGLevel>, cfg: &mut AMGConfig) -> Result<(), KError> {
9157 if let Some(target) = cfg.non_galerkin.oc_target {
9158 let allow_safeguard = cfg
9159 .grid_relax_type
9160 .contains(&RelaxType::SafeguardedGaussSeidel)
9161 || cfg.grid_relax_type.contains(&RelaxType::ChebyshevSafe)
9162 || cfg.grid_relax_type.contains(&RelaxType::Ilu0)
9163 || cfg.grid_relax_type.contains(&RelaxType::Ras);
9164 let mut trials_current = make_trial_matrix(cfg, levels[0].a.nrows())?;
9165 for _ in 0..cfg.non_galerkin.oc_max_iter {
9166 let oc = operator_complexity_estimate(levels);
9167 if oc <= target {
9168 break;
9169 }
9170 cfg.non_galerkin.drop_abs *= 1.25;
9171 cfg.non_galerkin.drop_rel = (cfg.non_galerkin.drop_rel * 1.1).min(0.95);
9172 for l in 0..levels.len() - 1 {
9173 if let Some(pat_full) = levels[l].a_next_pat.clone() {
9174 let r_tmp_storage = if cfg.keep_transpose {
9175 None
9176 } else {
9177 Some(build_r_from_p(&mut levels[l]))
9178 };
9179 let r_for_ops = r_tmp_storage.as_ref().unwrap_or(&levels[l].r);
9180 let trials_next = if let Some(ref trials) = trials_current {
9181 let mut next = Mat::<f64>::zeros(r_for_ops.nrows(), trials.ncols());
9182 restrict_trials(r_for_ops, trials.as_ref(), next.as_mut())?;
9183 Some(next)
9184 } else {
9185 None
9186 };
9187 if l + 1 < cfg.non_galerkin.start_level {
9188 trials_current = trials_next;
9189 continue;
9190 }
9191 let mut vals_full = vec![0.0; pat_full.col_idx.len()];
9192 rap_numeric(
9193 &pat_full,
9194 r_for_ops,
9195 &levels[l].a,
9196 &levels[l].p,
9197 &mut vals_full,
9198 );
9199 {
9200 let mut rf = |row: usize| RowFilter {
9201 tau_abs: cfg.rap_truncation_abs,
9202 tau_rel: cfg.truncation_factor,
9203 k_max: cfg.rap_max_elements_per_row,
9204 must_keep: if cfg.keep_pivot_in_rap {
9205 Some(row)
9206 } else {
9207 None
9208 },
9209 };
9210 apply_filter_to_csr_values_in_place(
9211 pat_full.nrows,
9212 &pat_full.row_ptr,
9213 &pat_full.col_idx,
9214 &mut vals_full,
9215 &mut rf,
9216 );
9217 }
9218 let block_size_next = levels[l].num_functions.max(1);
9219 let mut a_full = CsrMatrix::from_csr(
9220 pat_full.nrows,
9221 pat_full.ncols,
9222 pat_full.row_ptr.clone(),
9223 pat_full.col_idx.clone(),
9224 vals_full,
9225 );
9226 if !cfg.filter_after_non_galerkin {
9227 apply_trial_compensation(
9228 cfg,
9229 &mut a_full,
9230 trials_next.as_ref(),
9231 block_size_next,
9232 )?;
9233 }
9234 let (ng_pat, ng_vals, full2ng) = non_galerkin_filter_coarse(
9235 &pat_full,
9236 a_full.values(),
9237 cfg.non_galerkin.symmetry,
9238 NgRowFilter {
9239 tau_abs: cfg.non_galerkin.drop_abs,
9240 tau_rel: cfg.non_galerkin.drop_rel,
9241 k_max: cfg.non_galerkin.cap_row,
9242 lump_diag: cfg.non_galerkin.lump_diagonal,
9243 },
9244 );
9245 let mut a_ng = CsrMatrix::from_csr(
9246 ng_pat.nrows,
9247 ng_pat.ncols,
9248 ng_pat.row_ptr.clone(),
9249 ng_pat.col_idx.clone(),
9250 ng_vals,
9251 );
9252 if cfg.filter_after_non_galerkin {
9253 apply_trial_compensation(
9254 cfg,
9255 &mut a_ng,
9256 trials_next.as_ref(),
9257 block_size_next,
9258 )?;
9259 }
9260 levels[l].a_next_pat_ng = Some(ng_pat.clone());
9261 levels[l].rap_full2ng_pos = Some(full2ng);
9262 levels[l + 1].a = a_ng;
9263 levels[l + 1].diag_inv =
9264 diag_inv_from_csr_cfg_fallback(&levels[l + 1].a, cfg, allow_safeguard)?;
9265 trials_current = trials_next;
9266 } else {
9267 trials_current = None;
9268 }
9269 }
9270 }
9271 }
9272 Ok(())
9273}
9274
9275fn orthonormalize_aggregate(
9278 rows: &[usize],
9279 basis_cols: &[Vec<f64>],
9280 weights: &[f64],
9281 mfun: usize,
9282 row_basis: &mut [f64],
9283) {
9284 if rows.is_empty() {
9285 return;
9286 }
9287 let mut q_cols: Vec<Vec<f64>> = Vec::new();
9288 let limit = basis_cols.len().min(mfun);
9289 for f in 0..limit {
9290 let mut col: Vec<f64> = rows.iter().map(|&row_idx| basis_cols[f][row_idx]).collect();
9291 for prev in 0..q_cols.len() {
9292 let mut dot = 0.0;
9293 for (local_ix, &row_idx) in rows.iter().enumerate() {
9294 let w = weights[row_idx];
9295 dot += w * col[local_ix] * q_cols[prev][local_ix];
9296 }
9297 for local_ix in 0..col.len() {
9298 col[local_ix] -= dot * q_cols[prev][local_ix];
9299 }
9300 }
9301 let mut norm_sq = 0.0;
9302 for (local_ix, &row_idx) in rows.iter().enumerate() {
9303 let w = weights[row_idx];
9304 let v = col[local_ix];
9305 norm_sq += w * v * v;
9306 }
9307 if norm_sq <= 1e-24 {
9308 continue;
9309 }
9310 let norm = norm_sq.sqrt();
9311 for val in &mut col {
9312 *val /= norm;
9313 }
9314 q_cols.push(col);
9315 }
9316 for (local_ix, &row_idx) in rows.iter().enumerate() {
9317 for f in 0..mfun {
9318 let val = q_cols.get(f).map(|col| col[local_ix]).unwrap_or(0.0);
9319 row_basis[row_idx * mfun + f] = val;
9320 }
9321 }
9322}
9323
9324fn l1_diag_inv(a: &CsrMatrix<f64>) -> Vec<f64> {
9325 let n = a.nrows();
9326 let mut inv = vec![0.0; n];
9327 for i in 0..n {
9328 let mut s = 0.0;
9329 for p in a.row_ptr()[i]..a.row_ptr()[i + 1] {
9330 s += a.values()[p].abs();
9331 }
9332 inv[i] = 1.0 / s.max(1e-30);
9333 }
9334 inv
9335}
9336
9337fn diag_inv_from_csr_with_floor(
9338 a: &CsrMatrix<f64>,
9339 floor: f64,
9340 require_positive: bool,
9341) -> Result<Vec<f64>, KError> {
9342 let n = a.nrows();
9343 let mut d = vec![0.0; n];
9344 for i in 0..n {
9345 let rs = a.row_ptr()[i];
9346 let re = a.row_ptr()[i + 1];
9347 let mut aii = 0.0;
9348 for p in rs..re {
9349 if a.col_idx()[p] == i {
9350 aii = a.values()[p];
9351 break;
9352 }
9353 }
9354 if floor > 0.0 && aii <= 0.0 {
9355 aii += floor;
9356 }
9357 if aii.abs() < 1e-14 {
9358 return Err(KError::SolveError(format!("near-zero diagonal at row {i}")));
9359 }
9360 if require_positive && aii <= 0.0 {
9361 return Err(KError::SolveError(format!(
9362 "non-positive diagonal at row {i} (value {aii})"
9363 )));
9364 }
9365 d[i] = 1.0 / aii;
9366 }
9367 Ok(d)
9368}
9369
9370fn diag_inv_from_csr_cfg(a: &CsrMatrix<f64>, cfg: &AMGConfig) -> Result<Vec<f64>, KError> {
9371 let floor = if cfg.require_spd {
9372 cfg.spd_diag_floor.max(0.0)
9373 } else {
9374 0.0
9375 };
9376 diag_inv_from_csr_with_floor(a, floor, cfg.require_spd)
9377}
9378
9379fn diag_inv_from_csr_safeguarded(a: &CsrMatrix<f64>) -> Vec<f64> {
9380 let n = a.nrows();
9381 let mut inv = vec![0.0; n];
9382 let row_ptr = a.row_ptr();
9383 let col_idx = a.col_idx();
9384 let vals = a.values();
9385 for i in 0..n {
9386 let rs = row_ptr[i];
9387 let re = row_ptr[i + 1];
9388 let mut diag = 0.0;
9389 let mut row_sum = 0.0;
9390 for p in rs..re {
9391 let v = vals[p];
9392 row_sum += v.abs();
9393 if col_idx[p] == i {
9394 diag = v.abs();
9395 }
9396 }
9397 let denom = diag.max(row_sum).max(1e-30);
9398 inv[i] = 1.0 / denom;
9399 }
9400 inv
9401}
9402
9403fn diag_inv_from_csr_cfg_fallback(
9404 a: &CsrMatrix<f64>,
9405 cfg: &AMGConfig,
9406 allow_safeguard: bool,
9407) -> Result<Vec<f64>, KError> {
9408 match diag_inv_from_csr_cfg(a, cfg) {
9409 Ok(v) => Ok(v),
9410 #[allow(unused_variables)]
9411 Err(err) if allow_safeguard && !cfg.require_spd => Ok(diag_inv_from_csr_safeguarded(a)),
9412 Err(err) => Err(err),
9413 }
9414}
9415
9416fn diag_inv_from_csr(a: &CsrMatrix<f64>) -> Result<Vec<f64>, KError> {
9417 diag_inv_from_csr_with_floor(a, 0.0, false)
9418}
9419
9420fn make_d_sqrt_inv(diag_inv: &[f64]) -> Vec<f64> {
9421 let mut out = vec![0.0; diag_inv.len()];
9422 for (dst, &d) in out.iter_mut().zip(diag_inv.iter()) {
9423 *dst = if d > 0.0 { d.sqrt() } else { 0.0 };
9424 }
9425 out
9426}
9427
9428fn compute_cheb_data(
9429 cfg: &AMGConfig,
9430 a: &CsrMatrix<f64>,
9431 _diag_inv: &[f64],
9432 d_sqrt_inv: &[f64],
9433) -> Result<ChebData, KError> {
9434 let mut lam_max = chebyshev::estimate_lmax_sym(a, d_sqrt_inv, cfg.chebyshev_power_steps)?;
9435 if !lam_max.is_finite() || lam_max <= 0.0 {
9436 let mut fallback: f64 = 0.0;
9437 let row_ptr = a.row_ptr();
9438 let col_idx = a.col_idx();
9439 let vals = a.values();
9440 for i in 0..a.nrows() {
9441 let rs = row_ptr[i];
9442 let re = row_ptr[i + 1];
9443 let mut diag = 0.0f64;
9444 let mut sum = 0.0f64;
9445 for p in rs..re {
9446 let j = col_idx[p];
9447 let v = vals[p].abs();
9448 if j == i {
9449 diag = v;
9450 } else {
9451 sum += v;
9452 }
9453 }
9454 if diag > 0.0 {
9455 fallback = fallback.max(sum / diag);
9456 }
9457 }
9458 lam_max = fallback;
9459 }
9460 if !lam_max.is_finite() || lam_max <= 0.0 {
9461 lam_max = 1.0;
9462 }
9463 let safety = cfg.chebyshev_safety.max(1.0);
9464 lam_max *= safety;
9465 if !lam_max.is_finite() || lam_max <= 0.0 {
9466 lam_max = safety.max(1.0);
9467 }
9468 let ratio = cfg.chebyshev_lower_ratio.clamp(1e-6, 0.99);
9469 let lam_min = (ratio * lam_max).max(1e-30);
9470 Ok(ChebData {
9471 lambda_max: lam_max,
9472 lambda_min: lam_min,
9473 })
9474}
9475
9476fn update_level_caches(
9477 cfg: &AMGConfig,
9478 level: &mut AMGLevel,
9479 need_l1: bool,
9480 need_cheb: bool,
9481 need_safe_diag: bool,
9482 need_cheb_safe: bool,
9483 need_ilu0: bool,
9484 need_ras: bool,
9485 recompute_cheb: bool,
9486) -> Result<(), KError> {
9487 if need_l1 {
9488 level.l1_inv = Some(l1_diag_inv(&level.a));
9489 } else {
9490 level.l1_inv = None;
9491 }
9492 if need_cheb {
9493 let d_sqrt = make_d_sqrt_inv(&level.diag_inv);
9494 if recompute_cheb || level.cheb.is_none() {
9495 let cheb = compute_cheb_data(cfg, &level.a, &level.diag_inv, &d_sqrt)?;
9496 level.cheb = Some(cheb);
9497 }
9498 level.d_sqrt_inv = Some(d_sqrt);
9499 } else {
9500 level.d_sqrt_inv = None;
9501 level.cheb = None;
9502 }
9503 if need_safe_diag {
9504 let diag_safe = diag_inv_from_csr_safeguarded(&level.a);
9505 level.diag_inv_safe = Some(diag_safe);
9506 } else {
9507 level.diag_inv_safe = None;
9508 }
9509 if need_cheb_safe {
9510 let diag_safe = level
9511 .diag_inv_safe
9512 .as_ref()
9513 .ok_or_else(|| KError::InvalidInput("ChebyshevSafe cache missing".into()))?;
9514 let mut d_sqrt = vec![0.0; diag_safe.len()];
9515 for (dst, &d) in d_sqrt.iter_mut().zip(diag_safe.iter()) {
9516 *dst = if d > 0.0 { d.sqrt() } else { 0.0 };
9517 }
9518 if recompute_cheb || level.cheb_safe.is_none() {
9519 let cheb = compute_cheb_data(cfg, &level.a, diag_safe, &d_sqrt)?;
9520 level.cheb_safe = Some(cheb);
9521 }
9522 level.d_sqrt_inv_safe = Some(d_sqrt);
9523 } else {
9524 level.d_sqrt_inv_safe = None;
9525 level.cheb_safe = None;
9526 }
9527 if need_ilu0 {
9528 build_ilu0_cache(level)?;
9529 } else {
9530 level.ilu0 = None;
9531 }
9532 if need_ras {
9533 build_ras_cache(level)?;
9534 } else {
9535 level.ras = None;
9536 }
9537 refresh_mixed_precision_shadows(cfg, level);
9538 Ok(())
9539}
9540
9541fn build_ilu0_cache(level: &mut AMGLevel) -> Result<(), KError> {
9542 #[cfg(feature = "complex")]
9543 {
9544 let _ = level;
9545 return Err(KError::Unsupported(
9546 "ILU0 cache is not supported in complex AMG mode".into(),
9547 ));
9548 }
9549 #[cfg(not(feature = "complex"))]
9550 {
9551 let mut cfg = IluCsrConfig::default();
9552 cfg.kind = IluKind::Ilu0;
9553 cfg.pivot = PivotStrategy::DiagonalPerturbation;
9554 cfg.pivot_threshold = 1e-12;
9555 cfg.diag_perturb_factor = 1e-10;
9556 cfg.level_sched = false;
9557 cfg.reordering = ReorderingOptions::default();
9558 cfg.conditioning = ConditioningOptions::default();
9559 let mut ilu = IluCsr::new_with_config(cfg);
9560 let op = crate::matrix::op::CsrOp::new(Arc::new(level.a.clone()));
9561 ilu.setup(&op)?;
9562 level.ilu0 = Some(Mutex::new(ilu));
9563 Ok(())
9564 }
9565}
9566
9567fn build_ras_cache(level: &mut AMGLevel) -> Result<(), KError> {
9568 #[cfg(feature = "complex")]
9569 {
9570 let _ = level;
9571 return Err(KError::Unsupported(
9572 "RAS cache is not supported in complex AMG mode".into(),
9573 ));
9574 }
9575 #[cfg(not(feature = "complex"))]
9576 {
9577 let cfg = AsmConfig {
9578 overlap: 1,
9579 combine: AsmCombine::Restricted,
9580 local_solver: AsmLocalSolver::ILU,
9581 local_sweeps: 1,
9582 weight_partition_of_unity: false,
9583 deterministic: true,
9584 nparts: None,
9585 };
9586 let mut ras = Asm::with_config(cfg);
9587 let op = crate::matrix::op::CsrOp::new(Arc::new(level.a.clone()));
9588 Preconditioner::setup(&mut ras, &op)?;
9589 level.ras = Some(Mutex::new(ras));
9590 Ok(())
9591 }
9592}
9593
9594fn cast_slice_to_f32(src: &[f64]) -> Vec<f32> {
9595 src.iter().map(|&v| v as f32).collect()
9596}
9597
9598fn refresh_mixed_precision_shadows(cfg: &AMGConfig, level: &mut AMGLevel) {
9599 let Some(mp) = cfg.mixed_precision else {
9600 level.a_vals_f32 = None;
9601 level.diag_inv_f32 = None;
9602 level.d_sqrt_inv_f32 = None;
9603 level.l1_inv_f32 = None;
9604 level.fsai_g_vals_f32 = None;
9605 level.fsai_gt_vals_f32 = None;
9606 return;
9607 };
9608 if cfg.mixed_storage != MixedStorage::Cached {
9609 level.a_vals_f32 = None;
9610 level.diag_inv_f32 = None;
9611 level.d_sqrt_inv_f32 = None;
9612 level.l1_inv_f32 = None;
9613 level.fsai_g_vals_f32 = None;
9614 level.fsai_gt_vals_f32 = None;
9615 return;
9616 }
9617 if mp.residual_enabled() || mp.smoothers_enabled() {
9618 level.a_vals_f32 = Some(cast_slice_to_f32(level.a.values()));
9619 } else {
9620 level.a_vals_f32 = None;
9621 }
9622 if mp.smoothers_enabled() {
9623 level.diag_inv_f32 = Some(cast_slice_to_f32(&level.diag_inv));
9624 if let Some(ref d) = level.d_sqrt_inv {
9625 level.d_sqrt_inv_f32 = Some(cast_slice_to_f32(d));
9626 } else {
9627 level.d_sqrt_inv_f32 = None;
9628 }
9629 if let Some(ref l1) = level.l1_inv {
9630 level.l1_inv_f32 = Some(cast_slice_to_f32(l1));
9631 } else {
9632 level.l1_inv_f32 = None;
9633 }
9634 if let Some(ref fsai) = level.fsai {
9635 level.fsai_g_vals_f32 = Some(cast_slice_to_f32(fsai.g.values()));
9636 level.fsai_gt_vals_f32 = Some(cast_slice_to_f32(fsai.gt.values()));
9637 } else {
9638 level.fsai_g_vals_f32 = None;
9639 level.fsai_gt_vals_f32 = None;
9640 }
9641 } else {
9642 level.diag_inv_f32 = None;
9643 level.d_sqrt_inv_f32 = None;
9644 level.l1_inv_f32 = None;
9645 level.fsai_g_vals_f32 = None;
9646 level.fsai_gt_vals_f32 = None;
9647 }
9648}
9649
9650fn csr_lookup(a: &CsrMatrix<f64>, row: usize, col: usize) -> f64 {
9651 let rp = a.row_ptr();
9652 let ci = a.col_idx();
9653 let vv = a.values();
9654 let (rs, re) = (rp[row], rp[row + 1]);
9655 match ci[rs..re].binary_search(&col) {
9656 Ok(pos) => vv[rs + pos],
9657 Err(_) => 0.0,
9658 }
9659}
9660
9661fn gather_dense_submatrix(a: &CsrMatrix<f64>, pattern: &[usize], buf: &mut Vec<f64>) {
9662 let m = pattern.len();
9663 buf.resize(m * m, 0.0);
9664 for (i_local, &i) in pattern.iter().enumerate() {
9665 for (j_local, &j) in pattern.iter().take(i_local + 1).enumerate() {
9666 let val = csr_lookup(a, i, j);
9667 buf[i_local * m + j_local] = val;
9668 buf[j_local * m + i_local] = val;
9669 }
9670 }
9671}
9672
9673fn cholesky_factor(mat: &mut [f64], n: usize) -> bool {
9674 for i in 0..n {
9675 for j in 0..i {
9676 let mut sum = mat[i * n + j];
9677 for k in 0..j {
9678 sum -= mat[i * n + k] * mat[j * n + k];
9679 }
9680 let diag = mat[j * n + j];
9681 if diag <= 0.0 {
9682 return false;
9683 }
9684 sum /= diag;
9685 mat[i * n + j] = sum;
9686 }
9687 let mut sum = mat[i * n + i];
9688 for k in 0..i {
9689 let v = mat[i * n + k];
9690 sum -= v * v;
9691 }
9692 if sum <= 0.0 {
9693 return false;
9694 }
9695 let diag = sum.sqrt();
9696 mat[i * n + i] = diag;
9697 for j in (i + 1)..n {
9698 mat[i * n + j] = 0.0;
9699 }
9700 }
9701 true
9702}
9703
9704fn cholesky_solve(mat: &[f64], rhs: &[f64], n: usize) -> Vec<f64> {
9705 let mut y = vec![0.0; n];
9706 for i in 0..n {
9707 let mut sum = rhs[i];
9708 for k in 0..i {
9709 sum -= mat[i * n + k] * y[k];
9710 }
9711 let diag = mat[i * n + i];
9712 y[i] = sum / diag;
9713 }
9714 let mut x = vec![0.0; n];
9715 for i in (0..n).rev() {
9716 let mut sum = y[i];
9717 for k in (i + 1)..n {
9718 sum -= mat[k * n + i] * x[k];
9719 }
9720 let diag = mat[i * n + i];
9721 x[i] = sum / diag;
9722 }
9723 x
9724}
9725
9726fn solve_fsai_system(base: &[f64], rhs: &[f64], n: usize, lambda: f64) -> Option<Vec<f64>> {
9727 let mut attempt = if lambda >= 0.0 { lambda } else { 0.0 };
9728 let mut mat = vec![0.0; n * n];
9729 let mut tries = 0;
9730 while tries < 5 {
9731 mat.copy_from_slice(base);
9732 for d in 0..n {
9733 mat[d * n + d] += attempt;
9734 }
9735 if cholesky_factor(&mut mat, n) {
9736 return Some(cholesky_solve(&mat, rhs, n));
9737 }
9738 attempt = if attempt == 0.0 {
9739 1e-12
9740 } else {
9741 attempt * 10.0
9742 };
9743 tries += 1;
9744 }
9745 None
9746}
9747
9748fn prune_pattern_row(a: &CsrMatrix<f64>, row: usize, pattern: &mut Vec<usize>, cap: usize) {
9749 if pattern.is_empty() {
9750 pattern.push(row);
9751 }
9752 if !pattern.contains(&row) {
9753 pattern.push(row);
9754 }
9755 if cap == 0 {
9756 pattern.clear();
9757 pattern.push(row);
9758 return;
9759 }
9760 if pattern.len() <= cap {
9761 pattern.sort_unstable();
9762 return;
9763 }
9764 let mut entries: Vec<(usize, f64)> = pattern
9765 .iter()
9766 .copied()
9767 .filter(|&col| col != row)
9768 .map(|col| (col, csr_lookup(a, row, col).abs()))
9769 .collect();
9770 entries.sort_unstable_by(
9771 |a, b| match b.1.partial_cmp(&a.1).unwrap_or(CmpOrdering::Equal) {
9772 CmpOrdering::Equal => a.0.cmp(&b.0),
9773 other => other,
9774 },
9775 );
9776 let mut kept = Vec::with_capacity(cap.max(1));
9777 kept.push(row);
9778 for (col, _) in entries.into_iter().take(cap.saturating_sub(1)) {
9779 kept.push(col);
9780 }
9781 kept.sort_unstable();
9782 pattern.clear();
9783 pattern.extend(kept);
9784}
9785
9786fn fsai_build_pattern(
9787 a: &CsrMatrix<f64>,
9788 strength: Option<&Strength>,
9789 dist: usize,
9790 cap: usize,
9791) -> Vec<Vec<usize>> {
9792 let n = a.nrows();
9793 let mut patterns: Vec<Vec<usize>> = Vec::with_capacity(n);
9794 let mut mark = vec![0usize; n];
9795 let mut stamp = 1usize;
9796 let mut frontier: Vec<usize> = Vec::new();
9797 let mut next: Vec<usize> = Vec::new();
9798 for i in 0..n {
9799 let mut acc: Vec<usize> = vec![i];
9800 frontier.clear();
9801 frontier.push(i);
9802 mark[i] = stamp;
9803 for _ in 0..dist {
9804 next.clear();
9805 for &u in &frontier {
9806 let neighbors: &[usize] = if let Some(g) = strength {
9807 let rs = g.row_ptr[u];
9808 let re = g.row_ptr[u + 1];
9809 &g.col_idx[rs..re]
9810 } else {
9811 let rp = a.row_ptr();
9812 let ci = a.col_idx();
9813 let (rs, re) = (rp[u], rp[u + 1]);
9814 &ci[rs..re]
9815 };
9816 for &v in neighbors {
9817 if v >= n {
9818 continue;
9819 }
9820 if mark[v] != stamp {
9821 mark[v] = stamp;
9822 acc.push(v);
9823 next.push(v);
9824 }
9825 }
9826 }
9827 if next.is_empty() {
9828 break;
9829 }
9830 std::mem::swap(&mut frontier, &mut next);
9831 }
9832 stamp = stamp.wrapping_add(1);
9833 acc.sort_unstable();
9834 acc.dedup();
9835 prune_pattern_row(a, i, &mut acc, cap);
9836 patterns.push(acc);
9837 }
9838 if cap > 0 {
9839 let mut additions: Vec<(usize, usize)> = Vec::new();
9840 for i in 0..n {
9841 let current = patterns[i].clone();
9842 for &j in ¤t {
9843 if j >= n || j == i {
9844 continue;
9845 }
9846 if patterns[j].binary_search(&i).is_err() {
9847 additions.push((j, i));
9848 }
9849 }
9850 }
9851 additions.sort_unstable();
9852 additions.dedup();
9853 for (row, col) in additions {
9854 let pat = &mut patterns[row];
9855 match pat.binary_search(&col) {
9856 Ok(_) => {}
9857 Err(pos) => pat.insert(pos, col),
9858 }
9859 prune_pattern_row(a, row, pat, cap);
9860 }
9861 }
9862 patterns
9863}
9864
9865fn fsai_enrich_pattern(
9866 a: &CsrMatrix<f64>,
9867 pattern: &mut Vec<usize>,
9868 sol: &[f64],
9869 cap: usize,
9870) -> usize {
9871 if pattern.len() >= cap {
9872 return 0;
9873 }
9874 let rp = a.row_ptr();
9875 let ci = a.col_idx();
9876 let vv = a.values();
9877 let mut accum: Vec<(usize, f64)> = Vec::new();
9878 for (local, &row) in pattern.iter().enumerate() {
9879 let coeff = sol[local];
9880 if coeff == 0.0 {
9881 continue;
9882 }
9883 let (rs, re) = (rp[row], rp[row + 1]);
9884 for idx in rs..re {
9885 let col = ci[idx];
9886 if pattern.binary_search(&col).is_ok() {
9887 continue;
9888 }
9889 accum.push((col, -coeff * vv[idx]));
9890 }
9891 }
9892 if accum.is_empty() {
9893 return 0;
9894 }
9895 accum.sort_unstable_by(|a, b| a.0.cmp(&b.0));
9896 let mut merged: Vec<(usize, f64)> = Vec::new();
9897 for (col, val) in accum {
9898 if let Some(last) = merged.last_mut()
9899 && last.0 == col
9900 {
9901 last.1 += val;
9902 continue;
9903 }
9904 merged.push((col, val));
9905 }
9906 merged.sort_unstable_by(|a, b| {
9907 match b
9908 .1
9909 .abs()
9910 .partial_cmp(&a.1.abs())
9911 .unwrap_or(CmpOrdering::Equal)
9912 {
9913 CmpOrdering::Equal => a.0.cmp(&b.0),
9914 other => other,
9915 }
9916 });
9917 let mut added = 0usize;
9918 let space = cap.saturating_sub(pattern.len());
9919 for (col, _) in merged.into_iter() {
9920 if added >= space {
9921 break;
9922 }
9923 if col >= a.ncols() {
9924 continue;
9925 }
9926 if let Err(pos) = pattern.binary_search(&col) {
9927 pattern.insert(pos, col);
9928 added += 1;
9929 }
9930 }
9931 added
9932}
9933
9934fn fsai_drop_entries(
9935 a: &CsrMatrix<f64>,
9936 row: usize,
9937 pattern: &[usize],
9938 sol: &[f64],
9939 drop_tol: f64,
9940 cap: usize,
9941) -> (Vec<usize>, Vec<f64>) {
9942 let mut norm = 0.0;
9943 for &v in sol {
9944 norm += v * v;
9945 }
9946 let norm = norm.sqrt();
9947 let thr = drop_tol.max(0.0) * norm.max(1e-32);
9948 let mut cols: Vec<usize> = Vec::new();
9949 let mut vals: Vec<f64> = Vec::new();
9950 for (col, &val) in pattern.iter().zip(sol.iter()) {
9951 if *col == row {
9952 let mut keep = val;
9953 if keep.abs() < thr {
9954 let diag = csr_lookup(a, row, row);
9955 keep = if diag.abs() > 0.0 { 1.0 / diag } else { 1.0 };
9956 }
9957 cols.push(*col);
9958 vals.push(keep);
9959 } else if val.abs() >= thr {
9960 cols.push(*col);
9961 vals.push(val);
9962 }
9963 }
9964 if cols.is_empty() {
9965 let diag = csr_lookup(a, row, row);
9966 cols.push(row);
9967 vals.push(if diag.abs() > 0.0 { 1.0 / diag } else { 1.0 });
9968 }
9969 let limit = cap.max(1);
9970 if cols.len() > limit {
9971 let diag_pos = cols.iter().position(|&c| c == row).unwrap_or(0);
9972 let mut others: Vec<(usize, f64, usize)> = cols
9973 .iter()
9974 .enumerate()
9975 .filter(|&(idx, &c)| idx != diag_pos && c != row)
9976 .map(|(idx, &c)| (c, vals[idx].abs(), idx))
9977 .collect();
9978 others.sort_unstable_by(
9979 |a, b| match b.1.partial_cmp(&a.1).unwrap_or(CmpOrdering::Equal) {
9980 CmpOrdering::Equal => a.0.cmp(&b.0),
9981 other => other,
9982 },
9983 );
9984 let mut keep = vec![diag_pos];
9985 for (_, _, idx) in others.into_iter().take(limit.saturating_sub(1)) {
9986 keep.push(idx);
9987 }
9988 keep.sort_unstable();
9989 let mut new_cols = Vec::with_capacity(keep.len());
9990 let mut new_vals = Vec::with_capacity(keep.len());
9991 for idx in keep {
9992 new_cols.push(cols[idx]);
9993 new_vals.push(vals[idx]);
9994 }
9995 cols = new_cols;
9996 vals = new_vals;
9997 }
9998 let mut pairs: Vec<(usize, f64)> = cols.into_iter().zip(vals).collect();
9999 pairs.sort_unstable_by(|a, b| a.0.cmp(&b.0));
10000 let mut out_cols = Vec::with_capacity(pairs.len());
10001 let mut out_vals = Vec::with_capacity(pairs.len());
10002 for (c, v) in pairs {
10003 out_cols.push(c);
10004 out_vals.push(v);
10005 }
10006 (out_cols, out_vals)
10007}
10008
10009fn fsai_factor_values(
10010 a: &CsrMatrix<f64>,
10011 patterns: &mut [Vec<usize>],
10012 lambda: f64,
10013 drop_tol: f64,
10014 adaptive_passes: usize,
10015 cap: usize,
10016) -> Result<CsrMatrix<f64>, KError> {
10017 let n = a.nrows();
10018 let mut row_ptr = Vec::with_capacity(n + 1);
10019 let mut col_idx: Vec<usize> = Vec::new();
10020 let mut values: Vec<f64> = Vec::new();
10021 let mut base = Vec::new();
10022 let mut rhs = Vec::new();
10023 row_ptr.push(0);
10024 for row in 0..n {
10025 let pat = &mut patterns[row];
10026 if pat.is_empty() {
10027 pat.push(row);
10028 }
10029 if !pat.contains(&row) {
10030 pat.push(row);
10031 pat.sort_unstable();
10032 }
10033 let mut pass = 0usize;
10034 loop {
10035 let m = pat.len();
10036 if m == 0 {
10037 break;
10038 }
10039 gather_dense_submatrix(a, pat, &mut base);
10040 rhs.resize(m, 0.0);
10041 if let Ok(pos) = pat.binary_search(&row) {
10042 rhs[pos] = 1.0;
10043 } else {
10044 rhs[0] = 1.0;
10045 }
10046 let solved = match solve_fsai_system(&base, &rhs, m, lambda) {
10047 Some(sol) => sol,
10048 None => {
10049 pat.clear();
10050 pat.push(row);
10051 let diag = csr_lookup(a, row, row);
10052 let val = if diag.abs() > 0.0 { 1.0 / diag } else { 1.0 };
10053 col_idx.push(row);
10054 values.push(val);
10055 row_ptr.push(col_idx.len());
10056 break;
10057 }
10058 };
10059 if adaptive_passes > 0 && pass < adaptive_passes {
10060 let added = fsai_enrich_pattern(a, pat, &solved, cap);
10061 if added > 0 {
10062 pass += 1;
10063 continue;
10064 }
10065 }
10066 let (cols, vals) = fsai_drop_entries(a, row, pat, &solved, drop_tol, cap);
10067 pat.clear();
10068 pat.extend(cols.iter().copied());
10069 col_idx.extend_from_slice(&cols);
10070 values.extend_from_slice(&vals);
10071 row_ptr.push(col_idx.len());
10072 break;
10073 }
10074 }
10075 Ok(CsrMatrix::from_csr(n, n, row_ptr, col_idx, values))
10076}
10077
10078fn fsai_transpose_with_pos(g: &CsrMatrix<f64>) -> (CsrMatrix<f64>, Vec<usize>) {
10079 let m = g.nrows();
10080 let n = g.ncols();
10081 let nnz = g.nnz();
10082 let mut row_counts = vec![0usize; n + 1];
10083 for &col in g.col_idx() {
10084 row_counts[col + 1] += 1;
10085 }
10086 for i in 0..n {
10087 row_counts[i + 1] += row_counts[i];
10088 }
10089 let mut col_idx = vec![0usize; nnz];
10090 let mut vals = vec![0.0f64; nnz];
10091 let mut next = row_counts.clone();
10092 let mut map = vec![0usize; nnz];
10093 for row in 0..m {
10094 let (rs, re) = (g.row_ptr()[row], g.row_ptr()[row + 1]);
10095 for idx in rs..re {
10096 let col = g.col_idx()[idx];
10097 let dest = next[col];
10098 col_idx[dest] = row;
10099 vals[dest] = g.values()[idx];
10100 map[idx] = dest;
10101 next[col] += 1;
10102 }
10103 }
10104 (CsrMatrix::from_csr(n, m, row_counts, col_idx, vals), map)
10105}
10106
10107fn fsai_build_for_level(
10108 cfg: &AMGConfig,
10109 a: &CsrMatrix<f64>,
10110 strength: Option<&Strength>,
10111) -> Result<FsaiData, KError> {
10112 let strength_owned = if cfg.fsai_use_strength {
10113 if let Some(s) = strength {
10114 Some(s.symmetrize())
10115 } else {
10116 Some(Strength::from_csr(a, cfg.strong_threshold, cfg.normalize_strength).symmetrize())
10117 }
10118 } else {
10119 None
10120 };
10121 let mut patterns = fsai_build_pattern(
10122 a,
10123 strength_owned.as_ref(),
10124 cfg.fsai_dist.max(1),
10125 cfg.fsai_max_per_row.max(1),
10126 );
10127 let g = fsai_factor_values(
10128 a,
10129 &mut patterns,
10130 cfg.fsai_lambda,
10131 cfg.fsai_drop_tol,
10132 cfg.fsai_adaptive_passes,
10133 cfg.fsai_max_per_row.max(1),
10134 )?;
10135 let (gt, map) = fsai_transpose_with_pos(&g);
10136 Ok(FsaiData {
10137 g,
10138 gt,
10139 g2gt_pos: map,
10140 })
10141}
10142
10143fn fsai_refresh_numeric(
10144 a: &CsrMatrix<f64>,
10145 data: &mut FsaiData,
10146 lambda: f64,
10147) -> Result<(), KError> {
10148 let n = a.nrows();
10149 let rp = data.g.row_ptr().to_vec();
10150 let ci = data.g.col_idx().to_vec();
10151 let mut base = Vec::new();
10152 let mut rhs = Vec::new();
10153 {
10154 let vals = data.g.values_mut();
10155 for row in 0..n {
10156 let start = rp[row];
10157 let end = rp[row + 1];
10158 if start == end {
10159 continue;
10160 }
10161 let pattern = &ci[start..end];
10162 gather_dense_submatrix(a, pattern, &mut base);
10163 rhs.resize(pattern.len(), 0.0);
10164 if let Ok(pos) = pattern.binary_search(&row) {
10165 rhs[pos] = 1.0;
10166 } else if !pattern.is_empty() {
10167 rhs[0] = 1.0;
10168 }
10169 if let Some(sol) = solve_fsai_system(&base, &rhs, pattern.len(), lambda) {
10170 for (dst, val) in vals[start..end].iter_mut().zip(sol.iter()) {
10171 *dst = *val;
10172 }
10173 } else {
10174 for (dst, &col) in vals[start..end].iter_mut().zip(pattern.iter()) {
10175 if col == row {
10176 let diag = csr_lookup(a, row, row);
10177 *dst = if diag.abs() > 0.0 { 1.0 / diag } else { 1.0 };
10178 } else {
10179 *dst = 0.0;
10180 }
10181 }
10182 }
10183 }
10184 }
10185 let g_vals = data.g.values();
10186 let gt_vals = data.gt.values_mut();
10187 for (src, &dst) in data.g2gt_pos.iter().enumerate() {
10188 gt_vals[dst] = g_vals[src];
10189 }
10190 Ok(())
10191}
10192
10193fn csr_mul(a: &CsrMatrix<f64>, b: &CsrMatrix<f64>) -> Result<CsrMatrix<f64>, KError> {
10195 if a.ncols() != b.nrows() {
10196 return Err(KError::InvalidInput("csr_mul: dimension mismatch".into()));
10197 }
10198 let m = a.nrows();
10199 let n = b.ncols();
10200
10201 let mut row_ptr = Vec::with_capacity(m + 1);
10202 let mut col_idx: Vec<usize> = Vec::new();
10203 let mut vals: Vec<f64> = Vec::new();
10204 row_ptr.push(0);
10205
10206 let mut tmp_cols: Vec<usize> = Vec::new();
10207 let mut tmp_vals: Vec<f64> = Vec::new();
10208 let mut order: Vec<usize> = Vec::new();
10209
10210 for i in 0..m {
10211 tmp_cols.clear();
10212 tmp_vals.clear();
10213
10214 let ars = a.row_ptr()[i];
10215 let are = a.row_ptr()[i + 1];
10216 for ap in ars..are {
10217 let k = a.col_idx()[ap];
10218 let aik = a.values()[ap];
10219 let brs = b.row_ptr()[k];
10220 let bre = b.row_ptr()[k + 1];
10221 for bp in brs..bre {
10222 let j = b.col_idx()[bp];
10223 tmp_cols.push(j);
10224 tmp_vals.push(aik * b.values()[bp]);
10225 }
10226 }
10227
10228 if tmp_cols.is_empty() {
10229 row_ptr.push(col_idx.len());
10230 continue;
10231 }
10232
10233 order.clear();
10234 order.extend(0..tmp_cols.len());
10235 order.sort_unstable_by(|&u, &v| match tmp_cols[u].cmp(&tmp_cols[v]) {
10236 std::cmp::Ordering::Equal => u.cmp(&v),
10237 o => o,
10238 });
10239
10240 let mut run_col = tmp_cols[order[0]];
10241 let mut acc = 0.0f64;
10242 for &idx in &order {
10243 let c = tmp_cols[idx];
10244 if c == run_col {
10245 acc += tmp_vals[idx];
10246 } else {
10247 if acc != 0.0 {
10248 col_idx.push(run_col);
10249 vals.push(acc);
10250 }
10251 run_col = c;
10252 acc = tmp_vals[idx];
10253 }
10254 }
10255 if acc != 0.0 {
10256 col_idx.push(run_col);
10257 vals.push(acc);
10258 }
10259
10260 row_ptr.push(col_idx.len());
10261 }
10262
10263 Ok(CsrMatrix::from_csr(m, n, row_ptr, col_idx, vals))
10264}
10265
10266fn rap(
10268 r: &CsrMatrix<f64>,
10269 a: &CsrMatrix<f64>,
10270 p: &CsrMatrix<f64>,
10271) -> Result<CsrMatrix<f64>, KError> {
10272 let ap = csr_mul(a, p)?;
10273 csr_mul(r, &ap)
10274}
10275
10276fn compute_anisotropy<M>(a: &M) -> Vec<f64>
10279where
10280 M: DenseMatRef<f64> + Sync,
10281{
10282 let n = a.nrows();
10283 #[cfg(feature = "rayon")]
10284 return (0..n)
10285 .into_par_iter()
10286 .map(|i| {
10287 let diag = a.get(i, i);
10288 let mut max_off: f64 = 0.0;
10289 for j in 0..n {
10290 if i != j {
10291 max_off = max_off.max(a.get(i, j).abs());
10292 }
10293 }
10294 if diag.abs() > 1e-14 {
10295 max_off / diag.abs()
10296 } else {
10297 0.0
10298 }
10299 })
10300 .collect();
10301 #[cfg(not(feature = "rayon"))]
10302 {
10303 let mut out = vec![0.0; n];
10304 for i in 0..n {
10305 let diag = a.get(i, i);
10306 let mut max_off: f64 = 0.0;
10307 for j in 0..n {
10308 if i != j {
10309 max_off = max_off.max(a.get(i, j).abs());
10310 }
10311 }
10312 out[i] = if diag.abs() > 1e-14 {
10313 max_off / diag.abs()
10314 } else {
10315 0.0
10316 };
10317 }
10318 out
10319 }
10320}
10321
10322fn compute_adaptive_threshold<M>(a: &M, base_threshold: f64) -> f64
10323where
10324 M: DenseMatRef<f64> + Sync,
10325{
10326 let anis = compute_anisotropy(a);
10327 let avg = if anis.is_empty() {
10328 1.0
10329 } else {
10330 anis.iter().sum::<f64>() / anis.len() as f64
10331 };
10332 base_threshold * (1.0 + avg.max(0.5))
10333}
10334
10335fn compute_strength_matrix<M>(a: &M, thr: f64) -> Mat<f64>
10337where
10338 M: DenseMatRef<f64>,
10339{
10340 let n = a.nrows();
10341 let mut s = Mat::<f64>::zeros(n, n);
10342 let mut diag = vec![0.0; n];
10343 for i in 0..n {
10344 diag[i] = a.get(i, i).abs();
10345 }
10346 for i in 0..n {
10347 for j in 0..n {
10348 if i == j {
10349 continue;
10350 }
10351 let denom = (diag[i] * diag[j]).sqrt();
10352 if denom > 1e-14 {
10353 let st = a.get(i, j).abs() / denom;
10354 if st > thr {
10355 s[(i, j)] = st;
10356 }
10357 }
10358 }
10359 }
10360 s
10361}
10362
10363fn pairwise_aggregation(s: &Mat<f64>) -> Vec<usize> {
10364 let n = s.nrows();
10365 let mut agg = vec![usize::MAX; n];
10366 let mut vis = vec![false; n];
10367 let mut id = 0usize;
10368 for i in 0..n {
10369 if vis[i] {
10370 continue;
10371 }
10372 let mut best = None;
10373 let mut bestv = 0.0;
10374 for j in 0..n {
10375 if i == j || vis[j] {
10376 continue;
10377 }
10378 let v = s[(i, j)];
10379 if v > bestv {
10380 bestv = v;
10381 best = Some(j);
10382 }
10383 }
10384 if let Some(j) = best {
10385 agg[i] = id;
10386 agg[j] = id;
10387 vis[i] = true;
10388 vis[j] = true;
10389 id += 1;
10390 } else {
10391 agg[i] = id;
10392 vis[i] = true;
10393 id += 1;
10394 }
10395 }
10396 agg
10397}
10398
10399fn build_coarse_graph(s: &Mat<f64>, agg: &[usize]) -> Mat<f64> {
10400 let max_id = *agg.iter().max().unwrap_or(&0);
10401 let cn = max_id + 1;
10402 let mut cg = Mat::<f64>::zeros(cn, cn);
10403 let n = s.nrows();
10404 for i in 0..n {
10405 for j in 0..n {
10406 let ai = agg[i];
10407 let aj = agg[j];
10408 let v = s[(i, j)];
10409 if v != 0.0 {
10410 cg[(ai, aj)] += v;
10411 }
10412 }
10413 }
10414 cg
10415}
10416
10417fn remap_aggregates(first: &[usize], second: &[usize]) -> Vec<usize> {
10418 #[cfg(feature = "rayon")]
10419 return first.par_iter().map(|&c| second[c]).collect();
10420 #[cfg(not(feature = "rayon"))]
10421 first.iter().map(|&c| second[c]).collect()
10422}
10423
10424fn double_pairwise_aggregation(s: &Mat<f64>) -> Vec<usize> {
10425 let pass1 = pairwise_aggregation(s);
10426 let coarse = build_coarse_graph(s, &pass1);
10427 let pass2 = pairwise_aggregation(&coarse);
10428 remap_aggregates(&pass1, &pass2)
10429}
10430
10431fn greedy_aggregation(s: &Mat<f64>) -> Vec<usize> {
10433 let n = s.nrows();
10434 let mut agg = vec![usize::MAX; n];
10435 let mut next = 0usize;
10436 let max_sz = 4usize;
10437
10438 let mut order: Vec<(f64, usize)> = (0..n)
10440 .map(|i| ((0..n).map(|j| s[(i, j)]).sum::<f64>(), i))
10441 .collect();
10442 order.sort_by(|a, b| match b.0.total_cmp(&a.0) {
10443 std::cmp::Ordering::Equal => a.1.cmp(&b.1),
10444 o => o,
10445 });
10446
10447 for &(_, seed) in &order {
10448 if agg[seed] != usize::MAX {
10449 continue;
10450 }
10451 agg[seed] = next;
10452 let mut neigh: Vec<(f64, usize)> = (0..n)
10454 .filter(|&j| j != seed && agg[j] == usize::MAX)
10455 .map(|j| (s[(seed, j)], j))
10456 .collect();
10457 neigh.sort_by(|a, b| match b.0.total_cmp(&a.0) {
10458 std::cmp::Ordering::Equal => a.1.cmp(&b.1),
10459 o => o,
10460 });
10461 for &(_, j) in neigh.iter() {
10462 if (0..n).filter(|&i| agg[i] == next).count() >= max_sz {
10463 break;
10464 }
10465 if s[(seed, j)] > 0.1 && agg[j] == usize::MAX {
10466 agg[j] = next;
10467 }
10468 }
10469 next += 1;
10470 }
10471 agg
10472}
10473
10474fn construct_prolongation(_a: &Mat<f64>, aggregates: &[usize]) -> Mat<f64> {
10476 let n = aggregates.len();
10477 let max_id = *aggregates.iter().max().unwrap_or(&0);
10478 let nc = max_id + 1;
10479 let mut p = Mat::<f64>::zeros(n, nc);
10480 for (i, &g) in aggregates.iter().enumerate() {
10481 p[(i, g)] = 1.0;
10482 }
10483 p
10484}
10485
10486fn smooth_interpolation(p: &mut Mat<f64>, a: &Mat<f64>, weight: f64) {
10488 let r = p.nrows().min(a.nrows());
10489 let c = p.ncols();
10490 for i in 0..r {
10491 for j in 0..c {
10492 p[(i, j)] -= weight * a[(i, j.min(a.ncols() - 1))];
10493 }
10494 }
10495}
10496
10497fn minimize_energy(p: &mut Mat<f64>, _a: &Mat<f64>) {
10499 let (m, n) = (p.nrows(), p.ncols());
10500 for i in 0..m {
10501 let mut norm2 = R::default();
10502 for j in 0..n {
10503 norm2 += p[(i, j)] * p[(i, j)];
10504 }
10505 let s = if norm2 > 1e-14 { norm2.sqrt() } else { 1.0 };
10506 for j in 0..n {
10507 p[(i, j)] /= s;
10508 }
10509 }
10510}
10511
10512fn cg_sparse(
10515 a: &CsrMatrix<f64>,
10516 b: &[f64],
10517 x: &mut [f64],
10518 tol: f64,
10519 maxit: usize,
10520) -> Result<(), KError> {
10521 let n = a.nrows();
10522 if n == 0 {
10523 return Ok(());
10524 }
10525 x.fill(R::default());
10526
10527 let mut r = b.to_vec();
10528 let mut p = r.clone();
10529 let mut ap = vec![R::default(); n];
10530
10531 let mut rsold = dot(&r, &r);
10532 let atol = tol.max(1e-12) * rsold.sqrt().max(1e-30);
10533
10534 for _ in 0..maxit {
10535 a.spmv_scaled(1.0, &p, 0.0, &mut ap)?;
10536 let denom = dot(&p, &ap);
10537 if denom.abs() < 1e-30 {
10538 break;
10539 }
10540 let alpha = rsold / denom;
10541
10542 #[cfg(feature = "rayon")]
10543 {
10544 x.par_iter_mut()
10545 .zip(p.par_iter())
10546 .for_each(|(xi, &pi)| *xi += alpha * pi);
10547 }
10548 #[cfg(not(feature = "rayon"))]
10549 for i in 0..n {
10550 x[i] += alpha * p[i];
10551 }
10552
10553 #[cfg(feature = "rayon")]
10554 {
10555 r.par_iter_mut()
10556 .zip(ap.par_iter())
10557 .for_each(|(ri, &api)| *ri -= alpha * api);
10558 }
10559 #[cfg(not(feature = "rayon"))]
10560 for i in 0..n {
10561 r[i] -= alpha * ap[i];
10562 }
10563
10564 let rsnew = dot(&r, &r);
10565 if rsnew.sqrt() < atol {
10566 break;
10567 }
10568 let beta = rsnew / rsold;
10569
10570 #[cfg(feature = "rayon")]
10571 {
10572 p.par_iter_mut()
10573 .zip(r.par_iter())
10574 .for_each(|(pi, &ri)| *pi = ri + beta * *pi);
10575 }
10576 #[cfg(not(feature = "rayon"))]
10577 for i in 0..n {
10578 p[i] = r[i] + beta * p[i];
10579 }
10580
10581 rsold = rsnew;
10582 }
10583 Ok(())
10584}
10585
10586fn pcg_left_precond<F>(
10587 a: &CsrMatrix<f64>,
10588 b: &[f64],
10589 x: &mut [f64],
10590 tol: f64,
10591 maxit: usize,
10592 mut apply_prec: F,
10593) -> Result<(), KError>
10594where
10595 F: FnMut(&[f64], &mut [f64]) -> Result<(), KError>,
10596{
10597 let n = a.nrows();
10598 if n == 0 {
10599 return Ok(());
10600 }
10601 x.fill(R::default());
10602
10603 let mut r = b.to_vec();
10604 let mut z = vec![R::default(); n];
10605 let mut p = vec![R::default(); n];
10606 let mut ap = vec![R::default(); n];
10607
10608 apply_prec(&r, &mut z)?;
10609 let mut rz_old = dot(&r, &z);
10610 let atol = tol.max(1e-12) * rz_old.abs().sqrt().max(1e-30);
10611 if rz_old.abs().sqrt() < atol {
10612 return Ok(());
10613 }
10614 p.copy_from_slice(&z);
10615
10616 for _ in 0..maxit {
10617 a.spmv_scaled(1.0, &p, 0.0, &mut ap)?;
10618 let denom = dot(&p, &ap);
10619 if denom.abs() < 1e-30 {
10620 break;
10621 }
10622 let alpha = rz_old / denom;
10623 for i in 0..n {
10624 x[i] += alpha * p[i];
10625 r[i] -= alpha * ap[i];
10626 }
10627 apply_prec(&r, &mut z)?;
10628 let rz_new = dot(&r, &z);
10629 if rz_new.abs().sqrt() < atol {
10630 break;
10631 }
10632 let beta = rz_new / rz_old;
10633 for i in 0..n {
10634 p[i] = z[i] + beta * p[i];
10635 }
10636 rz_old = rz_new;
10637 }
10638 Ok(())
10639}
10640
10641#[inline]
10642fn dot(x: &[R], y: &[R]) -> R {
10643 let mut s = R::default();
10644 for i in 0..x.len() {
10645 s += x[i] * y[i];
10646 }
10647 s
10648}
10649
10650#[derive(Clone, Debug)]
10653pub struct LevelStats {
10654 pub level: usize,
10655 pub n: usize,
10656 pub nnz_a: usize,
10657 pub nnz_p: usize,
10658 pub nnz_r: usize,
10659 pub max_row_sum_a: f64,
10660 pub eff_nnz_a: Option<usize>,
10661 pub pre_sweeps: usize,
10662 pub post_sweeps: usize,
10663 pub pre_work_estimate: f64,
10664 pub post_work_estimate: f64,
10665 pub selected_relax_pre: String,
10666 pub selected_relax_post: String,
10667 pub coarse_solver: Option<String>,
10668}
10669
10670#[derive(Clone, Debug, Default)]
10671pub struct AmgLevelStats {
10672 pub p_min_col_norm: f64,
10673 pub p_cond_sketched: f64,
10674 pub galerkin_worst_rel: f64,
10675}
10676
10677#[derive(Clone, Debug, Default)]
10678pub struct LevelSetupTiming {
10679 pub strength: Duration,
10680 pub aggregate: Duration,
10681 pub prolong: Duration,
10682 pub restrict: Duration,
10683 pub rap_symbolic: Duration,
10684 pub rap_numeric: Duration,
10685 pub diag: Duration,
10686 pub total: Duration,
10687}
10688
10689#[derive(Clone, Debug, Default)]
10690pub struct SetupTimings {
10691 pub per_level: Vec<LevelSetupTiming>,
10692 pub total_setup: Duration,
10693 pub total_symbolic: Duration,
10694 pub total_numeric: Duration,
10695}
10696
10697#[derive(Clone, Debug, Default)]
10698pub struct CycleLevelTiming {
10699 pub level: usize,
10700 pub pre_smooth: Duration,
10701 pub matvec: Duration,
10702 pub residual_axpy: Duration,
10703 pub restrict: Duration,
10704 pub coarse_solve: Duration,
10705 pub prolong: Duration,
10706 pub post_smooth: Duration,
10707}
10708
10709#[derive(Clone, Debug)]
10710pub struct CycleTimings {
10711 pub per_level: Vec<CycleLevelTiming>,
10712 pub total_cycle: Duration,
10713 pub cycle_type: CycleType,
10714 pub kcycle: Option<KCycle>,
10715}
10716
10717impl Default for CycleTimings {
10718 fn default() -> Self {
10719 Self {
10720 per_level: Vec::new(),
10721 total_cycle: Duration::default(),
10722 cycle_type: CycleType::V,
10723 kcycle: None,
10724 }
10725 }
10726}
10727
10728#[derive(Clone, Debug)]
10729pub struct DistApplyStats {
10730 pub mode: DistCoarseStrategy,
10731 pub coarse_solver_route: DistCoarseSolverRoute,
10732 pub coarse_repartition: DistCoarseRepartition,
10733 pub local_apply: Duration,
10734 pub gather: Duration,
10735 pub scatter: Duration,
10736 pub halo_exchange: Duration,
10737 pub setup_total: Duration,
10738 pub per_level_apply: Vec<Duration>,
10739 pub comm_bytes: usize,
10740 pub per_level_comm_bytes: Vec<usize>,
10741 pub reductions: usize,
10742 pub setup_gathered_fine_matrix: bool,
10743 pub true_distributed_hierarchy: bool,
10744}
10745
10746impl Default for DistApplyStats {
10747 fn default() -> Self {
10748 Self {
10749 mode: DistCoarseStrategy::RootGather,
10750 coarse_solver_route: DistCoarseSolverRoute::Auto,
10751 coarse_repartition: DistCoarseRepartition::Keep,
10752 local_apply: Duration::default(),
10753 gather: Duration::default(),
10754 scatter: Duration::default(),
10755 halo_exchange: Duration::default(),
10756 setup_total: Duration::default(),
10757 per_level_apply: Vec::new(),
10758 comm_bytes: 0,
10759 per_level_comm_bytes: Vec::new(),
10760 reductions: 0,
10761 setup_gathered_fine_matrix: false,
10762 true_distributed_hierarchy: false,
10763 }
10764 }
10765}
10766
10767impl DistApplyStats {
10768 pub fn mode_label(&self) -> &'static str {
10773 dist_strategy_label(self.mode)
10774 }
10775
10776 pub fn coarse_solver_route_label(&self) -> &'static str {
10778 dist_route_label(self.coarse_solver_route, self.mode)
10779 }
10780
10781 pub fn uses_root_gather(&self) -> bool {
10783 matches!(self.mode, DistCoarseStrategy::RootGather)
10784 }
10785
10786 pub fn reports_distributed_support(&self) -> bool {
10788 self.true_distributed_hierarchy
10789 || amg_dist_route_reports_distributed(self.coarse_solver_route, self.mode)
10790 }
10791
10792 pub fn uses_distributed_fine_spmv(&self) -> bool {
10800 matches!(self.mode, DistCoarseStrategy::DistributedCsr)
10801 && !self.setup_uses_fine_matrix_gather()
10802 && !self.apply_uses_root_vector_gather()
10803 }
10804
10805 pub fn setup_uses_fine_matrix_gather(&self) -> bool {
10807 self.setup_gathered_fine_matrix
10808 }
10809
10810 pub fn apply_uses_root_vector_gather(&self) -> bool {
10812 self.gather > Duration::default() || self.scatter > Duration::default()
10813 }
10814}
10815
10816#[derive(Clone, Debug)]
10817pub struct AmgStats {
10818 pub grid_complexity: f64,
10819 pub operator_complexity: f64,
10820 pub total_nnz: usize,
10821 pub total_smoothing_work: f64,
10822 pub num_levels: usize,
10823 pub levels: Vec<LevelStats>,
10824 pub diagnostics: Vec<AmgLevelStats>,
10825 pub setup: SetupTimings,
10826 pub last_cycle: Option<CycleTimings>,
10827 pub selected_dist_coarse_route: Option<String>,
10828 pub dist_route_fallback: Vec<String>,
10829 #[cfg(feature = "complex")]
10830 pub complex_setup_mode: AmgComplexSetupMode,
10831 #[cfg(feature = "complex")]
10832 pub complex_setup_fallback_reason: Option<String>,
10833}
10834
10835impl AmgStats {
10836 fn from_dist_apply(stats: &DistApplyStats) -> Self {
10837 Self {
10838 grid_complexity: 0.0,
10839 operator_complexity: 0.0,
10840 total_nnz: 0,
10841 total_smoothing_work: 0.0,
10842 num_levels: 0,
10843 levels: Vec::new(),
10844 diagnostics: Vec::new(),
10845 setup: SetupTimings::default(),
10846 last_cycle: None,
10847 selected_dist_coarse_route: Some(stats.coarse_solver_route_label().to_string()),
10848 dist_route_fallback: dist_route_fallback_labels(stats.coarse_solver_route, stats.mode),
10849 #[cfg(feature = "complex")]
10850 complex_setup_mode: AmgComplexSetupMode::Unset,
10851 #[cfg(feature = "complex")]
10852 complex_setup_fallback_reason: None,
10853 }
10854 }
10855
10856 fn from_hierarchy(h: &AmgHierarchy) -> Self {
10857 let n0 = h.levels.first().map(|l| l.a.nrows() as f64).unwrap_or(1.0);
10858 let nnz0 = h.levels.first().map(|l| l.a.nnz() as f64).unwrap_or(1.0);
10859 let mut ng_sum = 0.0;
10860 let mut nnz_sum = 0.0;
10861 for l in &h.levels {
10862 ng_sum += l.a.nrows() as f64;
10863 nnz_sum += l.a.nnz() as f64;
10864 }
10865 Self {
10866 grid_complexity: ng_sum / n0,
10867 operator_complexity: nnz_sum / nnz0,
10868 total_nnz: h.levels.iter().map(|l| l.a.nnz()).sum(),
10869 total_smoothing_work: 0.0,
10870 num_levels: h.levels.len(),
10871 levels: Vec::new(),
10872 diagnostics: Vec::new(),
10873 setup: SetupTimings::default(),
10874 last_cycle: None,
10875 selected_dist_coarse_route: None,
10876 dist_route_fallback: Vec::new(),
10877 #[cfg(feature = "complex")]
10878 complex_setup_mode: AmgComplexSetupMode::Unset,
10879 #[cfg(feature = "complex")]
10880 complex_setup_fallback_reason: None,
10881 }
10882 }
10883}
10884
10885#[derive(Default)]
10886struct AmgRuntime {
10887 last_cycle: Option<CycleTimings>,
10888 last_dist_apply: Option<DistApplyStats>,
10889}
10890
10891fn operator_complexity_estimate(levels: &[AMGLevel]) -> f64 {
10892 if levels.is_empty() {
10893 return 0.0;
10894 }
10895 let nnz0 = levels[0].a.nnz() as f64;
10896 let nnz_sum: f64 = levels.iter().map(|l| l.a.nnz() as f64).sum();
10897 nnz_sum / nnz0
10898}
10899
10900fn collect_level_stats(
10901 h: &AmgHierarchy,
10902 cfg: &AMGConfig,
10903 relax_overrides: Option<&BTreeMap<usize, RelaxType>>,
10904 sweep_overrides: Option<&BTreeMap<usize, (usize, usize)>>,
10905) -> Vec<LevelStats> {
10906 let mut out = Vec::with_capacity(h.levels.len());
10907 let relax_overrides = relax_overrides.cloned().unwrap_or_default();
10908 let sweep_overrides = sweep_overrides.cloned().unwrap_or_default();
10909 for (i, lvl) in h.levels.iter().enumerate() {
10910 let coarsest = i == h.coarsest_ix();
10911 let default_pre_phase = if i == 0 {
10912 RelaxPhase::Fine
10913 } else {
10914 RelaxPhase::Down
10915 };
10916 let default_post_phase = if i == 0 {
10917 RelaxPhase::Fine
10918 } else {
10919 RelaxPhase::Up
10920 };
10921 let (pre_sweeps, post_sweeps, relax_pre, relax_post) = if coarsest {
10922 let sweeps = sweep_overrides
10923 .get(&i)
10924 .map(|(pre, _)| *pre)
10925 .unwrap_or(h.policy.sweeps[RelaxPhase::Coarsest.ix()]);
10926 let relax = relax_overrides
10927 .get(&i)
10928 .copied()
10929 .unwrap_or(h.policy.kind[RelaxPhase::Coarsest.ix()]);
10930 (sweeps, sweeps, relax, relax)
10931 } else {
10932 let (pre, post) = sweep_overrides.get(&i).copied().unwrap_or((
10933 h.policy.sweeps[default_pre_phase.ix()],
10934 h.policy.sweeps[default_post_phase.ix()],
10935 ));
10936 let relax = relax_overrides
10937 .get(&i)
10938 .copied()
10939 .unwrap_or(h.policy.kind[default_pre_phase.ix()]);
10940 (pre, post, relax, relax)
10941 };
10942 out.push(LevelStats {
10943 level: i,
10944 n: lvl.a.nrows(),
10945 nnz_a: lvl.a.nnz(),
10946 nnz_p: if i < h.coarsest_ix() { lvl.p.nnz() } else { 0 },
10947 nnz_r: if i < h.coarsest_ix() {
10948 if cfg.keep_transpose {
10949 lvl.r.nnz()
10950 } else {
10951 lvl.r_row_ptr
10952 .as_ref()
10953 .and_then(|rp| rp.last().copied())
10954 .unwrap_or(0)
10955 }
10956 } else {
10957 0
10958 },
10959 max_row_sum_a: max_row_sum_abs(&lvl.a),
10960 eff_nnz_a: Some(eff_nnz(&lvl.a, cfg.stats_eps)),
10961 pre_sweeps,
10962 post_sweeps,
10963 pre_work_estimate: pre_sweeps as f64 * lvl.a.nnz() as f64,
10964 post_work_estimate: post_sweeps as f64 * lvl.a.nnz() as f64,
10965 selected_relax_pre: format!("{:?}", relax_pre),
10966 selected_relax_post: format!("{:?}", relax_post),
10967 coarse_solver: if i == h.coarsest_ix() {
10968 Some(format!("{:?}", cfg.coarse_solve))
10969 } else {
10970 None
10971 },
10972 });
10973 }
10974 out
10975}
10976
10977fn print_setup_tables(stats: &AmgStats) {
10978 if stats.levels.is_empty() {
10979 return;
10980 }
10981 println!(
10982 "AMG hierarchy: {} levels\nGrid complexity: {:.3}, Operator complexity: {:.3}",
10983 stats.num_levels, stats.grid_complexity, stats.operator_complexity
10984 );
10985 println!(
10986 "Total nnz: {}, smoothing work estimate: {:.1}, dist route: {}, fallback: {}",
10987 stats.total_nnz,
10988 stats.total_smoothing_work,
10989 stats.selected_dist_coarse_route.as_deref().unwrap_or("n/a"),
10990 if stats.dist_route_fallback.is_empty() {
10991 "n/a".to_string()
10992 } else {
10993 stats.dist_route_fallback.join(" -> ")
10994 }
10995 );
10996 println!(
10997 "{:>5} {:>10} {:>10} {:>10} {:>10} {:>10} {:>10} {:>12} {:>14} {:>14}",
10998 "lev",
10999 "n",
11000 "nnz(A)",
11001 "nnz(P)",
11002 "nnz(R)",
11003 "pre_sw",
11004 "post_sw",
11005 "max_row_sum",
11006 "relax(pre)",
11007 "relax(post)"
11008 );
11009 for ls in &stats.levels {
11010 println!(
11011 "{:>5} {:>10} {:>10} {:>10} {:>10} {:>10} {:>10} {:>12.4e} {:>14} {:>14}",
11012 ls.level,
11013 ls.n,
11014 ls.nnz_a,
11015 ls.nnz_p,
11016 ls.nnz_r,
11017 ls.pre_sweeps,
11018 ls.post_sweeps,
11019 ls.max_row_sum_a,
11020 ls.selected_relax_pre,
11021 ls.selected_relax_post,
11022 );
11023 }
11024 if !stats.setup.per_level.is_empty() {
11025 println!("Setup timings (ms): level | strength agg prolon restr symRAP numRAP diag total");
11026 let ms = |d: Duration| (d.as_secs_f64() * 1e3).round() as u64;
11027 for (i, lt) in stats.setup.per_level.iter().enumerate() {
11028 println!(
11029 "{:>5} {:>9} {:>3} {:>6} {:>5} {:>7} {:>8} {:>4} {:>6}",
11030 i,
11031 ms(lt.strength),
11032 ms(lt.aggregate),
11033 ms(lt.prolong),
11034 ms(lt.restrict),
11035 ms(lt.rap_symbolic),
11036 ms(lt.rap_numeric),
11037 ms(lt.diag),
11038 ms(lt.total)
11039 );
11040 }
11041 println!(
11042 "Total setup: {} ms (symbolic {} ms, numeric {} ms)",
11043 ms(stats.setup.total_setup),
11044 ms(stats.setup.total_symbolic),
11045 ms(stats.setup.total_numeric)
11046 );
11047 }
11048}
11049
11050fn print_cycle_table(c: &CycleTimings) {
11051 let mut desc = match c.cycle_type {
11052 CycleType::V => "V-cycle".to_string(),
11053 CycleType::W { gamma } => format!("W-cycle(gamma={gamma})"),
11054 };
11055 if c.kcycle.is_some() {
11056 desc.push_str(" + K");
11057 }
11058 println!("{desc} timings (ms): level | pre mv axpy R coarse P post");
11059 let ms = |d: Duration| (d.as_secs_f64() * 1e3).round() as u64;
11060 for lv in &c.per_level {
11061 println!(
11062 "{:>5} {:>4} {:>2} {:>4} {:>1} {:>6} {:>1} {:>4}",
11063 lv.level,
11064 ms(lv.pre_smooth),
11065 ms(lv.matvec),
11066 ms(lv.residual_axpy),
11067 ms(lv.restrict),
11068 ms(lv.coarse_solve),
11069 ms(lv.prolong),
11070 ms(lv.post_smooth)
11071 );
11072 }
11073 println!("Total cycle: {} ms", ms(c.total_cycle));
11074}
11075
11076#[inline]
11077fn tic() -> Instant {
11078 Instant::now()
11079}
11080#[inline]
11081fn toc(t0: Instant) -> Duration {
11082 t0.elapsed()
11083}
11084
11085fn max_row_sum_abs<T: KrystScalar<Real = f64>>(a: &CsrMatrix<T>) -> f64 {
11086 let n = a.nrows();
11087 let rp = a.row_ptr();
11088 let vv = a.values();
11089 #[cfg(feature = "rayon")]
11090 {
11091 (0..n)
11092 .into_par_iter()
11093 .map(|i| {
11094 let mut s = 0.0;
11095 for p in rp[i]..rp[i + 1] {
11096 s += vv[p].abs();
11097 }
11098 s
11099 })
11100 .reduce(|| 0.0, |x, y| x.max(y))
11101 }
11102 #[cfg(not(feature = "rayon"))]
11103 {
11104 let mut m = 0.0;
11105 for i in 0..n {
11106 let mut s = 0.0;
11107 for p in rp[i]..rp[i + 1] {
11108 s += vv[p].abs();
11109 }
11110 if s > m {
11111 m = s;
11112 }
11113 }
11114 m
11115 }
11116}
11117
11118fn eff_nnz<T: KrystScalar<Real = f64>>(a: &CsrMatrix<T>, eps: f64) -> usize {
11119 if eps <= 0.0 {
11120 return a.nnz();
11121 }
11122 a.values().iter().filter(|&&v| v.abs() >= eps).count()
11123}
11124
11125#[inline]
11126fn with_timing<F, R>(enabled: bool, acc: &mut Duration, f: F) -> R
11127where
11128 F: FnOnce() -> R,
11129{
11130 if enabled {
11131 let t = tic();
11132 let out = f();
11133 *acc += toc(t);
11134 out
11135 } else {
11136 f()
11137 }
11138}
11139
11140fn transpose_csr_with_pos(p: &Pcsr) -> (Vec<usize>, Vec<usize>, Vec<f64>, Vec<usize>) {
11141 let p_csr = CsrMatrix::from_csr(
11142 p.m,
11143 p.n,
11144 p.row_ptr.clone(),
11145 p.col_idx.clone(),
11146 p.vals.clone(),
11147 );
11148 let (pt, p2r_pos) = adjoint_csr_with_pos(&p_csr);
11149 (
11150 pt.row_ptr().to_vec(),
11151 pt.col_idx().to_vec(),
11152 pt.values().to_vec(),
11153 p2r_pos,
11154 )
11155}
11156
11157#[cfg(all(test, not(feature = "complex")))]
11158mod tests {
11159 use super::*;
11160 use faer::Mat;
11161 use std::any::Any;
11162 use std::cmp::Ordering;
11163 use std::sync::{Mutex, OnceLock};
11164
11165 fn relax_lock() -> &'static Mutex<()> {
11166 static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
11167 LOCK.get_or_init(|| Mutex::new(()))
11168 }
11169
11170 mod chebyshev;
11171 #[cfg(feature = "complex")]
11172 mod complex_cycles;
11173 mod cycle_policy;
11174 mod fsai_smoother;
11175 mod mixed_precision;
11176 mod nodal_nns;
11177 mod nodal_strength;
11178 mod rank_galerkin;
11179
11180 #[inline]
11181 fn feq(a: f64, b: f64, atol: f64, rtol: f64) -> bool {
11182 let diff = (a - b).abs();
11183 diff <= atol.max(rtol * a.abs()).max(rtol * b.abs())
11184 }
11185
11186 fn assert_dense_eq(a: &Mat<f64>, b: &Mat<f64>, atol: f64, rtol: f64) {
11187 assert_eq!(a.nrows(), b.nrows());
11188 assert_eq!(a.ncols(), b.ncols());
11189 for i in 0..a.nrows() {
11190 for j in 0..a.ncols() {
11191 assert!(
11192 feq(a[(i, j)], b[(i, j)], atol, rtol),
11193 "dense mismatch at ({},{}): {} vs {}",
11194 i,
11195 j,
11196 a[(i, j)],
11197 b[(i, j)]
11198 );
11199 }
11200 }
11201 }
11202
11203 fn ready_hierarchy(amg: &AMG) -> &AmgHierarchy {
11204 match &amg.state {
11205 AmgState::Ready { hierarchy, .. } => hierarchy,
11206 _ => panic!("AMG not ready"),
11207 }
11208 }
11209
11210 fn reset_symbolic_counter() {
11211 BUILD_SYMBOLIC_COUNT.with(|c| c.set(0));
11212 }
11213
11214 fn symbolic_counter() -> usize {
11215 BUILD_SYMBOLIC_COUNT.with(|c| c.get())
11216 }
11217
11218 struct TestLinOp {
11219 mat: CsrMatrix<f64>,
11220 sid: StructureId,
11221 vid: ValuesId,
11222 }
11223
11224 impl TestLinOp {
11225 fn new(mat: CsrMatrix<f64>, sid: StructureId, vid: ValuesId) -> Self {
11226 Self { mat, sid, vid }
11227 }
11228
11229 fn with_values(&self, vid: ValuesId) -> Self {
11230 Self {
11231 mat: self.mat.clone(),
11232 sid: self.sid,
11233 vid,
11234 }
11235 }
11236 }
11237
11238 impl LinOp for TestLinOp {
11239 type S = f64;
11240
11241 fn dims(&self) -> (usize, usize) {
11242 (self.mat.nrows(), self.mat.ncols())
11243 }
11244
11245 fn matvec(&self, x: &[f64], y: &mut [f64]) {
11246 crate::matrix::spmv::csr_matvec(&self.mat, x, y).unwrap();
11247 }
11248
11249 fn as_any(&self) -> &dyn Any {
11250 &self.mat
11251 }
11252
11253 fn structure_id(&self) -> StructureId {
11254 self.sid
11255 }
11256
11257 fn values_id(&self) -> ValuesId {
11258 self.vid
11259 }
11260 }
11261
11262 fn csr_from_triples(m: usize, n: usize, mut trip: Vec<(usize, usize, f64)>) -> CsrMatrix<f64> {
11263 trip.sort_by(|a, b| match a.0.cmp(&b.0) {
11264 Ordering::Equal => a.1.cmp(&b.1),
11265 o => o,
11266 });
11267 let mut row_ptr = vec![0usize; m + 1];
11268 let mut col_idx = Vec::<usize>::new();
11269 let mut vals = Vec::<f64>::new();
11270 let mut i_cur = 0usize;
11271 let mut j_prev = usize::MAX;
11272 let mut acc = 0.0;
11273
11274 let push_acc = |row: usize,
11275 col: usize,
11276 v: f64,
11277 row_ptr: &mut [usize],
11278 col_idx: &mut Vec<usize>,
11279 vals: &mut Vec<f64>| {
11280 if v != 0.0 {
11281 col_idx.push(col);
11282 vals.push(v);
11283 }
11284 row_ptr[row + 1] = col_idx.len();
11285 };
11286
11287 for (r, c, v) in trip {
11288 while i_cur < r {
11289 if j_prev != usize::MAX {
11290 push_acc(i_cur, j_prev, acc, &mut row_ptr, &mut col_idx, &mut vals);
11291 j_prev = usize::MAX;
11292 acc = 0.0;
11293 }
11294 i_cur += 1;
11295 row_ptr[i_cur] = col_idx.len();
11296 }
11297 if j_prev == c {
11298 acc += v;
11299 } else {
11300 if j_prev != usize::MAX {
11301 push_acc(i_cur, j_prev, acc, &mut row_ptr, &mut col_idx, &mut vals);
11302 }
11303 j_prev = c;
11304 acc = v;
11305 }
11306 }
11307 while i_cur < m {
11308 if j_prev != usize::MAX {
11309 push_acc(i_cur, j_prev, acc, &mut row_ptr, &mut col_idx, &mut vals);
11310 j_prev = usize::MAX;
11311 acc = 0.0;
11312 }
11313 i_cur += 1;
11314 row_ptr[i_cur] = col_idx.len();
11315 }
11316
11317 CsrMatrix::from_csr(m, n, row_ptr, col_idx, vals)
11318 }
11319
11320 fn l2_norm(v: &[f64]) -> f64 {
11321 v.iter().map(|x| x * x).sum::<f64>().sqrt()
11322 }
11323
11324 fn identity_level() -> AMGLevel {
11325 let level = AMGLevel {
11326 a: CsrMatrix::identity(1),
11327 p: CsrMatrix::identity(1),
11328 r: CsrMatrix::identity(1),
11329 diag_inv: vec![1.0],
11330 d_sqrt_inv: None,
11331 l1_inv: None,
11332 diag_inv_safe: None,
11333 d_sqrt_inv_safe: None,
11334 cheb: None,
11335 cheb_safe: None,
11336 agg_of: vec![0],
11337 is_c: Vec::new(),
11338 cf: None,
11339 p2r_pos: vec![],
11340 num_functions: 1,
11341 row_basis: None,
11342 layout: None,
11343 nns: None,
11344 a_next_pat: None,
11345 a_next_pat_ng: None,
11346 rap_full2ng_pos: None,
11347 r_row_ptr: None,
11348 r_col_idx: None,
11349 r_vals_scratch: None,
11350 coarse_solver: None,
11351 ilu0: None,
11352 ras: None,
11353 fsai: None,
11354 a_vals_f32: None,
11355 diag_inv_f32: None,
11356 d_sqrt_inv_f32: None,
11357 l1_inv_f32: None,
11358 fsai_g_vals_f32: None,
11359 fsai_gt_vals_f32: None,
11360 };
11361 #[cfg(feature = "simd")]
11362 {
11363 let tuning = utils::default_spmv_tuning();
11364 build_level_spmv_plans(&mut level, &tuning);
11365 }
11366 level
11367 }
11368
11369 fn level_from_matrix(a: &CsrMatrix<f64>) -> AMGLevel {
11370 let n = a.nrows();
11371 let level = AMGLevel {
11372 a: a.clone(),
11373 p: CsrMatrix::identity(n),
11374 r: CsrMatrix::identity(n),
11375 diag_inv: diag_inv_from_csr(a).unwrap(),
11376 d_sqrt_inv: None,
11377 l1_inv: None,
11378 diag_inv_safe: None,
11379 d_sqrt_inv_safe: None,
11380 cheb: None,
11381 cheb_safe: None,
11382 agg_of: vec![0; n.max(1)],
11383 is_c: Vec::new(),
11384 cf: None,
11385 p2r_pos: vec![],
11386 num_functions: 1,
11387 row_basis: None,
11388 layout: None,
11389 nns: None,
11390 a_next_pat: None,
11391 a_next_pat_ng: None,
11392 rap_full2ng_pos: None,
11393 r_row_ptr: None,
11394 r_col_idx: None,
11395 r_vals_scratch: None,
11396 coarse_solver: None,
11397 ilu0: None,
11398 ras: None,
11399 fsai: None,
11400 a_vals_f32: None,
11401 diag_inv_f32: None,
11402 d_sqrt_inv_f32: None,
11403 l1_inv_f32: None,
11404 fsai_g_vals_f32: None,
11405 fsai_gt_vals_f32: None,
11406 };
11407 #[cfg(feature = "simd")]
11408 {
11409 let tuning = utils::default_spmv_tuning();
11410 build_level_spmv_plans(&mut level, &tuning);
11411 }
11412 level
11413 }
11414
11415 #[test]
11416 fn phase_selection_logic() {
11417 let _guard = relax_lock()
11418 .lock()
11419 .unwrap_or_else(|poison| poison.into_inner());
11420 reset_relax_counts();
11421 let levels = vec![identity_level(), identity_level(), identity_level()];
11422 let policy = RelaxPolicy {
11423 kind: [RelaxType::Jacobi; 4],
11424 sweeps: [1, 1, 1, 0],
11425 omega: 1.0,
11426 };
11427 let hier = AmgHierarchy {
11428 levels,
11429 policy,
11430 coarse_solve: CoarseSolve::DirectDense,
11431 };
11432 let mut amg = AMG {
11433 state: AmgState::Ready {
11434 hierarchy: Box::new(hier),
11435 last_structure_id: StructureId(0),
11436 last_values_id: ValuesId(0),
11437 pattern_hash: 0,
11438 },
11439 ..Default::default()
11440 };
11441 amg.cfg.require_spd = false;
11442 let rhs = [1.0];
11443 let mut sol = [0.0];
11444 amg.apply(PcSide::Left, &rhs, &mut sol).unwrap();
11445 let counts = get_relax_counts();
11446 assert_eq!(counts[RelaxPhase::Fine.ix()], 2);
11447 assert_eq!(counts[RelaxPhase::Down.ix()], 1);
11448 assert_eq!(counts[RelaxPhase::Up.ix()], 1);
11449 assert_eq!(counts[RelaxPhase::Coarsest.ix()], 0);
11450 }
11451
11452 #[test]
11453 fn fcg_presmooth_reduces_residual() {
11454 let n = 10;
11455 let a = poisson1d(n);
11456 let lvl = level_from_matrix(&a);
11457 let rhs = vec![1.0; n];
11458 let mut sol = vec![0.0; n];
11459 let mut ws = AMGWorkspace::new(n);
11460 let mut amg = AMG::default();
11461 amg.cfg.require_spd = false;
11462
11463 amg.fcg_presmooth(
11464 &lvl,
11465 &rhs,
11466 &mut sol,
11467 5,
11468 0.0,
11469 1,
11470 RelaxType::Jacobi,
11471 amg.cfg.jacobi_omega,
11472 &mut ws,
11473 )
11474 .unwrap();
11475
11476 let mut work = vec![0.0; n];
11477 a.spmv_scaled(1.0, &sol, 0.0, &mut work).unwrap();
11478 let mut r_out = 0.0;
11479 for i in 0..n {
11480 let ri = rhs[i] - work[i];
11481 r_out += ri * ri;
11482 }
11483 let r0 = (rhs.iter().map(|v| v * v).sum::<f64>()).sqrt();
11484 assert!(r_out.sqrt() < r0);
11485 }
11486
11487 #[test]
11488 fn flexible_presmooth_reduces_relax_calls() {
11489 let make_amg = || AMG {
11490 state: AmgState::Ready {
11491 hierarchy: Box::new(AmgHierarchy {
11492 levels: vec![identity_level(), identity_level()],
11493 policy: RelaxPolicy {
11494 kind: [RelaxType::Jacobi; 4],
11495 sweeps: [1, 1, 1, 0],
11496 omega: 1.0,
11497 },
11498 coarse_solve: CoarseSolve::DirectDense,
11499 }),
11500 last_structure_id: StructureId(0),
11501 last_values_id: ValuesId(0),
11502 pattern_hash: 0,
11503 },
11504 ..Default::default()
11505 };
11506
11507 let mut amg_std = make_amg();
11508 amg_std.cfg.require_spd = false;
11509 reset_relax_counts();
11510 let rhs = [1.0];
11511 let mut sol = [0.0];
11512 amg_std.apply(PcSide::Left, &rhs, &mut sol).unwrap();
11513 let baseline = get_relax_counts()[RelaxPhase::Fine.ix()];
11514
11515 let mut amg_flex = make_amg();
11516 amg_flex.cfg.require_spd = false;
11517 amg_flex.cfg.flexible_level = Some(0);
11518 amg_flex.cfg.flexible_iters = 3;
11519 amg_flex.cfg.flexible_pc_sweeps = 1;
11520 reset_relax_counts();
11521 sol[0] = 0.0;
11522 amg_flex.apply(PcSide::Left, &rhs, &mut sol).unwrap();
11523 let flex_counts = get_relax_counts()[RelaxPhase::Fine.ix()];
11524 assert!(flex_counts < baseline);
11525 }
11526
11527 #[test]
11528 fn flexible_presmooth_spd_guard() {
11529 let a = poisson1d(4);
11530 let mut amg = AMGBuilder::new()
11531 .grid_relax_type_all(RelaxType::GaussSeidel)
11532 .flexible_level(0)
11533 .flexible_iters(2)
11534 .build(&Mat::<f64>::zeros(0, 0))
11535 .unwrap();
11536 let err = amg.setup(&a).unwrap_err();
11537 match err {
11538 KError::InvalidInput(msg) => {
11539 assert!(msg.contains("flexible presmoothing"));
11540 }
11541 other => panic!("unexpected error: {other:?}"),
11542 }
11543 }
11544
11545 #[test]
11546 fn validation_failures() {
11547 let mut cfg = AMGConfig::default();
11548 cfg.grid_relax_type = [RelaxType::HybridGaussSeidel; 4];
11549 let err = validate_relax_policy(&cfg, cfg.coarse_solve).unwrap_err();
11550 assert!(matches!(err, KError::InvalidInput(_)));
11551
11552 let mut cfg = AMGConfig::default();
11553 cfg.coarse_solve = CoarseSolve::DirectDense;
11554 cfg.num_grid_sweeps[RelaxPhase::Coarsest.ix()] = 1;
11555 let err = validate_relax_policy(&cfg, cfg.coarse_solve).unwrap_err();
11556 assert!(matches!(err, KError::InvalidInput(_)));
11557
11558 let mut cfg = AMGConfig::default();
11559 cfg.truncation_factor = 1.2;
11560 assert!(validate_truncation_and_caps(&cfg).is_err());
11561 cfg.truncation_factor = -0.1;
11562 assert!(validate_truncation_and_caps(&cfg).is_err());
11563 cfg.truncation_factor = 0.0;
11564 cfg.interpolation_truncation = -1.0;
11565 assert!(validate_truncation_and_caps(&cfg).is_err());
11566 cfg.interpolation_truncation = 0.0;
11567 cfg.rap_truncation_abs = -1.0;
11568 assert!(validate_truncation_and_caps(&cfg).is_err());
11569 }
11570
11571 #[test]
11572 fn legacy_shim_populates_arrays() {
11573 let amg = AMG::builder()
11574 .smoothing_sweeps(2, 3)
11575 .build(&Mat::<f64>::zeros(0, 0))
11576 .unwrap();
11577 assert_eq!(amg.cfg.num_grid_sweeps, [2, 2, 3, 1]);
11578 }
11579
11580 #[test]
11581 fn rap_numeric_matches_dense_small() {
11582 let a = csr_from_triples(
11583 3,
11584 3,
11585 vec![
11586 (0, 0, 4.0),
11587 (0, 1, -1.0),
11588 (1, 0, -1.0),
11589 (1, 1, 4.0),
11590 (1, 2, -1.0),
11591 (2, 1, -1.0),
11592 (2, 2, 4.0),
11593 ],
11594 );
11595 let p = csr_from_triples(3, 2, vec![(0, 0, 1.0), (1, 0, 1.0), (2, 1, 1.0)]);
11596 let r = csr_from_triples(2, 3, vec![(0, 0, 1.0), (0, 1, 1.0), (1, 2, 1.0)]);
11597
11598 let pat = rap_ops::rap_symbolic(&r, &a, &p);
11599 let mut vals = vec![0.0; pat.col_idx.len()];
11600 rap_ops::rap_numeric(&pat, &r, &a, &p, &mut vals);
11601
11602 let ad = a.to_dense().unwrap();
11603 let pd = p.to_dense().unwrap();
11604 let rd = r.to_dense().unwrap();
11605 let cd = &rd * &ad * &pd;
11606
11607 let mut cpat = Mat::<f64>::zeros(pat.nrows, pat.ncols);
11608 for i in 0..pat.nrows {
11609 for k in pat.row_ptr[i]..pat.row_ptr[i + 1] {
11610 let j = pat.col_idx[k];
11611 cpat[(i, j)] = vals[k];
11612 }
11613 }
11614 assert_dense_eq(&cpat, &cd, 1e-12, 1e-12);
11615 }
11616
11617 #[test]
11618 fn transpose_bijection_and_values_small() {
11619 let m = 3;
11620 let n = 4;
11621 let p = prolong::Pcsr {
11622 m,
11623 n,
11624 row_ptr: vec![0, 2, 3, 5],
11625 col_idx: vec![0, 2, 1, 1, 3],
11626 vals: vec![1.0, 2.0, 3.0, 4.0, 5.0],
11627 };
11628 let (rr, rc, rv, p2r) = super::transpose_csr_with_pos(&p);
11629 assert_eq!(rr.len(), n + 1);
11630 assert_eq!(rc.len(), p.col_idx.len());
11631 assert_eq!(rv.len(), p.vals.len());
11632
11633 let nnz = p.vals.len();
11634 let mut seen = vec![false; nnz];
11635 for &q in &p2r {
11636 assert!(q < nnz);
11637 assert!(!seen[q]);
11638 seen[q] = true;
11639 }
11640 assert!(seen.into_iter().all(|b| b));
11641
11642 for (pi, &ri) in p2r.iter().enumerate() {
11643 assert!(feq(p.vals[pi], rv[ri], 0.0, 0.0));
11644 }
11645
11646 let mut p_dense = Mat::<f64>::zeros(m, n);
11647 for i in 0..m {
11648 for k in p.row_ptr[i]..p.row_ptr[i + 1] {
11649 p_dense[(i, p.col_idx[k])] = p.vals[k];
11650 }
11651 }
11652 let r_dense = p_dense.transpose().to_owned();
11653 let mut r_pat = Mat::<f64>::zeros(n, m);
11654 for i in 0..n {
11655 for k in rr[i]..rr[i + 1] {
11656 r_pat[(i, rc[k])] = rv[k];
11657 }
11658 }
11659 assert_dense_eq(&r_pat, &r_dense, 0.0, 0.0);
11660 }
11661
11662 #[test]
11663 #[cfg(not(feature = "complex"))]
11664 fn filter_enforces_row_sums() {
11665 let a = poisson1d(64);
11666 let mut filtered = AMGBuilder::new()
11667 .rap_drop_abs(0.05)
11668 .require_spd(false)
11669 .filter_omega(1.0)
11670 .build(&Mat::<f64>::zeros(0, 0))
11671 .unwrap();
11672 let mut baseline = AMGBuilder::new()
11673 .rap_drop_abs(0.05)
11674 .require_spd(false)
11675 .filter_omega(0.0)
11676 .build(&Mat::<f64>::zeros(0, 0))
11677 .unwrap();
11678 filtered.setup(&a).unwrap();
11679 baseline.setup(&a).unwrap();
11680 let h_filtered = ready_hierarchy(&filtered);
11681 let h_baseline = ready_hierarchy(&baseline);
11682 assert_eq!(h_filtered.levels.len(), h_baseline.levels.len());
11683 let mut best_ratio = f64::INFINITY;
11684 let mut any_significant = false;
11685 for (lvl_ix, (lvl_filtered, lvl_baseline)) in
11686 h_filtered.levels.iter().zip(&h_baseline.levels).enumerate()
11687 {
11688 if lvl_ix == 0 {
11689 continue;
11690 }
11691 let n = lvl_filtered.a.nrows();
11692 assert_eq!(n, lvl_baseline.a.nrows());
11693 let ones = vec![1.0; n];
11694 let mut sums_filtered = vec![0.0; n];
11695 let mut sums_baseline = vec![0.0; n];
11696 lvl_filtered
11697 .a
11698 .spmv_scaled(1.0, &ones, 0.0, &mut sums_filtered)
11699 .unwrap();
11700 lvl_baseline
11701 .a
11702 .spmv_scaled(1.0, &ones, 0.0, &mut sums_baseline)
11703 .unwrap();
11704 let max_filtered = sums_filtered.iter().fold(0.0f64, |acc, v| acc.max(v.abs()));
11705 let max_baseline = sums_baseline.iter().fold(0.0f64, |acc, v| acc.max(v.abs()));
11706 if max_baseline > 1e-14 {
11707 let ratio = max_filtered / max_baseline;
11708 best_ratio = best_ratio.min(ratio);
11709 any_significant = true;
11710 } else {
11711 assert!(
11712 max_filtered <= 1e-14,
11713 "level {lvl_ix} filtered row sum {} exceeds tight tolerance",
11714 max_filtered
11715 );
11716 }
11717 }
11718 assert!(
11719 any_significant,
11720 "no levels with meaningful baseline row sums"
11721 );
11722 assert!(
11723 best_ratio <= 0.1 + 1e-12,
11724 "expected at least one level to reduce max row sum by an order of magnitude: best_ratio={}",
11725 best_ratio
11726 );
11727 }
11728 #[test]
11729 #[cfg(not(feature = "complex"))]
11730 fn filter_non_galerkin_preserves_row_sums() {
11731 let a = poisson1d(64);
11732 let mut cfg_filtered = AMGConfig::default();
11733 cfg_filtered.require_spd = false;
11734 cfg_filtered.rap_truncation_abs = 0.02;
11735 cfg_filtered.filter_omega = 1.0;
11736 cfg_filtered.filter_after_non_galerkin = true;
11737 cfg_filtered.non_galerkin.enabled = true;
11738 cfg_filtered.non_galerkin.drop_abs = 0.05;
11739 cfg_filtered.non_galerkin.drop_rel = 0.0;
11740 cfg_filtered.non_galerkin.cap_row = 0;
11741 let mut cfg_baseline = cfg_filtered.clone();
11742 cfg_baseline.filter_omega = 0.0;
11743 let mut amg_filtered = AMG::with_config(cfg_filtered);
11744 let mut amg_baseline = AMG::with_config(cfg_baseline);
11745 amg_filtered.setup(&a).unwrap();
11746 amg_baseline.setup(&a).unwrap();
11747 let h_filtered = ready_hierarchy(&amg_filtered);
11748 let h_baseline = ready_hierarchy(&amg_baseline);
11749 assert_eq!(h_filtered.levels.len(), h_baseline.levels.len());
11750 let mut best_ratio = f64::INFINITY;
11751 let mut any_significant = false;
11752 for (lvl_ix, (lvl_filtered, lvl_baseline)) in
11753 h_filtered.levels.iter().zip(&h_baseline.levels).enumerate()
11754 {
11755 if lvl_ix == 0 {
11756 continue;
11757 }
11758 let n = lvl_filtered.a.nrows();
11759 assert_eq!(n, lvl_baseline.a.nrows());
11760 let ones = vec![1.0; n];
11761 let mut sums_filtered = vec![0.0; n];
11762 let mut sums_baseline = vec![0.0; n];
11763 lvl_filtered
11764 .a
11765 .spmv_scaled(1.0, &ones, 0.0, &mut sums_filtered)
11766 .unwrap();
11767 lvl_baseline
11768 .a
11769 .spmv_scaled(1.0, &ones, 0.0, &mut sums_baseline)
11770 .unwrap();
11771 let max_filtered = sums_filtered.iter().fold(0.0f64, |acc, v| acc.max(v.abs()));
11772 let max_baseline = sums_baseline.iter().fold(0.0f64, |acc, v| acc.max(v.abs()));
11773 if max_baseline > 1e-14 {
11774 let ratio = max_filtered / max_baseline;
11775 best_ratio = best_ratio.min(ratio);
11776 any_significant = true;
11777 } else {
11778 assert!(
11779 max_filtered <= 1e-14,
11780 "level {lvl_ix} filtered row sum {} exceeds tight tolerance",
11781 max_filtered
11782 );
11783 }
11784 }
11785 assert!(
11786 any_significant,
11787 "no levels with meaningful baseline row sums"
11788 );
11789 assert!(
11790 best_ratio <= 0.1 + 1e-12,
11791 "expected at least one level to reduce max row sum by an order of magnitude: best_ratio={}",
11792 best_ratio
11793 );
11794 }
11795 #[test]
11796 #[cfg(not(feature = "complex"))]
11797 fn filter_reduces_constant_mode_residual() {
11798 let a = poisson1d(64);
11799 let mut cfg_off = AMGConfig::default();
11800 cfg_off.rap_truncation_abs = 0.02;
11801 cfg_off.require_spd = false;
11802 cfg_off.filter_omega = 0.0;
11803 let mut cfg_on = cfg_off.clone();
11804 cfg_on.filter_omega = 1.0;
11805 let mut amg_off = AMG::with_config(cfg_off);
11806 let mut amg_on = AMG::with_config(cfg_on);
11807 amg_off.setup(&a).unwrap();
11808 amg_on.setup(&a).unwrap();
11809 let h_on = ready_hierarchy(&amg_on);
11810 let h_off = ready_hierarchy(&amg_off);
11811 assert!(h_on.coarsest_ix() >= 1);
11812 let coarse_ix = 1;
11813 let n_fine = h_on.levels[0].a.nrows();
11814 let ones_fine = vec![1.0; n_fine];
11815 let mut rhs_on = vec![0.0; h_on.levels[coarse_ix].a.nrows()];
11816 h_on.levels[0]
11817 .r
11818 .spmv_scaled(1.0, &ones_fine, 0.0, &mut rhs_on)
11819 .unwrap();
11820 let mut prod_on = vec![0.0; rhs_on.len()];
11821 h_on.levels[coarse_ix]
11822 .a
11823 .spmv_scaled(1.0, &rhs_on, 0.0, &mut prod_on)
11824 .unwrap();
11825 let norm_on = l2_norm(&prod_on);
11826
11827 let mut rhs_off = vec![0.0; h_off.levels[coarse_ix].a.nrows()];
11828 h_off.levels[0]
11829 .r
11830 .spmv_scaled(1.0, &ones_fine, 0.0, &mut rhs_off)
11831 .unwrap();
11832 let mut prod_off = vec![0.0; rhs_off.len()];
11833 h_off.levels[coarse_ix]
11834 .a
11835 .spmv_scaled(1.0, &rhs_off, 0.0, &mut prod_off)
11836 .unwrap();
11837 let norm_off = l2_norm(&prod_off);
11838 assert!(norm_off > 1e-12);
11839 assert!(norm_on <= norm_off * 0.1 + 1e-10);
11840 }
11841
11842 fn poisson1d(n: usize) -> CsrMatrix<f64> {
11843 let mut row_ptr = Vec::with_capacity(n + 1);
11844 let mut col_idx = Vec::new();
11845 let mut vals = Vec::new();
11846 row_ptr.push(0);
11847 for i in 0..n {
11848 if i > 0 {
11849 col_idx.push(i - 1);
11850 vals.push(-1.0);
11851 }
11852 col_idx.push(i);
11853 vals.push(2.0);
11854 if i + 1 < n {
11855 col_idx.push(i + 1);
11856 vals.push(-1.0);
11857 }
11858 row_ptr.push(col_idx.len());
11859 }
11860 CsrMatrix::from_csr(n, n, row_ptr, col_idx, vals)
11861 }
11862
11863 #[test]
11864 fn gs_symgs_residual() {
11865 let a = poisson1d(3);
11866 let d = diag_inv_from_csr(&a).unwrap();
11867 let rhs = vec![1.0; 3];
11868 let mut zf = vec![0.0; 3];
11869 AMG::gs_forward(1.0, &a, &d, &rhs, &mut zf, 1).unwrap();
11870 let mut work = vec![0.0; 3];
11871 a.spmv_scaled(1.0, &zf, 0.0, &mut work).unwrap();
11872 let mut res_f = vec![0.0; 3];
11873 for i in 0..3 {
11874 res_f[i] = rhs[i] - work[i];
11875 }
11876 let norm_f = dot(&res_f, &res_f);
11877
11878 let mut zs = vec![0.0; 3];
11879 AMG::sym_gs(1.0, &a, &d, &rhs, &mut zs, 1).unwrap();
11880 a.spmv_scaled(1.0, &zs, 0.0, &mut work).unwrap();
11881 let mut res_s = vec![0.0; 3];
11882 for i in 0..3 {
11883 res_s[i] = rhs[i] - work[i];
11884 }
11885 let norm_s = dot(&res_s, &res_s);
11886 assert!(norm_s < norm_f);
11887 }
11888
11889 #[test]
11890 fn l1_jacobi_no_worse_than_jacobi() {
11891 let a = poisson1d(3);
11892 let d = diag_inv_from_csr(&a).unwrap();
11893 let l1 = l1_diag_inv(&a);
11894 let rhs = vec![1.0; 3];
11895 let mut z_j = vec![0.0; 3];
11896 let mut z_l1 = vec![0.0; 3];
11897 let mut ws_j = AMGWorkspace::new(3);
11898 let mut ws_l1 = AMGWorkspace::new(3);
11899 AMG::jacobi_smooth_sparse(1.0, &a, &d, &rhs, &mut z_j, 1, &mut ws_j).unwrap();
11900 AMG::l1_jacobi(1.0, &a, &l1, &rhs, &mut z_l1, 1, &mut ws_l1).unwrap();
11901 a.spmv_scaled(1.0, &z_j, 0.0, &mut ws_j.work[..3]).unwrap();
11902 let mut rj = 0.0;
11903 for i in 0..3 {
11904 let ri = rhs[i] - ws_j.work[i];
11905 rj += ri * ri;
11906 }
11907 a.spmv_scaled(1.0, &z_l1, 0.0, &mut ws_l1.work[..3])
11908 .unwrap();
11909 let mut rl1 = 0.0;
11910 for i in 0..3 {
11911 let ri = rhs[i] - ws_l1.work[i];
11912 rl1 += ri * ri;
11913 }
11914 let r0 = 3.0; assert!(rj < r0);
11916 assert!(rl1 < r0);
11917 }
11918
11919 #[test]
11920 #[cfg(not(feature = "complex"))]
11921 fn refresh_updates_caches() {
11922 let a = poisson1d(4);
11923 let mut amg_l1 = AMGBuilder::new()
11924 .grid_relax_type_all(RelaxType::L1Jacobi)
11925 .build(&Mat::<f64>::zeros(0, 0))
11926 .unwrap();
11927 amg_l1.setup(&a).unwrap();
11928 let old = ready_hierarchy(&amg_l1).levels[0].l1_inv.as_ref().unwrap()[0];
11929 let mut a2 = a.clone();
11930 let rp = a2.row_ptr();
11931 for p in rp[0]..rp[1] {
11932 a2.values_mut()[p] *= 2.0;
11933 }
11934 amg_l1.update_numeric(&a2).unwrap();
11935 let new = ready_hierarchy(&amg_l1).levels[0].l1_inv.as_ref().unwrap()[0];
11936 assert!((new - old).abs() > 1e-12);
11937
11938 let mut amg_ch = AMGBuilder::new()
11939 .grid_relax_type_all(RelaxType::Chebyshev)
11940 .chebyshev_recompute_esteig(true)
11941 .build(&Mat::<f64>::zeros(0, 0))
11942 .unwrap();
11943 amg_ch.setup(&a).unwrap();
11944 let old_l = ready_hierarchy(&amg_ch).levels[0]
11945 .cheb
11946 .as_ref()
11947 .unwrap()
11948 .lambda_max;
11949 let old_ds = ready_hierarchy(&amg_ch).levels[0]
11950 .d_sqrt_inv
11951 .as_ref()
11952 .unwrap()[0];
11953 let mut a3 = a.clone();
11954 let rp3 = a3.row_ptr();
11955 for p in rp3[0]..rp3[1] {
11956 if a3.col_idx()[p] == 0 {
11957 a3.values_mut()[p] *= 1.5;
11958 }
11959 }
11960 amg_ch.update_numeric(&a3).unwrap();
11961 let new_l = ready_hierarchy(&amg_ch).levels[0]
11962 .cheb
11963 .as_ref()
11964 .unwrap()
11965 .lambda_max;
11966 let new_ds = ready_hierarchy(&amg_ch).levels[0]
11967 .d_sqrt_inv
11968 .as_ref()
11969 .unwrap()[0];
11970 assert!((new_l - old_l).abs() > 1e-6);
11971 assert!((new_ds - old_ds).abs() > 1e-12);
11972 }
11973
11974 #[test]
11975 fn coarse_dense_factorization_reused_across_cycles() {
11976 let a = poisson1d(16);
11977 let mut amg = AMGBuilder::new()
11978 .relaxation_type(RelaxType::Jacobi)
11979 .grid_relax_type_all(RelaxType::Jacobi)
11980 .coarse_solve(CoarseSolve::DirectDense)
11981 .num_grid_sweeps(RelaxPhase::Coarsest, 0)
11982 .build(&Mat::<f64>::zeros(0, 0))
11983 .unwrap();
11984 amg.setup(&a).unwrap();
11985
11986 let setups_before = {
11987 let h = ready_hierarchy(&amg);
11988 let lvl = &h.levels[h.coarsest_ix()];
11989 lvl.coarse_solver
11990 .as_ref()
11991 .expect("dense coarse solver cache")
11992 .lock()
11993 .unwrap()
11994 .nsetups()
11995 };
11996
11997 let rhs = vec![1.0; a.nrows()];
11998 let mut z = vec![0.0; a.nrows()];
11999 amg.apply(PcSide::Left, &rhs, &mut z).unwrap();
12000 amg.apply(PcSide::Left, &rhs, &mut z).unwrap();
12001
12002 let setups_after = {
12003 let h = ready_hierarchy(&amg);
12004 let lvl = &h.levels[h.coarsest_ix()];
12005 lvl.coarse_solver
12006 .as_ref()
12007 .expect("dense coarse solver cache")
12008 .lock()
12009 .unwrap()
12010 .nsetups()
12011 };
12012 assert!(setups_before > 0);
12013 assert_eq!(setups_after, setups_before);
12014 }
12015
12016 #[test]
12017 #[cfg(not(feature = "complex"))]
12018 fn apply_reuses_v_cycle_workspace() {
12019 let a = poisson1d(32);
12020 let mut amg = AMGBuilder::new()
12021 .relaxation_type(RelaxType::Jacobi)
12022 .grid_relax_type_all(RelaxType::Jacobi)
12023 .build(&Mat::<f64>::zeros(0, 0))
12024 .unwrap();
12025 amg.setup(&a).unwrap();
12026
12027 let rhs = vec![1.0; a.nrows()];
12028 let mut z = vec![0.0; a.nrows()];
12029 amg.apply(PcSide::Left, &rhs, &mut z).unwrap();
12030 let (temp_ptr, temp_cap, coarse_ptrs, coarse_caps) = {
12031 let pool = amg.workspace_pool.lock().unwrap();
12032 let ws = pool.last().expect("workspace returned to pool");
12033 (
12034 ws.temp.as_ptr(),
12035 ws.temp.capacity(),
12036 ws.coarse_sol.iter().map(Vec::as_ptr).collect::<Vec<_>>(),
12037 ws.coarse_sol.iter().map(Vec::capacity).collect::<Vec<_>>(),
12038 )
12039 };
12040
12041 amg.apply(PcSide::Left, &rhs, &mut z).unwrap();
12042 let pool = amg.workspace_pool.lock().unwrap();
12043 let ws = pool.last().expect("workspace returned to pool");
12044 assert_eq!(ws.temp.as_ptr(), temp_ptr);
12045 assert_eq!(ws.temp.capacity(), temp_cap);
12046 assert_eq!(
12047 ws.coarse_sol.iter().map(Vec::as_ptr).collect::<Vec<_>>(),
12048 coarse_ptrs
12049 );
12050 assert_eq!(
12051 ws.coarse_sol.iter().map(Vec::capacity).collect::<Vec<_>>(),
12052 coarse_caps
12053 );
12054 }
12055
12056 #[test]
12057 fn uniform_partition_detection_rejects_irregular_or_empty_parts() {
12058 assert_eq!(uniform_positive_partition_len(&[0, 3, 6, 9]), Some(3));
12059 assert_eq!(uniform_positive_partition_len(&[2, 5, 8]), Some(3));
12060 assert_eq!(uniform_positive_partition_len(&[0, 2, 5]), None);
12061 assert_eq!(uniform_positive_partition_len(&[0, 0, 0]), None);
12062 assert_eq!(uniform_positive_partition_len(&[4, 2]), None);
12063 assert_eq!(uniform_positive_partition_len(&[0]), None);
12064 }
12065
12066 #[test]
12067 fn root_vector_gather_scatter_uniform_fast_path_single_rank() {
12068 let comm = UniverseComm::NoComm(crate::parallel::NoComm);
12069 let row_part = vec![0usize, 4];
12070 let local = vec![1.0, 2.0, 3.0, 4.0];
12071
12072 let gathered = gather_vector(&comm, &row_part, 0, &local)
12073 .expect("uniform gather")
12074 .expect("root receives gathered vector");
12075 assert_eq!(gathered, local);
12076
12077 let mut scattered = vec![0.0; 4];
12078 scatter_vector(&comm, &row_part, 0, Some(&gathered), &mut scattered)
12079 .expect("uniform scatter");
12080 assert_eq!(scattered, gathered);
12081 }
12082
12083 #[test]
12084 fn root_vector_gather_scatter_reject_partition_length_mismatch() {
12085 let comm = UniverseComm::NoComm(crate::parallel::NoComm);
12086 let row_part = vec![0usize, 2, 4];
12087 let local = vec![1.0, 2.0];
12088
12089 assert!(gather_vector(&comm, &row_part, 0, &local).is_err());
12090 let mut out = vec![0.0; 2];
12091 assert!(scatter_vector(&comm, &row_part, 0, Some(&[1.0, 2.0]), &mut out).is_err());
12092 }
12093
12094 #[test]
12095 #[cfg(not(feature = "complex"))]
12096 fn coarse_ilu_reused() {
12097 let n = 8;
12098 let mut row_ptr = vec![0usize; n + 1];
12099 let mut col_idx = Vec::new();
12100 let mut vals = Vec::new();
12101 for i in 0..n {
12102 row_ptr[i] = col_idx.len();
12103 if i > 0 {
12104 col_idx.push(i - 1);
12105 vals.push(-1.0);
12106 }
12107 col_idx.push(i);
12108 vals.push(2.0);
12109 if i + 1 < n {
12110 col_idx.push(i + 1);
12111 vals.push(-1.0);
12112 }
12113 }
12114 row_ptr[n] = col_idx.len();
12115 let a = CsrMatrix::from_csr(n, n, row_ptr, col_idx, vals);
12116
12117 let mut amg = AMGBuilder::new()
12118 .relaxation_type(RelaxType::Jacobi)
12119 .grid_relax_type_all(RelaxType::Jacobi)
12120 .coarse_solve(CoarseSolve::ILU)
12121 .require_spd(false)
12122 .build(&Mat::<f64>::zeros(0, 0))
12123 .unwrap();
12124 amg.setup(&a).unwrap();
12125
12126 let setups_before = {
12127 let h = ready_hierarchy(&amg);
12128 let lvl = &h.levels[h.coarsest_ix()];
12129 lvl.coarse_solver
12130 .as_ref()
12131 .unwrap()
12132 .lock()
12133 .unwrap()
12134 .nsetups()
12135 };
12136
12137 let rhs = vec![1.0; n];
12138 let mut z = vec![0.0; n];
12139 amg.apply(PcSide::Left, &rhs, &mut z).unwrap();
12140 amg.apply(PcSide::Left, &rhs, &mut z).unwrap();
12141
12142 let setups_after = {
12143 let h = ready_hierarchy(&amg);
12144 let lvl = &h.levels[h.coarsest_ix()];
12145 lvl.coarse_solver
12146 .as_ref()
12147 .unwrap()
12148 .lock()
12149 .unwrap()
12150 .nsetups()
12151 };
12152
12153 assert_eq!(setups_before, 1);
12154 assert_eq!(setups_after, 1);
12155 }
12156
12157 #[test]
12158 #[cfg(not(feature = "complex"))]
12159 fn preserves_num_functions_across_levels() {
12160 let n = 16;
12162 let mut row_ptr = Vec::with_capacity(n + 1);
12163 let mut col_idx = Vec::new();
12164 let mut vals = Vec::new();
12165 row_ptr.push(0);
12166 for i in 0..n {
12167 if i > 0 {
12168 col_idx.push(i - 1);
12169 vals.push(-1.0);
12170 }
12171 col_idx.push(i);
12172 vals.push(2.0);
12173 if i + 1 < n {
12174 col_idx.push(i + 1);
12175 vals.push(-1.0);
12176 }
12177 row_ptr.push(col_idx.len());
12178 }
12179 let a = CsrMatrix::from_csr(n, n, row_ptr, col_idx, vals);
12180
12181 let t0 = vec![1.0; n];
12183 let t1: Vec<f64> = (0..n).map(|i| i as f64).collect();
12184
12185 let mut amg = AMGBuilder::new()
12186 .coarse_threshold(1)
12187 .max_coarse_size(1)
12188 .near_nullspace(vec![t0, t1])
12189 .build(&Mat::<f64>::zeros(0, 0))
12190 .unwrap();
12191 amg.setup(&a).unwrap();
12192
12193 let h = ready_hierarchy(&amg);
12194 for l in 0..h.coarsest_ix() {
12195 assert_eq!(h.levels[l].num_functions, 2);
12196 }
12197 }
12198
12199 #[test]
12200 fn apply_before_setup_returns_err() {
12201 let amg = AMG::default();
12202 let rhs = vec![1.0, 2.0];
12203 let mut sol = vec![0.0, 0.0];
12204 assert!(matches!(
12205 amg.apply(PcSide::Left, &rhs, &mut sol),
12206 Err(KError::InvalidInput(_))
12207 ));
12208 }
12209
12210 #[test]
12211 fn setup_is_idempotent_when_ids_unchanged() {
12212 reset_symbolic_counter();
12213 let mat = csr_from_triples(
12214 3,
12215 3,
12216 vec![
12217 (0, 0, 2.0),
12218 (0, 1, -1.0),
12219 (1, 0, -1.0),
12220 (1, 1, 2.0),
12221 (1, 2, -1.0),
12222 (2, 1, -1.0),
12223 (2, 2, 2.0),
12224 ],
12225 );
12226 let op = TestLinOp::new(mat.clone(), StructureId(1), ValuesId(1));
12227 let mut amg = AMG::default();
12228 amg.setup(&op).unwrap();
12229 assert_eq!(symbolic_counter(), 1);
12230 amg.setup(&op).unwrap();
12231 assert_eq!(symbolic_counter(), 1);
12232 }
12233
12234 #[test]
12235 fn values_id_change_refreshes_numeric_only() {
12236 reset_symbolic_counter();
12237 let mat = csr_from_triples(
12238 3,
12239 3,
12240 vec![
12241 (0, 0, 2.0),
12242 (0, 1, -1.0),
12243 (1, 0, -1.0),
12244 (1, 1, 2.0),
12245 (1, 2, -1.0),
12246 (2, 1, -1.0),
12247 (2, 2, 2.0),
12248 ],
12249 );
12250 let mut amg = AMG::default();
12251 let op1 = TestLinOp::new(mat.clone(), StructureId(3), ValuesId(3));
12252 amg.setup(&op1).unwrap();
12253 assert_eq!(symbolic_counter(), 1);
12254 let op2 = op1.with_values(ValuesId(4));
12255 amg.setup(&op2).unwrap();
12256 assert_eq!(symbolic_counter(), 1);
12257 }
12258
12259 #[test]
12260 fn structure_id_change_rebuilds_symbolic() {
12261 reset_symbolic_counter();
12262 let mat1 = csr_from_triples(
12263 3,
12264 3,
12265 vec![
12266 (0, 0, 2.0),
12267 (0, 1, -1.0),
12268 (1, 0, -1.0),
12269 (1, 1, 2.0),
12270 (1, 2, -1.0),
12271 (2, 1, -1.0),
12272 (2, 2, 2.0),
12273 ],
12274 );
12275 let mat2 = csr_from_triples(
12276 3,
12277 3,
12278 vec![
12279 (0, 0, 3.0),
12280 (0, 1, -1.0),
12281 (1, 0, -1.0),
12282 (1, 1, 3.0),
12283 (1, 2, -1.0),
12284 (2, 1, -1.0),
12285 (2, 2, 3.0),
12286 (2, 0, -0.5),
12287 ],
12288 );
12289 let mut amg = AMG::default();
12290 let op1 = TestLinOp::new(mat1, StructureId(5), ValuesId(5));
12291 amg.setup(&op1).unwrap();
12292 assert_eq!(symbolic_counter(), 1);
12293 let op2 = TestLinOp::new(mat2, StructureId(6), ValuesId(6));
12294 amg.setup(&op2).unwrap();
12295 assert_eq!(symbolic_counter(), 2);
12296 }
12297
12298 #[test]
12299 fn transpose_mapping_updates_values() {
12300 let p = prolong::Pcsr {
12301 m: 4,
12302 n: 5,
12303 row_ptr: vec![0, 2, 4, 6, 7],
12304 col_idx: vec![0, 2, 1, 3, 0, 4, 2],
12305 vals: vec![1.0, 2.0, 3.0, 4.0, -1.0, 5.0, 6.0],
12306 };
12307 let (rr, rc, _, p2r) = transpose_csr_with_pos(&p);
12308 let r = CsrMatrix::from_csr(p.n, p.m, rr.clone(), rc.clone(), vec![0.0; p.vals.len()]);
12309 #[cfg(debug_assertions)]
12310 debug_check_csr(&r, "test keep transpose");
12311 let updated_rc = r.col_idx().to_vec();
12312 let mut updated_vals = vec![0.0; p.vals.len()];
12313 for (pi, &ri) in p2r.iter().enumerate() {
12314 updated_vals[ri] = p.vals[pi] * 2.0;
12315 }
12316 let r_updated = CsrMatrix::from_csr(
12317 p.n,
12318 p.m,
12319 rr.clone(),
12320 updated_rc.to_vec(),
12321 updated_vals.clone(),
12322 );
12323 let mut p_dense = Mat::<f64>::zeros(p.m, p.n);
12324 for i in 0..p.m {
12325 for k in p.row_ptr[i]..p.row_ptr[i + 1] {
12326 p_dense[(i, p.col_idx[k])] = p.vals[k] * 2.0;
12327 }
12328 }
12329 let r_dense = p_dense.transpose().to_owned();
12330 let mut r_from_updated = Mat::<f64>::zeros(r_updated.nrows(), r_updated.ncols());
12331 for i in 0..r_updated.nrows() {
12332 for k in r_updated.row_ptr()[i]..r_updated.row_ptr()[i + 1] {
12333 r_from_updated[(i, r_updated.col_idx()[k])] = r_updated.values()[k];
12334 }
12335 }
12336 assert_dense_eq(&r_from_updated, &r_dense, 1e-12, 1e-12);
12337 }
12338
12339 #[test]
12340 fn galerkin_sample_check_spot() {
12341 let a = csr_from_triples(
12342 2,
12343 2,
12344 vec![(0, 0, 2.0), (0, 1, -1.0), (1, 0, -1.0), (1, 1, 2.0)],
12345 );
12346 let p = CsrMatrix::identity(2);
12347 let r = CsrMatrix::identity(2);
12348 let a_coarse = CsrMatrix::identity(2);
12349 let (ok, worst) = galerkin_sample_check(&a, &p, &r, &a_coarse, 4, 1e-12, 0xFEED).unwrap();
12350 assert!(ok);
12351 assert!(worst <= 1e-12);
12352 }
12353
12354 #[test]
12355 fn near_zero_diagonal_aborts_setup() {
12356 let a = csr_from_triples(2, 2, vec![(0, 1, -1.0), (1, 0, -1.0), (1, 1, 1.0)]);
12357 let op = TestLinOp::new(a, StructureId(7), ValuesId(7));
12358 let mut amg = AMG::default();
12359 let err = amg.setup(&op).unwrap_err();
12360 match err {
12361 KError::SolveError(msg) => assert!(msg.contains("near-zero diagonal")),
12362 other => panic!("unexpected error: {other:?}"),
12363 }
12364 }
12365
12366 #[cfg(feature = "complex")]
12367 mod bridge {
12368 use super::*;
12369 use crate::algebra::bridge::BridgeScratch;
12370 use crate::ops::kpc::KPreconditioner;
12371 use crate::preconditioner::PcSide;
12372
12373 fn poisson_1d(n: usize) -> CsrMatrix<f64> {
12374 let mut row_ptr = Vec::with_capacity(n + 1);
12375 let mut col_idx = Vec::new();
12376 let mut values = Vec::new();
12377 row_ptr.push(0);
12378 for i in 0..n {
12379 if i > 0 {
12380 col_idx.push(i - 1);
12381 values.push(-1.0);
12382 }
12383 col_idx.push(i);
12384 values.push(2.0);
12385 if i + 1 < n {
12386 col_idx.push(i + 1);
12387 values.push(-1.0);
12388 }
12389 row_ptr.push(col_idx.len());
12390 }
12391 CsrMatrix::from_csr(n, n, row_ptr, col_idx, values)
12392 }
12393
12394 #[test]
12395 #[cfg(not(feature = "complex"))]
12396 fn apply_s_matches_real_path() {
12397 let a = poisson_1d(12);
12398 let mut amg = AMGBuilder::new()
12399 .relaxation_type(RelaxType::Jacobi)
12400 .grid_relax_type_all(RelaxType::Jacobi)
12401 .build(&Mat::<f64>::zeros(0, 0))
12402 .expect("amg build");
12403 amg.setup(&a).expect("amg setup");
12404
12405 let rhs: Vec<f64> = (0..a.nrows()).map(|i| (i as f64).sin()).collect();
12406 let mut out_real = vec![0.0; rhs.len()];
12407 amg.apply(PcSide::Left, &rhs, &mut out_real)
12408 .expect("real amg apply");
12409
12410 let rhs_s: Vec<S> = rhs.iter().copied().map(S::from_real).collect();
12411 let mut out_s = vec![S::zero(); rhs_s.len()];
12412 let mut scratch = BridgeScratch::default();
12413 amg.apply_s(PcSide::Left, &rhs_s, &mut out_s, &mut scratch)
12414 .expect("scalar amg apply");
12415
12416 for (yr, ys) in out_real.iter().zip(out_s.iter()) {
12417 assert!((ys.real() - yr).abs() < 1e-10, "real mismatch");
12418 assert!(ys.imag().abs() < 1e-12, "imag component drift");
12419 }
12420 }
12421 }
12422}
12423
12424#[cfg(all(test, feature = "complex"))]
12425mod tests_complex {
12426 use super::{
12427 RowScaleMode, csr_pattern_hash, eff_nnz, local_qr, max_row_sum_abs, p_column_norms2,
12428 row_scaling, sync_adjoint_values_from_forward,
12429 };
12430 use crate::algebra::prelude::*;
12431 use crate::matrix::sparse::CsrMatrix;
12432 use crate::preconditioner::approxinv_csr::{ApproxInvBuilder, ApproxInvKind};
12433 use crate::preconditioner::{PcSide, Preconditioner};
12434
12435 fn poisson_1d_complex() -> CsrMatrix<S> {
12436 CsrMatrix::from_csr(
12437 3,
12438 3,
12439 vec![0, 2, 5, 7],
12440 vec![0, 1, 0, 1, 2, 1, 2],
12441 vec![
12442 S::from_parts(2.0, 0.1),
12443 S::from_parts(-1.0, 0.0),
12444 S::from_parts(-1.0, 0.0),
12445 S::from_parts(2.0, -0.2),
12446 S::from_parts(-1.0, 0.0),
12447 S::from_parts(-1.0, 0.0),
12448 S::from_parts(2.0, 0.1),
12449 ],
12450 )
12451 }
12452
12453 #[test]
12454 fn approxinv_complex_in_amg_scope_produces_finite_apply() {
12455 let a = poisson_1d_complex();
12456 let a_real = CsrMatrix::from_csr(
12457 a.nrows(),
12458 a.ncols(),
12459 a.row_ptr().to_vec(),
12460 a.col_idx().to_vec(),
12461 a.values().iter().map(|v| v.real()).collect(),
12462 );
12463 let mut fsai = ApproxInvBuilder::new(ApproxInvKind::FSAI)
12464 .levels(1)
12465 .build_fsai(&a_real)
12466 .unwrap();
12467
12468 let rhs = vec![S::from_parts(1.0, -0.4); 3];
12469 let mut y = vec![S::zero(); 3];
12470 fsai.setup(&a).unwrap();
12471 fsai.apply(PcSide::Left, &rhs, &mut y).unwrap();
12472 assert!(y.iter().all(|v| v.is_finite()));
12473 }
12474
12475 #[test]
12476 fn row_scaling_sum_to_one_preserves_complex_direction() {
12477 let row_ptr = vec![0, 2];
12478 let col_idx = vec![0, 1];
12479 let mut vals = vec![S::from_parts(1.0, 1.0), S::from_parts(1.0, -1.0)];
12480
12481 row_scaling(
12482 RowScaleMode::SumToOne,
12483 1,
12484 None,
12485 &[0],
12486 None,
12487 &row_ptr,
12488 &col_idx,
12489 &mut vals,
12490 )
12491 .unwrap();
12492
12493 let sum = vals[0] + vals[1];
12494 assert!((sum.real() - 1.0).abs() < 1e-12);
12495 assert!(sum.imag().abs() < 1e-12);
12496 assert!((vals[0].imag() - 0.5).abs() < 1e-12);
12497 assert!((vals[1].imag() + 0.5).abs() < 1e-12);
12498 }
12499
12500 #[test]
12501 fn local_qr_uses_hermitian_inner_products_for_complex_values() {
12502 let row_ptr = vec![0, 2, 4];
12503 let col_idx = vec![0, 1, 0, 1];
12504 let mut vals = vec![
12505 S::from_parts(1.0, 0.0),
12506 S::from_parts(1.0, 0.0),
12507 S::from_parts(0.0, 1.0),
12508 S::from_parts(1.0, 0.0),
12509 ];
12510
12511 local_qr(2, &[0, 0], &row_ptr, &col_idx, &mut vals).unwrap();
12512
12513 let q0 = [vals[0], vals[2]];
12514 let q1 = [vals[1], vals[3]];
12515 let dot = q0[0].conj() * q1[0] + q0[1].conj() * q1[1];
12516 let n0 = q0.iter().map(|v| v.abs2()).sum::<f64>().sqrt();
12517 let n1 = q1.iter().map(|v| v.abs2()).sum::<f64>().sqrt();
12518 assert!(
12519 dot.abs() < 1e-12,
12520 "columns are not Hermitian-orthogonal: {dot:?}"
12521 );
12522 assert!((n0 - 1.0).abs() < 1e-12);
12523 assert!((n1 - 1.0).abs() < 1e-12);
12524 }
12525
12526 #[test]
12527 fn csr_diagnostics_use_complex_magnitudes_and_pattern_only_hash() {
12528 let a = CsrMatrix::from_csr(
12529 2,
12530 2,
12531 vec![0, 2, 4],
12532 vec![0, 1, 0, 1],
12533 vec![
12534 S::from_parts(3.0, 4.0),
12535 S::from_parts(0.0, 2.0),
12536 S::from_parts(0.0, -1.0),
12537 S::from_parts(1.0, 1.0),
12538 ],
12539 );
12540 let same_pattern = CsrMatrix::from_csr(
12541 2,
12542 2,
12543 a.row_ptr().to_vec(),
12544 a.col_idx().to_vec(),
12545 vec![
12546 S::from_parts(10.0, -2.0),
12547 S::from_parts(0.0, 0.25),
12548 S::from_parts(7.0, 0.0),
12549 S::from_parts(-1.0, -3.0),
12550 ],
12551 );
12552
12553 let col_norms = p_column_norms2(&a);
12554 assert!((col_norms[0] - 26.0).abs() < 1e-12);
12555 assert!((col_norms[1] - 6.0).abs() < 1e-12);
12556 assert!((max_row_sum_abs(&a) - 7.0).abs() < 1e-12);
12557 assert_eq!(eff_nnz(&a, 2.0), 2);
12558 assert_eq!(csr_pattern_hash(&a), csr_pattern_hash(&same_pattern));
12559 }
12560
12561 #[test]
12562 fn restriction_value_sync_uses_adjoint_conjugates() {
12563 let forward = vec![
12564 S::from_parts(1.0, 2.0),
12565 S::from_parts(-3.0, 0.5),
12566 S::from_parts(0.0, -4.0),
12567 ];
12568 let p2r = vec![2, 0, 1];
12569 let mut adjoint = vec![S::zero(); 3];
12570
12571 sync_adjoint_values_from_forward(&forward, &p2r, &mut adjoint);
12572
12573 assert_eq!(adjoint[0], forward[1].conj());
12574 assert_eq!(adjoint[1], forward[2].conj());
12575 assert_eq!(adjoint[2], forward[0].conj());
12576 }
12577}