Skip to main content

ruda_runtime/runtime/tune/
validation.rs

1//! Numerical validation shared by tensor and graph autotuning, independent of a GPU backend.
2use alloc::{format, string::String};
3
4/// Check finite values using |actual-reference| <= atol + rtol*|reference|.
5/// Scaling before subtraction avoids overflow for large, finite values of opposite sign.
6pub fn validate_finite_values(
7    expected: impl IntoIterator<Item = f64>, actual: impl IntoIterator<Item = f64>,
8    absolute: f64, relative: f64,
9) -> Result<(), String> {
10    if !absolute.is_finite() || absolute < 0.0 || !relative.is_finite() || relative < 0.0 {
11        return Err("invalid numerical tolerance".into());
12    }
13    let mut expected = expected.into_iter(); let mut actual = actual.into_iter(); let mut index = 0usize;
14    loop {
15        match (expected.next(), actual.next()) {
16            (None, None) => return Ok(()),
17            (Some(a), Some(b)) => {
18                if !a.is_finite() || !b.is_finite() { return Err(format!("non-finite autotune output at element {index}")); }
19                let scale = a.abs().max(b.abs()).max(1.0);
20                let delta = (a/scale - b/scale).abs();
21                let limit = absolute/scale + relative*(a.abs()/scale);
22                if (absolute == 0.0 && relative == 0.0 && a != b) || delta > limit { return Err(format!("autotune output mismatch at element {index}: reference={a}, actual={b}")); }
23            }
24            _ => return Err("autotune output element counts differ".into()),
25        }
26        index = index.checked_add(1).ok_or_else(|| String::from("validation element count overflow"))?;
27    }
28}
29#[cfg(test)]
30mod tests {
31    use super::*;
32    #[test] fn valid_values() { assert!(validate_finite_values([1., 0.], [1.00001, 0.00001], 1e-4, 1e-3).is_ok()); }
33    #[test] fn nan_is_never_equivalent() { assert!(validate_finite_values([f64::NAN], [f64::NAN], 1., 1.).is_err()); }
34    #[test] fn infinity_is_rejected() { assert!(validate_finite_values([f64::INFINITY], [f64::INFINITY], 0., 0.).is_err()); }
35    #[test] fn opposite_extremes_do_not_pass_via_inf_comparison() { assert!(validate_finite_values([f64::MAX], [-f64::MAX], 0., 1e-3).is_err()); }
36    #[test] fn equal_extremes_pass() { assert!(validate_finite_values([f64::MAX], [f64::MAX], 0., 0.).is_ok()); }
37    #[test] fn exact_comparison_retains_subnormals() { assert!(validate_finite_values([f64::from_bits(1)], [0.], 0., 0.).is_err()); }
38    #[test] fn lengths_must_match() { assert!(validate_finite_values([1.,2.], [1.], 0., 0.).is_err()); }
39}