pub use laddu_expr::NumberClass;
use laddu_expr::{
ExprGraph, ExprId, ExprNode, ExprNodeSemantics, ValueKind,
parameters::{ParamState, Parameter},
};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct GraphFacts {
nodes: Vec<NodeFacts>,
}
impl GraphFacts {
pub fn analyze(graph: &ExprGraph) -> Self {
let mut nodes = Vec::with_capacity(graph.nodes().len());
let mut semantics = Vec::with_capacity(graph.nodes().len());
for node in graph.nodes() {
let node_semantics = node.semantics(&semantics);
nodes.push(NodeFacts::for_node(node, &nodes, node_semantics));
semantics.push(node_semantics);
}
Self { nodes }
}
pub fn get(&self, id: ExprId) -> Option<&NodeFacts> {
self.nodes.get(id.index())
}
pub fn nodes(&self) -> &[NodeFacts] {
&self.nodes
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub struct NodeFacts {
pub value_kind: ValueKind,
pub number_class: NumberClass,
pub dependency: DependencyFacts,
}
impl NodeFacts {
pub(crate) fn for_node(
node: &ExprNode,
facts: &[NodeFacts],
semantics: ExprNodeSemantics,
) -> Self {
let dependency = dependency(node, facts);
Self {
value_kind: semantics.value_kind,
number_class: semantics.number_class,
dependency,
}
}
pub fn evaluation_class(self) -> EvaluationClass {
self.dependency.evaluation_class()
}
}
#[derive(Copy, Clone, Debug, Default, PartialEq, Eq)]
pub struct DependencyFacts {
pub depends_on_free_params: bool,
pub depends_on_fixed_params: bool,
pub depends_on_event: bool,
}
impl DependencyFacts {
pub fn per_compile() -> Self {
Self::default()
}
pub fn from_parameter(parameter: &Parameter) -> Self {
match parameter.state() {
ParamState::Free => Self {
depends_on_free_params: true,
..Self::default()
},
ParamState::Fixed(_) => Self {
depends_on_fixed_params: true,
..Self::default()
},
}
}
pub fn from_event() -> Self {
Self {
depends_on_event: true,
..Self::default()
}
}
pub fn union(self, other: Self) -> Self {
Self {
depends_on_free_params: self.depends_on_free_params || other.depends_on_free_params,
depends_on_fixed_params: self.depends_on_fixed_params || other.depends_on_fixed_params,
depends_on_event: self.depends_on_event || other.depends_on_event,
}
}
pub fn evaluation_class(self) -> EvaluationClass {
if self.depends_on_free_params {
EvaluationClass::PerEvaluation
} else if self.depends_on_event {
EvaluationClass::PerEvent
} else {
EvaluationClass::PerCompile
}
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum EvaluationClass {
PerCompile,
PerEvent,
PerEvaluation,
}
fn dependency(node: &ExprNode, facts: &[NodeFacts]) -> DependencyFacts {
match node {
ExprNode::RealConst(_) | ExprNode::ComplexConst(_) => DependencyFacts::per_compile(),
ExprNode::ScalarParam(parameter) => DependencyFacts::from_parameter(parameter),
ExprNode::EventScalar(_) | ExprNode::EventP4Component { .. } => {
DependencyFacts::from_event()
}
ExprNode::Unary { .. }
| ExprNode::Binary { .. }
| ExprNode::NaryAdd { .. }
| ExprNode::NaryMul { .. }
| ExprNode::Complex { .. }
| ExprNode::Vector { .. }
| ExprNode::Matrix { .. }
| ExprNode::Component { .. }
| ExprNode::MatrixElement { .. }
| ExprNode::MatMul { .. }
| ExprNode::MatVec { .. }
| ExprNode::Dot { .. }
| ExprNode::Solve { .. } => union_children(node, facts),
}
}
fn union_children(node: &ExprNode, facts: &[NodeFacts]) -> DependencyFacts {
node.children()
.fold(DependencyFacts::per_compile(), |dependency, child| {
dependency.union(facts[child.index()].dependency)
})
}
#[cfg(test)]
mod tests {
use laddu_expr::{
Expr, ExprId, ExprNode, ValueKind, event_scalar, parameter, parameters::Parameter,
};
use crate::{CompileOptions, CompiledModel, EvaluationClass};
#[test]
fn facts_track_number_class_and_dependencies() {
let model =
event_scalar("mass") * Expr::from(Parameter::fixed("scale", 2.0)) + parameter!("x");
let compiled =
CompiledModel::from_expr_with_options(&model, &CompileOptions::without_optimizations())
.unwrap();
let event_id = compiled
.graph()
.nodes()
.iter()
.position(|node| matches!(node, ExprNode::EventScalar(name) if name.as_ref() == "mass"))
.map(ExprId::from_index)
.unwrap();
let root_facts = compiled.node_facts(compiled.graph().root()).unwrap();
assert_eq!(
compiled.node_facts(event_id).unwrap().value_kind,
ValueKind::Real
);
assert!(
compiled
.node_facts(event_id)
.unwrap()
.dependency
.depends_on_event
);
assert!(!compiled.graph().nodes().iter().any(
|node| matches!(node, ExprNode::ScalarParam(parameter) if parameter.name() == "scale")
));
assert!(
compiled
.graph()
.nodes()
.iter()
.any(|node| matches!(node, ExprNode::RealConst(value) if *value == 2.0))
);
assert!(root_facts.dependency.depends_on_free_params);
assert!(root_facts.dependency.depends_on_event);
assert_eq!(
root_facts.evaluation_class(),
EvaluationClass::PerEvaluation
);
}
}