radiate_gp/collections/graphs/
state.rs1use crate::{
2 Eval, EvalIntoMut, EvalMut, Graph, GraphEvaluator, GraphIterator, graphs::GraphEvalCache,
3};
4#[cfg(feature = "serde")]
5use serde::{Deserialize, Serialize};
6
7#[derive(Clone, PartialEq)]
8#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
9pub struct StatefulGraph<T, V> {
10 inner: Graph<T>,
11 state: Option<GraphEvalCache<V>>,
12}
13
14impl<T, V> StatefulGraph<T, V> {
15 pub fn new(inner: Graph<T>) -> Self {
16 StatefulGraph { inner, state: None }
17 }
18
19 pub fn reset(&mut self) {
20 self.state = None;
21 }
22
23 pub fn input_dim(&self) -> usize {
24 self.inner
25 .get_nodes_of_type(crate::NodeType::Input)
26 .collect::<Vec<_>>()
27 .len()
28 }
29
30 pub fn output_dim(&self) -> usize {
31 self.inner
32 .get_nodes_of_type(crate::NodeType::Output)
33 .collect::<Vec<_>>()
34 .len()
35 }
36
37 pub fn eval_scoped<F, O>(&mut self, eval_fn: F) -> O
38 where
39 F: FnOnce(&mut Self) -> O,
40 {
41 let current_state = self.state.take();
42 let output = eval_fn(self);
43 self.state = current_state;
44 output
45 }
46}
47
48impl<T, V> EvalIntoMut<[V], [V]> for StatefulGraph<T, V>
49where
50 T: Eval<[V], V>,
51 V: Copy + Default,
52{
53 fn eval_into_mut(&mut self, input: &[V], output: &mut [V]) {
54 let mut evaluator = match self.state.take() {
55 Some(c) => GraphEvaluator::from((&self.inner, c)),
56 None => GraphEvaluator::new(&self.inner),
57 };
58
59 evaluator.eval_into_mut(input, output);
60 self.state = Some(evaluator.take_cache());
61 }
62}
63
64impl<T, V> EvalMut<[V], Vec<V>> for StatefulGraph<T, V>
65where
66 T: Eval<[V], V>,
67 V: Copy + Default,
68{
69 fn eval_mut(&mut self, input: &[V]) -> Vec<V> {
70 let mut evaluator = match self.state.take() {
71 Some(c) => GraphEvaluator::from((&self.inner, c)),
72 None => GraphEvaluator::new(&self.inner),
73 };
74
75 let result = evaluator.eval_mut(input);
76 self.state = Some(evaluator.take_cache());
77 result
78 }
79}
80
81impl<T, V> AsRef<Graph<T>> for StatefulGraph<T, V> {
82 fn as_ref(&self) -> &Graph<T> {
83 &self.inner
84 }
85}
86
87impl<T, V> From<Graph<T>> for StatefulGraph<T, V>
88where
89 T: Eval<[V], V>,
90{
91 fn from(inner: Graph<T>) -> Self {
92 StatefulGraph { inner, state: None }
93 }
94}
95
96#[cfg(test)]
97mod tests {
98 #[test]
101 fn test_stateful_graph() {}
102}