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 anyhow::Context;
use fnv::FnvHashMap;
use petgraph::algo;
use petgraph::prelude::DiGraphMap;
use proptest::prelude::*;
fn build_dependency_graph<'a>(
assignments: impl IntoIterator<Item = (VariableID, &'a Function)>,
) -> Result<(DiGraphMap<VariableID, ()>, Vec<VariableID>), SubstitutionError> {
let assignments = assignments.into_iter().collect::<Vec<_>>();
let assigned_variables = assignments
.iter()
.map(|(id, _)| *id)
.collect::<std::collections::BTreeSet<_>>();
let mut dependency = DiGraphMap::new();
for &var_id in &assigned_variables {
dependency.add_node(var_id);
}
for &(assigned_var, function) in &assignments {
for required_var in function.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| assigned_variables.contains(var_id))
.collect();
Ok((dependency, topological_order))
}
#[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 (dependency, topological_order) =
build_dependency_graph(assignments.iter().map(|(&id, function)| (id, function)))?;
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(crate) fn substitute_acyclic_in_place_atomic(
&mut self,
acyclic: &AcyclicAssignments,
) -> Result<(), SubstitutionError> {
if acyclic.is_empty() {
return Ok(());
}
if self.is_empty() {
*self = acyclic.clone();
return Ok(());
}
let substituted_variables = acyclic.keys().collect::<std::collections::BTreeSet<_>>();
let mut replacements = FnvHashMap::default();
for (&id, function) in &self.assignments {
if substituted_variables.contains(&id) {
continue;
}
if !function.required_ids().is_disjoint(&substituted_variables) {
replacements.insert(id, function.clone().substitute_acyclic(acyclic)?);
}
}
let incoming = acyclic
.assignments
.iter()
.map(|(&id, function)| Ok((id, function.clone().substitute_acyclic(acyclic)?)))
.collect::<Result<FnvHashMap<_, _>, SubstitutionError>>()?;
let mut final_assignments = Vec::with_capacity(self.assignments.len() + incoming.len());
for (&id, function) in &self.assignments {
if incoming.contains_key(&id) {
continue;
}
final_assignments.push((id, replacements.get(&id).unwrap_or(function)));
}
final_assignments.extend(incoming.iter().map(|(&id, function)| (id, function)));
let (dependency, topological_order) = build_dependency_graph(final_assignments)?;
self.assignments.extend(replacements);
self.assignments.extend(incoming);
self.dependency = dependency;
self.topological_order = topological_order;
Ok(())
}
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(crate) fn into_evaluation_order(self) -> impl Iterator<Item = (VariableID, Function)> {
let Self {
mut assignments,
dependency: _,
topological_order,
} = self;
topological_order.into_iter().rev().map(move |id| {
let function = assignments
.remove(&id)
.expect("topological_order only contains assigned variables");
(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).with_context(|| {
format!("failed to evaluate assignment for variable {var_id:?}")
})?;
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> {
if self.is_empty() {
return Ok(FnvHashMap::default());
}
let assigned_ids = self.keys().collect::<Vec<_>>();
let evaluated_values = samples.try_map_ref(|state| {
let extended = self.evaluate(state, atol)?;
assigned_ids
.iter()
.map(|id| {
extended
.entries
.get(&id.into_inner())
.copied()
.ok_or_else(|| {
crate::error!(
{ id = ?id },
"evaluated assignment state is missing variable {id:?}"
)
})
})
.collect::<crate::Result<Vec<_>>>()
})?;
let mut result = FnvHashMap::default();
for (index, id) in assigned_ids.into_iter().enumerate() {
let values = evaluated_values.try_map_ref(|values| {
values.get(index).copied().ok_or_else(|| {
crate::error!(
{ id = ?id },
"evaluated assignment values are missing variable {id:?}"
)
})
})?;
result.insert(id, 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, SampleID, Sampled};
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 substitute_acyclic_in_place_atomic_matches_consuming_result() {
let initial = assign! {
1 <- linear!(2) + linear!(5),
2 <- linear!(6),
9 <- linear!(10)
};
let substitution = assign! {
2 <- linear!(3) + coeff!(1.0),
3 <- linear!(4)
};
let expected = initial.clone().substitute_acyclic(&substitution).unwrap();
let mut actual = initial;
actual
.substitute_acyclic_in_place_atomic(&substitution)
.unwrap();
assert_eq!(actual, expected);
}
#[test]
fn substitute_acyclic_in_place_atomic_preserves_value_on_error() {
let huge = crate::Coefficient::try_from(f64::MAX).unwrap();
let mut initial = AcyclicAssignments::new([(
VariableID::from(1),
Function::from((huge * linear!(2)).unwrap()),
)])
.unwrap();
let substitution = assign! {
2 <- (coeff!(2.0) * linear!(3)).unwrap()
};
let before = initial.clone();
let err = initial
.substitute_acyclic_in_place_atomic(&substitution)
.unwrap_err();
assert!(err.to_string().contains("Coefficient must be finite"));
assert_eq!(initial, before);
}
#[test]
fn substitute_acyclic_in_place_atomic_skips_overwritten_rhs() {
let huge = crate::Coefficient::try_from(f64::MAX).unwrap();
let mut initial = AcyclicAssignments::new([(
VariableID::from(1),
Function::from((huge * linear!(2)).unwrap()),
)])
.unwrap();
let substitution = assign! {
1 <- linear!(3),
2 <- (coeff!(2.0) * linear!(4)).unwrap()
};
initial
.substitute_acyclic_in_place_atomic(&substitution)
.unwrap();
assert_eq!(initial, substitution);
}
#[test]
fn into_evaluation_order_moves_every_assignment_in_dependency_order() {
let assignments = assign! {
1 <- linear!(2) + linear!(3),
4 <- linear!(1) + coeff!(2.0),
8 <- linear!(9)
};
let expected = assignments
.evaluation_order_iter()
.map(|(id, function)| (id, function.clone()))
.collect::<Vec<_>>();
let actual = assignments.into_evaluation_order().collect::<Vec<_>>();
assert_eq!(actual, expected);
}
#[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 evaluate_samples_matches_scalar_dag_evaluation_and_preserves_groups() {
let assignments = assign! {
1 <- linear!(2) + coeff!(1.0),
2 <- linear!(3)
};
let samples = Sampled::new(
vec![
vec![SampleID::from(10), SampleID::from(11)],
vec![SampleID::from(20)],
],
[
State::from_iter([(1, 100.0), (2, 200.0), (3, 2.0)]),
State::from_iter([(3, -2.0)]),
],
)
.unwrap();
let evaluated = assignments
.evaluate_samples(&samples, ATol::default())
.unwrap();
for values in evaluated.values() {
assert!(values.has_same_ids(&samples.ids()));
}
let mut groups = evaluated[&VariableID::from(1)]
.clone()
.chunk()
.into_iter()
.map(|(_, ids)| ids.into_iter().collect::<BTreeSet<_>>())
.collect::<Vec<_>>();
groups.sort();
assert_eq!(
groups,
vec![
BTreeSet::from([SampleID::from(10), SampleID::from(11)]),
BTreeSet::from([SampleID::from(20)]),
]
);
assert_eq!(
evaluated[&VariableID::from(1)].get(SampleID::from(10)),
Some(&3.0)
);
assert_eq!(
evaluated[&VariableID::from(1)].get(SampleID::from(11)),
Some(&3.0)
);
assert_eq!(
evaluated[&VariableID::from(1)].get(SampleID::from(20)),
Some(&-1.0)
);
assert_eq!(
evaluated[&VariableID::from(2)].get(SampleID::from(20)),
Some(&-2.0)
);
for (sample_id, state) in samples.iter() {
let scalar = assignments.evaluate(state, ATol::default()).unwrap();
for id in assignments.keys() {
assert_eq!(
evaluated[&id].get(*sample_id),
scalar.entries.get(&id.into_inner())
);
}
}
}
#[test]
fn evaluate_samples_preserves_source_groups_for_constant_assignment() {
let assignments = assign! {
1 <- coeff!(5.0)
};
let samples = Sampled::new(
vec![
vec![SampleID::from(10), SampleID::from(11)],
vec![SampleID::from(20)],
],
[State::from_iter([(9, 1.0)]), State::from_iter([(9, 2.0)])],
)
.unwrap();
let chunks = assignments
.evaluate_samples(&samples, ATol::default())
.unwrap()[&VariableID::from(1)]
.clone()
.chunk();
assert_eq!(chunks.len(), 2);
assert!(chunks.iter().all(|(value, _)| *value == 5.0));
let mut groups = chunks
.into_iter()
.map(|(_, ids)| ids.into_iter().collect::<BTreeSet<_>>())
.collect::<Vec<_>>();
groups.sort();
assert_eq!(
groups,
vec![
BTreeSet::from([SampleID::from(10), SampleID::from(11)]),
BTreeSet::from([SampleID::from(20)]),
]
);
}
#[test]
fn evaluate_samples_propagates_scalar_function_error() {
let assignments = assign! {
1 <- linear!(2),
2 <- linear!(3)
};
let state = State::from_iter([(3, f64::INFINITY)]);
let samples = Sampled::from((SampleID::from(10), state.clone()));
let scalar_error = assignments.evaluate(&state, ATol::default()).unwrap_err();
let sampled_error = assignments
.evaluate_samples(&samples, ATol::default())
.unwrap_err();
let sampled_message = sampled_error.to_string();
assert_eq!(sampled_message, scalar_error.to_string());
assert!(matches!(
sampled_error.downcast_ref::<crate::FunctionEvaluationError>(),
Some(crate::FunctionEvaluationError::NonFiniteResult { .. })
));
}
#[test]
fn evaluate_samples_reports_scalar_missing_external_dependency() {
let assignments = assign! {
1 <- linear!(2),
2 <- linear!(3)
};
let state = State::default();
let samples = Sampled::from((SampleID::from(10), state.clone()));
let scalar_error = assignments.evaluate(&state, ATol::default()).unwrap_err();
let sampled_error = assignments
.evaluate_samples(&samples, ATol::default())
.unwrap_err();
let scalar_missing = scalar_error
.downcast_ref::<crate::MissingStateEntries>()
.unwrap();
let sampled_missing = sampled_error
.downcast_ref::<crate::MissingStateEntries>()
.unwrap();
assert_eq!(sampled_missing.ids, scalar_missing.ids);
assert_eq!(sampled_missing.ids, BTreeSet::from([VariableID::from(3)]));
}
#[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("failed to evaluate assignment for variable VariableID(1)"));
assert!(err.is::<crate::FunctionEvaluationError>());
}
}