Aluode/PerceptionLabPortable
0
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 