diff --git a/sklearn/decomposition/tests/test_pca.py b/sklearn/decomposition/tests/test_pca.py index c0d1060217f..83f71381c0b 100644 --- a/sklearn/decomposition/tests/test_pca.py +++ b/sklearn/decomposition/tests/test_pca.py @@ -817,9 +817,9 @@ def test_variance_correctness(copy): np.testing.assert_allclose(pca_var, true_var) -def check_array_api_get_precision(name, estimator, array_namespace, device, dtype): +def check_array_api_get_precision(name, estimator, array_namespace, device, dtype_name): xp = _array_api_for_tests(array_namespace, device) - iris_np = iris.data.astype(dtype) + iris_np = iris.data.astype(dtype_name) iris_xp = xp.asarray(iris_np, device=device) estimator.fit(iris_np) @@ -835,7 +835,7 @@ def check_array_api_get_precision(name, estimator, array_namespace, device, dtyp assert_allclose( _convert_to_numpy(precision_xp, xp=xp), precision_np, - atol=_atol_for_type(dtype), + atol=_atol_for_type(dtype_name), ) covariance_xp = estimator_xp.get_covariance() assert covariance_xp.shape == (4, 4) @@ -844,12 +844,12 @@ def check_array_api_get_precision(name, estimator, array_namespace, device, dtyp assert_allclose( _convert_to_numpy(covariance_xp, xp=xp), covariance_np, - atol=_atol_for_type(dtype), + atol=_atol_for_type(dtype_name), ) @pytest.mark.parametrize( - "array_namespace, device, dtype", yield_namespace_device_dtype_combinations() + "array_namespace, device, dtype_name", yield_namespace_device_dtype_combinations() ) @pytest.mark.parametrize( "check", @@ -870,13 +870,15 @@ def check_array_api_get_precision(name, estimator, array_namespace, device, dtyp ], ids=_get_check_estimator_ids, ) -def test_pca_array_api_compliance(estimator, check, array_namespace, device, dtype): +def test_pca_array_api_compliance( + estimator, check, array_namespace, device, dtype_name +): name = estimator.__class__.__name__ - check(name, estimator, array_namespace, device=device, dtype=dtype) + check(name, estimator, array_namespace, device=device, dtype_name=dtype_name) @pytest.mark.parametrize( - "array_namespace, device, dtype", yield_namespace_device_dtype_combinations() + "array_namespace, device, dtype_name", yield_namespace_device_dtype_combinations() ) @pytest.mark.parametrize( "check", @@ -892,9 +894,11 @@ def test_pca_array_api_compliance(estimator, check, array_namespace, device, dty ], ids=_get_check_estimator_ids, ) -def test_pca_mle_array_api_compliance(estimator, check, array_namespace, device, dtype): +def test_pca_mle_array_api_compliance( + estimator, check, array_namespace, device, dtype_name +): name = estimator.__class__.__name__ - check(name, estimator, array_namespace, device=device, dtype=dtype) + check(name, estimator, array_namespace, device=device, dtype_name=dtype_name) def test_array_api_error_and_warnings_on_unsupported_params(): diff --git a/sklearn/metrics/tests/test_common.py b/sklearn/metrics/tests/test_common.py index 4f5b10a51a4..8fad63870e4 100644 --- a/sklearn/metrics/tests/test_common.py +++ b/sklearn/metrics/tests/test_common.py @@ -1733,67 +1733,87 @@ def test_metrics_pos_label_error_str(metric, y_pred_threshold, dtype_y_str): def check_array_api_metric( - metric, array_namespace, device, dtype, y_true_np, y_pred_np, sample_weight=None + metric, array_namespace, device, dtype_name, y_true_np, y_pred_np, sample_weight ): xp = _array_api_for_tests(array_namespace, device) + y_true_xp = xp.asarray(y_true_np, device=device) y_pred_xp = xp.asarray(y_pred_np, device=device) metric_np = metric(y_true_np, y_pred_np, sample_weight=sample_weight) + if sample_weight is not None: + sample_weight = xp.asarray(sample_weight, device=device) + with config_context(array_api_dispatch=True): - if sample_weight is not None: - sample_weight = xp.asarray(sample_weight, device=device) metric_xp = metric(y_true_xp, y_pred_xp, sample_weight=sample_weight) assert_allclose( metric_xp, metric_np, - atol=_atol_for_type(dtype), + atol=_atol_for_type(dtype_name), ) def check_array_api_binary_classification_metric( - metric, array_namespace, device, dtype + metric, array_namespace, device, dtype_name ): y_true_np = np.array([0, 0, 1, 1]) y_pred_np = np.array([0, 1, 0, 1]) + check_array_api_metric( - metric, array_namespace, device, dtype, y_true_np=y_true_np, y_pred_np=y_pred_np + metric, + array_namespace, + device, + dtype_name, + y_true_np=y_true_np, + y_pred_np=y_pred_np, + sample_weight=None, + ) + + sample_weight = np.array([0.0, 0.1, 2.0, 1.0], dtype=dtype_name) + + check_array_api_metric( + metric, + array_namespace, + device, + dtype_name, + y_true_np=y_true_np, + y_pred_np=y_pred_np, + sample_weight=sample_weight, ) - if "sample_weight" in signature(metric).parameters: - check_array_api_metric( - metric, - array_namespace, - device, - dtype, - y_true_np=y_true_np, - y_pred_np=y_pred_np, - sample_weight=np.array([0.0, 0.1, 2.0, 1.0]), - ) def check_array_api_multiclass_classification_metric( - metric, array_namespace, device, dtype + metric, array_namespace, device, dtype_name ): y_true_np = np.array([0, 1, 2, 3]) y_pred_np = np.array([0, 1, 0, 2]) + check_array_api_metric( - metric, array_namespace, device, dtype, y_true_np=y_true_np, y_pred_np=y_pred_np + metric, + array_namespace, + device, + dtype_name, + y_true_np=y_true_np, + y_pred_np=y_pred_np, + sample_weight=None, + ) + + sample_weight = np.array([0.0, 0.1, 2.0, 1.0], dtype=dtype_name) + + check_array_api_metric( + metric, + array_namespace, + device, + dtype_name, + y_true_np=y_true_np, + y_pred_np=y_pred_np, + sample_weight=sample_weight, ) - if "sample_weight" in signature(metric).parameters: - check_array_api_metric( - metric, - array_namespace, - device, - dtype, - y_true_np=y_true_np, - y_pred_np=y_pred_np, - sample_weight=np.array([0.0, 0.1, 2.0, 1.0]), - ) -metric_checkers = { +array_api_metric_checkers = { accuracy_score: [ check_array_api_binary_classification_metric, check_array_api_multiclass_classification_metric, @@ -1805,15 +1825,15 @@ metric_checkers = { } -def yield_metric_checker_combinations(metric_checkers=metric_checkers): +def yield_metric_checker_combinations(metric_checkers=array_api_metric_checkers): for metric, checkers in metric_checkers.items(): for checker in checkers: yield metric, checker @pytest.mark.parametrize( - "array_namespace, device, dtype", yield_namespace_device_dtype_combinations() + "array_namespace, device, dtype_name", yield_namespace_device_dtype_combinations() ) @pytest.mark.parametrize("metric, check_func", yield_metric_checker_combinations()) -def test_array_api_compliance(metric, array_namespace, device, dtype, check_func): - check_func(metric, array_namespace, device, dtype) +def test_array_api_compliance(metric, array_namespace, device, dtype_name, check_func): + check_func(metric, array_namespace, device, dtype_name) diff --git a/sklearn/model_selection/tests/test_split.py b/sklearn/model_selection/tests/test_split.py index 94a33cf1a81..57bc6b22351 100644 --- a/sklearn/model_selection/tests/test_split.py +++ b/sklearn/model_selection/tests/test_split.py @@ -1267,7 +1267,7 @@ def test_train_test_split_default_test_size(train_size, exp_train, exp_test): @pytest.mark.parametrize( - "array_namespace, device, dtype", yield_namespace_device_dtype_combinations() + "array_namespace, device, dtype_name", yield_namespace_device_dtype_combinations() ) @pytest.mark.parametrize( "shuffle,stratify", @@ -1278,16 +1278,18 @@ def test_train_test_split_default_test_size(train_size, exp_train, exp_test): (False, None), ), ) -def test_array_api_train_test_split(shuffle, stratify, array_namespace, device, dtype): +def test_array_api_train_test_split( + shuffle, stratify, array_namespace, device, dtype_name +): xp = _array_api_for_tests(array_namespace, device) X = np.arange(100).reshape((10, 10)) y = np.arange(10) - X_np = X.astype(dtype) + X_np = X.astype(dtype_name) X_xp = xp.asarray(X_np, device=device) - y_np = y.astype(dtype) + y_np = y.astype(dtype_name) y_xp = xp.asarray(y_np, device=device) X_train_np, X_test_np, y_train_np, y_test_np = train_test_split( diff --git a/sklearn/preprocessing/tests/test_data.py b/sklearn/preprocessing/tests/test_data.py index 5a70c3091a8..2896e729e97 100644 --- a/sklearn/preprocessing/tests/test_data.py +++ b/sklearn/preprocessing/tests/test_data.py @@ -682,7 +682,7 @@ def test_standard_check_array_of_inverse_transform(): @pytest.mark.parametrize( - "array_namespace, device, dtype", yield_namespace_device_dtype_combinations() + "array_namespace, device, dtype_name", yield_namespace_device_dtype_combinations() ) @pytest.mark.parametrize( "check", @@ -701,9 +701,11 @@ def test_standard_check_array_of_inverse_transform(): ], ids=_get_check_estimator_ids, ) -def test_scaler_array_api_compliance(estimator, check, array_namespace, device, dtype): +def test_scaler_array_api_compliance( + estimator, check, array_namespace, device, dtype_name +): name = estimator.__class__.__name__ - check(name, estimator, array_namespace, device=device, dtype=dtype) + check(name, estimator, array_namespace, device=device, dtype_name=dtype_name) def test_min_max_scaler_iris(): diff --git a/sklearn/utils/_array_api.py b/sklearn/utils/_array_api.py index 6072c0fab85..0c386a843bf 100644 --- a/sklearn/utils/_array_api.py +++ b/sklearn/utils/_array_api.py @@ -24,7 +24,7 @@ def yield_namespace_device_dtype_combinations(): The name of the device on which to allocate the arrays. Can be None to indicate that the default value should be used. - dtype : str + dtype_name : str The name of the data type to use for arrays. Can be None to indicate that the default value should be used. """ @@ -444,7 +444,9 @@ def _weighted_sum(sample_score, sample_weight, normalize=False, xp=None): sample_score = xp.astype(xp.asarray(sample_score, device="cpu"), xp.float64) if sample_weight is not None: - sample_weight = xp.asarray(sample_weight, dtype=sample_score.dtype) + sample_weight = xp.asarray( + sample_weight, dtype=sample_score.dtype, device=device(sample_score) + ) if not xp.isdtype(sample_weight.dtype, "real floating"): sample_weight = xp.astype(sample_weight, xp.float64) diff --git a/sklearn/utils/estimator_checks.py b/sklearn/utils/estimator_checks.py index 28c099441e1..4d87357d688 100644 --- a/sklearn/utils/estimator_checks.py +++ b/sklearn/utils/estimator_checks.py @@ -311,11 +311,15 @@ def _yield_outliers_checks(estimator): def _yield_array_api_checks(estimator): - for array_namespace, device, dtype in yield_namespace_device_dtype_combinations(): + for ( + array_namespace, + device, + dtype_name, + ) in yield_namespace_device_dtype_combinations(): yield partial( check_array_api_input, array_namespace=array_namespace, - dtype=dtype, + dtype_name=dtype_name, device=device, ) @@ -864,7 +868,7 @@ def check_array_api_input( estimator_orig, array_namespace, device=None, - dtype="float64", + dtype_name="float64", check_values=False, ): """Check that the estimator can work consistently with the Array API @@ -878,7 +882,7 @@ def check_array_api_input( xp = _array_api_for_tests(array_namespace, device) X, y = make_classification(random_state=42) - X = X.astype(dtype, copy=False) + X = X.astype(dtype_name, copy=False) X = _enforce_estimator_tags_X(estimator_orig, X) y = _enforce_estimator_tags_y(estimator_orig, y) @@ -1007,14 +1011,14 @@ def check_array_api_input_and_values( estimator_orig, array_namespace, device=None, - dtype="float64", + dtype_name="float64", ): return check_array_api_input( name, estimator_orig, array_namespace=array_namespace, device=device, - dtype=dtype, + dtype_name=dtype_name, check_values=True, ) diff --git a/sklearn/utils/tests/test_array_api.py b/sklearn/utils/tests/test_array_api.py index 01b1f2bf1ad..1df81cf823b 100644 --- a/sklearn/utils/tests/test_array_api.py +++ b/sklearn/utils/tests/test_array_api.py @@ -129,7 +129,7 @@ def test_asarray_with_order_ignored(): @pytest.mark.parametrize( - "array_namespace, device, dtype", yield_namespace_device_dtype_combinations() + "array_namespace, device, dtype_name", yield_namespace_device_dtype_combinations() ) @pytest.mark.parametrize( "sample_weight, normalize, expected", @@ -143,20 +143,20 @@ def test_asarray_with_order_ignored(): ], ) def test_weighted_sum( - array_namespace, device, dtype, sample_weight, normalize, expected + array_namespace, device, dtype_name, sample_weight, normalize, expected ): xp = _array_api_for_tests(array_namespace, device) - sample_score = numpy.asarray([1, 2, 3, 4], dtype=dtype) + sample_score = numpy.asarray([1, 2, 3, 4], dtype=dtype_name) sample_score = xp.asarray(sample_score, device=device) if sample_weight is not None: - sample_weight = numpy.asarray(sample_weight, dtype=dtype) + sample_weight = numpy.asarray(sample_weight, dtype=dtype_name) sample_weight = xp.asarray(sample_weight, device=device) with config_context(array_api_dispatch=True): result = _weighted_sum(sample_score, sample_weight, normalize) assert isinstance(result, float) - assert_allclose(result, expected, atol=_atol_for_type(dtype)) + assert_allclose(result, expected, atol=_atol_for_type(dtype_name)) @skip_if_array_api_compat_not_configured diff --git a/sklearn/utils/tests/test_multiclass.py b/sklearn/utils/tests/test_multiclass.py index d7702ba35cf..6603aca206e 100644 --- a/sklearn/utils/tests/test_multiclass.py +++ b/sklearn/utils/tests/test_multiclass.py @@ -379,17 +379,17 @@ def test_is_multilabel(): @pytest.mark.parametrize( - "array_namespace, device, dtype", + "array_namespace, device, dtype_name", yield_namespace_device_dtype_combinations(), ) -def test_is_multilabel_array_api_compliance(array_namespace, device, dtype): +def test_is_multilabel_array_api_compliance(array_namespace, device, dtype_name): xp = _array_api_for_tests(array_namespace, device) for group, group_examples in ARRAY_API_EXAMPLES.items(): dense_exp = group == "multilabel-indicator" for example in group_examples: if np.asarray(example).dtype.kind == "f": - example = np.asarray(example, dtype=dtype) + example = np.asarray(example, dtype=dtype_name) else: example = np.asarray(example) example = xp.asarray(example, device=device)