Skip to content

Commit 5d756d8

Browse files
committed
Fix: accept numpy scalars and 0/1 for solver params
1 parent 64fb294 commit 5d756d8

2 files changed

Lines changed: 46 additions & 9 deletions

File tree

‎python/cupdlpx/model.py‎

Lines changed: 21 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -488,21 +488,36 @@ def _resolve_param_key(self, name: str) -> str:
488488

489489
def _validate_param_value(self, key: str, value: Any) -> Any:
490490
if key in _BOOL_PARAMS:
491-
if not isinstance(value, bool):
492-
raise TypeError(f"Parameter '{key}' must be bool.")
493-
return value
491+
# accept a real bool, numpy bool, or an integer 0/1; store a Python bool
492+
if isinstance(value, (bool, np.bool_)):
493+
return bool(value)
494+
if isinstance(value, (int, np.integer)):
495+
if int(value) not in (0, 1):
496+
raise ValueError(f"Parameter '{key}' must be 0 or 1 when given as an int.")
497+
return bool(value)
498+
raise TypeError(f"Parameter '{key}' must be a bool (or 0/1).")
494499

495500
if key in _INT_PARAMS:
496-
if isinstance(value, bool) or not isinstance(value, int):
497-
raise TypeError(f"Parameter '{key}' must be int.")
501+
# accept any Python/numpy integer or an integer-valued float; store a Python int
502+
if isinstance(value, (bool, np.bool_)):
503+
raise TypeError(f"Parameter '{key}' must be an int.")
504+
if isinstance(value, (int, np.integer)):
505+
value = int(value)
506+
elif isinstance(value, (float, np.floating)) and float(value).is_integer():
507+
value = int(value)
508+
else:
509+
raise TypeError(f"Parameter '{key}' must be an int.")
498510
if key in _POSITIVE_INT_PARAMS and value <= 0:
499511
raise ValueError(f"Parameter '{key}' must be positive.")
500512
if key not in _POSITIVE_INT_PARAMS and value < 0:
501513
raise ValueError(f"Parameter '{key}' must be nonnegative.")
502514
return value
503515

504516
if key in _FLOAT_PARAMS:
505-
if isinstance(value, bool):
517+
# accept any real number (Python/numpy int or float); store a Python float
518+
if isinstance(value, (bool, np.bool_)):
519+
raise TypeError(f"Parameter '{key}' must be a number.")
520+
if not isinstance(value, (int, float, np.integer, np.floating)):
506521
raise TypeError(f"Parameter '{key}' must be a number.")
507522
value = float(value)
508523
if not np.isfinite(value):

‎test/test_api_surface.py‎

Lines changed: 25 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -206,15 +206,37 @@ def test_set_params_is_transactional(base_lp_data):
206206

207207
def test_param_value_validation(base_lp_data):
208208
model = _model(base_lp_data)
209+
# float params accept ints and numpy numbers
209210
model.setParam("TimeLimit", 12)
210211
assert model.getParam("TimeLimit") == 12.0
212+
model.setParam("TimeLimit", np.float64(30.0))
213+
assert model.getParam("TimeLimit") == 30.0
211214
model.setParam("OptimalityNorm", "LINF")
212215
assert model.getParam("OptimalityNorm") == "linf"
213216

217+
# bool params accept real/numpy bools and 0/1 (coerced to a Python bool)
218+
model.setParam("OutputFlag", 1)
219+
assert model.getParam("OutputFlag") is True
220+
model.setParam("OutputFlag", 0)
221+
assert model.getParam("OutputFlag") is False
222+
model.setParam("OutputFlag", np.bool_(True))
223+
assert model.getParam("OutputFlag") is True
224+
225+
# int params accept numpy ints and integer-valued floats (coerced to a Python int)
226+
model.setParam("IterationLimit", np.int64(1000))
227+
assert model.getParam("IterationLimit") == 1000 and isinstance(model.getParam("IterationLimit"), int)
228+
model.setParam("IterationLimit", 2000.0)
229+
assert model.getParam("IterationLimit") == 2000
230+
231+
# still-invalid inputs
214232
with pytest.raises(TypeError):
215-
model.setParam("OutputFlag", 1)
233+
model.setParam("IterationLimit", False) # bool is not an int here
216234
with pytest.raises(TypeError):
217-
model.setParam("IterationLimit", False)
235+
model.setParam("IterationLimit", 1.5) # non-integer float
236+
with pytest.raises(TypeError):
237+
model.setParam("OutputFlag", "yes") # string is not a bool
238+
with pytest.raises(ValueError):
239+
model.setParam("OutputFlag", 2) # only 0/1 allowed as int
218240
with pytest.raises(ValueError):
219241
model.setParam("IterationLimit", -1)
220242
with pytest.raises(ValueError):
@@ -261,7 +283,7 @@ def test_set_params_value_validation_is_transactional(base_lp_data):
261283
model = _model(base_lp_data)
262284
old_time_limit = model.getParam("TimeLimit")
263285
with pytest.raises(TypeError):
264-
model.setParams(TimeLimit=123.0, OutputFlag=1)
286+
model.setParams(TimeLimit=123.0, OutputFlag="yes") # OutputFlag invalid
265287
assert model.getParam("TimeLimit") == old_time_limit
266288

267289

0 commit comments

Comments
 (0)