Skip to content

Commit 6775954

Browse files
committed
Simplify wrappers
1 parent 802074c commit 6775954

1 file changed

Lines changed: 6 additions & 27 deletions

File tree

metric_learn/_util.py

Lines changed: 6 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -27,43 +27,22 @@ def vector_norm(X):
2727
'force_all_finite' in signature(check_array).parameters)
2828
_CHECK_X_Y_SUPPORTS_FORCE_ALL_FINITE = (
2929
'force_all_finite' in signature(check_X_y).parameters)
30-
_MISSING = object() # sentinel value to check if an argument is given or not
31-
32-
def _normalize_force_all_finite_arg(kwargs, supports_force_all_finite):
33-
"""Take a dictionary of arguments, and make sure that the argument that
34-
controls finite values check in `check_array` and `check_X_y` is named
35-
correctly for the version of scikit-learn used (`force_all_finite` for
36-
older version, `ensure_all_finite` for newer ones).
37-
"""
38-
kwargs = kwargs.copy()
39-
force_all_finite = kwargs.pop('force_all_finite', _MISSING)
40-
ensure_all_finite = kwargs.pop('ensure_all_finite', _MISSING)
41-
if supports_force_all_finite:
42-
if ensure_all_finite is not _MISSING:
43-
kwargs['force_all_finite'] = ensure_all_finite
44-
elif force_all_finite is not _MISSING:
45-
kwargs['force_all_finite'] = force_all_finite
46-
else:
47-
if force_all_finite is not _MISSING:
48-
kwargs['ensure_all_finite'] = force_all_finite
49-
elif ensure_all_finite is not _MISSING:
50-
kwargs['ensure_all_finite'] = ensure_all_finite
51-
return kwargs
5230

5331

5432
def _check_array(*args, **kwargs):
5533
"""Local wrapper around `sklearn.utils.check_array` to deal with the change
5634
from `force_all_finite` to `ensure_all_finite` in scikit-learn."""
57-
kwargs = _normalize_force_all_finite_arg(
58-
kwargs, _CHECK_ARRAY_SUPPORTS_FORCE_ALL_FINITE)
35+
if not _CHECK_ARRAY_SUPPORTS_FORCE_ALL_FINITE and "force_all_finite" in kwargs:
36+
kwargs = kwargs.copy()
37+
kwargs["ensure_all_finite"] = kwargs.pop("force_all_finite")
5938
return check_array(*args, **kwargs)
6039

61-
6240
def _check_X_y(*args, **kwargs):
6341
"""Local wrapper around `sklearn.utils.check_X_y` to deal with the change
6442
from `force_all_finite` to `ensure_all_finite` in scikit-learn."""
65-
kwargs = _normalize_force_all_finite_arg(
66-
kwargs, _CHECK_X_Y_SUPPORTS_FORCE_ALL_FINITE)
43+
if not _CHECK_X_Y_SUPPORTS_FORCE_ALL_FINITE and "force_all_finite" in kwargs:
44+
kwargs = kwargs.copy()
45+
kwargs["ensure_all_finite"] = kwargs.pop("force_all_finite")
6746
return check_X_y(*args, **kwargs)
6847

6948

0 commit comments

Comments
 (0)