Skip to main content

solve_matrix/
solve_matrix.rs

1extern crate lapack_src;
2use solver::linalg::{csr::CsrMatrix, eigen::eigen_symmetric_tridiagonal};
3use std::env;
4use std::fs::File;
5use std::io::{BufRead, BufReader};
6use std::path::Path;
7use std::time::Instant;
8
9#[allow(clippy::needless_range_loop)]
10fn read_matrix(filename: &str) -> (CsrMatrix, Vec<f64>) {
11    let path = Path::new("benches").join("data").join(filename);
12    let file = File::open(&path).unwrap_or_else(|_| panic!("Failed to open file {:?}", path));
13    let mut reader = BufReader::new(file);
14    let mut line = String::new();
15
16    // Size
17    reader.read_line(&mut line).unwrap();
18    let size: usize = line.trim().parse().unwrap();
19    line.clear();
20
21    // row_ptr
22    reader.read_line(&mut line).unwrap();
23    let row_ptr: Vec<usize> = line
24        .split_whitespace()
25        .map(|s| s.parse().unwrap())
26        .collect();
27    line.clear();
28
29    // col_indices
30    reader.read_line(&mut line).unwrap();
31    let col_indices: Vec<usize> = line
32        .split_whitespace()
33        .map(|s| s.parse().unwrap())
34        .collect();
35    line.clear();
36
37    // values
38    reader.read_line(&mut line).unwrap();
39    let values: Vec<f64> = line
40        .split_whitespace()
41        .map(|s| s.parse().unwrap())
42        .collect();
43    line.clear();
44
45    // b
46    reader.read_line(&mut line).unwrap();
47    let b: Vec<f64> = line
48        .split_whitespace()
49        .map(|s| s.parse().unwrap())
50        .collect();
51
52    // Compute diag_indices
53    let mut diag_indices = Vec::with_capacity(size);
54    for i in 0..size {
55        let start = row_ptr[i];
56        let end = row_ptr[i + 1];
57        let mut found = false;
58        for idx in start..end {
59            if col_indices[idx] == i {
60                diag_indices.push(idx);
61                found = true;
62                break;
63            }
64        }
65        if !found {
66            diag_indices.push(usize::MAX);
67        }
68    }
69
70    (
71        CsrMatrix {
72            values,
73            col_indices,
74            row_ptr,
75            size,
76            diag_indices,
77        },
78        b,
79    )
80}
81
82fn main() {
83    let args: Vec<String> = env::args().collect();
84    if args.len() < 3 {
85        eprintln!("Usage: solve_matrix <filename> <method>");
86        return;
87    }
88    let filename = &args[1];
89    let method = &args[2];
90
91    let (matrix, b) = read_matrix(filename);
92
93    if method == "sor" {
94        let mut x = vec![0.0; matrix.size];
95        let start = Instant::now();
96        // Use same params as bench: tol=1e-6, max_iter=10000, omega=1.5
97        let iters = matrix.solve_sor(&b, &mut x, 1e-6, 10000, 1.5);
98        let duration = start.elapsed();
99
100        eprintln!(
101            "Rust SOR time: {:.4} ms (iters: {})",
102            duration.as_secs_f64() * 1000.0,
103            iters
104        );
105
106        // Print result to stdout for Python to read
107        for val in x {
108            print!("{} ", val);
109        }
110        println!();
111    } else if method == "eigen" {
112        // Extract diagonal and off-diagonal
113        let mut d = vec![0.0; matrix.size];
114        let mut e = vec![0.0; matrix.size];
115
116        for i in 0..matrix.size {
117            let row_start = matrix.row_ptr[i];
118            let row_end = matrix.row_ptr[i + 1];
119
120            for idx in row_start..row_end {
121                let col = matrix.col_indices[idx];
122                if col == i {
123                    d[i] = matrix.values[idx];
124                } else if col == i + 1 {
125                    e[i] = matrix.values[idx];
126                }
127            }
128        }
129
130        let start = Instant::now();
131        let (eigenvalues, _) = eigen_symmetric_tridiagonal(d, e);
132        let duration = start.elapsed();
133
134        eprintln!("Rust Eigen time: {:.4} ms", duration.as_secs_f64() * 1000.0);
135
136        // Print eigenvalues to stdout
137        for val in eigenvalues {
138            print!("{} ", val);
139        }
140        println!();
141    } else if method == "faer" {
142        use faer::{Mat, Side};
143
144        // Convert CSR to Dense Mat
145        let mut dense_mat = Mat::<f64>::zeros(matrix.size, matrix.size);
146        for i in 0..matrix.size {
147            let row_start = matrix.row_ptr[i];
148            let row_end = matrix.row_ptr[i + 1];
149            for idx in row_start..row_end {
150                let col = matrix.col_indices[idx];
151                let val = matrix.values[idx];
152                dense_mat[(i, col)] = val;
153            }
154        }
155
156        let start = Instant::now();
157        // High-level API for eigenvalues only
158        // Assuming signature: selfadjoint_eigenvalues(Side)
159        let s = dense_mat.selfadjoint_eigenvalues(Side::Lower);
160        let duration = start.elapsed();
161
162        eprintln!("Rust Faer time: {:.4} ms", duration.as_secs_f64() * 1000.0);
163
164        for val in s.iter() {
165            print!("{} ", val);
166        }
167        println!();
168    } else if method == "lapack" {
169        use lapack::dstev;
170
171        // Extract diagonal and off-diagonal
172        let mut d = vec![0.0; matrix.size];
173        let mut e = vec![0.0; matrix.size - 1]; // dstev expects e to be len n-1
174
175        for i in 0..matrix.size {
176            let row_start = matrix.row_ptr[i];
177            let row_end = matrix.row_ptr[i + 1];
178
179            for idx in row_start..row_end {
180                let col = matrix.col_indices[idx];
181                if col == i {
182                    d[i] = matrix.values[idx];
183                } else if col == i + 1 {
184                    e[i] = matrix.values[idx];
185                }
186            }
187        }
188
189        let n = matrix.size as i32;
190        // 'V' means compute eigenvectors, 'N' means eigenvalues only
191        // We want eigenvalues only for fair comparison with faer/scipy eigh_tridiagonal default?
192        // Scipy eigh_tridiagonal returns (w, v) by default.
193        // My custom eigen returns (w, v).
194        // So let's compute eigenvectors too ('V').
195        // But wait, bench_solvers.py only compares eigenvalues (w).
196        // And faer benchmark above used selfadjoint_eigenvalues (only w).
197        // Let's stick to eigenvalues only ('N') for speed if we only compare w,
198        // BUT bench_solvers.py compares w.
199        // However, if I want to beat Scipy which computes both (usually), I should check what Scipy does.
200        // Scipy eigh_tridiagonal(d, e) returns w, v.
201        // So I should use 'V'.
202        // But wait, if I use 'N', I don't need Z.
203
204        // Let's use 'N' (eigenvalues only) first to see raw speed of the solver,
205        // as we are comparing against faer which was only eigenvalues.
206        // And bench_solvers.py only checks w.
207
208        let jobz = b'N';
209        let mut z = vec![0.0; 1]; // Not referenced if jobz == 'N'
210        let ldz = 1;
211        let mut work = vec![0.0; std::cmp::max(1, 2 * matrix.size - 2)];
212        let mut info = 0;
213
214        let start = Instant::now();
215        unsafe {
216            dstev(jobz, n, &mut d, &mut e, &mut z, ldz, &mut work, &mut info);
217        }
218        let duration = start.elapsed();
219
220        if info != 0 {
221            eprintln!("LAPACK dstev failed with info={}", info);
222        } else {
223            eprintln!(
224                "Rust LAPACK time: {:.4} ms",
225                duration.as_secs_f64() * 1000.0
226            );
227            for val in d.iter() {
228                print!("{} ", val);
229            }
230            println!();
231        }
232    } else if method == "lapack_solve" {
233        use lapack::dgtsv;
234
235        // Extract diagonal and off-diagonal
236        let mut d = vec![0.0; matrix.size];
237        let mut du = vec![0.0; matrix.size - 1];
238        let mut dl = vec![0.0; matrix.size - 1];
239
240        for i in 0..matrix.size {
241            let row_start = matrix.row_ptr[i];
242            let row_end = matrix.row_ptr[i + 1];
243
244            for idx in row_start..row_end {
245                let col = matrix.col_indices[idx];
246                if col == i {
247                    d[i] = matrix.values[idx];
248                } else if col == i + 1 {
249                    du[i] = matrix.values[idx];
250                    dl[i] = matrix.values[idx]; // Symmetric
251                }
252            }
253        }
254
255        let n = matrix.size as i32;
256        let nrhs = 1;
257        let ldb = n;
258        let mut info = 0;
259
260        // b is consumed and replaced by solution
261        let mut x = b.clone();
262
263        let start = Instant::now();
264        unsafe {
265            dgtsv(n, nrhs, &mut dl, &mut d, &mut du, &mut x, ldb, &mut info);
266        }
267        let duration = start.elapsed();
268
269        if info != 0 {
270            eprintln!("LAPACK dgtsv failed with info={}", info);
271        } else {
272            eprintln!(
273                "Rust LAPACK Solve time: {:.4} ms",
274                duration.as_secs_f64() * 1000.0
275            );
276            for val in x {
277                print!("{} ", val);
278            }
279            println!();
280        }
281    } else {
282        eprintln!("Unknown method: {}", method);
283    }
284}