Skip to main content

runmat_analysis_fea/solve/backend/
runtime_tensor.rs

1use futures::executor::block_on;
2
3use runmat_accelerate_api::{provider, HostTensorView};
4
5use super::{cpu_reference::CpuReferenceBackend, linear_algebra::LinearAlgebraBackend};
6
7#[derive(Debug, Clone, Copy, Default)]
8pub struct RuntimeTensorBackend;
9
10impl LinearAlgebraBackend for RuntimeTensorBackend {
11    fn dot(&self, a: &[f64], b: &[f64]) -> f64 {
12        if let Some(value) = try_dot_gpu(a, b) {
13            return value;
14        }
15        CpuReferenceBackend.dot(a, b)
16    }
17
18    fn axpy(&self, alpha: f64, x: &[f64], y: &mut [f64]) {
19        if let Some(out) = try_axpy_gpu(alpha, x, y) {
20            if out.len() == y.len() {
21                y.copy_from_slice(&out);
22                return;
23            }
24        }
25        CpuReferenceBackend.axpy(alpha, x, y)
26    }
27
28    fn vec_sub(&self, a: &[f64], b: &[f64]) -> Vec<f64> {
29        if let Some(out) = try_vec_sub_gpu(a, b) {
30            return out;
31        }
32        CpuReferenceBackend.vec_sub(a, b)
33    }
34}
35
36fn try_dot_gpu(a: &[f64], b: &[f64]) -> Option<f64> {
37    if a.len() != b.len() {
38        return None;
39    }
40    let provider = provider()?;
41    let shape = [a.len()];
42    let ah = provider
43        .upload(&HostTensorView {
44            data: a,
45            shape: &shape,
46        })
47        .ok()?;
48    let bh = provider
49        .upload(&HostTensorView {
50            data: b,
51            shape: &shape,
52        })
53        .ok()?;
54    let mul = block_on(provider.elem_mul(&ah, &bh)).ok()?;
55    let sum = block_on(provider.reduce_sum(&mul)).ok()?;
56
57    let scalar = match provider.read_scalar(&sum, 0) {
58        Ok(value) => Some(value),
59        Err(_) => block_on(provider.download(&sum))
60            .ok()
61            .and_then(|host| host.data.first().copied()),
62    };
63
64    let _ = provider.free(&sum);
65    let _ = provider.free(&mul);
66    let _ = provider.free(&bh);
67    let _ = provider.free(&ah);
68    scalar
69}
70
71fn try_axpy_gpu(alpha: f64, x: &[f64], y: &[f64]) -> Option<Vec<f64>> {
72    if x.len() != y.len() {
73        return None;
74    }
75    let provider = provider()?;
76    let shape = [x.len()];
77    let xh = provider
78        .upload(&HostTensorView {
79            data: x,
80            shape: &shape,
81        })
82        .ok()?;
83    let yh = provider
84        .upload(&HostTensorView {
85            data: y,
86            shape: &shape,
87        })
88        .ok()?;
89    let scaled = provider.scalar_mul(&xh, alpha).ok()?;
90    let out_h = block_on(provider.elem_add(&scaled, &yh)).ok()?;
91    let host = block_on(provider.download(&out_h)).ok()?;
92
93    let _ = provider.free(&out_h);
94    let _ = provider.free(&scaled);
95    let _ = provider.free(&yh);
96    let _ = provider.free(&xh);
97    Some(host.data)
98}
99
100fn try_vec_sub_gpu(a: &[f64], b: &[f64]) -> Option<Vec<f64>> {
101    if a.len() != b.len() {
102        return None;
103    }
104    let provider = provider()?;
105    let shape = [a.len()];
106    let ah = provider
107        .upload(&HostTensorView {
108            data: a,
109            shape: &shape,
110        })
111        .ok()?;
112    let bh = provider
113        .upload(&HostTensorView {
114            data: b,
115            shape: &shape,
116        })
117        .ok()?;
118    let out_h = block_on(provider.elem_sub(&ah, &bh)).ok()?;
119    let host = block_on(provider.download(&out_h)).ok()?;
120
121    let _ = provider.free(&out_h);
122    let _ = provider.free(&bh);
123    let _ = provider.free(&ah);
124    Some(host.data)
125}