CoolFace
Apppublic

Aluode/PerceptionLabPortable

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
_encode.py377 linesDownload Raw Back to utils
1# Authors: The scikit-learn developers
2# SPDX-License-Identifier: BSD-3-Clause
3
4from collections import Counter
5from contextlib import suppress
6from typing import NamedTuple
7
8import numpy as np
9
10from ._array_api import (
11    _isin,
12    _searchsorted,
13    device,
14    get_namespace,
15    xpx,
16)
17from ._missing import is_scalar_nan
18
19
20def _unique(values, *, return_inverse=False, return_counts=False):
21    """Helper function to find unique values with support for python objects.
22
23    Uses pure python method for object dtype, and numpy method for
24    all other dtypes.
25
26    Parameters
27    ----------
28    values : ndarray
29        Values to check for unknowns.
30
31    return_inverse : bool, default=False
32        If True, also return the indices of the unique values.
33
34    return_counts : bool, default=False
35        If True, also return the number of times each unique item appears in
36        values.
37
38    Returns
39    -------
40    unique : ndarray
41        The sorted unique values.
42
43    unique_inverse : ndarray
44        The indices to reconstruct the original array from the unique array.
45        Only provided if `return_inverse` is True.
46
47    unique_counts : ndarray
48        The number of times each of the unique values comes up in the original
49        array. Only provided if `return_counts` is True.
50    """
51    if values.dtype == object:
52        return _unique_python(
53            values, return_inverse=return_inverse, return_counts=return_counts
54        )
55    # numerical
56    return _unique_np(
57        values, return_inverse=return_inverse, return_counts=return_counts
58    )
59
60
61def _unique_np(values, return_inverse=False, return_counts=False):
62    """Helper function to find unique values for numpy arrays that correctly
63    accounts for nans. See `_unique` documentation for details."""
64    xp, _ = get_namespace(values)
65
66    inverse, counts = None, None
67
68    if return_inverse and return_counts:
69        uniques, _, inverse, counts = xp.unique_all(values)
70    elif return_inverse:
71        uniques, inverse = xp.unique_inverse(values)
72    elif return_counts:
73        uniques, counts = xp.unique_counts(values)
74    else:
75        uniques = xp.unique_values(values)
76
77    # np.unique will have duplicate missing values at the end of `uniques`
78    # here we clip the nans and remove it from uniques
79    if uniques.size and is_scalar_nan(uniques[-1]):
80        nan_idx = _searchsorted(uniques, xp.nan, xp=xp)
81        uniques = uniques[: nan_idx + 1]
82        if return_inverse:
83            inverse[inverse > nan_idx] = nan_idx
84
85        if return_counts:
86            counts[nan_idx] = xp.sum(counts[nan_idx:])
87            counts = counts[: nan_idx + 1]
88
89    ret = (uniques,)
90
91    if return_inverse:
92        ret += (inverse,)
93
94    if return_counts:
95        ret += (counts,)
96
97    return ret[0] if len(ret) == 1 else ret
98
99
100class MissingValues(NamedTuple):
101    """Data class for missing data information"""
102
103    nan: bool
104    none: bool
105
106    def to_list(self):
107        """Convert tuple to a list where None is always first."""
108        output = []
109        if self.none:
110            output.append(None)
111        if self.nan:
112            output.append(np.nan)
113        return output
114
115
116def _extract_missing(values):
117    """Extract missing values from `values`.
118
119    Parameters
120    ----------
121    values: set
122        Set of values to extract missing from.
123
124    Returns
125    -------
126    output: set
127        Set with missing values extracted.
128
129    missing_values: MissingValues
130        Object with missing value information.
131    """
132    missing_values_set = {
133        value for value in values if value is None or is_scalar_nan(value)
134    }
135
136    if not missing_values_set:
137        return values, MissingValues(nan=False, none=False)
138
139    if None in missing_values_set:
140        if len(missing_values_set) == 1:
141            output_missing_values = MissingValues(nan=False, none=True)
142        else:
143            # If there is more than one missing value, then it has to be
144            # float('nan') or np.nan
145            output_missing_values = MissingValues(nan=True, none=True)
146    else:
147        output_missing_values = MissingValues(nan=True, none=False)
148
149    # create set without the missing values
150    output = values - missing_values_set
151    return output, output_missing_values
152
153
154class _nandict(dict):
155    """Dictionary with support for nans."""
156
157    def __init__(self, mapping):
158        super().__init__(mapping)
159        for key, value in mapping.items():
160            if is_scalar_nan(key):
161                self.nan_value = value
162                break
163
164    def __missing__(self, key):
165        if hasattr(self, "nan_value") and is_scalar_nan(key):
166            return self.nan_value
167        raise KeyError(key)
168
169
170def _map_to_integer(values, uniques):
171    """Map values based on its position in uniques."""
172    xp, _ = get_namespace(values, uniques)
173    table = _nandict({val: i for i, val in enumerate(uniques)})
174    return xp.asarray([table[v] for v in values], device=device(values))
175
176
177def _unique_python(values, *, return_inverse, return_counts):
178    # Only used in `_uniques`, see docstring there for details
179    try:
180        uniques_set = set(values)
181        uniques_set, missing_values = _extract_missing(uniques_set)
182
183        uniques = sorted(uniques_set)
184        uniques.extend(missing_values.to_list())
185        uniques = np.array(uniques, dtype=values.dtype)
186    except TypeError:
187        types = sorted(t.__qualname__ for t in set(type(v) for v in values))
188        raise TypeError(
189            "Encoders require their input argument must be uniformly "
190            f"strings or numbers. Got {types}"
191        )
192    ret = (uniques,)
193
194    if return_inverse:
195        ret += (_map_to_integer(values, uniques),)
196
197    if return_counts:
198        ret += (_get_counts(values, uniques),)
199
200    return ret[0] if len(ret) == 1 else ret
201
202
203def _encode(values, *, uniques, check_unknown=True):
204    """Helper function to encode values into [0, n_uniques - 1].
205
206    Uses pure python method for object dtype, and numpy method for
207    all other dtypes.
208    The numpy method has the limitation that the `uniques` need to
209    be sorted. Importantly, this is not checked but assumed to already be
210    the case. The calling method needs to ensure this for all non-object
211    values.
212
213    Parameters
214    ----------
215    values : ndarray
216        Values to encode.
217    uniques : ndarray
218        The unique values in `values`. If the dtype is not object, then
219        `uniques` needs to be sorted.
220    check_unknown : bool, default=True
221        If True, check for values in `values` that are not in `unique`
222        and raise an error. This is ignored for object dtype, and treated as
223        True in this case. This parameter is useful for
224        _BaseEncoder._transform() to avoid calling _check_unknown()
225        twice.
226
227    Returns
228    -------
229    encoded : ndarray
230        Encoded values
231    """
232    xp, _ = get_namespace(values, uniques)
233    if not xp.isdtype(values.dtype, "numeric"):
234        try:
235            return _map_to_integer(values, uniques)
236        except KeyError as e:
237            raise ValueError(f"y contains previously unseen labels: {e}")
238    else:
239        if check_unknown:
240            diff = _check_unknown(values, uniques)
241            if diff:
242                raise ValueError(f"y contains previously unseen labels: {diff}")
243        return _searchsorted(uniques, values, xp=xp)
244
245
246def _check_unknown(values, known_values, return_mask=False):
247    """
248    Helper function to check for unknowns in values to be encoded.
249
250    Uses pure python method for object dtype, and numpy method for
251    all other dtypes.
252
253    Parameters
254    ----------
255    values : array
256        Values to check for unknowns.
257    known_values : array
258        Known values. Must be unique.
259    return_mask : bool, default=False
260        If True, return a mask of the same shape as `values` indicating
261        the valid values.
262
263    Returns
264    -------
265    diff : list
266        The unique values present in `values` and not in `know_values`.
267    valid_mask : boolean array
268        Additionally returned if ``return_mask=True``.
269
270    """
271    xp, _ = get_namespace(values, known_values)
272    valid_mask = None
273
274    if not xp.isdtype(values.dtype, "numeric"):
275        values_set = set(values)
276        values_set, missing_in_values = _extract_missing(values_set)
277
278        uniques_set = set(known_values)
279        uniques_set, missing_in_uniques = _extract_missing(uniques_set)
280        diff = values_set - uniques_set
281
282        nan_in_diff = missing_in_values.nan and not missing_in_uniques.nan
283        none_in_diff = missing_in_values.none and not missing_in_uniques.none
284
285        def is_valid(value):
286            return (
287                value in uniques_set
288                or (missing_in_uniques.none and value is None)
289                or (missing_in_uniques.nan and is_scalar_nan(value))
290            )
291
292        if return_mask:
293            if diff or nan_in_diff or none_in_diff:
294                valid_mask = xp.array([is_valid(value) for value in values])
295            else:
296                valid_mask = xp.ones(len(values), dtype=xp.bool)
297
298        diff = list(diff)
299        if none_in_diff:
300            diff.append(None)
301        if nan_in_diff:
302            diff.append(np.nan)
303    else:
304        unique_values = xp.unique_values(values)
305        diff = xpx.setdiff1d(unique_values, known_values, assume_unique=True, xp=xp)
306        if return_mask:
307            if diff.size:
308                valid_mask = _isin(values, known_values, xp)
309            else:
310                valid_mask = xp.ones(len(values), dtype=xp.bool)
311
312        # check for nans in the known_values
313        if xp.any(xp.isnan(known_values)):
314            diff_is_nan = xp.isnan(diff)
315            if xp.any(diff_is_nan):
316                # removes nan from valid_mask
317                if diff.size and return_mask:
318                    is_nan = xp.isnan(values)
319                    valid_mask[is_nan] = 1
320
321                # remove nan from diff
322                diff = diff[~diff_is_nan]
323        diff = list(diff)
324
325    if return_mask:
326        return diff, valid_mask
327    return diff
328
329
330class _NaNCounter(Counter):
331    """Counter with support for nan values."""
332
333    def __init__(self, items):
334        super().__init__(self._generate_items(items))
335
336    def _generate_items(self, items):
337        """Generate items without nans. Stores the nan counts separately."""
338        for item in items:
339            if not is_scalar_nan(item):
340                yield item
341                continue
342            if not hasattr(self, "nan_count"):
343                self.nan_count = 0
344            self.nan_count += 1
345
346    def __missing__(self, key):
347        if hasattr(self, "nan_count") and is_scalar_nan(key):
348            return self.nan_count
349        raise KeyError(key)
350
351
352def _get_counts(values, uniques):
353    """Get the count of each of the `uniques` in `values`.
354
355    The counts will use the order passed in by `uniques`. For non-object dtypes,
356    `uniques` is assumed to be sorted and `np.nan` is at the end.
357    """
358    if values.dtype.kind in "OU":
359        counter = _NaNCounter(values)
360        output = np.zeros(len(uniques), dtype=np.int64)
361        for i, item in enumerate(uniques):
362            with suppress(KeyError):
363                output[i] = counter[item]
364        return output
365
366    unique_values, counts = _unique_np(values, return_counts=True)
367
368    # Recorder unique_values based on input: `uniques`
369    uniques_in_values = np.isin(uniques, unique_values, assume_unique=True)
370    if np.isnan(unique_values[-1]) and np.isnan(uniques[-1]):
371        uniques_in_values[-1] = True
372
373    unique_valid_indices = np.searchsorted(unique_values, uniques[uniques_in_values])
374    output = np.zeros_like(uniques, dtype=np.int64)
375    output[uniques_in_values] = counts[unique_valid_indices]
376    return output
377