Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 14 additions & 12 deletions cereeberus/cereeberus/reeb/reebgraph.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import numpy as np

from ..draw import draw
from collections import Counter

# from build.lib.cereeberus.reeb import graph

Expand Down Expand Up @@ -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]
Expand All @@ -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])
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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):
Expand Down
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ dependencies = ["numpy",
"matplotlib",
"pandas",
"scipy",
"pulp",
"pulp<4",
"gudhi",
"scikit-learn"
]
Expand All @@ -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/"
documentation = "https://munchlab.github.io/ceREEBerus/"
2 changes: 1 addition & 1 deletion requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ networkx
matplotlib
pandas
scipy
pulp
pulp<4
gudhi
sphinx
nbsphinx
Expand Down
26 changes: 24 additions & 2 deletions tests/test_reeb_class.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down
Loading