diff --git a/cereeberus/cereeberus/reeb/reebgraph.py b/cereeberus/cereeberus/reeb/reebgraph.py index 7ead05a..fea1141 100644 --- a/cereeberus/cereeberus/reeb/reebgraph.py +++ b/cereeberus/cereeberus/reeb/reebgraph.py @@ -5,6 +5,7 @@ import numpy as np from ..draw import draw +from collections import Counter # from build.lib.cereeberus.reeb import graph @@ -710,9 +711,9 @@ def slice(self, a, b, type="open", verbose=False): # (a, a) is an empty interval for open type return ReebGraph() if type == "open": - v_list = [v for v in self.nodes() if self.f[v] > a and self.f[v] < b] + v_list = set(v for v in self.nodes() if self.f[v] > a and self.f[v] < b) elif type == "closed": - v_list = [v for v in self.nodes() if self.f[v] >= a and self.f[v] <= b] + v_list = set(v for v in self.nodes() if self.f[v] >= a and self.f[v] <= b) # Keep the edges where either endpoint (or both) is in (a,b) e_list = [e for e in self.edges() if e[0] in v_list or e[1] in v_list] @@ -732,7 +733,7 @@ def slice(self, a, b, type="open", verbose=False): ) # Make a dictionary of counts to deal with multiedges - e_dict = {e: e_list.count(e) for e in e_list} + e_dict = Counter(e_list) if verbose: print("Vertices (v,f(v)):", [(v, self.f[v]) for v in v_list]) @@ -743,7 +744,7 @@ def slice(self, a, b, type="open", verbose=False): H = ReebGraph() for v in v_list: - H.add_node(v, self.f[v]) + H.add_node(v, self.f[v], reset_pos=False) for e in e_dict: if e[0] in v_list and e[1] in v_list: @@ -752,7 +753,7 @@ def slice(self, a, b, type="open", verbose=False): print(f"Adding {e_dict[e]} of edge {e} entirely inside slice:") for i in range(e_dict[e]): # Add an edge for each copy in the list - H.add_edge(e[0], e[1]) + H.add_edge(e[0], e[1], reset_pos=False) elif e[0] not in v_list and e[1] not in v_list: # The edge is entirely crossing the slice, so we add two vertices and an edge @@ -764,9 +765,9 @@ def slice(self, a, b, type="open", verbose=False): for i in range(e_dict[e]): v1 = "-".join([str(v) for v in e]) + "_" + str(i) + "_lower" v2 = "-".join([str(v) for v in e]) + "_" + str(i) + "_upper" - H.add_node(v1, a) - H.add_node(v2, b) - H.add_edge(v1, v2) + H.add_node(v1, a, reset_pos=False) + H.add_node(v2, b, reset_pos=False) + H.add_edge(v1, v2, reset_pos=False) else: # One vertex is in the set and one is out. # Need to check (for the closed case) that this isn't an edge going up from the top bound or down from the bottom bound @@ -805,8 +806,8 @@ def slice(self, a, b, type="open", verbose=False): func_val = b # Add a new vertex called edge_name with value b - H.add_node(edge_name, func_val) - H.add_edge(e[0], edge_name) + H.add_node(edge_name, func_val, reset_pos=False) + H.add_edge(e[0], edge_name, reset_pos=False) else: # The higher edge is in the set, so the other vertex must have # value below the min @@ -823,8 +824,9 @@ def slice(self, a, b, type="open", verbose=False): func_val = a # Add a new vertex called edge_name with value a - H.add_node(edge_name, func_val) - H.add_edge(edge_name, e[1]) + H.add_node(edge_name, func_val, reset_pos=False) + H.add_edge(edge_name, e[1], reset_pos=False) + H.set_pos_from_f() return H def connected_components(self): diff --git a/pyproject.toml b/pyproject.toml index d70d00e..f21a1d9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,7 +20,7 @@ dependencies = ["numpy", "matplotlib", "pandas", "scipy", - "pulp", + "pulp<4", "gudhi", "scikit-learn" ] @@ -34,4 +34,4 @@ cereeberus = ["data"] [project.urls] repository = "https://github.com/MunchLab/ceREEBerus" homepage = "https://munchlab.github.io/ceREEBerus/" -documentation = "https://munchlab.github.io/ceREEBerus/" \ No newline at end of file +documentation = "https://munchlab.github.io/ceREEBerus/" diff --git a/requirements.txt b/requirements.txt index e1fa66f..a462920 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,7 +3,7 @@ networkx matplotlib pandas scipy -pulp +pulp<4 gudhi sphinx nbsphinx diff --git a/tests/test_reeb_class.py b/tests/test_reeb_class.py index de3b189..527ad5f 100644 --- a/tests/test_reeb_class.py +++ b/tests/test_reeb_class.py @@ -299,8 +299,8 @@ def test_set_pos_from_f_preserves_y_function_values(self): self.assertEqual(set(R.nodes), set(R.pos_f.keys())) for v in R.nodes: self.assertEqual(R.pos_f[v][1], R.f[v]) - - + + def test_remove_node_deferred_pos_cleanup(self): # Regression test: remove_node(reset_pos=False) must still drop the # removed vertex's pos_f entry immediately. Previously this cleanup @@ -411,7 +411,29 @@ def test_slice_closed_unchanged(self): self.assertEqual(self._shape(T.slice(1, 4, type='closed')), (2, 2, 1)) self.assertEqual(self._shape(T.slice(2, 2, type='closed')), (2, 0, 2)) + + + + def test_slice_multiedge(self): + # Slicing across a multiedge should produce one new lower/upper vertex + # pair per parallel copy, and the result should still be a well-formed + # Reeb graph (positions computed, edges pointing upward, etc). + R = ex_rg.torus() # nodes a=0, b=1, c=4, d=5, with a double edge b-c + + # Interval (2,3) falls strictly between b and c, so v_list is empty and + # both copies of the b-c multiedge cross the slice entirely. + H = R.slice(2, 3) + + self.assertEqual(len(H.nodes), 4) # 2 lower + 2 upper subdivision vertices + self.assertEqual(len(H.edges), 2) # one edge per copy of the multiedge + self.assertEqual(H.number_connected_components(), 2) + self.check_reeb(H) + # Same check for the closed-interval case + H = R.slice(2, 3, type='closed') + self.assertEqual(len(H.nodes), 4) + self.assertEqual(len(H.edges), 2) + self.check_reeb(H) def test_boundary_map_parallel_edges(self):