Skip to content

Numerical stability of confusion_matrix_at_thresholds on float32-only capable array API devices #34813

Description

@ogrisel

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_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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Type

    No type

    Projects

    • Status
      No status

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions