laddu-compile 0.25.0

Amplitude analysis tools for Rust
Documentation
use super::*;

#[test]
fn cse_merges_duplicate_subtrees() {
    let x = Expr::from(parameter!("x"));
    let y = Expr::from(parameter!("y"));
    let sum = x + y;
    let model = sum.clone() * sum;
    let compiled = CompiledModel::from_expr_with_options(&model, &exact_options()).unwrap();

    assert_eq!(count_nary_add(&compiled), 1);
}

#[test]
fn cse_canonicalizes_commutative_binary_operands() {
    let x = Expr::from(parameter!("x"));
    let y = Expr::from(parameter!("y"));
    let model = (x.clone() + y.clone()) * (y + x);
    let compiled = CompiledModel::from_expr(&model).unwrap();

    assert_eq!(count_nary_add(&compiled), 1);
    assert!(matches!(
        compiled.graph().node(compiled.graph().root()),
        Some(
            ExprNode::Unary {
                op: UnaryOp::PowI(2),
                ..
            } | ExprNode::NaryMul { .. }
        )
    ));
}

#[test]
fn cse_canonicalizes_associative_addition_trees() {
    let x = Expr::from(parameter!("x"));
    let y = Expr::from(parameter!("y"));
    let z = Expr::from(parameter!("z"));
    let lhs = (x.clone() + y.clone()) + z.clone();
    let rhs = x + (z + y);
    let compiled = CompiledModel::from_expr(&(lhs * rhs)).unwrap();

    assert!(matches!(
        compiled.graph().node(compiled.graph().root()),
        Some(
            ExprNode::Unary {
                op: UnaryOp::PowI(2),
                ..
            } | ExprNode::NaryMul { .. }
        )
    ));
    assert_eq!(count_nary_add(&compiled), 1);
}

#[test]
fn cse_canonicalizes_associative_multiplication_trees() {
    let x = Expr::from(parameter!("x"));
    let y = Expr::from(parameter!("y"));
    let z = Expr::from(parameter!("z"));
    let lhs = (x.clone() * y.clone()) * z.clone();
    let rhs = z * (y * x);
    let compiled = CompiledModel::from_expr(&(lhs + rhs)).unwrap();
    assert!(compiled.cost().weighted_ops() <= 6);
}

#[test]
fn cse_ignores_metadata_when_merging_duplicate_subtrees() {
    let x = Expr::from(parameter!("x"));
    let y = Expr::from(parameter!("y"));
    let lhs = (x.clone() + y.clone()).named("lhs");
    let rhs = (x + y).tagged("rhs");
    let compiled = CompiledModel::from_expr(&(lhs * rhs)).unwrap();

    assert_eq!(count_nary_add(&compiled), 1);
}

#[test]
fn rewritten_subexpression_keeps_source_annotation() {
    let x = Expr::from(parameter!("x"));
    let marked = (x + 0.0).named("inner").tagged("retain");
    let compiled = CompiledModel::from_expr(&marked.sin()).unwrap();
    let parameter = compiled
        .graph()
        .nodes()
        .iter()
        .position(|node| matches!(node, ExprNode::ScalarParam(_)))
        .unwrap();
    let metadata = compiled
        .graph()
        .metadata(laddu_expr::ExprId::from_index(parameter))
        .unwrap();
    assert_eq!(metadata.name(), Some("inner"));
    assert!(metadata.has_tag("retain"));
}