use crate::da::Da;
use crate::error::{codes, dace_panic};
use crate::norm::Interval;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ToleranceKind {
Absolute,
Relative,
}
#[derive(Debug, Clone)]
pub struct AdsConfig {
pub tolerances: Vec<f64>,
pub tolerance_kind: ToleranceKind,
pub targets: Vec<usize>,
pub max_splits_per_var: u32,
pub max_leaves: usize,
}
impl Default for AdsConfig {
fn default() -> Self {
AdsConfig {
tolerances: vec![1e-8],
tolerance_kind: ToleranceKind::Absolute,
targets: Vec::new(),
max_splits_per_var: 32,
max_leaves: 1024,
}
}
}
#[derive(Debug, Clone)]
pub struct AdsLeaf {
pub center: Vec<f64>,
pub half_width: Vec<f64>,
pub values: Vec<Da>,
pub bounds: Vec<Interval>,
pub met: bool,
}
#[derive(Debug, Clone)]
pub struct AdsResult {
pub leaves: Vec<AdsLeaf>,
pub splits_per_var: Vec<u32>,
pub met_leaves: usize,
}
pub fn split<F>(f: F, domain: &[Interval], config: &AdsConfig) -> AdsResult
where
F: Fn(&[Da]) -> Vec<Da>,
{
validate_config(config);
if domain.is_empty() {
dace_panic(
codes::OUT_OF_DOMAIN,
"ADS: the domain must cover at least one variable",
);
}
if domain
.iter()
.any(|iv| iv.lo.is_nan() || iv.hi.is_nan() || iv.lo > iv.hi)
{
dace_panic(
codes::OUT_OF_DOMAIN,
"ADS: domain interval with lo > hi or a NaN endpoint",
);
}
let center: Vec<f64> = domain.iter().map(|iv| 0.5 * (iv.lo + iv.hi)).collect();
let half_width: Vec<f64> = domain.iter().map(|iv| 0.5 * (iv.hi - iv.lo)).collect();
let mut leaves: Vec<AdsLeaf> = Vec::new();
let mut splits_per_var = vec![0u32; half_width.len()];
let mut met_leaves = 0usize;
let mut stack = vec![Node {
center,
half_width,
path_splits: vec![0; domain.len()],
}];
while let Some(Node {
center,
half_width,
path_splits,
}) = stack.pop()
{
let values = expand(&f, ¢er, &half_width);
let targets = resolve_targets(&config.targets, values.len());
if config.tolerances.len() != targets.len() {
dace_panic(
codes::OUT_OF_DOMAIN,
"ADS: tolerances must align with targets (or all components when targets is empty)",
);
}
let bounds: Vec<Interval> = values.iter().map(Da::bound).collect();
let met = targets
.iter()
.zip(&config.tolerances)
.all(|(&j, &tol)| within_tol(bounds[j], tol, config.tolerance_kind));
let direction = if met || leaves.len() + stack.len() + 2 > config.max_leaves {
None
} else {
choose_direction(&values, &targets, &half_width, &path_splits, config)
};
if let Some(dir) = direction {
splits_per_var[dir] += 1;
let h = half_width[dir] / 2.0;
let mut upper = Node {
center: center.clone(),
half_width: half_width.clone(),
path_splits: path_splits.clone(),
};
upper.center[dir] += h;
upper.half_width[dir] = h;
upper.path_splits[dir] += 1;
let mut lower = Node {
center,
half_width,
path_splits,
};
lower.center[dir] -= h;
lower.half_width[dir] = h;
lower.path_splits[dir] += 1;
stack.push(upper);
stack.push(lower);
} else {
met_leaves += usize::from(met);
leaves.push(AdsLeaf {
center,
half_width,
values,
bounds,
met,
});
}
}
AdsResult {
leaves,
splits_per_var,
met_leaves,
}
}
struct Node {
center: Vec<f64>,
half_width: Vec<f64>,
path_splits: Vec<u32>,
}
fn choose_direction(
values: &[Da],
targets: &[usize],
half_width: &[f64],
path_splits: &[u32],
config: &AdsConfig,
) -> Option<usize> {
let mut best: Option<(usize, f64)> = None;
for (i, &h) in half_width.iter().enumerate() {
if path_splits[i] >= config.max_splits_per_var || h <= 0.0 || h.is_nan() {
continue;
}
let var = i as u32 + 1;
let mut contrib = 0.0f64;
for &j in targets {
let b = values[j].deriv(var).bound();
contrib = contrib.max(b.lo.abs()).max(b.hi.abs());
}
if contrib > 0.0 && best.is_none_or(|(_, c)| contrib > c) {
best = Some((i, contrib));
}
}
best.map(|(i, _)| i)
}
fn expand<F>(f: &F, center: &[f64], half_width: &[f64]) -> Vec<Da>
where
F: Fn(&[Da]) -> Vec<Da>,
{
let vars: Vec<Da> = (0..center.len())
.map(|i| {
let var = i as u32 + 1;
Da::variable(var).translate_variable(var, half_width[i], center[i])
})
.collect();
let values = f(&vars);
if values.is_empty() {
dace_panic(codes::OUT_OF_DOMAIN, "ADS: the map returned no components");
}
values
}
fn resolve_targets(targets: &[usize], nvals: usize) -> Vec<usize> {
if targets.is_empty() {
(0..nvals).collect()
} else {
if targets.iter().any(|&j| j >= nvals) {
dace_panic(codes::OUT_OF_DOMAIN, "ADS: target index out of range");
}
targets.to_vec()
}
}
fn within_tol(b: Interval, tol: f64, kind: ToleranceKind) -> bool {
let width = b.hi - b.lo;
match kind {
ToleranceKind::Absolute => width <= tol,
ToleranceKind::Relative => width <= tol * b.lo.abs().max(b.hi.abs()),
}
}
fn validate_config(config: &AdsConfig) {
if config.max_leaves == 0 {
dace_panic(codes::OUT_OF_DOMAIN, "ADS: max_leaves must be at least 1");
}
if config.tolerances.is_empty() {
dace_panic(codes::OUT_OF_DOMAIN, "ADS: tolerances must not be empty");
}
if config
.tolerances
.iter()
.any(|&t| !t.is_finite() || t <= 0.0)
{
dace_panic(
codes::OUT_OF_DOMAIN,
"ADS: tolerances must be positive and finite",
);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::CONTEXT_LOCK;
fn cfg(tolerances: Vec<f64>) -> AdsConfig {
AdsConfig {
tolerances,
..Default::default()
}
}
#[test]
fn constant_map_is_one_met_leaf() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(4, 2).unwrap();
let result = split(
|_: &[Da]| vec![Da::constant(0.5)],
&[
Interval { lo: -1.0, hi: 1.0 },
Interval { lo: -2.0, hi: 2.0 },
],
&cfg(vec![1e-6]),
);
assert_eq!(result.leaves.len(), 1);
assert_eq!(result.met_leaves, 1);
assert!(result.leaves[0].met);
assert_eq!(result.leaves[0].bounds[0], Interval { lo: 0.5, hi: 0.5 });
assert_eq!(result.leaves[0].center, vec![0.0, 0.0]);
assert_eq!(result.leaves[0].half_width, vec![1.0, 2.0]);
assert_eq!(result.splits_per_var, vec![0, 0]);
}
#[test]
fn linear_map_meets_tolerance_without_splitting() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(4, 1).unwrap();
let result = split(
|x: &[Da]| vec![0.1 * x[0].clone()],
&[Interval { lo: -1.0, hi: 1.0 }],
&cfg(vec![0.3]),
);
assert_eq!(result.leaves.len(), 1);
assert!(result.leaves[0].met);
assert_eq!(result.met_leaves, 1);
let b = result.leaves[0].bounds[0];
assert!((b.hi - b.lo - 0.2).abs() < 1e-15, "width {b:?}");
}
#[test]
fn relative_tolerance_uses_bound_scale() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(4, 1).unwrap();
let domain = &[Interval { lo: -1.0, hi: 1.0 }];
let map = |x: &[Da]| vec![Da::constant(1000.0) + 0.001 * x[0].clone()];
let abs = AdsConfig {
tolerances: vec![1e-3],
max_splits_per_var: 0,
..Default::default()
};
let r = split(map, domain, &abs);
assert_eq!(r.leaves.len(), 1);
assert!(!r.leaves[0].met);
assert_eq!(r.met_leaves, 0);
let rel = AdsConfig {
tolerances: vec![1e-3],
tolerance_kind: ToleranceKind::Relative,
max_splits_per_var: 0,
..Default::default()
};
let r = split(map, domain, &rel);
assert_eq!(r.leaves.len(), 1);
assert!(r.leaves[0].met);
assert_eq!(r.met_leaves, 1);
}
#[test]
fn sin_on_big_box_splits_with_rigorous_leaf_bounds() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(16, 1).unwrap();
let cfg = AdsConfig {
tolerances: vec![1e-2],
..Default::default()
};
let result = split(
|x: &[Da]| vec![crate::elementary::sin(&x[0])],
&[Interval { lo: -1.5, hi: 1.5 }],
&cfg,
);
assert!(result.leaves.len() > 1, "a near-half-period box must split");
assert_eq!(result.met_leaves, result.leaves.len());
assert!(result.splits_per_var[0] > 0);
for leaf in &result.leaves {
let (lo, hi) = (
leaf.center[0] - leaf.half_width[0],
leaf.center[0] + leaf.half_width[0],
);
let (true_lo, true_hi) = sin_range(lo, hi);
assert!(
leaf.bounds[0].lo <= true_lo + 1e-10,
"leaf bound {leaf:?} below true range on [{lo}, {hi}]"
);
assert!(
leaf.bounds[0].hi >= true_hi - 1e-10,
"leaf bound {leaf:?} above true range on [{lo}, {hi}]"
);
}
}
fn sin_range(lo: f64, hi: f64) -> (f64, f64) {
let mut lo_v = lo.sin().min(hi.sin());
let mut hi_v = lo.sin().max(hi.sin());
let k_min = ((lo - std::f64::consts::FRAC_PI_2) / std::f64::consts::PI).ceil() as i64;
let k_max = ((hi - std::f64::consts::FRAC_PI_2) / std::f64::consts::PI).floor() as i64;
for k in k_min..=k_max {
let x = std::f64::consts::FRAC_PI_2 + k as f64 * std::f64::consts::PI;
if (lo..=hi).contains(&x) {
let v = x.sin();
lo_v = lo_v.min(v);
hi_v = hi_v.max(v);
}
}
(lo_v, hi_v)
}
#[test]
fn leaf_budget_caps_the_tree_and_flags_unmet() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(16, 1).unwrap();
let cfg = AdsConfig {
tolerances: vec![1e-6],
max_leaves: 8,
..Default::default()
};
let result = split(
|x: &[Da]| vec![crate::elementary::sin(&x[0])],
&[Interval { lo: -1.5, hi: 1.5 }],
&cfg,
);
assert_eq!(result.leaves.len(), 8);
assert_eq!(result.met_leaves, 0);
assert!(result.leaves.iter().all(|l| !l.met));
assert!((result.leaves[0].center[0] - result.leaves[0].half_width[0] + 1.5).abs() < 1e-12);
for w in result.leaves.windows(2) {
let (a, b) = (&w[0], &w[1]);
let a_hi = a.center[0] + a.half_width[0];
let b_lo = b.center[0] - b.half_width[0];
assert!((a_hi - b_lo).abs() < 1e-12, "leaves must tile in order");
}
let last = result.leaves.last().unwrap();
assert!((last.center[0] + last.half_width[0] - 1.5).abs() < 1e-12);
}
#[test]
fn per_direction_cap_limits_half_width() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(16, 1).unwrap();
let cfg = AdsConfig {
tolerances: vec![1e-9],
max_splits_per_var: 3,
..Default::default()
};
let result = split(
|x: &[Da]| vec![crate::elementary::sin(&x[0])],
&[Interval { lo: -1.5, hi: 1.5 }],
&cfg,
);
assert_eq!(result.leaves.len(), 8);
assert_eq!(result.splits_per_var, vec![7]);
assert_eq!(result.met_leaves, 0);
for leaf in &result.leaves {
assert!(
(leaf.half_width[0] - 1.5 / 8.0).abs() < 1e-12,
"cap must bound the half-width: {leaf:?}"
);
}
}
#[test]
fn direction_choice_follows_the_largest_contribution() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(8, 2).unwrap();
let cfg = AdsConfig {
tolerances: vec![1.0],
..Default::default()
};
let result = split(
|x: &[Da]| vec![0.001 * x[0].clone() + 10.0 * x[1].clone()],
&[
Interval { lo: -1.0, hi: 1.0 },
Interval { lo: -1.0, hi: 1.0 },
],
&cfg,
);
assert_eq!(result.splits_per_var, vec![0, 31]);
assert_eq!(result.leaves.len(), 32);
assert_eq!(result.met_leaves, 32);
for leaf in &result.leaves {
assert!((leaf.half_width[0] - 1.0).abs() < 1e-12);
assert!((leaf.half_width[1] - 1.0 / 32.0).abs() < 1e-12);
}
}
#[test]
fn invalid_inputs_panic_with_code_650() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(4, 2).unwrap();
fn expect_code(f: impl FnOnce() -> AdsResult + std::panic::UnwindSafe, code: u32) {
let err = std::panic::catch_unwind(f).expect_err("must panic");
let e = err
.downcast_ref::<crate::error::DaceError>()
.expect("DaceError payload");
assert_eq!(e.code, code, "{}", e);
}
let unit = &[Interval { lo: -1.0, hi: 1.0 }];
let ok = || AdsConfig {
tolerances: vec![1e-3],
..Default::default()
};
let lin = |x: &[Da]| vec![x[0].clone()];
expect_code(|| split(lin, &[], &ok()), codes::OUT_OF_DOMAIN);
expect_code(
|| split(lin, &[Interval { lo: 1.0, hi: -1.0 }], &ok()),
codes::OUT_OF_DOMAIN,
);
expect_code(
|| {
split(
lin,
&[Interval {
lo: f64::NAN,
hi: 1.0,
}],
&ok(),
)
},
codes::OUT_OF_DOMAIN,
);
expect_code(
|| {
split(
lin,
unit,
&AdsConfig {
tolerances: vec![],
..Default::default()
},
)
},
codes::OUT_OF_DOMAIN,
);
for bad in [0.0, -1e-3, f64::NAN, f64::INFINITY] {
expect_code(
|| {
split(
lin,
unit,
&AdsConfig {
tolerances: vec![bad],
..Default::default()
},
)
},
codes::OUT_OF_DOMAIN,
);
}
expect_code(
|| {
split(
lin,
unit,
&AdsConfig {
max_leaves: 0,
..Default::default()
},
)
},
codes::OUT_OF_DOMAIN,
);
expect_code(
|| {
split(
lin,
unit,
&AdsConfig {
tolerances: vec![1e-3, 1e-3],
..Default::default()
},
)
},
codes::OUT_OF_DOMAIN,
);
expect_code(
|| {
split(
lin,
unit,
&AdsConfig {
tolerances: vec![1e-3],
targets: vec![1],
..Default::default()
},
)
},
codes::OUT_OF_DOMAIN,
);
expect_code(
|| split(|_: &[Da]| vec![], unit, &ok()),
codes::OUT_OF_DOMAIN,
);
}
#[test]
fn identical_runs_produce_identical_results() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(12, 2).unwrap();
let cfg = AdsConfig {
tolerances: vec![1e-3, 1e-3],
..Default::default()
};
let domain = &[
Interval { lo: -1.2, hi: 1.0 },
Interval { lo: -0.8, hi: 1.4 },
];
let map = |x: &[Da]| {
vec![
crate::elementary::sin(&x[0]) * crate::elementary::cos(&x[1]),
x[0].clone() * x[1].clone(),
]
};
let a = split(map, domain, &cfg);
let b = split(map, domain, &cfg);
assert_eq!(a.leaves.len(), b.leaves.len());
assert_eq!(a.splits_per_var, b.splits_per_var);
assert_eq!(a.met_leaves, b.met_leaves);
for (la, lb) in a.leaves.iter().zip(&b.leaves) {
assert_eq!(la.center, lb.center);
assert_eq!(la.half_width, lb.half_width);
assert_eq!(la.bounds, lb.bounds);
assert_eq!(la.met, lb.met);
}
}
#[test]
fn zero_width_variables_never_split() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(8, 2).unwrap();
let cfg = AdsConfig {
tolerances: vec![1e-2],
..Default::default()
};
let result = split(
|x: &[Da]| vec![x[0].clone() * crate::elementary::sin(&x[1])],
&[
Interval { lo: 0.5, hi: 0.5 },
Interval { lo: -1.5, hi: 1.5 },
],
&cfg,
);
assert!(result.leaves.len() > 1);
assert_eq!(result.splits_per_var[0], 0);
assert!(result.leaves.iter().all(|l| l.half_width[0] == 0.0));
assert_eq!(result.met_leaves, result.leaves.len());
}
}