Skip to content
Merged
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
21 changes: 20 additions & 1 deletion namedisl/set_like.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,6 +189,7 @@ class _NamedIslSetOrMapLike(NamedIslObject[IslSetOrMapLikeT_co]):
.. automethod:: project_out
.. automethod:: project_out_except
.. automethod:: gist
.. automethod:: detect_equalities
.. automethod:: remove_divs
.. automethod:: compute_divs
.. automethod:: __and__
Expand Down Expand Up @@ -312,6 +313,12 @@ def gist(self, context: Self) -> Self:
self_aligned.space,
)

def detect_equalities(self) -> Self:
return type(self)(
cast("IslSetOrMapLikeT_co", self._obj.detect_equalities()),
self.space,
)

def compute_divs(self) -> Self:
return type(self)(
cast("IslSetOrMapLikeT_co", self._obj.compute_divs()),
Expand Down Expand Up @@ -738,6 +745,9 @@ def make_map_from_domain_and_range(
class Map(_NamedIslMapLike[isl.Map], _NamedIslUnbasic[isl.Map]):
"""
.. automethod:: is_bijective
.. automethod:: is_injective
.. automethod:: is_single_valued
.. automethod:: lexmin
.. automethod:: complement
.. automethod:: simple_hull
.. automethod:: convex_hull
Expand All @@ -756,9 +766,18 @@ class Map(_NamedIslMapLike[isl.Map], _NamedIslUnbasic[isl.Map]):

_isl_type: ClassVar[type[IslObject]] = isl.Map

def is_bijective(self):
def is_bijective(self) -> bool:
return self._obj.is_bijective()

def is_injective(self) -> bool:
return self._obj.is_injective()

def is_single_valued(self) -> bool:
return self._obj.is_single_valued()

def lexmin(self) -> Map:
return Map(self._obj.lexmin(), self.space)

def complement(self) -> Map:
return Map(self._obj.complement(), self.space)

Expand Down
35 changes: 35 additions & 0 deletions namedisl/test/test_set_like.py
Original file line number Diff line number Diff line change
Expand Up @@ -315,6 +315,41 @@ def test_map_coalesce() -> None:
assert len(map_.coalesce().basic_maps()) == 1


def test_map_detect_equalities() -> None:
map_ = nisl.make_map("{ [i] -> [j] : 0 <= i <= j < 5 }")
detected = map_.detect_equalities()

assert isinstance(detected, nisl.Map)
assert detected.space == map_.space
assert detected.equals(map_)


def test_map_function_properties() -> None:
bijective = nisl.make_map("{ [i] -> [j = i] : 0 <= i < 2 }")
many_to_one = nisl.make_map("{ [i] -> [j = 0] : 0 <= i < 2 }")
one_to_many = nisl.make_map("{ [i = 0] -> [j] : 0 <= j < 2 }")

assert bijective.is_bijective()
assert bijective.is_injective()
assert bijective.is_single_valued()

assert not many_to_one.is_bijective()
assert not many_to_one.is_injective()
assert many_to_one.is_single_valued()

assert not one_to_many.is_bijective()
assert one_to_many.is_injective()
assert not one_to_many.is_single_valued()


def test_map_lexmin_preserves_named_space() -> None:
map_ = nisl.make_map(
"{ [target] -> [source] : target = 0 and 1 <= source <= 2 }"
)

assert map_.lexmin() == nisl.make_map("{ [target = 0] -> [source = 1] }")


@pytest.mark.parametrize("ndims_domain", [2, 3, 4, 5])
@pytest.mark.parametrize("ndims_range", [2, 3, 4, 5])
@pytest.mark.parametrize("has_params", [True, False])
Expand Down
Loading