#![cfg(feature = "parallel")]
use crate::deriv::log::{DerivationLog, DerivedExpr};
use crate::kernel::{ExprData, ExprId, ExprPool};
use crate::simplify::engine::SimplifyConfig;
use crate::simplify::rules::RewriteRule;
use dashmap::DashMap;
use rayon::prelude::*;
use std::sync::Arc;
const PAR_THRESHOLD: usize = 4;
const SEGMENT_STACK_BYTES: usize = 16 * 1024 * 1024;
const FOREIGN_STACK_BUDGET: usize = 512 * 1024;
const OWNED_STACK_BUDGET: usize = SEGMENT_STACK_BYTES - 4 * 1024 * 1024;
thread_local! {
static SEGMENT_BASE: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
static SEGMENT_BUDGET: std::cell::Cell<usize> =
const { std::cell::Cell::new(FOREIGN_STACK_BUDGET) };
}
type Memo = DashMap<ExprId, ExprId>;
type Rules = Arc<Vec<Box<dyn RewriteRule>>>;
pub fn simplify_par(expr: ExprId, pool: &ExprPool) -> DerivedExpr<ExprId> {
simplify_par_with_config(expr, pool, &SimplifyConfig::default())
}
pub fn simplify_par_with_config(
expr: ExprId,
pool: &ExprPool,
config: &SimplifyConfig,
) -> DerivedExpr<ExprId> {
let rules: Rules = Arc::new(crate::simplify::rules_for_config(config));
let mut current = expr;
let mut full_log = DerivationLog::new();
for _ in 0..config.max_iterations {
let memo = Memo::new();
let result = simplify_node_par(current, pool, &rules, &memo);
full_log = full_log.merge(result.log);
if result.value == current {
break;
}
current = result.value;
}
let mut assumptions = config.assumptions.clone();
super::assumptions::collect_static_domain_facts(current, pool, &mut assumptions);
if !assumptions.is_empty() {
let colored = super::colored_egraph::apply_colored_if_needed(current, pool, &assumptions);
return DerivedExpr::with_log(colored.value, full_log.merge(colored.log));
}
DerivedExpr::with_log(current, full_log)
}
fn simplify_node_par(
expr: ExprId,
pool: &ExprPool,
rules: &Rules,
memo: &Memo,
) -> DerivedExpr<ExprId> {
if let Some(cached) = memo.get(&expr) {
return DerivedExpr::new(*cached);
}
let result = with_stack_segment(|| {
let (rebuilt, child_log) = pool.with(expr, |data| {
simplify_children_par(expr, data, pool, rules, memo)
});
let (current, rule_log) =
crate::simplify::engine::apply_rules(rebuilt, pool, rules.as_ref());
let limit_log = crate::simplify::engine::expand_limit_log();
DerivedExpr::with_log(current, child_log.merge(rule_log).merge(limit_log))
});
memo.insert(expr, result.value);
result
}
fn with_stack_segment<R: Send>(f: impl FnOnce() -> R + Send) -> R {
if stack_used() < SEGMENT_BUDGET.with(|b| b.get()) {
return f();
}
std::thread::scope(|scope| {
std::thread::Builder::new()
.stack_size(SEGMENT_STACK_BYTES)
.spawn_scoped(scope, || {
SEGMENT_BUDGET.with(|b| b.set(OWNED_STACK_BUDGET));
f()
})
.expect("failed to spawn stack segment for deep recursion")
.join()
.unwrap_or_else(|payload| std::panic::resume_unwind(payload))
})
}
fn stack_used() -> usize {
let probe = 0u8;
let here = &probe as *const u8 as usize;
SEGMENT_BASE.with(|base| {
if base.get() == 0 || here >= base.get() {
base.set(here);
0
} else {
base.get() - here
}
})
}
fn simplify_children_par(
expr: ExprId,
data: &ExprData,
pool: &ExprPool,
rules: &Rules,
memo: &Memo,
) -> (ExprId, DerivationLog) {
match data {
ExprData::Add(args) if args.len() >= PAR_THRESHOLD => {
let (new_args, log) = par_children(args, pool, rules, memo);
(rebuild_nary(expr, args, new_args, pool, NAry::Add), log)
}
ExprData::Mul(args) if args.len() >= PAR_THRESHOLD => {
let (new_args, log) = par_children(args, pool, rules, memo);
(rebuild_nary(expr, args, new_args, pool, NAry::Mul), log)
}
ExprData::Add(args) => {
let (new_args, log) = seq_children(args, pool, rules, memo);
(rebuild_nary(expr, args, new_args, pool, NAry::Add), log)
}
ExprData::Mul(args) => {
let (new_args, log) = seq_children(args, pool, rules, memo);
(rebuild_nary(expr, args, new_args, pool, NAry::Mul), log)
}
ExprData::Pow { base, exp } => {
let rb = simplify_node_par(*base, pool, rules, memo);
let re = simplify_node_par(*exp, pool, rules, memo);
let log = rb.log.merge(re.log);
let id = if rb.value == *base && re.value == *exp {
expr
} else {
pool.pow(rb.value, re.value)
};
(id, log)
}
ExprData::Func { name, args } => {
let (new_args, log) = seq_children(args, pool, rules, memo);
let id = if new_args == *args {
expr
} else {
pool.func(name.as_str(), new_args)
};
(id, log)
}
ExprData::Piecewise { branches, default } => {
let mut log = DerivationLog::new();
let mut changed = false;
let new_branches: Vec<(ExprId, ExprId)> = branches
.iter()
.map(|&(cond, val)| {
let rv = simplify_node_par(val, pool, rules, memo);
log = std::mem::take(&mut log).merge(rv.log);
changed |= rv.value != val;
(cond, rv.value)
})
.collect();
let rd = simplify_node_par(*default, pool, rules, memo);
log = log.merge(rd.log);
let id = if !changed && rd.value == *default {
expr
} else {
pool.piecewise(new_branches, rd.value)
};
(id, log)
}
ExprData::Predicate { kind, args } => {
let (new_args, log) = seq_children(args, pool, rules, memo);
let id = if new_args == *args {
expr
} else {
pool.predicate(kind.clone(), new_args)
};
(id, log)
}
ExprData::Forall { var, body } => {
let rb = simplify_node_par(*body, pool, rules, memo);
let id = if rb.value == *body {
expr
} else {
pool.forall(*var, rb.value)
};
(id, rb.log)
}
ExprData::Exists { var, body } => {
let rb = simplify_node_par(*body, pool, rules, memo);
let id = if rb.value == *body {
expr
} else {
pool.exists(*var, rb.value)
};
(id, rb.log)
}
ExprData::BigO(arg) => {
let r = simplify_node_par(*arg, pool, rules, memo);
let id = if r.value == *arg {
expr
} else {
pool.big_o(r.value)
};
(id, r.log)
}
_ => (expr, DerivationLog::new()),
}
}
fn par_children(
args: &[ExprId],
pool: &ExprPool,
rules: &Rules,
memo: &Memo,
) -> (Vec<ExprId>, DerivationLog) {
let results: Vec<DerivedExpr<ExprId>> = args
.par_iter()
.map(|&a| simplify_node_par(a, pool, rules, memo))
.collect();
let new_args: Vec<ExprId> = results.iter().map(|r| r.value).collect();
let mut log = DerivationLog::new();
for r in results {
log = log.merge(r.log);
}
(new_args, log)
}
fn seq_children(
args: &[ExprId],
pool: &ExprPool,
rules: &Rules,
memo: &Memo,
) -> (Vec<ExprId>, DerivationLog) {
let mut log = DerivationLog::new();
let new_args: Vec<ExprId> = args
.iter()
.map(|&a| {
let r = simplify_node_par(a, pool, rules, memo);
log = std::mem::take(&mut log).merge(r.log);
r.value
})
.collect();
(new_args, log)
}
enum NAry {
Add,
Mul,
}
fn rebuild_nary(
expr: ExprId,
old_args: &[ExprId],
new_args: Vec<ExprId>,
pool: &ExprPool,
kind: NAry,
) -> ExprId {
if new_args == old_args && old_args.windows(2).all(|w| w[0] <= w[1]) {
return expr;
}
match kind {
NAry::Add => pool.add(new_args),
NAry::Mul => pool.mul(new_args),
}
}
pub fn rules_for_config_par(config: &SimplifyConfig) -> Vec<Box<dyn RewriteRule + Send + Sync>> {
crate::simplify::engine::rules_for_config(config)
.into_iter()
.map(|rule| rule as Box<dyn RewriteRule + Send + Sync>)
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernel::{Domain, ExprPool};
use crate::simplify::simplify;
fn p() -> ExprPool {
ExprPool::new()
}
#[test]
fn par_matches_sequential_add() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let zero = pool.integer(0_i32);
let expr = pool.add(vec![x, zero, zero, zero, zero, zero]);
let seq = simplify(expr, &pool);
let par = simplify_par(expr, &pool);
assert_eq!(seq.value, par.value);
}
#[test]
fn par_matches_sequential_mul() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let one = pool.integer(1_i32);
let expr = pool.mul(vec![x, one, one, one, one, one]);
let seq = simplify(expr, &pool);
let par = simplify_par(expr, &pool);
assert_eq!(seq.value, par.value);
}
#[test]
fn par_constant_folding() {
let pool = p();
let a = pool.integer(2_i32);
let b = pool.integer(3_i32);
let c = pool.integer(4_i32);
let d = pool.integer(5_i32);
let expr = pool.add(vec![a, b, c, d]);
let par = simplify_par(expr, &pool);
let expected = pool.integer(14_i32);
assert_eq!(par.value, expected);
}
#[test]
fn ruleset_matches_sequential_exactly() {
for expand in [false, true] {
let config = SimplifyConfig {
expand,
..Default::default()
};
let seq: Vec<&str> = crate::simplify::rules_for_config(&config)
.iter()
.map(|r| r.name())
.collect();
let par: Vec<&str> = rules_for_config_par(&config)
.iter()
.map(|r| r.name())
.collect();
assert_eq!(seq, par, "rule lists diverged for expand={expand}");
}
}
#[test]
fn par_expand_matches_sequential() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let two = pool.integer(2_i32);
let sum = pool.add(vec![x, y]);
let sq = pool.pow(sum, two);
let config = SimplifyConfig {
expand: true,
..Default::default()
};
let seq = crate::simplify::simplify_with(
sq,
&pool,
&crate::simplify::rules_for_config(&config),
config.clone(),
);
let par = simplify_par_with_config(sq, &pool, &config);
assert_eq!(seq.value, par.value, "(x + y)^2 expanded differently");
assert_ne!(par.value, sq);
}
#[test]
fn par_honours_static_domains() {
let pool = p();
let x = pool.symbol("x", Domain::Positive);
let two = pool.integer(2_i32);
let sq = pool.pow(x, two);
let sqrt = pool.func("sqrt", vec![sq]);
let seq = simplify(sqrt, &pool);
let par = simplify_par(sqrt, &pool);
assert_eq!(seq.value, par.value);
}
#[test]
fn par_memoises_shared_subexpressions() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let one = pool.integer(1_i32);
let zero = pool.integer(0_i32);
let mut shared = x;
for _ in 0..8 {
shared = pool.mul(vec![shared, one]);
shared = pool.add(vec![shared, zero]);
}
let args: Vec<ExprId> = (1..=64)
.map(|i| {
let c = pool.integer(i);
pool.mul(vec![shared, c])
})
.collect();
let expr = pool.add(args);
let seq = simplify(expr, &pool);
let tp = rayon::ThreadPoolBuilder::new()
.num_threads(1)
.build()
.unwrap();
let par = tp.install(|| simplify_par(expr, &pool));
assert_eq!(seq.value, par.value);
assert_eq!(
seq.log.len(),
par.log.len(),
"parallel path re-simplified shared subexpressions"
);
}
fn deep_chain(pool: &ExprPool, depth: usize) -> ExprId {
let x = pool.symbol("x", Domain::Real);
let one = pool.integer(1_i32);
let zero = pool.integer(0_i32);
let mut e = x;
for _ in 0..depth {
e = pool.mul(vec![e, one]);
e = pool.add(vec![e, zero]);
e = pool.pow(e, one);
}
e
}
#[test]
fn par_survives_deep_chain_on_worker_thread() {
let pool = p();
let deep = deep_chain(&pool, 1000);
let x = pool.symbol("x", Domain::Real);
let tp = rayon::ThreadPoolBuilder::new()
.num_threads(2)
.build()
.unwrap();
let par = tp.install(|| simplify_par(deep, &pool));
assert_eq!(par.value, x);
}
#[inline(never)]
fn probe_at_depth(frames: u32) -> usize {
let mut pad = [0u8; 256];
pad[0] = frames as u8;
std::hint::black_box(&pad);
if frames == 0 {
stack_used()
} else {
probe_at_depth(frames - 1)
}
}
#[test]
fn stack_probe_rebaselines_after_unwinding() {
let deep = probe_at_depth(400);
assert_eq!(deep, 0, "the first probe on a thread establishes the base");
assert_eq!(stack_used(), 0, "a probe above the old base must re-base");
let used = probe_at_depth(64);
assert!(
used > 0,
"stack usage under-read as {used} after re-baselining"
);
}
#[test]
fn par_matches_sequential_on_moderate_chain() {
let pool = p();
let deep = deep_chain(&pool, 100);
let seq = simplify(deep, &pool);
let tp = rayon::ThreadPoolBuilder::new()
.num_threads(2)
.build()
.unwrap();
let par = tp.install(|| simplify_par(deep, &pool));
assert_eq!(seq.value, par.value);
}
#[test]
fn par_large_sum() {
let pool = p();
let args: Vec<ExprId> = (1..=20).map(|i| pool.integer(i)).collect();
let expr = pool.add(args);
let par = simplify_par(expr, &pool);
let seq = simplify(expr, &pool);
assert_eq!(par.value, seq.value);
}
#[test]
fn parallel_expansion_declines_are_recorded() {
let pool = p();
let vars: Vec<_> = (0..4)
.map(|i| pool.symbol(format!("v{i}"), Domain::Complex))
.collect();
let sum = pool.add(vars.clone());
let twelve = pool.integer(12);
let big = pool.pow(sum, twelve);
let config = SimplifyConfig {
expand: true,
..SimplifyConfig::default()
};
let out = simplify_par_with_config(big, &pool, &config);
assert_eq!(
out.value, big,
"the power is over budget, so it must not expand"
);
assert!(
out.log
.steps()
.iter()
.any(|s| s.rule_name == crate::simplify::rules::EXPAND_POW_LIMIT_RULE),
"the decline must be recorded, not silent"
);
}
}