Skip to main content

radiate_gp/collections/graphs/
state.rs

1use 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    // use super::*;
99
100    #[test]
101    fn test_stateful_graph() {}
102}