Skip to main content

kryst/preconditioner/
amg.rs

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
66// New sparse SA/RS submodules
67mod 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// ===== Public enums (kept compatible with your old file) =====================
103
104/// Coarsening strategies.
105#[derive(Clone, Copy, Debug, PartialEq, Eq)]
106pub enum CoarsenType {
107    RS,
108    HMIS,
109    PMIS,
110    Falgout,
111}
112
113/// Interpolation strategies.
114#[derive(Clone, Copy, Debug, PartialEq, Eq)]
115pub enum InterpType {
116    Classical,
117    Direct,
118    Multipass,
119    Extended,
120    Standard,
121    HE,
122}
123
124/// Response when rank diagnostics flag interpolation issues.
125#[derive(Clone, Copy, Debug, PartialEq, Eq)]
126pub enum RankFallback {
127    RetryLooserInterp,
128    SwitchInterpKind,
129    Reaggregate,
130    Abort,
131}
132
133/// Relaxation/smoothing choices.
134#[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/// Per-phase relaxation controls mirroring Boomer semantics.
151#[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// ===== Config + Builder ======================================================
198
199#[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    /// Levels (0 = finest) where the Krylov smoother is enabled.
216    pub levels: Vec<usize>,
217    /// Number of Krylov iterations when enabled.
218    pub iters: usize,
219    pub algo: KrylovAlgo,
220    /// Apply the Krylov smoother after the standard post-smoothing step.
221    pub place_post: bool,
222    /// Apply the Krylov smoother before the standard post-smoothing step.
223    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    /// For scalar or when near-nullspace is provided: enforce per-row reproduction
265    /// target by scaling the subset of columns for each function α.
266    ToNearNullspace,
267    /// Fallback for scalar r=1 without NNS: enforce ∑ P\[i,*\] = 1.
268    SumToOne,
269    /// Make each row’s α-slice unit L2 norm (robust when T is noisy).
270    L2Unit,
271    /// Diagonal-weighted norm: sqrt(p^T D p) = 1 using D = diag(A).
272    DUnit,
273}
274
275#[derive(Clone, Copy, Debug, PartialEq)]
276pub enum PostInterpType {
277    None,
278    /// Pure row-scaling family (stable, cheap, deterministic)
279    RowScaling(RowScaleMode),
280    /// Optional orthonormalization of per-aggregate blocks
281    LocalQR,
282    /// Extra SA-like smoothing passes on P values (fixed pattern).
283    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,           // HYPRE default: 25
310    pub strong_threshold: f64,       // HYPRE default: 0.25
311    pub coarse_threshold: usize,     // HYPRE default: 9
312    pub max_coarse_size: usize,      // HYPRE default: 9
313    pub min_coarse_size: usize,      // HYPRE minimum: 1
314    pub truncation_factor: f64,      // 0 => no truncation
315    pub max_elements_per_row: usize, // 0 => unlimited
316    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], // [Fine, Down, Up, Coarsest]
325    pub num_grid_sweeps: [usize; 4],     // [Fine, Down, Up, Coarsest]
326    // legacy shims
327    pub pre_sweeps: usize,         // HYPRE default: 1
328    pub post_sweeps: usize,        // HYPRE default: 1
329    pub coarsen_type: CoarsenType, // HYPRE default: HMIS
330    pub interp_type: InterpType,   // robust: Extended/Standard
331    pub relax_type: RelaxType,     // HYPRE default: Gauss-Seidel (we implement Jacobi)
332    pub logging_level: usize,
333    pub print_level: usize,
334    pub tolerance: f64, // for coarse direct solve (CG)
335    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,  // NEW: used for dense->CSR conversion
354    pub stats_eps: f64, // threshold for effective nnz reporting
355    // P rank diagnostics
356    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    // Galerkin sampling controls
361    pub verify_galerkin: bool,
362    pub galerkin_samples: usize,
363    pub galerkin_rel_tol: f64,
364    // Rank fallback policy
365    pub on_rank_failure: RankFallback,
366    // FSAI smoother controls
367    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    // New SA/RS controls
375    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    /// In complex builds, reject genuinely complex non-diagonal operators until AMG
411    /// carries complex values through the full hierarchy.
412    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
1326/// Builder for `AMG` (preserves old chaining API).
1327pub 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        // leave Coarsest as-is
1448        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
1739// ===== Workspace, levels & hierarchy ========================================
1740
1741fn 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_l (coarse operator at this level, l = 0 is finest)
2124    a: CsrMatrix<f64>,
2125    /// P_l (interpolation to next coarser level)
2126    p: CsrMatrix<f64>,
2127    /// R_l (restriction to next coarser level)
2128    r: CsrMatrix<f64>,
2129    /// diag(A_l)^{-1}
2130    diag_inv: Vec<f64>,
2131    /// diag(A_l)^{-1/2} (optional)
2132    d_sqrt_inv: Option<Vec<f64>>,
2133    /// 1 / (\sum_j |a_ij|) for L1-Jacobi
2134    l1_inv: Option<Vec<f64>>,
2135    /// diag(A_l)^{-1} with safeguards for poor diagonals (optional)
2136    diag_inv_safe: Option<Vec<f64>>,
2137    /// diag(A_l)^{-1/2} derived from safeguarded diagonal (optional)
2138    d_sqrt_inv_safe: Option<Vec<f64>>,
2139    /// Cached spectral bounds for Chebyshev smoother
2140    cheb: Option<ChebData>,
2141    /// Cached spectral bounds for safeguarded Chebyshev smoother
2142    cheb_safe: Option<ChebData>,
2143    /// fine->coarse aggregate id used to rebuild P values (SA numeric refresh)
2144    agg_of: Vec<usize>,
2145    /// coarse/fine flags for classical interpolation
2146    is_c: Vec<bool>,
2147    /// CF metadata for classical interpolation
2148    cf: Option<CFInfo>,
2149    /// Mapping from P entry index -> index in R (transpose) values array
2150    p2r_pos: Vec<usize>,
2151    /// number of coarse basis functions per aggregate
2152    num_functions: usize,
2153    /// Row-local orthonormalized basis coefficients for nodal SA, length = nrows * num_functions.
2154    row_basis: Option<Vec<f64>>,
2155    /// DOF layout for nodal aggregation
2156    layout: Option<DofLayout>,
2157    /// stored near-nullspace basis (optional)
2158    nns: Option<Vec<Vec<f64>>>,
2159    /// Symbolic pattern for A_{l+1}
2160    a_next_pat: Option<CsrPattern>,
2161    /// Filtered NG pattern for A_{l+1}
2162    a_next_pat_ng: Option<CsrPattern>,
2163    /// Mapping from full RAP nnz index -> NG nnz index
2164    rap_full2ng_pos: Option<Vec<Option<usize>>>,
2165    /// Optional transpose storage when keep_transpose is disabled
2166    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>, // 0..L ; L is coarsest
2276    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")]
3024// Compatibility cache for APIs that still store a real CSR for dimensions,
3025// pattern checks, or legacy real-only hierarchy setup. Native complex setup and
3026// apply use the original complex CSR values.
3027fn 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    // Root-centric gather: every rank sends its owned slice to the root, which
3358    // assembles the full global vector. Non-root ranks return None.
3359    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    // Root-centric scatter: the root slices the global vector by row partition
3411    // and sends each owned segment to its rank. Non-root ranks only receive.
3412    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// ===== Main AMG object =======================================================
3868
3869#[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    /// Build transfer operators with the canonical AMG restriction `R = P^H`.
3926    ///
3927    /// AMG setup derives its restriction from the prolongation in scalar-generic
3928    /// Galerkin paths. This constructor makes that contract explicit for callers
3929    /// that only need to provide a custom prolongation.
3930    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    // ---- Setup paths --------------------------------------------------------
4089
4090    #[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        // Build the full hierarchy from scratch (symbolic + numeric)
4177        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            // best-effort: pick the deepest matching override for the built hierarchy.
4218            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        // Distributed setup is root-centric: we gather the global matrix, build the
4262        // hierarchy on rank 0, and keep non-root ranks in an uninitialized state.
4263        // Validation that depends on distributed apply (e.g., SPD probing) is disabled
4264        // under MPI because apply_dist is collective. Errors are synchronized so all
4265        // ranks return the same failure before any iteration starts.
4266        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        // Update finest A_0 and diag(A_0)^{-1}
5188        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        // Recompute P_l values, R_l values, and A_{l+1} values using fixed patterns
5213        for l in 0..h.coarsest_ix() {
5214            // Recompute P_l values in-place using SA smoother with fixed pattern
5215            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                    &params,
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            // Update R values from P via precomputed transpose mapping
5412            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            // Recompute A_{l+1} values by RAP numeric using fixed pattern
5433            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                // Safety fallback: full RAP (structure + values)
5556                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    // ---- Smoother -----------------------------------------------------------
5722
5723    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            // work = A * temp
5744            a.spmv_scaled(1.0, &ws.temp[..n], 0.0, &mut ws.work[..n])?;
5745            // temp += omega * D^{-1} * (r - work)
5746            #[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                &gt_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    // single dispatch point for all relaxation strategies
6655    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    // ---- Multigrid cycle ----------------------------------------------------
6786
6787    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        // Pre-smooth
7314        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        // residual = rhs - A * sol
7354        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        // r_c = R * residual
7402        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        // Post-smooth
7476        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    // Convenience to avoid trait ambiguity in examples
7539    #[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    /// True when the latest distributed setup used a non-scalable gather of the
7580    /// fine matrix.
7581    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// ===== Preconditioner trait (new API) =======================================
7648
7649#[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// ===== Legacy adapter (unchanged external signature) ========================
8089
8090#[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
8100// ===== Hierarchy construction (symbolic + numeric) ==========================
8101
8102fn 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    // Level 0 (finest)
8130    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    // Drive coarsening: build levels 0..L (inclusive L is coarsest)
8251    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        // 1) Strength of connection (sparse)
8265        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        // 2) Aggregates
8286        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                        &params,
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        // R = P^T pattern and values
8634        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        // 4) Coarse operator A_c symbolic and numeric
8666        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        // Replace previous temporary P/R by actual inter-level transfers and agg mapping
8790        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        // Next level (coarser)
8820        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        } // stalled
8926        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
9275// ===== Sparse utilities (local; avoid dense on hot path) ====================
9276
9277fn 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 &current {
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
10193/// CSR * CSR using per-row growing maps (Gustavson-style, simple).
10194fn 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
10266/// RAP = R * A * P
10267fn 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
10276// ===== Coarsening & interpolation (dense helpers, same as old) ==============
10277
10278fn 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
10335/// S(i,j) = |A_ij| / sqrt(|A_ii| |A_jj|) if above threshold.
10336fn 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
10431/// Greedy aggregation (balanced, small aggregates).
10432fn 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    // Order by total strength descending
10439    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        // pick strongest distinct neighbors
10453        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
10474/// Piecewise-constant prolongation from aggregates (dense).
10475fn 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
10486/// Simple smoothing of P with weight* A (Jacobi-like).
10487fn 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
10497/// Row 2-norm normalization.
10498fn 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
10512// ===== Small sparse CG for coarsest level ===================================
10513
10514fn 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// ===== Helpers for transpose mapping and stats ==============================
10651
10652#[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    /// User-facing label for the distributed AMG apply/setup route.
10769    ///
10770    /// These labels are stable diagnostics strings and intentionally differ from
10771    /// Rust enum debug names.
10772    pub fn mode_label(&self) -> &'static str {
10773        dist_strategy_label(self.mode)
10774    }
10775
10776    /// User-facing label for the selected distributed coarse solver route.
10777    pub fn coarse_solver_route_label(&self) -> &'static str {
10778        dist_route_label(self.coarse_solver_route, self.mode)
10779    }
10780
10781    /// True when the selected route is the non-scalable root gather path.
10782    pub fn uses_root_gather(&self) -> bool {
10783        matches!(self.mode, DistCoarseStrategy::RootGather)
10784    }
10785
10786    /// True when this resolved route provides native distributed execution.
10787    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    /// True when the route uses rank-local `DistCsrOp` SpMV for the fine-grid
10793    /// residual instead of gathering the fine matrix or vectors to a root rank.
10794    ///
10795    /// This is weaker than [`Self::reports_distributed_support`]: the current
10796    /// `distributed_csr` milestone uses the canonical distributed CSR operator
10797    /// for fine residual correction, while the multilevel AMG hierarchy itself
10798    /// is still local-block based.
10799    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    /// True when setup used a non-scalable gather of the fine distributed matrix.
10806    pub fn setup_uses_fine_matrix_gather(&self) -> bool {
10807        self.setup_gathered_fine_matrix
10808    }
10809
10810    /// True when the latest apply used root-centric vector gather/scatter.
10811    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; // initial residual norm squared for rhs=[1,1,1]
11915        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        // Build 1D Poisson matrix of size 16
12161        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        // Two-function near nullspace: constant and linear
12182        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}