Skip to main content

runmat_analysis_fea/solve/runtime_tensor_solver/
mod.rs

1use futures::executor::block_on;
2use runmat_accelerate_api::{provider, GpuTensorHandle, HostTensorView};
3
4use crate::{
5    assembly::AssemblySummary,
6    diagnostics::{FeaDiagnostic, FeaDiagnosticSeverity},
7    operator::{csr_stiffness, dense_stiffness},
8    solve::{linear::LinearSolveResult, preconditioner::SpdPreconditionerKind},
9};
10
11mod operator_impl;
12mod preconditioner_impl;
13
14use operator_impl::{
15    apply_k_device, apply_k_host_from_prepared, dot_handle, linear_shift_indices,
16    DeviceOperatorContext,
17};
18use preconditioner_impl::{
19    apply_preconditioner_device, build_ilu0_factors, PreconditionerDeviceContext,
20};
21
22#[derive(Debug, Clone)]
23pub struct RuntimeTensorPreparedLinearSystem {
24    pub(crate) dof_count: usize,
25    pub(crate) shape: Vec<usize>,
26    pub(crate) diag: Vec<f64>,
27    pub(crate) upper_left: Vec<f64>,
28    pub(crate) upper_right: Vec<f64>,
29    pub(crate) inv_diag: Vec<f64>,
30    pub(crate) ilu_l_subdiag: Vec<f64>,
31    pub(crate) ilu_upper_superdiag: Vec<f64>,
32    pub(crate) ilu_inv_u_diag: Vec<f64>,
33    pub(crate) constrained_mask: Vec<f64>,
34    pub(crate) unconstrained_mask: Vec<f64>,
35    pub(crate) prev_indices: Vec<u32>,
36    pub(crate) next_indices: Vec<u32>,
37}
38
39pub(crate) struct RuntimeTensorWorkspace {
40    pub(crate) full_indices: Vec<u32>,
41    pub(crate) precond_y: Option<GpuTensorHandle>,
42    pub(crate) precond_z: Option<GpuTensorHandle>,
43}
44
45impl RuntimeTensorWorkspace {
46    fn new(n: usize) -> Option<Self> {
47        if n > u32::MAX as usize {
48            return None;
49        }
50        Some(Self {
51            full_indices: (0..n as u32).collect(),
52            precond_y: None,
53            precond_z: None,
54        })
55    }
56
57    fn release(&mut self, provider: &dyn runmat_accelerate_api::AccelProvider) {
58        if let Some(handle) = self.precond_y.take() {
59            let _ = provider.free(&handle);
60        }
61        if let Some(handle) = self.precond_z.take() {
62            let _ = provider.free(&handle);
63        }
64    }
65}
66
67pub fn solve_linear_system_runtime_tensor(
68    summary: &AssemblySummary,
69    preconditioner_kind: SpdPreconditionerKind,
70) -> Option<LinearSolveResult> {
71    solve_linear_system_runtime_tensor_with_initial_guess(summary, preconditioner_kind, None)
72}
73
74pub fn solve_linear_system_runtime_tensor_with_initial_guess(
75    summary: &AssemblySummary,
76    preconditioner_kind: SpdPreconditionerKind,
77    initial_guess: Option<&[f64]>,
78) -> Option<LinearSolveResult> {
79    solve_runtime_tensor_linear_system_internal(
80        summary,
81        None,
82        &summary.operator.rhs,
83        preconditioner_kind,
84        initial_guess,
85    )
86}
87
88pub fn prepare_runtime_tensor_linear_system(
89    summary: &AssemblySummary,
90) -> Option<RuntimeTensorPreparedLinearSystem> {
91    let n = summary.dof_count;
92    if n == 0
93        || dense_stiffness(&summary.operator).is_some()
94        || csr_stiffness(&summary.operator).is_some()
95    {
96        return None;
97    }
98
99    let inv_diag: Vec<f64> = (0..n)
100        .map(|i| {
101            if summary.operator.constrained[i] {
102                1.0
103            } else {
104                1.0 / summary.operator.stiffness_diag[i].abs().max(1.0e-12)
105            }
106        })
107        .collect();
108    let (ilu_l_subdiag, ilu_upper_superdiag, ilu_inv_u_diag) = build_ilu0_factors(summary);
109    let diag = summary.operator.stiffness_diag.clone();
110    let mut upper_left = vec![0.0; n];
111    let mut upper_right = vec![0.0; n];
112    for i in 0..n {
113        if i > 0 && !summary.operator.constrained[i - 1] && !summary.operator.constrained[i] {
114            upper_left[i] = summary.operator.stiffness_upper[i - 1];
115        }
116        if i + 1 < n && !summary.operator.constrained[i + 1] && !summary.operator.constrained[i] {
117            upper_right[i] = summary.operator.stiffness_upper[i];
118        }
119    }
120    let constrained_mask: Vec<f64> = summary
121        .operator
122        .constrained
123        .iter()
124        .map(|&value| if value { 1.0 } else { 0.0 })
125        .collect();
126    let unconstrained_mask: Vec<f64> = summary
127        .operator
128        .constrained
129        .iter()
130        .map(|&value| if value { 0.0 } else { 1.0 })
131        .collect();
132    let prev_indices = linear_shift_indices(n, -1)?;
133    let next_indices = linear_shift_indices(n, 1)?;
134
135    Some(RuntimeTensorPreparedLinearSystem {
136        dof_count: n,
137        shape: vec![n],
138        diag,
139        upper_left,
140        upper_right,
141        inv_diag,
142        ilu_l_subdiag,
143        ilu_upper_superdiag,
144        ilu_inv_u_diag,
145        constrained_mask,
146        unconstrained_mask,
147        prev_indices,
148        next_indices,
149    })
150}
151
152pub fn solve_prepared_linear_system_runtime_tensor(
153    summary: &AssemblySummary,
154    prepared: &RuntimeTensorPreparedLinearSystem,
155    rhs: &[f64],
156    preconditioner_kind: SpdPreconditionerKind,
157    initial_guess: Option<&[f64]>,
158) -> Option<LinearSolveResult> {
159    solve_runtime_tensor_linear_system_internal(
160        summary,
161        Some(prepared),
162        rhs,
163        preconditioner_kind,
164        initial_guess,
165    )
166}
167
168fn solve_runtime_tensor_linear_system_internal(
169    summary: &AssemblySummary,
170    prepared: Option<&RuntimeTensorPreparedLinearSystem>,
171    rhs: &[f64],
172    preconditioner_kind: SpdPreconditionerKind,
173    initial_guess: Option<&[f64]>,
174) -> Option<LinearSolveResult> {
175    let provider = provider()?;
176    let dof_count = prepared
177        .map(|value| value.dof_count)
178        .unwrap_or(summary.dof_count);
179    if dof_count == 0 || rhs.len() != dof_count {
180        return None;
181    }
182
183    let shape_storage = prepared
184        .map(|value| value.shape.clone())
185        .unwrap_or_else(|| vec![dof_count]);
186    let shape = shape_storage.as_slice();
187    let zeros = vec![0.0; dof_count];
188    let inv_diag = prepared
189        .map(|value| value.inv_diag.clone())
190        .unwrap_or_else(|| {
191            (0..dof_count)
192                .map(|i| {
193                    if summary.operator.constrained[i] {
194                        1.0
195                    } else {
196                        1.0 / summary.operator.stiffness_diag[i].abs().max(1.0e-12)
197                    }
198                })
199                .collect()
200        });
201    let (ilu_l_subdiag, ilu_upper_superdiag, ilu_inv_u_diag) = prepared
202        .map(|value| {
203            (
204                value.ilu_l_subdiag.clone(),
205                value.ilu_upper_superdiag.clone(),
206                value.ilu_inv_u_diag.clone(),
207            )
208        })
209        .unwrap_or_else(|| build_ilu0_factors(summary));
210    let (diag, upper_left, upper_right) = prepared
211        .map(|value| {
212            (
213                value.diag.clone(),
214                value.upper_left.clone(),
215                value.upper_right.clone(),
216            )
217        })
218        .unwrap_or_else(|| {
219            let diag = summary.operator.stiffness_diag.clone();
220            let mut upper_left = vec![0.0; dof_count];
221            let mut upper_right = vec![0.0; dof_count];
222            for i in 0..dof_count {
223                if i > 0 && !summary.operator.constrained[i - 1] && !summary.operator.constrained[i]
224                {
225                    upper_left[i] = summary.operator.stiffness_upper[i - 1];
226                }
227                if i + 1 < dof_count
228                    && !summary.operator.constrained[i + 1]
229                    && !summary.operator.constrained[i]
230                {
231                    upper_right[i] = summary.operator.stiffness_upper[i];
232                }
233            }
234            (diag, upper_left, upper_right)
235        });
236    let (constrained_mask, unconstrained_mask) = prepared
237        .map(|value| {
238            (
239                value.constrained_mask.clone(),
240                value.unconstrained_mask.clone(),
241            )
242        })
243        .unwrap_or_else(|| {
244            let constrained_mask: Vec<f64> = summary
245                .operator
246                .constrained
247                .iter()
248                .map(|&value| if value { 1.0 } else { 0.0 })
249                .collect();
250            let unconstrained_mask: Vec<f64> = summary
251                .operator
252                .constrained
253                .iter()
254                .map(|&value| if value { 0.0 } else { 1.0 })
255                .collect();
256            (constrained_mask, unconstrained_mask)
257        });
258    let prev_indices = prepared
259        .map(|value| value.prev_indices.clone())
260        .unwrap_or(linear_shift_indices(dof_count, -1)?);
261    let next_indices = prepared
262        .map(|value| value.next_indices.clone())
263        .unwrap_or(linear_shift_indices(dof_count, 1)?);
264    let mut workspace = RuntimeTensorWorkspace::new(dof_count)?;
265
266    let initial_x = match initial_guess {
267        Some(values) if values.len() == dof_count => values.to_vec(),
268        _ => zeros.clone(),
269    };
270    let mut x = provider
271        .upload(&HostTensorView {
272            data: &initial_x,
273            shape,
274        })
275        .ok()?;
276    let zero_h = provider
277        .upload(&HostTensorView {
278            data: &zeros,
279            shape,
280        })
281        .ok()?;
282    let rhs_h = provider.upload(&HostTensorView { data: rhs, shape }).ok()?;
283    let inv = provider
284        .upload(&HostTensorView {
285            data: &inv_diag,
286            shape,
287        })
288        .ok()?;
289    let ilu_l_subdiag_h = provider
290        .upload(&HostTensorView {
291            data: &ilu_l_subdiag,
292            shape,
293        })
294        .ok()?;
295    let ilu_upper_superdiag_h = provider
296        .upload(&HostTensorView {
297            data: &ilu_upper_superdiag,
298            shape,
299        })
300        .ok()?;
301    let ilu_inv_u_diag_h = provider
302        .upload(&HostTensorView {
303            data: &ilu_inv_u_diag,
304            shape,
305        })
306        .ok()?;
307    let diag_h = provider
308        .upload(&HostTensorView { data: &diag, shape })
309        .ok()?;
310    let upper_left_h = provider
311        .upload(&HostTensorView {
312            data: &upper_left,
313            shape,
314        })
315        .ok()?;
316    let upper_right_h = provider
317        .upload(&HostTensorView {
318            data: &upper_right,
319            shape,
320        })
321        .ok()?;
322    let constrained_mask_h = provider
323        .upload(&HostTensorView {
324            data: &constrained_mask,
325            shape,
326        })
327        .ok()?;
328    let unconstrained_mask_h = provider
329        .upload(&HostTensorView {
330            data: &unconstrained_mask,
331            shape,
332        })
333        .ok()?;
334    let device_operator = DeviceOperatorContext {
335        provider,
336        diag: &diag_h,
337        upper_left: &upper_left_h,
338        upper_right: &upper_right_h,
339        constrained_mask: &constrained_mask_h,
340        unconstrained_mask: &unconstrained_mask_h,
341        prev_indices: &prev_indices,
342        next_indices: &next_indices,
343        shape,
344    };
345
346    let mut r = if initial_guess.is_some() {
347        let ax = apply_k_device(&device_operator, &x)?;
348        let residual = block_on(provider.elem_sub(&rhs_h, &ax)).ok()?;
349        let _ = provider.free(&ax);
350        residual
351    } else {
352        block_on(provider.elem_add(&rhs_h, &zero_h)).ok()?
353    };
354
355    let preconditioner_ctx = PreconditionerDeviceContext {
356        provider,
357        inv_diag: &inv,
358        ilu_l_subdiag: &ilu_l_subdiag_h,
359        ilu_upper_superdiag: &ilu_upper_superdiag_h,
360        ilu_inv_u_diag: &ilu_inv_u_diag_h,
361        constrained_mask: &constrained_mask_h,
362        unconstrained_mask: &unconstrained_mask_h,
363        prev_indices: &prev_indices,
364        next_indices: &next_indices,
365        shape,
366        zero_like: &zero_h,
367    };
368
369    let mut z =
370        apply_preconditioner_device(&preconditioner_ctx, preconditioner_kind, &r, &mut workspace)?;
371    let mut p = block_on(provider.elem_add(&z, &x)).ok()?;
372
373    let mut host_sync_count: u32 = 0;
374    let mut rz_old = dot_handle(provider, &r, &z, &mut host_sync_count)?;
375    let b_norm = rhs.iter().map(|v| v * v).sum::<f64>().sqrt().max(1.0);
376    let tol = 1.0e-8;
377    let max_iters = 64;
378    let mut converged = false;
379    let mut iterations = 0u32;
380    let mut last_rr: Option<f64> = None;
381    let mut device_apply_k_count: u32 = 0;
382    let mut device_apply_k_attempt_count: u32 = 0;
383
384    for _ in 0..max_iters {
385        device_apply_k_attempt_count = device_apply_k_attempt_count.saturating_add(1);
386        let ap = match apply_k_device(&device_operator, &p) {
387            Some(value) => {
388                device_apply_k_count = device_apply_k_count.saturating_add(1);
389                value
390            }
391            None => {
392                host_sync_count = host_sync_count.saturating_add(1);
393                let p_host = block_on(provider.download(&p)).ok()?;
394                let ap_host = apply_k_host_from_prepared(
395                    &diag,
396                    &upper_left,
397                    &upper_right,
398                    &constrained_mask,
399                    &unconstrained_mask,
400                    &p_host.data,
401                );
402                provider
403                    .upload(&HostTensorView {
404                        data: &ap_host,
405                        shape,
406                    })
407                    .ok()?
408            }
409        };
410
411        let denom = dot_handle(provider, &p, &ap, &mut host_sync_count)?;
412        if denom.abs() <= 1.0e-18 {
413            let _ = provider.free(&ap);
414            break;
415        }
416        let alpha = rz_old / denom;
417
418        let scaled_p = provider.scalar_mul(&p, alpha).ok()?;
419        let new_x = block_on(provider.elem_add(&x, &scaled_p)).ok()?;
420        let _ = provider.free(&x);
421        let _ = provider.free(&scaled_p);
422        x = new_x;
423
424        let scaled_ap = provider.scalar_mul(&ap, alpha).ok()?;
425        let new_r = block_on(provider.elem_sub(&r, &scaled_ap)).ok()?;
426        let _ = provider.free(&r);
427        let _ = provider.free(&scaled_ap);
428        let _ = provider.free(&ap);
429        r = new_r;
430
431        let rr = dot_handle(provider, &r, &r, &mut host_sync_count)?;
432        last_rr = Some(rr);
433        let residual_norm = rr.sqrt();
434        iterations += 1;
435        if residual_norm / b_norm <= tol {
436            converged = true;
437            break;
438        }
439
440        let new_z = apply_preconditioner_device(
441            &preconditioner_ctx,
442            preconditioner_kind,
443            &r,
444            &mut workspace,
445        )?;
446        let rz_new = dot_handle(provider, &r, &new_z, &mut host_sync_count)?;
447        if rz_old.abs() <= 1.0e-18 {
448            let _ = provider.free(&z);
449            z = new_z;
450            break;
451        }
452        let beta = rz_new / rz_old;
453
454        let beta_p = provider.scalar_mul(&p, beta).ok()?;
455        let new_p = block_on(provider.elem_add(&new_z, &beta_p)).ok()?;
456        let _ = provider.free(&p);
457        let _ = provider.free(&beta_p);
458        let _ = provider.free(&z);
459        p = new_p;
460        z = new_z;
461        rz_old = rz_new;
462    }
463
464    host_sync_count = host_sync_count.saturating_add(1);
465    let x_host = block_on(provider.download(&x)).ok()?;
466    let residual_norm = if let Some(rr) = last_rr {
467        rr.sqrt()
468    } else {
469        dot_handle(provider, &r, &r, &mut host_sync_count)?.sqrt()
470    };
471    let mut diagnostics = vec![FeaDiagnostic {
472        code: "FEA_SOLVER_METHOD".to_string(),
473        severity: FeaDiagnosticSeverity::Info,
474        message: format!(
475            "solver=pcg preconditioner={} matrix_free=true backend=runtime_tensor",
476            preconditioner_kind.as_str()
477        ),
478    }];
479    if !converged {
480        diagnostics.push(FeaDiagnostic {
481            code: "FEA_CG_MAX_ITERS".to_string(),
482            severity: FeaDiagnosticSeverity::Warning,
483            message: format!(
484                "runtime_tensor pcg reached max iterations ({max_iters}) with residual_norm={residual_norm}"
485            ),
486        });
487    }
488
489    let _ = provider.free(&z);
490    let _ = provider.free(&p);
491    let _ = provider.free(&r);
492    let _ = provider.free(&x);
493    let _ = provider.free(&zero_h);
494    let _ = provider.free(&inv);
495    let _ = provider.free(&ilu_l_subdiag_h);
496    let _ = provider.free(&ilu_upper_superdiag_h);
497    let _ = provider.free(&ilu_inv_u_diag_h);
498    let _ = provider.free(&diag_h);
499    let _ = provider.free(&upper_left_h);
500    let _ = provider.free(&upper_right_h);
501    let _ = provider.free(&constrained_mask_h);
502    let _ = provider.free(&unconstrained_mask_h);
503    let _ = provider.free(&rhs_h);
504    workspace.release(provider);
505
506    Some(LinearSolveResult {
507        iterations,
508        residual_norm,
509        converged,
510        host_sync_count,
511        solver_backend: "runtime_tensor".to_string(),
512        device_apply_k_count,
513        device_apply_k_attempt_count,
514        solution: x_host.data,
515        solver_method: "matrix_free_pcg".to_string(),
516        preconditioner: preconditioner_kind.as_str().to_string(),
517        diagnostics,
518    })
519}
520
521pub fn estimate_runtime_tensor_pcg_host_syncs(max_iters: u32) -> u32 {
522    let _ = max_iters;
523    1
524}
525
526#[cfg(test)]
527mod tests {
528    use super::*;
529
530    #[test]
531    fn host_sync_estimator_matches_formula() {
532        assert_eq!(estimate_runtime_tensor_pcg_host_syncs(0), 1);
533        assert_eq!(estimate_runtime_tensor_pcg_host_syncs(1), 1);
534        assert_eq!(estimate_runtime_tensor_pcg_host_syncs(64), 1);
535    }
536}