Skip to main content

runmat_analysis_fea/solve/modal/
mod.rs

1use runmat_accelerate_api::provider as accel_provider;
2use serde::{Deserialize, Serialize};
3use std::time::Instant;
4
5use crate::{
6    assembly::AssemblySummary,
7    diagnostics::{FeaDiagnostic, FeaDiagnosticSeverity},
8    operator::{apply_k, apply_m},
9    solve::runtime_tensor_solver::prepare_runtime_tensor_linear_system,
10    ComputeBackend,
11};
12
13mod diagnostics;
14mod linear_solve;
15mod math;
16
17use diagnostics::push_modal_quality_diagnostics;
18use linear_solve::{solve_k_system_cg, CgSolveOptions};
19use math::{dot, normalize_mass, orthonormalize_mass, relative_l2_update};
20
21#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
22pub struct ModalSolveResult {
23    pub converged: bool,
24    pub eigenvalues_hz: Vec<f64>,
25    pub mode_shapes: Vec<Vec<f64>>,
26    pub residual_norms: Vec<f64>,
27    pub diagnostics: Vec<FeaDiagnostic>,
28    pub solver_method: String,
29    pub solver_backend: String,
30    pub solver_host_sync_count: u32,
31    pub device_apply_k_count: u32,
32    pub device_apply_k_attempt_count: u32,
33}
34
35pub fn solve_modal_system(
36    summary: &AssemblySummary,
37    mode_count: usize,
38    backend: ComputeBackend,
39) -> ModalSolveResult {
40    if summary.dof_count == 0 || mode_count == 0 {
41        return ModalSolveResult {
42            converged: false,
43            eigenvalues_hz: Vec::new(),
44            mode_shapes: Vec::new(),
45            residual_norms: Vec::new(),
46            diagnostics: vec![FeaDiagnostic {
47                code: "FEA_MODAL_EMPTY_SYSTEM".to_string(),
48                severity: FeaDiagnosticSeverity::Warning,
49                message:
50                    "modal solve skipped because assembled system has zero DOFs or requested mode_count is zero"
51                        .to_string(),
52            }],
53            solver_method: "matrix_free_subspace_iteration".to_string(),
54            solver_backend: "cpu_reference".to_string(),
55            solver_host_sync_count: 0,
56            device_apply_k_count: 0,
57            device_apply_k_attempt_count: 0,
58        };
59    }
60
61    let use_runtime_tensor = backend == ComputeBackend::Gpu;
62    let mut solver_backend = "cpu_reference".to_string();
63    let mut solver_host_sync_count = 0u32;
64    let mut device_apply_k_count = 0u32;
65    let mut device_apply_k_attempt_count = 0u32;
66
67    let unconstrained: Vec<usize> = summary
68        .operator
69        .constrained
70        .iter()
71        .enumerate()
72        .filter_map(|(i, is_constrained)| if *is_constrained { None } else { Some(i) })
73        .collect();
74    let target_mode_count = mode_count.min(unconstrained.len());
75
76    let mut basis: Vec<Vec<f64>> = Vec::with_capacity(target_mode_count);
77    let mut modes: Vec<(f64, Vec<f64>, f64)> = Vec::with_capacity(target_mode_count);
78    let has_accel_provider = use_runtime_tensor && accel_provider().is_some();
79    let mut prepared_build_ms = 0.0_f64;
80    let prepared_runtime_system = if has_accel_provider {
81        let prepared_start = Instant::now();
82        let prepared = prepare_runtime_tensor_linear_system(summary);
83        prepared_build_ms = prepared_start.elapsed().as_secs_f64() * 1_000.0;
84        prepared
85    } else {
86        None
87    };
88    let mut solve_ms = 0.0_f64;
89    let mut fallback_apply_count = 0u32;
90    let (max_inverse_iters, min_inverse_iters, update_tol) = if has_accel_provider {
91        (6usize, 2usize, 5.0e-4)
92    } else {
93        (8usize, 3usize, 1.0e-4)
94    };
95
96    for mode_idx in 0..target_mode_count {
97        let mut q = vec![0.0; summary.operator.dof_count];
98        q[unconstrained[mode_idx]] = 1.0;
99        normalize_mass(&summary.operator, &mut q);
100        let mut linear_guess: Option<Vec<f64>> = None;
101
102        for iter in 0..max_inverse_iters {
103            let q_prev = q.clone();
104            let mq = apply_m(&summary.operator, &q);
105            let solve_start = Instant::now();
106            let z = solve_k_system_cg(
107                summary,
108                &summary.operator,
109                &mq,
110                CgSolveOptions {
111                    max_iters: 64,
112                    tol: 1.0e-10,
113                    use_runtime_tensor,
114                    prepared_runtime_system: prepared_runtime_system.as_ref(),
115                    initial_guess: linear_guess.as_deref(),
116                },
117            );
118            solve_ms += solve_start.elapsed().as_secs_f64() * 1_000.0;
119            if let Some(solve) = z.runtime_tensor {
120                solver_backend = solve.solver_backend;
121                solver_host_sync_count =
122                    solver_host_sync_count.saturating_add(solve.host_sync_count);
123                device_apply_k_count =
124                    device_apply_k_count.saturating_add(solve.device_apply_k_count);
125                device_apply_k_attempt_count =
126                    device_apply_k_attempt_count.saturating_add(solve.device_apply_k_attempt_count);
127                fallback_apply_count = fallback_apply_count.saturating_add(
128                    solve
129                        .device_apply_k_attempt_count
130                        .saturating_sub(solve.device_apply_k_count),
131                );
132            }
133            let mut z_vec = z.vector;
134            linear_guess = Some(z_vec.clone());
135            orthonormalize_mass(&summary.operator, &mut z_vec, &basis);
136            normalize_mass(&summary.operator, &mut z_vec);
137            q = z_vec;
138
139            if iter + 1 >= min_inverse_iters {
140                let rel_update = relative_l2_update(&q_prev, &q);
141                if rel_update <= update_tol {
142                    break;
143                }
144            }
145        }
146
147        let kq = apply_k(&summary.operator, &q);
148        let mq = apply_m(&summary.operator, &q);
149        let lambda = (dot(&q, &kq) / dot(&q, &mq).abs().max(1.0e-12)).max(0.0);
150        let freq_hz = lambda.sqrt() / (2.0 * std::f64::consts::PI);
151
152        let residual = kq
153            .iter()
154            .zip(mq.iter())
155            .map(|(k_value, m_value)| {
156                let diff = *k_value - lambda * *m_value;
157                diff * diff
158            })
159            .sum::<f64>()
160            .sqrt();
161        let kq_norm = kq
162            .iter()
163            .map(|value| value * value)
164            .sum::<f64>()
165            .sqrt()
166            .max(1.0e-12);
167
168        basis.push(q.clone());
169        modes.push((freq_hz, q, residual / kq_norm));
170    }
171
172    modes.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal));
173    let mut eigenvalues_hz = Vec::with_capacity(modes.len());
174    let mut mode_shapes = Vec::with_capacity(modes.len());
175    let mut residual_norms = Vec::with_capacity(modes.len());
176    for (freq_hz, shape, residual) in modes {
177        eigenvalues_hz.push(freq_hz);
178        mode_shapes.push(shape);
179        residual_norms.push(residual);
180    }
181
182    let converged = !eigenvalues_hz.is_empty();
183    let mut diagnostics = vec![FeaDiagnostic {
184        code: "FEA_MODAL_METHOD".to_string(),
185        severity: FeaDiagnosticSeverity::Info,
186        message: "solver=matrix_free_subspace_iteration inverse_k=true".to_string(),
187    }];
188    diagnostics.push(FeaDiagnostic {
189        code: "FEA_MODAL_CONVERGENCE".to_string(),
190        severity: if converged {
191            FeaDiagnosticSeverity::Info
192        } else {
193            FeaDiagnosticSeverity::Warning
194        },
195        message: format!(
196            "mode_count_requested={} mode_count_solved={} converged={}",
197            mode_count,
198            eigenvalues_hz.len(),
199            converged
200        ),
201    });
202
203    push_modal_quality_diagnostics(
204        &mut diagnostics,
205        &summary.operator,
206        &eigenvalues_hz,
207        &mode_shapes,
208        &residual_norms,
209    );
210    diagnostics.push(FeaDiagnostic {
211        code: "FEA_MODAL_COST".to_string(),
212        severity: FeaDiagnosticSeverity::Info,
213        message: format!(
214            "prepared_build_ms={} solve_ms={} fallback_apply_count={}",
215            prepared_build_ms, solve_ms, fallback_apply_count
216        ),
217    });
218
219    ModalSolveResult {
220        converged,
221        eigenvalues_hz,
222        mode_shapes,
223        residual_norms,
224        diagnostics,
225        solver_method: "matrix_free_subspace_iteration".to_string(),
226        solver_backend,
227        solver_host_sync_count,
228        device_apply_k_count,
229        device_apply_k_attempt_count,
230    }
231}