Skip to main content

generate_matrix/
generate_matrix.rs

1use rand::prelude::*;
2use std::env;
3use std::fs::{self, File};
4use std::io::{BufWriter, Write};
5use std::path::Path;
6
7fn main() -> std::io::Result<()> {
8    let args: Vec<String> = env::args().collect();
9    let size = if args.len() > 1 {
10        args[1].parse::<usize>().unwrap_or(1000)
11    } else {
12        1000
13    };
14
15    let filename_arg = if args.len() > 2 {
16        &args[2]
17    } else {
18        "bench_matrix.txt"
19    };
20
21    let data_dir = Path::new("benches").join("data");
22    if !data_dir.exists() {
23        fs::create_dir_all(&data_dir)?;
24    }
25    let file_path = data_dir.join(filename_arg);
26
27    println!("Generating {}x{} matrix to {:?}", size, size, file_path);
28
29    let mut rng = rand::rng();
30
31    // Generate Tridiagonal Matrix (CSR format)
32    let mut values = Vec::new();
33    let mut col_indices = Vec::new();
34    let mut row_ptr = Vec::new();
35    row_ptr.push(0);
36
37    let mut next_lower_diag: f64 = 0.0;
38
39    for i in 0..size {
40        let mut row_sum = 0.0;
41
42        // Lower diagonal
43        if i > 0 {
44            let val = next_lower_diag;
45            values.push(val);
46            col_indices.push(i - 1);
47            row_sum += val.abs();
48        }
49
50        // Diagonal placeholder (will insert later to ensure diagonal dominance)
51        let diag_idx = values.len();
52        values.push(0.0);
53        col_indices.push(i);
54
55        // Upper diagonal
56        if i < size - 1 {
57            let val = -rng.random::<f64>();
58            next_lower_diag = val; // Store for next row's lower diagonal to ensure symmetry
59            values.push(val);
60            col_indices.push(i + 1);
61            row_sum += val.abs();
62        }
63
64        // Set diagonal value to be slightly larger than sum of off-diagonals
65        values[diag_idx] = row_sum + 1.0 + rng.random::<f64>();
66
67        row_ptr.push(values.len());
68    }
69
70    // Generate RHS vector b
71    let b: Vec<f64> = (0..size).map(|_| rng.random::<f64>()).collect();
72
73    // Write to file
74    let file = File::create(file_path)?;
75    let mut writer = BufWriter::new(file);
76
77    // Format:
78    // Size
79    // row_ptr (space separated)
80    // col_indices (space separated)
81    // values (space separated)
82    // b (space separated)
83
84    writeln!(writer, "{}", size)?;
85
86    for (i, val) in row_ptr.iter().enumerate() {
87        if i > 0 {
88            write!(writer, " ")?;
89        }
90        write!(writer, "{}", val)?;
91    }
92    writeln!(writer)?;
93
94    for (i, val) in col_indices.iter().enumerate() {
95        if i > 0 {
96            write!(writer, " ")?;
97        }
98        write!(writer, "{}", val)?;
99    }
100    writeln!(writer)?;
101
102    for (i, val) in values.iter().enumerate() {
103        if i > 0 {
104            write!(writer, " ")?;
105        }
106        write!(writer, "{}", val)?;
107    }
108    writeln!(writer)?;
109
110    for (i, val) in b.iter().enumerate() {
111        if i > 0 {
112            write!(writer, " ")?;
113        }
114        write!(writer, "{}", val)?;
115    }
116    writeln!(writer)?;
117
118    println!("Done.");
119    Ok(())
120}