use std::env;
use std::error::Error;
use std::fs::{self, File};
use std::io::{BufRead, BufReader};
use std::path::{Path, PathBuf};
use std::time::Instant;
use hybit::{
analyze_csr32, pcg_with_workspace, read_matrix_market, recommend_rigid_body_aggregate_nodes,
Csr32Matrix, ParallelCsr32Operator, ParallelRigidBodyTwoLevelPreconditioner, PcgWorkspace,
Preconditioner, RigidBodyAggregation, RigidBodyTwoLevelBlockJacobiPreconditioner,
SolverOptions,
};
#[derive(Debug)]
struct Args {
matrix: PathBuf,
coordinates: PathBuf,
rhs: Option<PathBuf>,
relative_tolerance: f64,
max_iterations: usize,
target_coarse_dimension: usize,
aggregation: RigidBodyAggregation,
kernel_repeats: usize,
}
impl Args {
fn parse() -> Result<Self, Box<dyn Error>> {
let mut matrix = None;
let mut coordinates = None;
let mut rhs = None;
let mut relative_tolerance: f64 = 1.0e-8;
let mut max_iterations = 3000usize;
let mut target_coarse_dimension = 1536usize;
let mut aggregation = RigidBodyAggregation::Graph;
let mut kernel_repeats = 20usize;
let mut it = env::args().skip(1);
while let Some(arg) = it.next() {
match arg.as_str() {
"--matrix" => matrix = Some(PathBuf::from(next_value(&mut it, "--matrix")?)),
"--coords" => coordinates = Some(PathBuf::from(next_value(&mut it, "--coords")?)),
"--rhs" => rhs = Some(PathBuf::from(next_value(&mut it, "--rhs")?)),
"--tol" => relative_tolerance = next_value(&mut it, "--tol")?.parse()?,
"--max-iters" => max_iterations = next_value(&mut it, "--max-iters")?.parse()?,
"--target-coarse-dim" => {
target_coarse_dimension = next_value(&mut it, "--target-coarse-dim")?.parse()?
}
"--kernel-repeats" => {
kernel_repeats = next_value(&mut it, "--kernel-repeats")?.parse()?
}
"--aggregation" => {
aggregation = match next_value(&mut it, "--aggregation")?
.to_ascii_lowercase()
.as_str()
{
"contiguous" => RigidBodyAggregation::Contiguous,
"graph" => RigidBodyAggregation::Graph,
other => {
return Err(format!(
"unknown aggregation '{other}'; use contiguous or graph"
)
.into())
}
}
}
"-h" | "--help" => {
print_usage();
std::process::exit(0);
}
other if !other.starts_with('-') && matrix.is_none() => {
matrix = Some(PathBuf::from(other))
}
other => return Err(format!("unknown argument '{other}'").into()),
}
}
let matrix = matrix.ok_or("missing matrix path; use --matrix FILE.mtx")?;
let coordinates = coordinates.unwrap_or_else(|| matrix.with_extension("coords"));
if !relative_tolerance.is_finite() || relative_tolerance <= 0.0 {
return Err("--tol must be finite and > 0".into());
}
if max_iterations == 0 {
return Err("--max-iters must be > 0".into());
}
if target_coarse_dimension < 6 {
return Err("--target-coarse-dim must be >= 6".into());
}
if kernel_repeats == 0 {
return Err("--kernel-repeats must be > 0".into());
}
Ok(Self {
matrix,
coordinates,
rhs,
relative_tolerance,
max_iterations,
target_coarse_dimension,
aggregation,
kernel_repeats,
})
}
}
fn next_value<I: Iterator<Item = String>>(
it: &mut I,
flag: &str,
) -> Result<String, Box<dyn Error>> {
it.next()
.ok_or_else(|| format!("missing value after {flag}").into())
}
fn print_usage() {
println!("HyBIT serial vs parallel rigid-body preconditioner benchmark");
println!("Usage: fem_structural_precond_parallel --matrix K.mtx [--coords K.coords] [--rhs b.txt] [--tol 1e-8] [--max-iters 3000] [--target-coarse-dim 1536] [--aggregation graph|contiguous] [--kernel-repeats 20]");
println!("The fine-grid operator is ParallelCsr32Operator in both PCG runs.");
println!("Set RAYON_NUM_THREADS before launch to control the shared Rayon pool.");
}
fn read_coordinates(path: &Path) -> Result<Vec<[f64; 3]>, Box<dyn Error>> {
let file = File::open(path)?;
let reader = BufReader::new(file);
let mut expected = None::<usize>;
let mut coordinates = Vec::new();
for (line_no, line) in reader.lines().enumerate() {
let line = line?;
let text = line.trim();
if text.is_empty() || text.starts_with('#') {
continue;
}
if expected.is_none() {
expected = Some(text.parse::<usize>().map_err(|e| {
format!(
"{}:{}: invalid coordinate count: {e}",
path.display(),
line_no + 1
)
})?);
coordinates.reserve(expected.unwrap());
continue;
}
let fields: Vec<&str> = text.split_whitespace().collect();
if fields.len() != 3 {
return Err(format!(
"{}:{}: expected three coordinates",
path.display(),
line_no + 1
)
.into());
}
let xyz = [
fields[0].parse::<f64>()?,
fields[1].parse::<f64>()?,
fields[2].parse::<f64>()?,
];
if xyz.iter().any(|v| !v.is_finite()) {
return Err(
format!("{}:{}: non-finite coordinate", path.display(), line_no + 1).into(),
);
}
coordinates.push(xyz);
}
let expected =
expected.ok_or_else(|| format!("{}: missing coordinate count", path.display()))?;
if coordinates.len() != expected {
return Err(format!(
"{}: coordinate count mismatch: header says {}, read {}",
path.display(),
expected,
coordinates.len()
)
.into());
}
Ok(coordinates)
}
fn load_rhs(path: &Path, n: usize) -> Result<Vec<f64>, Box<dyn Error>> {
let text = fs::read_to_string(path)?;
let mut values = Vec::with_capacity(n);
for (line_index, raw_line) in text.lines().enumerate() {
let line = raw_line.trim_start_matches('\u{feff}').trim();
if line.is_empty() || line.starts_with('#') || line.starts_with('%') {
continue;
}
for (token_index, token) in line.split_whitespace().enumerate() {
let value = token.parse::<f64>().map_err(|e| {
format!(
"{}: invalid RHS float at line {}, token {}: {:?} ({e})",
path.display(),
line_index + 1,
token_index + 1,
token
)
})?;
if !value.is_finite() {
return Err(format!(
"{}: non-finite RHS value at line {}, token {}",
path.display(),
line_index + 1,
token_index + 1
)
.into());
}
values.push(value);
}
}
if values.len() != n {
return Err(format!(
"{}: RHS length mismatch: expected {n}, got {}",
path.display(),
values.len()
)
.into());
}
Ok(values)
}
fn norm2(x: &[f64]) -> f64 {
x.iter().map(|v| v * v).sum::<f64>().sqrt()
}
fn mib(bytes: usize) -> f64 {
bytes as f64 / (1024.0 * 1024.0)
}
fn ms_per(seconds: f64, repeats: usize) -> f64 {
seconds * 1.0e3 / repeats as f64
}
fn verified_relative_residual(
a: &Csr32Matrix,
b: &[f64],
x: &[f64],
) -> Result<f64, Box<dyn Error>> {
let ax = a.spmv(x)?;
let rr = b
.iter()
.zip(&ax)
.map(|(&bi, &ai)| {
let r = bi - ai;
r * r
})
.sum::<f64>()
.sqrt();
let bn = norm2(b);
Ok(if bn == 0.0 { rr } else { rr / bn })
}
fn main() -> Result<(), Box<dyn Error>> {
let args = Args::parse()?;
println!(
"HyBIT {} serial vs parallel rigid-body preconditioner benchmark",
env!("CARGO_PKG_VERSION")
);
println!("matrix : {}", args.matrix.display());
println!("coordinates : {}", args.coordinates.display());
let load_start = Instant::now();
let (matrix, mm) = read_matrix_market(&args.matrix)?;
let matrix_load = load_start.elapsed().as_secs_f64();
let coord_start = Instant::now();
let coordinates = read_coordinates(&args.coordinates)?;
let coord_load = coord_start.elapsed().as_secs_f64();
let profile = analyze_csr32(&matrix)?;
println!(
"Matrix Market : {:?}, {} input entries -> {} CSR nnz",
mm.symmetry, mm.input_entries, mm.csr_nnz
);
println!("dimensions : {} x {}", profile.nrows, profile.ncols);
println!("nnz : {}", profile.nnz);
println!(
"CSR storage : {:.3} MiB",
mib(matrix.storage_bytes())
);
println!("matrix load : {:.3} ms", matrix_load * 1.0e3);
println!("coordinate nodes : {}", coordinates.len());
println!("coordinate load : {:.3} ms", coord_load * 1.0e3);
println!("aggregation : {:?}", args.aggregation);
println!("target coarse dim : {}", args.target_coarse_dimension);
println!("kernel repeats : {}", args.kernel_repeats);
if !profile.square || !profile.full_diagonal || !profile.positive_diagonal {
return Err(
"structural PCG benchmark requires a square matrix with a complete positive diagonal"
.into(),
);
}
if matrix.nrows() != coordinates.len() * 3 {
return Err(format!(
"matrix/coordinate mismatch: {} matrix rows != {} coordinate nodes * 3",
matrix.nrows(),
coordinates.len()
)
.into());
}
let b = if let Some(path) = args.rhs.as_deref() {
println!("RHS : {}", path.display());
load_rhs(path, matrix.nrows())?
} else {
println!("RHS : generated as b=A*1 (known exact solution)");
matrix.spmv(&vec![1.0; matrix.ncols()])?
};
let aggregate_nodes =
recommend_rigid_body_aggregate_nodes(coordinates.len(), args.target_coarse_dimension)?;
let setup_start = Instant::now();
let preconditioner = match args.aggregation {
RigidBodyAggregation::Graph => {
RigidBodyTwoLevelBlockJacobiPreconditioner::from_csr32_graph(
&matrix,
&coordinates,
aggregate_nodes,
)?
}
RigidBodyAggregation::Contiguous => RigidBodyTwoLevelBlockJacobiPreconditioner::from_csr32(
&matrix,
&coordinates,
aggregate_nodes,
)?,
RigidBodyAggregation::Auto => unreachable!(),
};
let setup_seconds = setup_start.elapsed().as_secs_f64();
let parallel_preconditioner = ParallelRigidBodyTwoLevelPreconditioner::new(&preconditioner)?;
let parallel_operator = ParallelCsr32Operator::new(&matrix);
println!("aggregate target : {} nodes", aggregate_nodes);
println!("aggregate count : {}", preconditioner.aggregate_count());
println!("coarse dimension : {}", preconditioner.coarse_dimension());
println!(
"prec storage : {:.3} MiB",
mib(preconditioner.factor_bytes())
);
println!(
"parallel index : {:.3} MiB",
mib(parallel_preconditioner.index_storage_bytes())
);
println!("setup : {:.3} ms", setup_seconds * 1.0e3);
println!(
"Rayon threads : {}",
parallel_preconditioner.rayon_threads()
);
let mut zs = vec![0.0; matrix.nrows()];
let mut zp = vec![0.0; matrix.nrows()];
preconditioner.apply(&b, &mut zs)?;
parallel_preconditioner.apply(&b, &mut zp)?;
let scale = zs.iter().fold(1.0f64, |m, &v| m.max(v.abs()));
let max_diff = zs
.iter()
.zip(&zp)
.fold(0.0f64, |m, (&a, &b)| m.max((a - b).abs()));
if max_diff > 1.0e-11 * scale {
return Err(format!(
"serial/parallel preconditioner mismatch: max diff={max_diff:e}, scale={scale:e}"
)
.into());
}
let serial_prec_start = Instant::now();
for _ in 0..args.kernel_repeats {
preconditioner.apply(&b, &mut zs)?;
}
let serial_prec = serial_prec_start.elapsed().as_secs_f64();
let parallel_prec_start = Instant::now();
for _ in 0..args.kernel_repeats {
parallel_preconditioner.apply(&b, &mut zp)?;
}
let parallel_prec = parallel_prec_start.elapsed().as_secs_f64();
println!();
println!("Preconditioner microbenchmark");
println!(
"serial / apply : {:.3} ms",
ms_per(serial_prec, args.kernel_repeats)
);
println!(
"parallel / apply : {:.3} ms",
ms_per(parallel_prec, args.kernel_repeats)
);
println!(
"precond speedup : {:.3}x",
serial_prec / parallel_prec.max(f64::MIN_POSITIVE)
);
let options = SolverOptions {
relative_tolerance: args.relative_tolerance,
absolute_tolerance: 0.0,
max_iterations: args.max_iterations,
};
let mut xs = vec![0.0; matrix.ncols()];
let mut ws = PcgWorkspace::new(matrix.nrows());
let serial_solve_start = Instant::now();
let serial_out = pcg_with_workspace(
¶llel_operator,
&preconditioner,
&b,
&mut xs,
options,
&mut ws,
)?;
let serial_solve = serial_solve_start.elapsed().as_secs_f64();
let serial_verified = verified_relative_residual(&matrix, &b, &xs)?;
let mut xp = vec![0.0; matrix.ncols()];
let mut wp = PcgWorkspace::new(matrix.nrows());
let parallel_solve_start = Instant::now();
let parallel_out = pcg_with_workspace(
¶llel_operator,
¶llel_preconditioner,
&b,
&mut xp,
options,
&mut wp,
)?;
let parallel_solve = parallel_solve_start.elapsed().as_secs_f64();
let parallel_verified = verified_relative_residual(&matrix, &b, &xp)?;
println!();
println!("[1/2] Parallel CSR + serial preconditioner");
println!("status : {:?}", serial_out.status);
println!("iterations : {}", serial_out.iterations);
println!(
"reported residual : {:.6e}",
serial_out.final_residual / norm2(&b).max(f64::MIN_POSITIVE)
);
println!("verified residual : {:.6e}", serial_verified);
println!("solve time : {:.3} ms", serial_solve * 1.0e3);
println!();
println!("[2/2] Parallel CSR + parallel preconditioner");
println!("status : {:?}", parallel_out.status);
println!("iterations : {}", parallel_out.iterations);
println!(
"reported residual : {:.6e}",
parallel_out.final_residual / norm2(&b).max(f64::MIN_POSITIVE)
);
println!("verified residual : {:.6e}", parallel_verified);
println!("solve time : {:.3} ms", parallel_solve * 1.0e3);
println!();
println!("Comparison");
println!(
"iteration ratio : {:.3} (parallel prec / serial prec)",
parallel_out.iterations as f64 / serial_out.iterations.max(1) as f64
);
println!(
"solve-time ratio : {:.3} (parallel prec / serial prec)",
parallel_solve / serial_solve.max(f64::MIN_POSITIVE)
);
println!(
"PCG speedup : {:.3}x",
serial_solve / parallel_solve.max(f64::MIN_POSITIVE)
);
if !serial_verified.is_finite() || !parallel_verified.is_finite() {
return Err("non-finite independently verified residual".into());
}
Ok(())
}