use super::*;
#[test]
fn vector_and_matrix_extraction_alias_selected_scalar() {
let x = Expr::from(parameter!("x"));
let y = Expr::from(parameter!("y"));
let component = CompiledModel::from_expr(&vector([x.clone(), y.clone()]).component(1)).unwrap();
let element = CompiledModel::from_expr(
&matrix([[x, y.clone()], [3.0.into(), 4.0.into()]]).matrix_element(0, 1),
)
.unwrap();
for compiled in [component, element] {
assert_eq!(compiled.graph().nodes().len(), 1);
assert!(matches!(
compiled.graph().node(ExprId::from_index(0)),
Some(ExprNode::ScalarParam(parameter)) if parameter.name() == "y"
));
}
}
#[test]
fn tensor_scalar_addition_then_subtraction_cancels_per_element() {
let x = event_scalar("x");
let scale = Expr::from(parameter!("scale"));
let vector = vector([x.clone(), 2.0.into()]);
let matrix = matrix([[x, 3.0.into()]]);
for expression in [
((&vector + &scale) - &scale).component(0),
((&matrix + &scale) - &scale).matrix_element(0, 0),
] {
let compiled = CompiledModel::from_expr(&expression).unwrap();
assert!(
matches!(compiled.graph().node(compiled.graph().root()), Some(ExprNode::EventScalar(name)) if name.as_ref() == "x")
);
}
}
#[test]
fn matrix_vector_identities_and_zeroes_simplify() {
let x = event_scalar("x");
let y = event_scalar("y");
let identity = matrix([[1.0, 0.0], [0.0, 1.0]]);
let zero_matrix = matrix([[0.0, 0.0], [0.0, 0.0]]);
let vector = vector([x, y]);
let identity_product = CompiledModel::from_expr(&matvec(identity, vector.clone())).unwrap();
assert!(matches!(
identity_product.graph().node(identity_product.graph().root()),
Some(ExprNode::Vector { elements }) if elements.len() == 2
));
let zero_product = CompiledModel::from_expr(&matvec(zero_matrix, vector)).unwrap();
assert!(matches!(
zero_product.graph().node(zero_product.graph().root()),
Some(ExprNode::Vector { elements }) if elements.len() == 2
&& elements.iter().all(|id| matches!(zero_product.graph().node(*id), Some(ExprNode::RealConst(0.0))))
));
}
#[test]
fn dot_and_matvec_lower_to_scalar_arithmetic_when_cheaper() {
let x = event_scalar("x");
let y = event_scalar("y");
let dot_product =
CompiledModel::from_expr(&dot(vector([x.clone(), y.clone()]), vector([2.0, 3.0]))).unwrap();
assert_eq!(
dot_product
.graph()
.nodes()
.iter()
.filter(|node| matches!(node, ExprNode::Dot { .. }))
.count(),
0
);
assert!(matches!(
dot_product.graph().node(dot_product.graph().root()),
Some(ExprNode::NaryAdd { terms }) if terms.len() == 2
));
let matrix_product =
CompiledModel::from_expr(&matvec(matrix([[1.0, 2.0], [3.0, 4.0]]), vector([x, y])))
.unwrap();
assert_eq!(
matrix_product
.graph()
.nodes()
.iter()
.filter(|node| matches!(node, ExprNode::MatVec { .. }))
.count(),
0
);
assert!(matches!(
matrix_product.graph().node(matrix_product.graph().root()),
Some(ExprNode::Vector { elements }) if elements.len() == 2
));
}
#[test]
fn selected_aggregate_outputs_only_lower_required_contractions() {
const N: usize = 8;
let matrix_values = matrix::<N, N, Expr>(std::array::from_fn(|row| {
std::array::from_fn(|col| event_scalar(format!("m{row}_{col}")))
}));
let vector_values = vector(std::array::from_fn::<Expr, N, _>(|index| {
event_scalar(format!("v{index}"))
}));
let selected_row =
CompiledModel::from_expr(&matvec(matrix_values, vector_values).component(3)).unwrap();
assert!(matches!(
selected_row.graph().node(selected_row.graph().root()),
Some(ExprNode::NaryAdd { terms }) if terms.len() == N
));
assert!(
!selected_row
.graph()
.nodes()
.iter()
.any(|node| matches!(node, ExprNode::MatVec { .. } | ExprNode::Component { .. }))
);
let selected_names = selected_row
.graph()
.nodes()
.iter()
.filter_map(|node| match node {
ExprNode::EventScalar(name) => Some(name.as_ref()),
_ => None,
})
.collect::<Vec<_>>();
assert_eq!(selected_names.len(), 2 * N);
assert!(
selected_names
.iter()
.all(|name| { name.starts_with("v") || name.starts_with("m3_") })
);
let lhs = matrix::<N, N, Expr>(std::array::from_fn(|row| {
std::array::from_fn(|col| event_scalar(format!("a{row}_{col}")))
}));
let rhs = matrix::<N, N, Expr>(std::array::from_fn(|row| {
std::array::from_fn(|col| event_scalar(format!("b{row}_{col}")))
}));
let selected_element =
CompiledModel::from_expr(&matmul(lhs, rhs).matrix_element(2, 5)).unwrap();
assert!(matches!(
selected_element.graph().node(selected_element.graph().root()),
Some(ExprNode::NaryAdd { terms }) if terms.len() == N
));
assert!(
!selected_element.graph().nodes().iter().any(|node| matches!(
node,
ExprNode::MatMul { .. } | ExprNode::MatrixElement { .. }
))
);
let selected_names = selected_element
.graph()
.nodes()
.iter()
.filter_map(|node| match node {
ExprNode::EventScalar(name) => Some(name.as_ref()),
_ => None,
})
.collect::<Vec<_>>();
assert_eq!(selected_names.len(), 2 * N);
assert!(
selected_names
.iter()
.all(|name| { name.starts_with("a2_") || name.ends_with("_5") })
);
}
#[test]
fn matrix_multiplication_identity_and_zero_simplify() {
let x = event_scalar("x");
let matrix_value = matrix([[x, 2.0.into()], [3.0.into(), 4.0.into()]]);
let identity = matrix([[1.0, 0.0], [0.0, 1.0]]);
let identity_product = CompiledModel::from_expr(&matmul(identity, matrix_value)).unwrap();
assert!(matches!(
identity_product
.graph()
.node(identity_product.graph().root()),
Some(ExprNode::Matrix {
rows: 2,
cols: 2,
..
})
));
let zero_product = CompiledModel::from_expr(&matmul(
matrix([[0.0, 0.0], [0.0, 0.0]]),
matrix([[1.0, 2.0], [3.0, 4.0]]),
))
.unwrap();
assert!(matches!(
zero_product.graph().node(zero_product.graph().root()),
Some(ExprNode::Matrix { rows: 2, cols: 2, elements }) if elements
.iter()
.all(|id| matches!(zero_product.graph().node(*id), Some(ExprNode::RealConst(0.0))))
));
}
#[test]
fn equation_rules_handle_right_identity_and_zero_dot() {
let value = matrix([[event_scalar("a"), 2.0.into()], [3.0.into(), 4.0.into()]]);
let identity = matrix([[1.0, 0.0], [0.0, 1.0]]);
let product = CompiledModel::from_expr(&matmul(value, identity)).unwrap();
assert!(matches!(
product.graph().node(product.graph().root()),
Some(ExprNode::Matrix {
rows: 2,
cols: 2,
..
})
));
let zero = vector([0.0, 0.0]);
let other = vector([event_scalar("x"), event_scalar("y")]);
let dot_product = CompiledModel::from_expr(&dot(other, zero)).unwrap();
assert!(matches!(
dot_product.graph().node(dot_product.graph().root()),
Some(ExprNode::RealConst(0.0))
));
}
#[test]
fn indexed_formula_distributes_diagonal_scalar_over_matrix_element() {
let scale = Expr::from(parameter!("scale"));
let value = matrix([
[event_scalar("a"), event_scalar("b")],
[event_scalar("c"), event_scalar("d")],
]);
let diagonal = matrix([[scale.clone(), 0.0.into()], [0.0.into(), scale]]);
let selected = CompiledModel::from_expr(&matmul(diagonal, value).matrix_element(0, 1)).unwrap();
assert!(
matches!(selected.graph().node(selected.graph().root()), Some(ExprNode::NaryMul { factors })
if factors.len() == 2
&& factors.iter().any(|id| matches!(selected.graph().node(*id), Some(ExprNode::ScalarParam(parameter)) if parameter.name() == "scale"))
&& factors.iter().any(|id| matches!(selected.graph().node(*id), Some(ExprNode::EventScalar(name)) if name.as_ref() == "b")))
);
}