As discussed in #33200, confusion_matrix_at_thresholds currently relies on an upcast to _max_precision_float_dtype(xp, device) which is float64 for numpy on CPU and many array API devices for can be limitted to float32 on common devices such as torch MPS (Apple GPUs) or some older Intel GPUs via torch XPI.
Possible solution
I think this code could be changed to improve numerical stability (and maybe processing speed) by implementing either or both of the following two strategies:
When sample_weight is None
Using an integer dtype to perform the call to xp.cumulative_sum instead and only after casting the results to the dtype of y_score.
When sample_weight is not None
Remove the upcast to _max_precision_float_dtype(xp, device) and instead compute the total weight sum and normalize the weight array before calling the cumulative sum. Here is the pseudo-code tps = xp.cumulative_sum(y_true * weight / weigh.sum(), dtype=y_score.dtype) * weight.sum().
Resolution plan
I think the first step is to do a quick empirical study on a float32 only device to confirm that this function actually has a numerical stability problem. Instead of MPS, we could leverage array-api-strict's float32-only devices and call that function on arrays with a large number of elements to try to trigger the problem and then check that the proposed solution actually work as intended. If this is the case we could open two independent PRs, one for the weighted case and one for the unweighted case.
Warning
This issue is not yet ready for a PR. If you are interested in contributing to scikit-learn, please have a look at our contributing guidelines, and in particular the sections for new contributors and the "Needs triage" label.
As discussed in #33200,
confusion_matrix_at_thresholdscurrently relies on an upcast to_max_precision_float_dtype(xp, device)which isfloat64for numpy on CPU and many array API devices for can be limitted tofloat32on common devices such as torch MPS (Apple GPUs) or some older Intel GPUs via torch XPI.Possible solution
I think this code could be changed to improve numerical stability (and maybe processing speed) by implementing either or both of the following two strategies:
When
sample_weight is NoneUsing an integer dtype to perform the call to
xp.cumulative_suminstead and only after casting the results to the dtype ofy_score.When
sample_weight is not NoneRemove the upcast to
_max_precision_float_dtype(xp, device)and instead compute the total weight sum and normalize theweightarray before calling the cumulative sum. Here is the pseudo-codetps = xp.cumulative_sum(y_true * weight / weigh.sum(), dtype=y_score.dtype) * weight.sum().Resolution plan
I think the first step is to do a quick empirical study on a float32 only device to confirm that this function actually has a numerical stability problem. Instead of MPS, we could leverage array-api-strict's float32-only devices and call that function on arrays with a large number of elements to try to trigger the problem and then check that the proposed solution actually work as intended. If this is the case we could open two independent PRs, one for the weighted case and one for the unweighted case.