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}