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
5 changes: 5 additions & 0 deletions coremltools/converters/mil/input_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -398,6 +398,11 @@ def __init__(
symbol:
Optional symbol name for the dim. Autogenerate a symbol name if not specified.
"""
if upper_bound > 0 and lower_bound > upper_bound:
raise ValueError(
f"Lower bound {lower_bound} is greater than upper bound ({upper_bound}) for range"
)

if symbol is None:
from coremltools.converters.mil.mil import get_new_symbol
self.symbol = get_new_symbol()
Expand Down
20 changes: 20 additions & 0 deletions coremltools/converters/mil/test/test_input_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,26 @@ def test_rangedim_raises_if_default_above_upper():
with pytest.raises(ValueError, match=r"greater than maximum value"):
RangeDim(lower_bound=0, upper_bound=5, default=6)

@pytest.mark.parametrize("default", [None, 4, 5])
def test_rangedim_raises_if_lower_above_positive_finite_upper(default):
with pytest.raises(ValueError, match=r"Lower bound.*greater than upper bound"):
RangeDim(lower_bound=5, upper_bound=3, default=default)

@pytest.mark.parametrize(
"lower_bound, upper_bound, default, expected_default",
[
(5, 5, None, 5),
(0, 5, None, 0),
(1, 5, 3, 3),
(5, -1, None, 5),
],
)
def test_rangedim_accepts_valid_bounds(lower_bound, upper_bound, default, expected_default):
dim = RangeDim(lower_bound=lower_bound, upper_bound=upper_bound, default=default)
assert dim.lower_bound == lower_bound
assert dim.upper_bound == upper_bound
assert dim.default == expected_default

def test_rangedim_ior_merges_bounds_and_adjusts_default():
dim1 = RangeDim(lower_bound=0, upper_bound=10, default=5)
dim2 = RangeDim(lower_bound=2, upper_bound=8, default=3)
Expand Down