Skip to main content

gammalooprs/cff/
tree.rs

1use std::collections::HashSet;
2
3use bincode_trait_derive::{Decode, Encode};
4use derive_more::{From, Into};
5use itertools::Itertools;
6use serde::{Deserialize, Serialize};
7use symbolica::{
8    atom::{Atom, AtomCore},
9    parse,
10};
11use typed_index_collections::TiVec;
12
13/// data structure for a tree
14
15#[derive(
16    Debug,
17    From,
18    Into,
19    Copy,
20    Clone,
21    Serialize,
22    Deserialize,
23    Encode,
24    Decode,
25    PartialEq,
26    Eq,
27    PartialOrd,
28    Ord,
29    Hash,
30)]
31pub struct NodeId(usize);
32
33impl NodeId {
34    pub const fn root() -> Self {
35        NodeId(0)
36    }
37}
38
39#[derive(Debug, Clone, Serialize, Deserialize, Encode, Decode, Eq, PartialEq)]
40pub struct TreeNode<T> {
41    pub data: T,
42    pub node_id: NodeId,
43    pub children: Vec<NodeId>,
44    pub parent: Option<NodeId>,
45}
46
47#[allow(dead_code)]
48fn determine_shifted_id(removed_ids_sorted: &[NodeId], original_id: NodeId) -> NodeId {
49    let shift = removed_ids_sorted
50        .iter()
51        .filter(|&&removed_id| removed_id < original_id)
52        .count();
53    NodeId(original_id.0 - shift)
54}
55
56impl<T> TreeNode<T> {
57    #[allow(dead_code)]
58    fn update_node_ids(&mut self, removed_ids: &[NodeId]) {
59        self.node_id = determine_shifted_id(removed_ids, self.node_id);
60
61        self.children
62            .retain(|child_id| !removed_ids.contains(child_id));
63        self.children.iter_mut().for_each(|child_id| {
64            *child_id = determine_shifted_id(removed_ids, *child_id);
65        });
66
67        if let Some(parent_id) = self.parent {
68            if removed_ids.contains(&parent_id) {
69                self.parent = None;
70            } else {
71                let shifted_parent_id = determine_shifted_id(removed_ids, parent_id);
72                self.parent = Some(shifted_parent_id);
73            }
74        }
75    }
76}
77
78#[derive(Debug, Clone, Serialize, Deserialize, Encode, Decode, Eq, PartialEq)]
79pub struct Tree<T> {
80    nodes: TiVec<NodeId, TreeNode<T>>,
81}
82
83impl<T> Tree<T> {
84    pub(crate) fn from_root(data: T) -> Self {
85        let node_id = NodeId(0);
86        let root_node = TreeNode {
87            data,
88            node_id,
89            children: Vec::new(),
90            parent: None,
91        };
92        Tree {
93            nodes: vec![root_node].into(),
94        }
95    }
96
97    pub(crate) fn insert_node(&mut self, parent_id: NodeId, data: T) {
98        let node_id = NodeId(self.nodes.len());
99        let new_node = TreeNode {
100            data,
101            node_id,
102            children: Vec::new(),
103            parent: Some(parent_id),
104        };
105        self.nodes.push(new_node);
106        self.nodes[parent_id].children.push(node_id);
107    }
108
109    pub(crate) fn get_node(&self, node_id: NodeId) -> &TreeNode<T> {
110        &self.nodes[node_id]
111    }
112
113    #[allow(dead_code)]
114    pub(crate) fn get_data_node_mut(&mut self, node_id: NodeId) -> &mut T {
115        &mut self.nodes[node_id].data
116    }
117
118    pub(crate) fn apply_mut_closure(&mut self, node_id: NodeId, closure: impl Fn(&mut T)) {
119        closure(&mut self.nodes[node_id].data);
120    }
121
122    pub(crate) fn get_bottom_layer(&self) -> Vec<NodeId> {
123        self.nodes
124            .iter()
125            .filter(|node| node.children.is_empty())
126            .map(|node| node.node_id)
127            .collect()
128    }
129
130    pub(crate) fn map<G>(self, f: impl Fn(T) -> G) -> Tree<G> {
131        Tree {
132            nodes: self
133                .nodes
134                .into_iter()
135                .map(|node| TreeNode {
136                    data: f(node.data),
137                    node_id: node.node_id,
138                    children: node.children,
139                    parent: node.parent,
140                })
141                .collect(),
142        }
143    }
144
145    pub(crate) fn map_mut(&mut self, f: impl Fn(&mut T)) {
146        for node in &mut self.nodes {
147            f(&mut node.data);
148        }
149    }
150
151    #[allow(dead_code)]
152    pub(crate) fn get_num_nodes(&self) -> usize {
153        self.nodes.len()
154    }
155
156    pub(crate) fn iter_nodes(&self) -> impl Iterator<Item = &TreeNode<T>> {
157        self.nodes.iter()
158    }
159
160    fn path_to_root(&self, mut node_id: NodeId) -> Vec<NodeId> {
161        let mut path = Vec::new();
162        loop {
163            path.push(node_id);
164            if let Some(parent_id) = self.nodes[node_id].parent {
165                node_id = parent_id;
166            } else {
167                break;
168            }
169        }
170        path.reverse();
171        path
172    }
173
174    #[allow(dead_code)]
175    fn obtain_subtree_node_ids_impl(&self, node_id: NodeId, collected_ids: &mut Vec<NodeId>) {
176        for &child_id in &self.nodes[node_id].children {
177            collected_ids.push(child_id);
178            self.obtain_subtree_node_ids_impl(child_id, collected_ids);
179        }
180    }
181
182    #[allow(dead_code)]
183    fn obtain_subtree_node_ids(&self, node_id: NodeId) -> Vec<NodeId> {
184        let mut collected_ids = vec![node_id];
185        self.obtain_subtree_node_ids_impl(node_id, &mut collected_ids);
186        collected_ids
187    }
188
189    #[allow(dead_code)]
190    pub(crate) fn remove_node(&mut self, node_id: NodeId) {
191        let mut subtree = self.obtain_subtree_node_ids(node_id);
192        subtree.sort();
193
194        // from largest to smallest, this way the shifting of indices does not affect the removal
195        for id in subtree.iter().rev() {
196            self.nodes.remove(*id);
197        }
198
199        for node in &mut self.nodes {
200            // shift all the node ids accordingly
201            node.update_node_ids(&subtree);
202        }
203    }
204
205    #[allow(dead_code)]
206    pub(crate) fn filter_mut(&mut self, predicate: impl Fn(&T) -> bool) {
207        let mut nodes_to_remove = HashSet::new();
208
209        for node in &self.nodes {
210            if !predicate(&node.data) {
211                nodes_to_remove.extend(self.obtain_subtree_node_ids(node.node_id));
212            }
213        }
214
215        let nodes_to_remove: Vec<NodeId> = nodes_to_remove.into_iter().sorted().collect();
216
217        for id in nodes_to_remove.iter().rev() {
218            self.nodes.remove(*id);
219        }
220
221        for node in &mut self.nodes {
222            // shift all the node ids accordingly
223            node.update_node_ids(&nodes_to_remove);
224        }
225    }
226
227    pub(crate) fn keep_branches_with_value_count_mut(&mut self, value: &T, n: usize)
228    where
229        T: Eq,
230    {
231        if self.nodes.is_empty() {
232            return;
233        }
234
235        let leaves = self.get_bottom_layer();
236        let mut nodes_to_keep = HashSet::new();
237        let mut has_match = false;
238
239        for leaf in leaves {
240            let path = self.path_to_root(leaf);
241            let count = path
242                .iter()
243                .filter(|&&node_id| self.nodes[node_id].data == *value)
244                .count();
245
246            if count == n {
247                has_match = true;
248                nodes_to_keep.extend(path);
249            }
250        }
251
252        if !has_match {
253            self.nodes.clear();
254            return;
255        }
256
257        let nodes_to_remove: Vec<NodeId> = self
258            .nodes
259            .iter()
260            .filter(|node| !nodes_to_keep.contains(&node.node_id))
261            .map(|node| node.node_id)
262            .sorted()
263            .collect();
264
265        for id in nodes_to_remove.iter().rev() {
266            self.nodes.remove(*id);
267        }
268
269        for node in &mut self.nodes {
270            node.update_node_ids(&nodes_to_remove);
271        }
272    }
273
274    pub(crate) fn max_value_count_on_branch(&self, value: &T) -> usize
275    where
276        T: Eq,
277    {
278        if self.nodes.is_empty() {
279            return 0;
280        }
281
282        self.get_bottom_layer()
283            .into_iter()
284            .map(|leaf| {
285                self.path_to_root(leaf)
286                    .into_iter()
287                    .filter(|&node_id| self.nodes[node_id].data == *value)
288                    .count()
289            })
290            .max()
291            .unwrap_or(0)
292    }
293}
294
295impl<T> Tree<T>
296where
297    Atom: From<T>,
298    T: Copy,
299{
300    fn to_atom_inv_impl(&self, cur_node: NodeId) -> Atom {
301        let node = &self.nodes[cur_node];
302        let inv_data_esurface = (Atom::num(1) / Atom::from(node.data))
303            .replace(parse!("η_inf^-1"))
304            .with(Atom::num(0));
305
306        let child_sum = node
307            .children
308            .iter()
309            .map(|&child| self.to_atom_inv_impl(child))
310            .reduce(|acc, x| acc + x)
311            .unwrap_or(Atom::num(1));
312
313        inv_data_esurface * child_sum
314    }
315
316    pub(crate) fn to_atom_inv(&self) -> Atom {
317        if self.nodes.is_empty() {
318            return Atom::num(0);
319        }
320        self.to_atom_inv_impl(NodeId::root())
321    }
322}
323
324#[cfg(test)]
325mod tests {
326
327    use crate::cff::tree::{NodeId, Tree};
328
329    #[test]
330    fn test_remove_node() {
331        let mut tree = Tree::from_root(0);
332        tree.insert_node(NodeId(0), 1);
333        tree.insert_node(NodeId(0), 2);
334        tree.insert_node(NodeId(2), 3);
335        tree.insert_node(NodeId(2), 4);
336        tree.insert_node(NodeId(4), 5);
337        tree.insert_node(NodeId(4), 6);
338        tree.insert_node(NodeId(4), 7);
339        tree.insert_node(NodeId(7), 8);
340        tree.insert_node(NodeId(7), 9);
341
342        tree.remove_node(NodeId(4));
343
344        let mut expected_tree = Tree::from_root(0);
345        expected_tree.insert_node(NodeId(0), 1);
346        expected_tree.insert_node(NodeId(0), 2);
347        expected_tree.insert_node(NodeId(2), 3);
348
349        assert_eq!(tree, expected_tree);
350
351        let mut tree = Tree::from_root(0);
352        tree.insert_node(NodeId(0), 1);
353        tree.insert_node(NodeId(1), 2);
354        tree.insert_node(NodeId(1), 3);
355        tree.insert_node(NodeId(1), 4);
356        tree.insert_node(NodeId(4), 5);
357        tree.insert_node(NodeId(4), 6);
358        tree.insert_node(NodeId(0), 7);
359        tree.insert_node(NodeId(7), 8);
360        tree.insert_node(NodeId(7), 9);
361
362        tree.remove_node(NodeId(1));
363
364        let mut expected_tree = Tree::from_root(0);
365        expected_tree.insert_node(NodeId(0), 7);
366        expected_tree.insert_node(NodeId(1), 8);
367        expected_tree.insert_node(NodeId(1), 9);
368
369        assert_eq!(tree, expected_tree);
370    }
371
372    #[test]
373    fn test_keep_branches_with_value_count_mut() {
374        let mut tree = Tree::from_root(0);
375        tree.insert_node(NodeId(0), 1);
376        tree.insert_node(NodeId(1), 2);
377        tree.insert_node(NodeId(2), 1);
378        tree.insert_node(NodeId(0), 1);
379        tree.insert_node(NodeId(4), 3);
380        tree.insert_node(NodeId(5), 4);
381
382        tree.keep_branches_with_value_count_mut(&1, 1);
383
384        let mut expected_tree = Tree::from_root(0);
385        expected_tree.insert_node(NodeId(0), 1);
386        expected_tree.insert_node(NodeId(1), 3);
387        expected_tree.insert_node(NodeId(2), 4);
388
389        assert_eq!(tree, expected_tree);
390    }
391
392    #[test]
393    fn test_keep_branches_with_value_count_mut_no_match() {
394        let mut tree = Tree::from_root(0);
395        tree.insert_node(NodeId(0), 1);
396        tree.insert_node(NodeId(1), 2);
397        tree.insert_node(NodeId(2), 1);
398
399        tree.keep_branches_with_value_count_mut(&1, 3);
400
401        assert_eq!(tree.get_num_nodes(), 0);
402    }
403
404    #[test]
405    fn test_keep_branches_with_value_count_mut_zero_occurrences() {
406        let mut tree = Tree::from_root(0);
407        tree.insert_node(NodeId(0), 2);
408        tree.insert_node(NodeId(1), 3);
409        tree.insert_node(NodeId(0), 1);
410        tree.insert_node(NodeId(3), 4);
411
412        tree.keep_branches_with_value_count_mut(&1, 0);
413
414        let mut expected_tree = Tree::from_root(0);
415        expected_tree.insert_node(NodeId(0), 2);
416        expected_tree.insert_node(NodeId(1), 3);
417
418        assert_eq!(tree, expected_tree);
419    }
420
421    #[test]
422    fn test_keep_branches_with_value_count_mut_keeps_shared_prefix() {
423        let mut tree = Tree::from_root(0);
424        tree.insert_node(NodeId(0), 9);
425        tree.insert_node(NodeId(1), 1);
426        tree.insert_node(NodeId(2), 5);
427        tree.insert_node(NodeId(1), 2);
428        tree.insert_node(NodeId(4), 6);
429
430        tree.keep_branches_with_value_count_mut(&1, 1);
431
432        let mut expected_tree = Tree::from_root(0);
433        expected_tree.insert_node(NodeId(0), 9);
434        expected_tree.insert_node(NodeId(1), 1);
435        expected_tree.insert_node(NodeId(2), 5);
436
437        assert_eq!(tree, expected_tree);
438    }
439
440    #[test]
441    fn test_max_value_count_on_branch() {
442        let mut tree = Tree::from_root(1);
443        tree.insert_node(NodeId(0), 1);
444        tree.insert_node(NodeId(1), 2);
445        tree.insert_node(NodeId(2), 1);
446        tree.insert_node(NodeId(0), 3);
447        tree.insert_node(NodeId(4), 1);
448
449        assert_eq!(tree.max_value_count_on_branch(&1), 3);
450    }
451
452    #[test]
453    fn test_max_value_count_on_branch_value_absent() {
454        let mut tree = Tree::from_root(0);
455        tree.insert_node(NodeId(0), 2);
456        tree.insert_node(NodeId(1), 3);
457        tree.insert_node(NodeId(0), 4);
458
459        assert_eq!(tree.max_value_count_on_branch(&1), 0);
460    }
461
462    #[test]
463    fn test_max_value_count_on_branch_single_branch() {
464        let mut tree = Tree::from_root(1);
465        tree.insert_node(NodeId(0), 1);
466
467        assert_eq!(tree.max_value_count_on_branch(&1), 2);
468    }
469}