radiate_gp/collections/trees/
iter.rs1use super::TreeChromosome;
2use crate::collections::{Tree, TreeNode};
3use std::{collections::VecDeque, marker::PhantomData};
4
5pub trait TreeIterator<T> {
57 fn iter_pre_order(&self) -> PreOrderIterator<'_, T>;
58 fn iter_post_order(&self) -> PostOrderIterator<'_, T>;
59 fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T>;
60 fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F);
61}
62
63impl<T> TreeIterator<T> for TreeNode<T> {
67 fn iter_pre_order(&self) -> PreOrderIterator<'_, T> {
68 PreOrderIterator { stack: vec![self] }
69 }
70
71 fn iter_post_order(&self) -> PostOrderIterator<'_, T> {
72 PostOrderIterator {
73 stack: vec![(self, false)],
74 }
75 }
76
77 fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T> {
78 TreeBreadthFirstIterator {
79 queue: vec![self].into_iter().collect(),
80 }
81 }
82
83 fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F) {
84 let visitor = TreeVisitor::new(visit_fn);
85 visitor.visit(self);
86 }
87}
88
89impl<T> TreeIterator<T> for Tree<T> {
93 fn iter_pre_order(&self) -> PreOrderIterator<'_, T> {
94 PreOrderIterator {
95 stack: self
96 .root()
97 .map_or(Vec::new(), |root| vec![root].into_iter().collect()),
98 }
99 }
100
101 fn iter_post_order(&self) -> PostOrderIterator<'_, T> {
102 PostOrderIterator {
103 stack: self
104 .root()
105 .map_or(Vec::new(), |root| vec![(root, false)].into_iter().collect()),
106 }
107 }
108 fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T> {
109 TreeBreadthFirstIterator {
110 queue: self
111 .root()
112 .map_or(VecDeque::new(), |root| vec![root].into_iter().collect()),
113 }
114 }
115
116 fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F) {
117 let visitor = TreeVisitor::new(visit_fn);
118 if let Some(root) = self.root_mut() {
119 visitor.visit(root);
120 }
121 }
122}
123
124impl<T> TreeIterator<T> for TreeChromosome<T> {
125 fn iter_pre_order(&self) -> PreOrderIterator<'_, T> {
126 self.root().iter_pre_order()
127 }
128
129 fn iter_post_order(&self) -> PostOrderIterator<'_, T> {
130 self.root().iter_post_order()
131 }
132
133 fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T> {
134 self.root().iter_breadth_first()
135 }
136
137 fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F) {
138 self.root_mut().apply(visit_fn);
139 }
140}
141
142pub struct PreOrderIterator<'a, T> {
143 stack: Vec<&'a TreeNode<T>>,
144}
145
146impl<'a, T> Iterator for PreOrderIterator<'a, T> {
147 type Item = &'a TreeNode<T>;
148
149 fn next(&mut self) -> Option<Self::Item> {
150 self.stack.pop().inspect(|node| {
151 if let Some(children) = node.children() {
152 for child in children.iter().rev() {
153 self.stack.push(child);
154 }
155 }
156 })
157 }
158}
159
160pub struct PostOrderIterator<'a, T> {
161 stack: Vec<(&'a TreeNode<T>, bool)>,
162}
163
164impl<'a, T> Iterator for PostOrderIterator<'a, T> {
165 type Item = &'a TreeNode<T>;
166
167 fn next(&mut self) -> Option<Self::Item> {
168 while let Some((node, visited)) = self.stack.pop() {
169 if visited {
170 return Some(node);
171 }
172 self.stack.push((node, true));
173 if let Some(children) = node.children() {
174 for child in children.iter().rev() {
175 self.stack.push((child, false));
176 }
177 }
178 }
179 None
180 }
181}
182
183pub struct TreeBreadthFirstIterator<'a, T> {
184 queue: VecDeque<&'a TreeNode<T>>,
185}
186
187impl<'a, T> Iterator for TreeBreadthFirstIterator<'a, T> {
188 type Item = &'a TreeNode<T>;
189
190 fn next(&mut self) -> Option<Self::Item> {
191 let node = self.queue.pop_front()?;
192 if let Some(children) = node.children() {
193 self.queue.extend(children.iter());
194 }
195 Some(node)
196 }
197}
198
199pub struct TreeVisitor<T, F>
200where
201 F: Fn(&mut TreeNode<T>),
202{
203 visitor: F,
204 _marker: PhantomData<T>,
205}
206
207impl<T, F> TreeVisitor<T, F>
208where
209 F: Fn(&mut TreeNode<T>),
210{
211 pub fn new(visitor: F) -> Self {
212 TreeVisitor {
213 visitor,
214 _marker: PhantomData,
215 }
216 }
217
218 pub fn visit(&self, node: &mut TreeNode<T>) {
219 (self.visitor)(node);
220
221 if let Some(children) = node.children_mut() {
222 for child in children.iter_mut() {
223 self.visit(child);
224 }
225 }
226 }
227}
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232 use crate::Op;
233 use crate::collections::{Tree, TreeNode};
234 use crate::node::Node;
235
236 #[test]
237 fn test_tree_traversal() {
238 let leaf = Op::constant(4.0);
245 let node2 = TreeNode::with_children(Op::constant(2.0), vec![TreeNode::new(leaf)]);
246
247 let node3 = TreeNode::new(Op::constant(3.0));
248
249 let root = Tree::new(TreeNode::with_children(
250 Op::constant(1.0),
251 vec![node2, node3],
252 ));
253
254 let pre_order: Vec<f32> = root
256 .iter_pre_order()
257 .map(|n| match &n.value() {
258 Op::Const(_, v) => *v,
259 _ => panic!("Expected constant but got {:?}", n.value()),
260 })
261 .collect();
262 assert_eq!(pre_order, vec![1.0, 2.0, 4.0, 3.0]);
263
264 let post_order: Vec<f32> = root
266 .iter_post_order()
267 .map(|n| match &n.value() {
268 Op::Const(_, v) => *v,
269 _ => panic!("Expected constant"),
270 })
271 .collect();
272 assert_eq!(post_order, vec![4.0, 2.0, 3.0, 1.0]);
273
274 let bfs: Vec<f32> = root
276 .iter_breadth_first()
277 .map(|n| match &n.value() {
278 Op::Const(_, v) => *v,
279 _ => panic!("Expected constant"),
280 })
281 .collect();
282 assert_eq!(bfs, vec![1.0, 2.0, 3.0, 4.0]);
283 }
284}