mod logical_memory;
use super::error::SubstitutionError;
use crate::{
check_self_assignment, decision_variable::VariableID, substitute_acyclic_via_one, v1::State,
ATol, Evaluate, Function, Substitute, VariableIDSet,
};
use fnv::FnvHashMap;
use petgraph::algo;
use petgraph::prelude::DiGraphMap;
use proptest::prelude::*;
#[derive(Debug, Clone, Default)]
pub struct AcyclicAssignments {
assignments: FnvHashMap<VariableID, Function>,
dependency: DiGraphMap<VariableID, ()>,
topological_order: Vec<VariableID>,
}
impl AcyclicAssignments {
pub fn new(
iter: impl IntoIterator<Item = (VariableID, Function)>,
) -> Result<Self, SubstitutionError> {
let assignments: FnvHashMap<VariableID, Function> = iter.into_iter().collect();
let mut dependency = DiGraphMap::new();
for &var_id in assignments.keys() {
dependency.add_node(var_id);
}
for (&assigned_var, linear) in &assignments {
for required_var in linear.required_ids() {
if required_var == assigned_var {
return Err(SubstitutionError::RecursiveAssignment {
var_id: assigned_var,
});
}
dependency.add_edge(assigned_var, required_var, ());
}
}
let topological_order = algo::toposort(&dependency, None)
.map_err(|_| SubstitutionError::CyclicAssignmentDetected)?
.into_iter()
.filter(|var_id| assignments.contains_key(var_id))
.collect();
Ok(Self {
assignments,
dependency,
topological_order,
})
}
pub fn is_empty(&self) -> bool {
self.assignments.is_empty()
}
pub fn len(&self) -> usize {
self.assignments.len()
}
pub fn get(&self, id: &VariableID) -> Option<&Function> {
self.assignments.get(id)
}
pub fn iter(&self) -> impl Iterator<Item = (&VariableID, &Function)> {
self.assignments.iter()
}
pub fn substitution_order_iter(&self) -> impl Iterator<Item = (VariableID, &Function)> {
self.topological_order.iter().copied().map(move |var_id| {
let function = self
.assignments
.get(&var_id)
.expect("topological_order only contains assigned variables");
(var_id, function)
})
}
pub fn evaluation_order_iter(&self) -> impl Iterator<Item = (VariableID, &Function)> {
self.topological_order
.iter()
.copied()
.rev()
.map(move |var_id| {
let function = self
.assignments
.get(&var_id)
.expect("topological_order only contains assigned variables");
(var_id, function)
})
}
pub fn keys(&self) -> impl Iterator<Item = VariableID> + '_ {
self.assignments.keys().copied()
}
}
impl PartialEq for AcyclicAssignments {
fn eq(&self, other: &Self) -> bool {
if self.assignments != other.assignments {
return false;
}
let self_nodes: std::collections::BTreeSet<_> = self.dependency.nodes().collect();
let other_nodes: std::collections::BTreeSet<_> = other.dependency.nodes().collect();
if self_nodes != other_nodes {
return false;
}
let self_edges: std::collections::BTreeSet<_> = self.dependency.all_edges().collect();
let other_edges: std::collections::BTreeSet<_> = other.dependency.all_edges().collect();
self_edges == other_edges
}
}
impl IntoIterator for AcyclicAssignments {
type Item = (VariableID, Function);
type IntoIter = <FnvHashMap<VariableID, Function> as IntoIterator>::IntoIter;
fn into_iter(self) -> Self::IntoIter {
self.assignments.into_iter()
}
}
impl Arbitrary for AcyclicAssignments {
type Parameters = AcyclicAssignmentsParameters;
type Strategy = BoxedStrategy<Self>;
fn arbitrary_with(p: Self::Parameters) -> Self::Strategy {
proptest::collection::vec(
(
(0..=p.function_parameters.max_id().into_inner()).prop_map(VariableID::from),
Function::arbitrary_with(p.function_parameters),
),
0..=p.max_assignments,
)
.prop_filter_map("Acyclic", |assignments| {
AcyclicAssignments::new(assignments).ok()
})
.boxed()
}
}
#[derive(Debug, Clone, Copy)]
pub struct AcyclicAssignmentsParameters {
pub max_assignments: usize,
pub function_parameters: <Function as Arbitrary>::Parameters,
}
impl Default for AcyclicAssignmentsParameters {
fn default() -> Self {
Self {
max_assignments: 10,
function_parameters: <Function as Arbitrary>::Parameters::default(),
}
}
}
impl Evaluate for AcyclicAssignments {
type Output = State;
type SampledOutput = FnvHashMap<VariableID, crate::Sampled<f64>>;
fn evaluate(&self, state: &State, atol: ATol) -> crate::Result<Self::Output> {
let mut extended_state = state.clone();
for (var_id, function) in self.evaluation_order_iter() {
let value = function.evaluate(&extended_state, atol)?;
if !value.is_finite() {
return Err(crate::error!(
"Assignment for variable {var_id:?} evaluated to non-finite value: {value}"
));
}
extended_state.entries.insert(var_id.into_inner(), value);
}
Ok(extended_state)
}
fn partial_evaluate(&mut self, state: &State, atol: ATol) -> crate::Result<()> {
let mut new_assignments = Vec::new();
for (var_id, function) in self.assignments.iter() {
let mut function_clone = function.clone();
function_clone.partial_evaluate(state, atol)?;
new_assignments.push((*var_id, function_clone));
}
*self = Self::new(new_assignments)?;
Ok(())
}
fn evaluate_samples(
&self,
samples: &crate::Sampled<State>,
atol: ATol,
) -> crate::Result<Self::SampledOutput> {
let mut result = FnvHashMap::default();
for (var_id, function) in self.substitution_order_iter() {
let sampled_values = function.evaluate_samples(samples, atol)?;
result.insert(var_id, sampled_values);
}
Ok(result)
}
fn required_ids(&self) -> VariableIDSet {
self.assignments
.values()
.flat_map(|function| function.required_ids())
.collect()
}
}
impl Substitute for AcyclicAssignments {
type Output = Self;
fn substitute_acyclic(
self,
acyclic: &crate::AcyclicAssignments,
) -> Result<Self::Output, crate::SubstitutionError> {
if self.is_empty() {
return Ok(acyclic.clone());
}
if acyclic.is_empty() {
return Ok(self);
}
substitute_acyclic_via_one(self, acyclic)
}
fn substitute_one(
self,
assigned: VariableID,
function: &Function,
) -> Result<Self::Output, SubstitutionError> {
check_self_assignment(assigned, function)?;
let mut new_assignments = Vec::new();
for (var_id, func) in self.assignments {
let substituted_func = func.substitute_one(assigned, function)?;
new_assignments.push((var_id, substituted_func));
}
new_assignments.push((assigned, function.clone()));
Self::new(new_assignments)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{assign, coeff, linear};
use ::approx::assert_abs_diff_eq;
use std::collections::BTreeSet;
#[test]
fn test_substitute_acyclic_success() {
let initial = assign! {
1 <- linear!(2) + linear!(3)
};
let substitution = assign! {
3 <- linear!(4) + coeff!(1.0)
};
let expected = assign! {
1 <- ((linear!(2) + linear!(4)).unwrap() + coeff!(1.0)).unwrap(),
3 <- linear!(4) + coeff!(1.0)
};
let result = initial.substitute_acyclic(&substitution).unwrap();
assert_eq!(result.assignments.len(), expected.assignments.len());
for (var_id, expected_function) in expected.iter() {
assert_abs_diff_eq!(result.get(var_id).unwrap(), expected_function);
}
assert_eq!(
result.dependency.all_edges().collect::<BTreeSet<_>>(),
expected.dependency.all_edges().collect::<BTreeSet<_>>()
);
}
#[test]
fn test_substitute_acyclic_fail() {
let initial = assign! {
1 <- linear!(2) + linear!(3)
};
let substitution = assign! {
3 <- linear!(1) + coeff!(1.0)
};
insta::assert_snapshot!(
initial.substitute_acyclic(&substitution).unwrap_err(),
@"Recursive assignment detected: variable 1 cannot be assigned to a function that depends on itself"
);
}
#[test]
fn test_evaluate_topological_order() {
let assignments = assign! {
1 <- linear!(2) + linear!(3), 4 <- linear!(1) + coeff!(2.0) };
let state = State::from_iter(vec![(2, 1.0), (3, 2.0)]);
let result = assignments.evaluate(&state, ATol::default()).unwrap();
assert_eq!(result.entries[&1], 3.0); assert_eq!(result.entries[&2], 1.0); assert_eq!(result.entries[&3], 2.0); assert_eq!(result.entries[&4], 5.0); }
#[test]
fn test_topological_order_contains_only_assigned_variables() {
let assignments = assign! {
1 <- linear!(2) + linear!(3),
4 <- linear!(1) + linear!(5)
};
assert_eq!(assignments.topological_order.len(), assignments.len());
assert!(assignments
.topological_order
.iter()
.all(|var_id| assignments.assignments.contains_key(var_id)));
}
#[test]
fn test_evaluate_rejects_non_finite_assignment_value() {
let assignments = assign! {
1 <- linear!(2)
};
let state = State::from_iter([(2, f64::INFINITY)]);
let err = assignments.evaluate(&state, ATol::default()).unwrap_err();
assert!(err.to_string().contains("evaluated to non-finite value"));
}
}