use crate::base::arena::Arena;
use crate::base::node::{ExprId, ExprNode};
use crate::base::walk;
use rustc_hash::{FxHashMap, FxHashSet};
use smallvec::SmallVec;
pub(crate) struct CseResult {
pub bindings: Vec<(ExprId, ExprId)>,
pub expr: ExprId,
}
pub(crate) struct CseMultiResult {
pub bindings: Vec<(ExprId, ExprId)>,
pub exprs: Vec<ExprId>,
}
pub(crate) fn cse(arena: &mut Arena, expr: ExprId) -> CseResult {
let multi = cse_multi(arena, &[expr]);
CseResult {
bindings: multi.bindings,
expr: multi.exprs[0],
}
}
pub(crate) fn cse_multi(arena: &mut Arena, exprs: &[ExprId]) -> CseMultiResult {
if exprs.is_empty() {
return CseMultiResult {
bindings: Vec::new(),
exprs: Vec::new(),
};
}
let mut ref_count: FxHashMap<ExprId, usize> = FxHashMap::default();
for &root in exprs {
count_refs(arena, root, &mut ref_count);
}
let all_ids = collect_all_ids(arena, exprs);
let overlap_rewrites = find_and_apply_partial_overlaps(arena, &all_ids);
let current_exprs: Vec<ExprId> = if overlap_rewrites.is_empty() {
exprs.to_vec()
} else {
exprs
.iter()
.map(|&e| apply_replacements(arena, e, &overlap_rewrites))
.collect()
};
let mut ref_count: FxHashMap<ExprId, usize> = FxHashMap::default();
for &root in ¤t_exprs {
count_refs(arena, root, &mut ref_count);
}
let mut combined_post_order: Vec<ExprId> = Vec::new();
let mut combined_visited: FxHashSet<ExprId> = FxHashSet::default();
for &root in ¤t_exprs {
let po = walk::post_order_ids(arena, root);
for id in po {
if combined_visited.insert(id) {
combined_post_order.push(id);
}
}
}
let extract: Vec<ExprId> = combined_post_order
.iter()
.copied()
.filter(|&id| {
let count = ref_count.get(&id).copied().unwrap_or(0);
count >= min_uses(arena, id)
})
.collect();
let mut bindings: Vec<(ExprId, ExprId)> = Vec::new();
let mut replacement_map: FxHashMap<ExprId, ExprId> = FxHashMap::default();
for (i, &id) in extract.iter().enumerate() {
let name = format!("__cse_{i}");
let name_id = arena.symbol(&name);
let replaced_value = apply_replacements(arena, id, &replacement_map);
bindings.push((name_id, replaced_value));
replacement_map.insert(id, name_id);
}
let final_exprs: Vec<ExprId> = current_exprs
.iter()
.map(|&e| apply_replacements(arena, e, &replacement_map))
.collect();
CseMultiResult {
bindings,
exprs: final_exprs,
}
}
fn min_uses(arena: &Arena, id: ExprId) -> usize {
let node = arena.node(id);
if node.is_atom() {
return usize::MAX;
}
let is_atom = |c: ExprId| arena.node(c).is_atom();
match node {
ExprNode::BoolTrue
| ExprNode::BoolFalse
| ExprNode::Gt(_, _)
| ExprNode::Ge(_, _)
| ExprNode::Eq_(_, _)
| ExprNode::Ne(_, _)
| ExprNode::And(_)
| ExprNode::Or(_)
| ExprNode::Not(_) => usize::MAX,
ExprNode::Neg(x) if is_atom(*x) => 3,
ExprNode::Mul(ch) if ch.len() == 2 && arena.as_num(ch[0]).is_some() && is_atom(ch[1]) => 3,
ExprNode::Pow(b, e) if is_atom(*b) => {
let two = num_bigint::BigUint::from(2u32);
let small = arena
.as_num(*e)
.is_some_and(|r| r.is_integer() && *r.numer().magnitude() <= two);
if small { 3 } else { 2 }
}
_ => 2,
}
}
fn collect_all_ids(arena: &Arena, roots: &[ExprId]) -> Vec<ExprId> {
let mut visited: FxHashSet<ExprId> = FxHashSet::default();
let mut result = Vec::new();
for &root in roots {
let po = walk::post_order_ids(arena, root);
for id in po {
if visited.insert(id) {
result.push(id);
}
}
}
result
}
fn find_and_apply_partial_overlaps(
arena: &mut Arena,
all_ids: &[ExprId],
) -> FxHashMap<ExprId, ExprId> {
let mut replacement_map: FxHashMap<ExprId, ExprId> = FxHashMap::default();
for is_add in [true, false] {
let mut node_children: Vec<(ExprId, Vec<ExprId>)> = Vec::new();
for &id in all_ids {
let dominated = match arena.node(id) {
ExprNode::Add(_) if is_add => true,
ExprNode::Mul(_) if !is_add => true,
_ => false,
};
if dominated {
let children: Vec<ExprId> = arena.node(id).children().to_vec();
if children.len() >= 2 {
node_children.push((id, children));
}
}
}
if node_children.len() < 2 {
continue;
}
let mut child_to_parents: FxHashMap<ExprId, Vec<usize>> = FxHashMap::default();
for (idx, (_id, children)) in node_children.iter().enumerate() {
for &child in children {
child_to_parents.entry(child).or_default().push(idx);
}
}
let mut pair_shared: FxHashMap<(usize, usize), Vec<ExprId>> = FxHashMap::default();
for (&child, parents) in &child_to_parents {
if parents.len() < 2 {
continue;
}
for i in 0..parents.len() {
for j in (i + 1)..parents.len() {
let a = parents[i].min(parents[j]);
let b = parents[i].max(parents[j]);
pair_shared.entry((a, b)).or_default().push(child);
}
}
}
let mut rewritten: FxHashSet<ExprId> = FxHashSet::default();
let mut pairs: Vec<((usize, usize), Vec<ExprId>)> = pair_shared
.into_iter()
.filter(|(_, shared)| shared.len() >= 2)
.collect();
pairs.sort_by_key(|p| (std::cmp::Reverse(p.1.len()), p.0));
for ((idx_a, idx_b), shared_children) in pairs {
let (parent_a, _) = &node_children[idx_a];
let (parent_b, _) = &node_children[idx_b];
if rewritten.contains(parent_a) || rewritten.contains(parent_b) {
continue;
}
let shared_set: FxHashSet<ExprId> = shared_children.iter().copied().collect();
if shared_set.len() < 2 {
continue;
}
let shared_vec: Vec<ExprId> = shared_set.iter().copied().collect();
let shared_expr = if is_add {
arena.add(&shared_vec)
} else {
arena.mul(&shared_vec)
};
let new_a = rewrite_nary_node(arena, *parent_a, &shared_set, shared_expr, is_add);
let new_b = rewrite_nary_node(arena, *parent_b, &shared_set, shared_expr, is_add);
if new_a != *parent_a {
replacement_map.insert(*parent_a, new_a);
rewritten.insert(*parent_a);
}
if new_b != *parent_b {
replacement_map.insert(*parent_b, new_b);
rewritten.insert(*parent_b);
}
}
}
replacement_map
}
fn rewrite_nary_node(
arena: &mut Arena,
parent: ExprId,
shared_children: &FxHashSet<ExprId>,
shared_expr: ExprId,
is_add: bool,
) -> ExprId {
let original_children: SmallVec<[ExprId; 6]> = arena.node(parent).children();
let mut remaining: Vec<ExprId> = original_children
.iter()
.filter(|c| !shared_children.contains(c))
.copied()
.collect();
remaining.push(shared_expr);
if remaining.len() == 1 {
return remaining[0];
}
if is_add {
arena.add(&remaining)
} else {
arena.mul(&remaining)
}
}
fn count_refs(arena: &Arena, root: ExprId, counts: &mut FxHashMap<ExprId, usize>) {
let mut stack = vec![root];
while let Some(id) = stack.pop() {
*counts.entry(id).or_insert(0) += 1;
if counts[&id] == 1 {
for &child in arena.node(id).children().iter() {
stack.push(child);
}
}
}
}
fn apply_replacements(arena: &mut Arena, expr: ExprId, map: &FxHashMap<ExprId, ExprId>) -> ExprId {
if map.is_empty() {
return expr;
}
walk::walk_and_rebuild(arena, expr, &|_arena, id| map.get(&id).copied())
}
#[cfg(test)]
mod tests {
use super::*;
fn sym(a: &mut Arena, name: &str) -> ExprId {
a.symbol(name)
}
fn display(a: &Arena, id: ExprId) -> String {
a.display(id).to_string()
}
#[test]
fn cse_no_common() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let expr = a.add(&[x, y]);
let result = cse(&mut a, expr);
assert!(result.bindings.is_empty(), "no common subexpressions");
assert_eq!(result.expr, expr);
}
#[test]
fn cse_with_common_subexpr() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let sin_x = a.sin(x);
let two = a.int(2);
let sin_sq = a.pow(sin_x, two);
let expr = a.add(&[sin_sq, sin_x]);
let result = cse(&mut a, expr);
assert!(!result.bindings.is_empty(), "should extract sin(x)");
let binding_val = display(&a, result.bindings[0].1);
assert!(
binding_val.contains("sin"),
"binding should be sin(x): {binding_val}"
);
}
#[test]
fn cse_deeply_shared() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let sin_x = a.sin(x);
let exp_sin = a.exp(sin_x);
let two = a.int(2);
let sq = a.pow(exp_sin, two);
let expr = a.add(&[exp_sin, sq]);
let result = cse(&mut a, expr);
assert!(
!result.bindings.is_empty(),
"should extract common subexprs"
);
}
#[test]
fn cse_atoms_not_extracted() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.add(&[x, x]);
let result = cse(&mut a, expr);
for (_, val) in &result.bindings {
assert!(!a.node(*val).is_atom(), "atoms should not be extracted");
}
}
#[test]
fn cse_preserves_value() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let three = a.int(3);
let x2 = a.pow(x, two);
let three_x2 = a.mul(&[three, x2]);
let expr = a.add(&[x2, three_x2]);
let result = cse(&mut a, expr);
let final_s = display(&a, result.expr);
assert!(!final_s.is_empty());
}
#[test]
fn cse_multi_empty() {
let mut a = Arena::new();
let result = cse_multi(&mut a, &[]);
assert!(result.bindings.is_empty());
assert!(result.exprs.is_empty());
}
#[test]
fn cse_multi_single_expr() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let sin_x = a.sin(x);
let two = a.int(2);
let sin_sq = a.pow(sin_x, two);
let expr = a.add(&[sin_sq, sin_x]);
let result = cse_multi(&mut a, &[expr]);
assert!(!result.bindings.is_empty(), "should extract sin(x)");
assert_eq!(result.exprs.len(), 1);
}
#[test]
fn cse_multi_shared_across_exprs() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let sin_x = a.sin(x);
let two = a.int(2);
let one = a.int(1);
let expr1 = a.add(&[sin_x, one]);
let expr2 = a.pow(sin_x, two);
let result = cse_multi(&mut a, &[expr1, expr2]);
assert!(
!result.bindings.is_empty(),
"should extract sin(x) shared between expr1 and expr2"
);
assert_eq!(result.exprs.len(), 2);
}
#[test]
fn cse_multi_no_shared() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let expr1 = a.sin(x);
let expr2 = a.cos(y);
let result = cse_multi(&mut a, &[expr1, expr2]);
assert!(result.bindings.is_empty(), "no common subexpressions");
assert_eq!(result.exprs.len(), 2);
}
#[test]
fn cse_multi_three_exprs() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let cos_x = a.cos(x);
let two = a.int(2);
let three = a.int(3);
let expr1 = a.pow(cos_x, two);
let expr2 = a.pow(cos_x, three);
let one = a.int(1);
let expr3 = a.add(&[cos_x, one]);
let result = cse_multi(&mut a, &[expr1, expr2, expr3]);
assert!(
!result.bindings.is_empty(),
"should extract cos(x) shared across all three"
);
assert_eq!(result.exprs.len(), 3);
}
}