Skip to content

Commit 6099369

Browse files
Fazel94ogrisel
andauthored
FEA Add array API support for matthews_corrcoef (scikit-learn#34423)
Co-authored-by: Olivier Grisel <olivier.grisel@ensta.org>
1 parent b6574d3 commit 6099369

4 files changed

Lines changed: 28 additions & 10 deletions

File tree

doc/modules/array_api.rst

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -217,6 +217,7 @@ Metrics
217217
- :func:`sklearn.metrics.hamming_loss`
218218
- :func:`sklearn.metrics.jaccard_score`
219219
- :func:`sklearn.metrics.log_loss`
220+
- :func:`sklearn.metrics.matthews_corrcoef` (see :ref:`device_support_for_float64`)
220221
- :func:`sklearn.metrics.max_error`
221222
- :func:`sklearn.metrics.mean_absolute_error`
222223
- :func:`sklearn.metrics.mean_absolute_percentage_error`
Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,2 @@
1+
- :func:`sklearn.metrics.matthews_corrcoef` now supports array API compatible inputs.
2+
By :user:`Mohamad Fazeli <Fazel94>`.

sklearn/metrics/_classification.py

Lines changed: 21 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1317,25 +1317,36 @@ def matthews_corrcoef(y_true, y_pred, *, sample_weight=None):
13171317
if y_type not in {"binary", "multiclass"}:
13181318
raise ValueError("%s is not supported" % y_type)
13191319

1320+
xp, _, device_ = get_namespace_and_device(y_true, y_pred)
1321+
13201322
lb = LabelEncoder()
1321-
lb.fit(np.hstack([y_true, y_pred]))
1323+
lb.fit(xp.concat([y_true, y_pred]))
13221324
y_true = lb.transform(y_true)
13231325
y_pred = lb.transform(y_pred)
13241326

13251327
C = confusion_matrix(y_true, y_pred, sample_weight=sample_weight)
1326-
t_sum = C.sum(axis=1, dtype=np.float64)
1327-
p_sum = C.sum(axis=0, dtype=np.float64)
1328-
n_correct = np.trace(C, dtype=np.float64)
1329-
n_samples = p_sum.sum()
1330-
cov_ytyp = n_correct * n_samples - np.dot(t_sum, p_sum)
1331-
cov_ypyp = n_samples**2 - np.dot(p_sum, p_sum)
1332-
cov_ytyt = n_samples**2 - np.dot(t_sum, t_sum)
1328+
# Cast the confusion matrix to the maximum-precision float dtype for two
1329+
# reasons:
1330+
# 1. The covariance terms below reach n_samples**4, which overflows int64
1331+
# for n_samples as small as ~55k (see issue #9622); floats keep the
1332+
# computation exact up to 2**53 and accurate beyond.
1333+
# 2. The array API standard leaves integer-dtype __truediv__
1334+
# implementation-defined, and array_api_strict intentionally raises
1335+
# for it.
1336+
C = xp.astype(C, _max_precision_float_dtype(xp, device=device_), copy=False)
1337+
t_sum = xp.sum(C, axis=1)
1338+
p_sum = xp.sum(C, axis=0)
1339+
n_correct = xp.linalg.trace(C)
1340+
n_samples = xp.sum(p_sum)
1341+
cov_ytyp = n_correct * n_samples - xp.sum(t_sum * p_sum)
1342+
cov_ypyp = n_samples**2 - xp.sum(p_sum * p_sum)
1343+
cov_ytyt = n_samples**2 - xp.sum(t_sum * t_sum)
13331344

13341345
cov_ypyp_ytyt = cov_ypyp * cov_ytyt
1335-
if cov_ypyp_ytyt == 0:
1346+
if float(cov_ypyp_ytyt) == 0.0:
13361347
return 0.0
13371348
else:
1338-
return float(cov_ytyp / np.sqrt(cov_ypyp_ytyt))
1349+
return float(cov_ytyp / xp.sqrt(cov_ypyp_ytyt))
13391350

13401351

13411352
@validate_params(

sklearn/metrics/tests/test_common.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2532,6 +2532,10 @@ def check_array_api_metric_pairwise(metric, array_namespace, device_name, dtype_
25322532
check_array_api_multiclass_classification_metric,
25332533
check_array_api_multilabel_classification_metric,
25342534
],
2535+
matthews_corrcoef: [
2536+
check_array_api_binary_classification_metric,
2537+
check_array_api_multiclass_classification_metric,
2538+
],
25352539
multilabel_confusion_matrix: [
25362540
check_array_api_binary_classification_metric,
25372541
check_array_api_multiclass_classification_metric,

0 commit comments

Comments
 (0)