Prompt48 commited on
Commit
ad17684
·
verified ·
1 Parent(s): dfdfaa7

Upload edit\Qwen3-TTS-test\.venv\Lib\site-packages\sklearn\externals\array_api_compat\common\_helpers.py with huggingface_hub

Browse files
edit//Qwen3-TTS-test//.venv//Lib//site-packages//sklearn//externals//array_api_compat//common//_helpers.py ADDED
@@ -0,0 +1,1058 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Various helper functions which are not part of the spec.
3
+
4
+ Functions which start with an underscore are for internal use only but helpers
5
+ that are in __all__ are intended as additional helper functions for use by end
6
+ users of the compat library.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import inspect
12
+ import math
13
+ import sys
14
+ import warnings
15
+ from collections.abc import Collection, Hashable
16
+ from functools import lru_cache
17
+ from typing import (
18
+ TYPE_CHECKING,
19
+ Any,
20
+ Final,
21
+ Literal,
22
+ SupportsIndex,
23
+ TypeAlias,
24
+ TypeGuard,
25
+ TypeVar,
26
+ cast,
27
+ overload,
28
+ )
29
+
30
+ from ._typing import Array, Device, HasShape, Namespace, SupportsArrayNamespace
31
+
32
+ if TYPE_CHECKING:
33
+
34
+ import dask.array as da
35
+ import jax
36
+ import ndonnx as ndx
37
+ import numpy as np
38
+ import numpy.typing as npt
39
+ import sparse # pyright: ignore[reportMissingTypeStubs]
40
+ import torch
41
+
42
+ # TODO: import from typing (requires Python >=3.13)
43
+ from typing_extensions import TypeIs, TypeVar
44
+
45
+ _SizeT = TypeVar("_SizeT", bound = int | None)
46
+
47
+ _ZeroGradientArray: TypeAlias = npt.NDArray[np.void]
48
+ _CupyArray: TypeAlias = Any # cupy has no py.typed
49
+
50
+ _ArrayApiObj: TypeAlias = (
51
+ npt.NDArray[Any]
52
+ | da.Array
53
+ | jax.Array
54
+ | ndx.Array
55
+ | sparse.SparseArray
56
+ | torch.Tensor
57
+ | SupportsArrayNamespace[Any]
58
+ | _CupyArray
59
+ )
60
+
61
+ _API_VERSIONS_OLD: Final = frozenset({"2021.12", "2022.12", "2023.12"})
62
+ _API_VERSIONS: Final = _API_VERSIONS_OLD | frozenset({"2024.12"})
63
+
64
+
65
+ @lru_cache(100)
66
+ def _issubclass_fast(cls: type, modname: str, clsname: str) -> bool:
67
+ try:
68
+ mod = sys.modules[modname]
69
+ except KeyError:
70
+ return False
71
+ parent_cls = getattr(mod, clsname)
72
+ return issubclass(cls, parent_cls)
73
+
74
+
75
+ def _is_jax_zero_gradient_array(x: object) -> TypeGuard[_ZeroGradientArray]:
76
+ """Return True if `x` is a zero-gradient array.
77
+
78
+ These arrays are a design quirk of Jax that may one day be removed.
79
+ See https://github.com/google/jax/issues/20620.
80
+ """
81
+ # Fast exit
82
+ try:
83
+ dtype = x.dtype # type: ignore[attr-defined]
84
+ except AttributeError:
85
+ return False
86
+ cls = cast(Hashable, type(dtype))
87
+ if not _issubclass_fast(cls, "numpy.dtypes", "VoidDType"):
88
+ return False
89
+
90
+ if "jax" not in sys.modules:
91
+ return False
92
+
93
+ import jax
94
+ # jax.float0 is a np.dtype([('float0', 'V')])
95
+ return dtype == jax.float0
96
+
97
+
98
+ def is_numpy_array(x: object) -> TypeGuard[npt.NDArray[Any]]:
99
+ """
100
+ Return True if `x` is a NumPy array.
101
+
102
+ This function does not import NumPy if it has not already been imported
103
+ and is therefore cheap to use.
104
+
105
+ This also returns True for `ndarray` subclasses and NumPy scalar objects.
106
+
107
+ See Also
108
+ --------
109
+
110
+ array_namespace
111
+ is_array_api_obj
112
+ is_cupy_array
113
+ is_torch_array
114
+ is_ndonnx_array
115
+ is_dask_array
116
+ is_jax_array
117
+ is_pydata_sparse_array
118
+ """
119
+ # TODO: Should we reject ndarray subclasses?
120
+ cls = cast(Hashable, type(x))
121
+ return (
122
+ _issubclass_fast(cls, "numpy", "ndarray")
123
+ or _issubclass_fast(cls, "numpy", "generic")
124
+ ) and not _is_jax_zero_gradient_array(x)
125
+
126
+
127
+ def is_cupy_array(x: object) -> bool:
128
+ """
129
+ Return True if `x` is a CuPy array.
130
+
131
+ This function does not import CuPy if it has not already been imported
132
+ and is therefore cheap to use.
133
+
134
+ This also returns True for `cupy.ndarray` subclasses and CuPy scalar objects.
135
+
136
+ See Also
137
+ --------
138
+
139
+ array_namespace
140
+ is_array_api_obj
141
+ is_numpy_array
142
+ is_torch_array
143
+ is_ndonnx_array
144
+ is_dask_array
145
+ is_jax_array
146
+ is_pydata_sparse_array
147
+ """
148
+ cls = cast(Hashable, type(x))
149
+ return _issubclass_fast(cls, "cupy", "ndarray")
150
+
151
+
152
+ def is_torch_array(x: object) -> TypeIs[torch.Tensor]:
153
+ """
154
+ Return True if `x` is a PyTorch tensor.
155
+
156
+ This function does not import PyTorch if it has not already been imported
157
+ and is therefore cheap to use.
158
+
159
+ See Also
160
+ --------
161
+
162
+ array_namespace
163
+ is_array_api_obj
164
+ is_numpy_array
165
+ is_cupy_array
166
+ is_dask_array
167
+ is_jax_array
168
+ is_pydata_sparse_array
169
+ """
170
+ cls = cast(Hashable, type(x))
171
+ return _issubclass_fast(cls, "torch", "Tensor")
172
+
173
+
174
+ def is_ndonnx_array(x: object) -> TypeIs[ndx.Array]:
175
+ """
176
+ Return True if `x` is a ndonnx Array.
177
+
178
+ This function does not import ndonnx if it has not already been imported
179
+ and is therefore cheap to use.
180
+
181
+ See Also
182
+ --------
183
+
184
+ array_namespace
185
+ is_array_api_obj
186
+ is_numpy_array
187
+ is_cupy_array
188
+ is_ndonnx_array
189
+ is_dask_array
190
+ is_jax_array
191
+ is_pydata_sparse_array
192
+ """
193
+ cls = cast(Hashable, type(x))
194
+ return _issubclass_fast(cls, "ndonnx", "Array")
195
+
196
+
197
+ def is_dask_array(x: object) -> TypeIs[da.Array]:
198
+ """
199
+ Return True if `x` is a dask.array Array.
200
+
201
+ This function does not import dask if it has not already been imported
202
+ and is therefore cheap to use.
203
+
204
+ See Also
205
+ --------
206
+
207
+ array_namespace
208
+ is_array_api_obj
209
+ is_numpy_array
210
+ is_cupy_array
211
+ is_torch_array
212
+ is_ndonnx_array
213
+ is_jax_array
214
+ is_pydata_sparse_array
215
+ """
216
+ cls = cast(Hashable, type(x))
217
+ return _issubclass_fast(cls, "dask.array", "Array")
218
+
219
+
220
+ def is_jax_array(x: object) -> TypeIs[jax.Array]:
221
+ """
222
+ Return True if `x` is a JAX array.
223
+
224
+ This function does not import JAX if it has not already been imported
225
+ and is therefore cheap to use.
226
+
227
+
228
+ See Also
229
+ --------
230
+
231
+ array_namespace
232
+ is_array_api_obj
233
+ is_numpy_array
234
+ is_cupy_array
235
+ is_torch_array
236
+ is_ndonnx_array
237
+ is_dask_array
238
+ is_pydata_sparse_array
239
+ """
240
+ cls = cast(Hashable, type(x))
241
+ return _issubclass_fast(cls, "jax", "Array") or _is_jax_zero_gradient_array(x)
242
+
243
+
244
+ def is_pydata_sparse_array(x: object) -> TypeIs[sparse.SparseArray]:
245
+ """
246
+ Return True if `x` is an array from the `sparse` package.
247
+
248
+ This function does not import `sparse` if it has not already been imported
249
+ and is therefore cheap to use.
250
+
251
+
252
+ See Also
253
+ --------
254
+
255
+ array_namespace
256
+ is_array_api_obj
257
+ is_numpy_array
258
+ is_cupy_array
259
+ is_torch_array
260
+ is_ndonnx_array
261
+ is_dask_array
262
+ is_jax_array
263
+ """
264
+ # TODO: Account for other backends.
265
+ cls = cast(Hashable, type(x))
266
+ return _issubclass_fast(cls, "sparse", "SparseArray")
267
+
268
+
269
+ def is_array_api_obj(x: object) -> TypeIs[_ArrayApiObj]: # pyright: ignore[reportUnknownParameterType]
270
+ """
271
+ Return True if `x` is an array API compatible array object.
272
+
273
+ See Also
274
+ --------
275
+
276
+ array_namespace
277
+ is_numpy_array
278
+ is_cupy_array
279
+ is_torch_array
280
+ is_ndonnx_array
281
+ is_dask_array
282
+ is_jax_array
283
+ """
284
+ return (
285
+ hasattr(x, '__array_namespace__')
286
+ or _is_array_api_cls(cast(Hashable, type(x)))
287
+ )
288
+
289
+
290
+ @lru_cache(100)
291
+ def _is_array_api_cls(cls: type) -> bool:
292
+ return (
293
+ # TODO: drop support for numpy<2 which didn't have __array_namespace__
294
+ _issubclass_fast(cls, "numpy", "ndarray")
295
+ or _issubclass_fast(cls, "numpy", "generic")
296
+ or _issubclass_fast(cls, "cupy", "ndarray")
297
+ or _issubclass_fast(cls, "torch", "Tensor")
298
+ or _issubclass_fast(cls, "dask.array", "Array")
299
+ or _issubclass_fast(cls, "sparse", "SparseArray")
300
+ # TODO: drop support for jax<0.4.32 which didn't have __array_namespace__
301
+ or _issubclass_fast(cls, "jax", "Array")
302
+ )
303
+
304
+
305
+ def _compat_module_name() -> str:
306
+ assert __name__.endswith(".common._helpers")
307
+ return __name__.removesuffix(".common._helpers")
308
+
309
+
310
+ @lru_cache(100)
311
+ def is_numpy_namespace(xp: Namespace) -> bool:
312
+ """
313
+ Returns True if `xp` is a NumPy namespace.
314
+
315
+ This includes both NumPy itself and the version wrapped by array-api-compat.
316
+
317
+ See Also
318
+ --------
319
+
320
+ array_namespace
321
+ is_cupy_namespace
322
+ is_torch_namespace
323
+ is_ndonnx_namespace
324
+ is_dask_namespace
325
+ is_jax_namespace
326
+ is_pydata_sparse_namespace
327
+ is_array_api_strict_namespace
328
+ """
329
+ return xp.__name__ in {"numpy", _compat_module_name() + ".numpy"}
330
+
331
+
332
+ @lru_cache(100)
333
+ def is_cupy_namespace(xp: Namespace) -> bool:
334
+ """
335
+ Returns True if `xp` is a CuPy namespace.
336
+
337
+ This includes both CuPy itself and the version wrapped by array-api-compat.
338
+
339
+ See Also
340
+ --------
341
+
342
+ array_namespace
343
+ is_numpy_namespace
344
+ is_torch_namespace
345
+ is_ndonnx_namespace
346
+ is_dask_namespace
347
+ is_jax_namespace
348
+ is_pydata_sparse_namespace
349
+ is_array_api_strict_namespace
350
+ """
351
+ return xp.__name__ in {"cupy", _compat_module_name() + ".cupy"}
352
+
353
+
354
+ @lru_cache(100)
355
+ def is_torch_namespace(xp: Namespace) -> bool:
356
+ """
357
+ Returns True if `xp` is a PyTorch namespace.
358
+
359
+ This includes both PyTorch itself and the version wrapped by array-api-compat.
360
+
361
+ See Also
362
+ --------
363
+
364
+ array_namespace
365
+ is_numpy_namespace
366
+ is_cupy_namespace
367
+ is_ndonnx_namespace
368
+ is_dask_namespace
369
+ is_jax_namespace
370
+ is_pydata_sparse_namespace
371
+ is_array_api_strict_namespace
372
+ """
373
+ return xp.__name__ in {"torch", _compat_module_name() + ".torch"}
374
+
375
+
376
+ def is_ndonnx_namespace(xp: Namespace) -> bool:
377
+ """
378
+ Returns True if `xp` is an NDONNX namespace.
379
+
380
+ See Also
381
+ --------
382
+
383
+ array_namespace
384
+ is_numpy_namespace
385
+ is_cupy_namespace
386
+ is_torch_namespace
387
+ is_dask_namespace
388
+ is_jax_namespace
389
+ is_pydata_sparse_namespace
390
+ is_array_api_strict_namespace
391
+ """
392
+ return xp.__name__ == "ndonnx"
393
+
394
+
395
+ @lru_cache(100)
396
+ def is_dask_namespace(xp: Namespace) -> bool:
397
+ """
398
+ Returns True if `xp` is a Dask namespace.
399
+
400
+ This includes both ``dask.array`` itself and the version wrapped by array-api-compat.
401
+
402
+ See Also
403
+ --------
404
+
405
+ array_namespace
406
+ is_numpy_namespace
407
+ is_cupy_namespace
408
+ is_torch_namespace
409
+ is_ndonnx_namespace
410
+ is_jax_namespace
411
+ is_pydata_sparse_namespace
412
+ is_array_api_strict_namespace
413
+ """
414
+ return xp.__name__ in {"dask.array", _compat_module_name() + ".dask.array"}
415
+
416
+
417
+ def is_jax_namespace(xp: Namespace) -> bool:
418
+ """
419
+ Returns True if `xp` is a JAX namespace.
420
+
421
+ This includes ``jax.numpy`` and ``jax.experimental.array_api`` which existed in
422
+ older versions of JAX.
423
+
424
+ See Also
425
+ --------
426
+
427
+ array_namespace
428
+ is_numpy_namespace
429
+ is_cupy_namespace
430
+ is_torch_namespace
431
+ is_ndonnx_namespace
432
+ is_dask_namespace
433
+ is_pydata_sparse_namespace
434
+ is_array_api_strict_namespace
435
+ """
436
+ return xp.__name__ in {"jax.numpy", "jax.experimental.array_api"}
437
+
438
+
439
+ def is_pydata_sparse_namespace(xp: Namespace) -> bool:
440
+ """
441
+ Returns True if `xp` is a pydata/sparse namespace.
442
+
443
+ See Also
444
+ --------
445
+
446
+ array_namespace
447
+ is_numpy_namespace
448
+ is_cupy_namespace
449
+ is_torch_namespace
450
+ is_ndonnx_namespace
451
+ is_dask_namespace
452
+ is_jax_namespace
453
+ is_array_api_strict_namespace
454
+ """
455
+ return xp.__name__ == "sparse"
456
+
457
+
458
+ def is_array_api_strict_namespace(xp: Namespace) -> bool:
459
+ """
460
+ Returns True if `xp` is an array-api-strict namespace.
461
+
462
+ See Also
463
+ --------
464
+
465
+ array_namespace
466
+ is_numpy_namespace
467
+ is_cupy_namespace
468
+ is_torch_namespace
469
+ is_ndonnx_namespace
470
+ is_dask_namespace
471
+ is_jax_namespace
472
+ is_pydata_sparse_namespace
473
+ """
474
+ return xp.__name__ == "array_api_strict"
475
+
476
+
477
+ def _check_api_version(api_version: str | None) -> None:
478
+ if api_version in _API_VERSIONS_OLD:
479
+ warnings.warn(
480
+ f"The {api_version} version of the array API specification was requested but the returned namespace is actually version 2024.12"
481
+ )
482
+ elif api_version is not None and api_version not in _API_VERSIONS:
483
+ raise ValueError(
484
+ "Only the 2024.12 version of the array API specification is currently supported"
485
+ )
486
+
487
+
488
+ def array_namespace(
489
+ *xs: Array | complex | None,
490
+ api_version: str | None = None,
491
+ use_compat: bool | None = None,
492
+ ) -> Namespace:
493
+ """
494
+ Get the array API compatible namespace for the arrays `xs`.
495
+
496
+ Parameters
497
+ ----------
498
+ xs: arrays
499
+ one or more arrays. xs can also be Python scalars (bool, int, float,
500
+ complex, or None), which are ignored.
501
+
502
+ api_version: str
503
+ The newest version of the spec that you need support for (currently
504
+ the compat library wrapped APIs support v2024.12).
505
+
506
+ use_compat: bool or None
507
+ If None (the default), the native namespace will be returned if it is
508
+ already array API compatible, otherwise a compat wrapper is used. If
509
+ True, the compat library wrapped library will be returned. If False,
510
+ the native library namespace is returned.
511
+
512
+ Returns
513
+ -------
514
+
515
+ out: namespace
516
+ The array API compatible namespace corresponding to the arrays in `xs`.
517
+
518
+ Raises
519
+ ------
520
+ TypeError
521
+ If `xs` contains arrays from different array libraries or contains a
522
+ non-array.
523
+
524
+
525
+ Typical usage is to pass the arguments of a function to
526
+ `array_namespace()` at the top of a function to get the corresponding
527
+ array API namespace:
528
+
529
+ .. code:: python
530
+
531
+ def your_function(x, y):
532
+ xp = array_api_compat.array_namespace(x, y)
533
+ # Now use xp as the array library namespace
534
+ return xp.mean(x, axis=0) + 2*xp.std(y, axis=0)
535
+
536
+
537
+ Wrapped array namespaces can also be imported directly. For example,
538
+ `array_namespace(np.array(...))` will return `array_api_compat.numpy`.
539
+ This function will also work for any array library not wrapped by
540
+ array-api-compat if it explicitly defines `__array_namespace__
541
+ <https://data-apis.org/array-api/latest/API_specification/generated/array_api.array.__array_namespace__.html>`__
542
+ (the wrapped namespace is always preferred if it exists).
543
+
544
+ See Also
545
+ --------
546
+
547
+ is_array_api_obj
548
+ is_numpy_array
549
+ is_cupy_array
550
+ is_torch_array
551
+ is_dask_array
552
+ is_jax_array
553
+ is_pydata_sparse_array
554
+
555
+ """
556
+ if use_compat not in [None, True, False]:
557
+ raise ValueError("use_compat must be None, True, or False")
558
+
559
+ _use_compat = use_compat in [None, True]
560
+
561
+ namespaces: set[Namespace] = set()
562
+ for x in xs:
563
+ if is_numpy_array(x):
564
+ import numpy as np
565
+
566
+ from .. import numpy as numpy_namespace
567
+
568
+ if use_compat is True:
569
+ _check_api_version(api_version)
570
+ namespaces.add(numpy_namespace)
571
+ elif use_compat is False:
572
+ namespaces.add(np)
573
+ else:
574
+ # numpy 2.0+ have __array_namespace__, however, they are not yet fully array API
575
+ # compatible.
576
+ namespaces.add(numpy_namespace)
577
+ elif is_cupy_array(x):
578
+ if _use_compat:
579
+ _check_api_version(api_version)
580
+ from .. import cupy as cupy_namespace
581
+
582
+ namespaces.add(cupy_namespace)
583
+ else:
584
+ import cupy as cp # pyright: ignore[reportMissingTypeStubs]
585
+
586
+ namespaces.add(cp)
587
+ elif is_torch_array(x):
588
+ if _use_compat:
589
+ _check_api_version(api_version)
590
+ from .. import torch as torch_namespace
591
+
592
+ namespaces.add(torch_namespace)
593
+ else:
594
+ import torch
595
+
596
+ namespaces.add(torch)
597
+ elif is_dask_array(x):
598
+ if _use_compat:
599
+ _check_api_version(api_version)
600
+ from ..dask import array as dask_namespace
601
+
602
+ namespaces.add(dask_namespace)
603
+ else:
604
+ import dask.array as da
605
+
606
+ namespaces.add(da)
607
+ elif is_jax_array(x):
608
+ if use_compat is True:
609
+ _check_api_version(api_version)
610
+ raise ValueError("JAX does not have an array-api-compat wrapper")
611
+ elif use_compat is False:
612
+ import jax.numpy as jnp
613
+ else:
614
+ # JAX v0.4.32 and newer implements the array API directly in jax.numpy.
615
+ # For older JAX versions, it is available via jax.experimental.array_api.
616
+ import jax.numpy
617
+
618
+ if hasattr(jax.numpy, "__array_api_version__"):
619
+ jnp = jax.numpy
620
+ else:
621
+ import jax.experimental.array_api as jnp # pyright: ignore[reportMissingImports]
622
+ namespaces.add(jnp)
623
+ elif is_pydata_sparse_array(x):
624
+ if use_compat is True:
625
+ _check_api_version(api_version)
626
+ raise ValueError("`sparse` does not have an array-api-compat wrapper")
627
+ else:
628
+ import sparse # pyright: ignore[reportMissingTypeStubs]
629
+ # `sparse` is already an array namespace. We do not have a wrapper
630
+ # submodule for it.
631
+ namespaces.add(sparse)
632
+ elif hasattr(x, "__array_namespace__"):
633
+ if use_compat is True:
634
+ raise ValueError(
635
+ "The given array does not have an array-api-compat wrapper"
636
+ )
637
+ x = cast("SupportsArrayNamespace[Any]", x)
638
+ namespaces.add(x.__array_namespace__(api_version=api_version))
639
+ elif isinstance(x, (bool, int, float, complex, type(None))):
640
+ continue
641
+ else:
642
+ # TODO: Support Python scalars?
643
+ raise TypeError(f"{type(x).__name__} is not a supported array type")
644
+
645
+ if not namespaces:
646
+ raise TypeError("Unrecognized array input")
647
+
648
+ if len(namespaces) != 1:
649
+ raise TypeError(f"Multiple namespaces for array inputs: {namespaces}")
650
+
651
+ (xp,) = namespaces
652
+
653
+ return xp
654
+
655
+
656
+ # backwards compatibility alias
657
+ get_namespace = array_namespace
658
+
659
+
660
+ def _check_device(bare_xp: Namespace, device: Device) -> None: # pyright: ignore[reportUnusedFunction]
661
+ """
662
+ Validate dummy device on device-less array backends.
663
+
664
+ Notes
665
+ -----
666
+ This function is also invoked by CuPy, which does have multiple devices
667
+ if there are multiple GPUs available.
668
+ However, CuPy multi-device support is currently impossible
669
+ without using the global device or a context manager:
670
+
671
+ https://github.com/data-apis/array-api-compat/pull/293
672
+ """
673
+ if bare_xp is sys.modules.get("numpy"):
674
+ if device not in ("cpu", None):
675
+ raise ValueError(f"Unsupported device for NumPy: {device!r}")
676
+
677
+ elif bare_xp is sys.modules.get("dask.array"):
678
+ if device not in ("cpu", _DASK_DEVICE, None):
679
+ raise ValueError(f"Unsupported device for Dask: {device!r}")
680
+
681
+
682
+ # Placeholder object to represent the dask device
683
+ # when the array backend is not the CPU.
684
+ # (since it is not easy to tell which device a dask array is on)
685
+ class _dask_device:
686
+ def __repr__(self) -> Literal["DASK_DEVICE"]:
687
+ return "DASK_DEVICE"
688
+
689
+
690
+ _DASK_DEVICE = _dask_device()
691
+
692
+
693
+ # device() is not on numpy.ndarray or dask.array and to_device() is not on numpy.ndarray
694
+ # or cupy.ndarray. They are not included in array objects of this library
695
+ # because this library just reuses the respective ndarray classes without
696
+ # wrapping or subclassing them. These helper functions can be used instead of
697
+ # the wrapper functions for libraries that need to support both NumPy/CuPy and
698
+ # other libraries that use devices.
699
+ def device(x: _ArrayApiObj, /) -> Device:
700
+ """
701
+ Hardware device the array data resides on.
702
+
703
+ This is equivalent to `x.device` according to the `standard
704
+ <https://data-apis.org/array-api/latest/API_specification/generated/array_api.array.device.html>`__.
705
+ This helper is included because some array libraries either do not have
706
+ the `device` attribute or include it with an incompatible API.
707
+
708
+ Parameters
709
+ ----------
710
+ x: array
711
+ array instance from an array API compatible library.
712
+
713
+ Returns
714
+ -------
715
+ out: device
716
+ a ``device`` object (see the `Device Support <https://data-apis.org/array-api/latest/design_topics/device_support.html>`__
717
+ section of the array API specification).
718
+
719
+ Notes
720
+ -----
721
+
722
+ For NumPy the device is always `"cpu"`. For Dask, the device is always a
723
+ special `DASK_DEVICE` object.
724
+
725
+ See Also
726
+ --------
727
+
728
+ to_device : Move array data to a different device.
729
+
730
+ """
731
+ if is_numpy_array(x):
732
+ return "cpu"
733
+ elif is_dask_array(x):
734
+ # Peek at the metadata of the Dask array to determine type
735
+ if is_numpy_array(x._meta): # pyright: ignore
736
+ # Must be on CPU since backed by numpy
737
+ return "cpu"
738
+ return _DASK_DEVICE
739
+ elif is_jax_array(x):
740
+ # FIXME Jitted JAX arrays do not have a device attribute
741
+ # https://github.com/jax-ml/jax/issues/26000
742
+ # Return None in this case. Note that this workaround breaks
743
+ # the standard and will result in new arrays being created on the
744
+ # default device instead of the same device as the input array(s).
745
+ x_device = getattr(x, "device", None)
746
+ # Older JAX releases had .device() as a method, which has been replaced
747
+ # with a property in accordance with the standard.
748
+ if inspect.ismethod(x_device):
749
+ return x_device()
750
+ else:
751
+ return x_device
752
+ elif is_pydata_sparse_array(x):
753
+ # `sparse` will gain `.device`, so check for this first.
754
+ x_device = getattr(x, "device", None)
755
+ if x_device is not None:
756
+ return x_device
757
+ # Everything but DOK has this attr.
758
+ try:
759
+ inner = x.data # pyright: ignore
760
+ except AttributeError:
761
+ return "cpu"
762
+ # Return the device of the constituent array
763
+ return device(inner) # pyright: ignore
764
+ return x.device # pyright: ignore
765
+
766
+
767
+ # Prevent shadowing, used below
768
+ _device = device
769
+
770
+
771
+ # Based on cupy.array_api.Array.to_device
772
+ def _cupy_to_device(
773
+ x: _CupyArray,
774
+ device: Device,
775
+ /,
776
+ stream: int | Any | None = None,
777
+ ) -> _CupyArray:
778
+ import cupy as cp
779
+
780
+ if device == "cpu":
781
+ # allowing us to use `to_device(x, "cpu")`
782
+ # is useful for portable test swapping between
783
+ # host and device backends
784
+ return x.get()
785
+ if not isinstance(device, cp.cuda.Device):
786
+ raise TypeError(f"Unsupported device type {device!r}")
787
+
788
+ if stream is None:
789
+ with device:
790
+ return cp.asarray(x)
791
+
792
+ # stream can be an int as specified in __dlpack__, or a CuPy stream
793
+ if isinstance(stream, int):
794
+ stream = cp.cuda.ExternalStream(stream)
795
+ elif not isinstance(stream, cp.cuda.Stream):
796
+ raise TypeError(f"Unsupported stream type {stream!r}")
797
+
798
+ with device, stream:
799
+ return cp.asarray(x)
800
+
801
+
802
+ def _torch_to_device(
803
+ x: torch.Tensor,
804
+ device: torch.device | str | int,
805
+ /,
806
+ stream: None = None,
807
+ ) -> torch.Tensor:
808
+ if stream is not None:
809
+ raise NotImplementedError
810
+ return x.to(device)
811
+
812
+
813
+ def to_device(x: Array, device: Device, /, *, stream: int | Any | None = None) -> Array:
814
+ """
815
+ Copy the array from the device on which it currently resides to the specified ``device``.
816
+
817
+ This is equivalent to `x.to_device(device, stream=stream)` according to
818
+ the `standard
819
+ <https://data-apis.org/array-api/latest/API_specification/generated/array_api.array.to_device.html>`__.
820
+ This helper is included because some array libraries do not have the
821
+ `to_device` method.
822
+
823
+ Parameters
824
+ ----------
825
+
826
+ x: array
827
+ array instance from an array API compatible library.
828
+
829
+ device: device
830
+ a ``device`` object (see the `Device Support <https://data-apis.org/array-api/latest/design_topics/device_support.html>`__
831
+ section of the array API specification).
832
+
833
+ stream: int | Any | None
834
+ stream object to use during copy. In addition to the types supported
835
+ in ``array.__dlpack__``, implementations may choose to support any
836
+ library-specific stream object with the caveat that any code using
837
+ such an object would not be portable.
838
+
839
+ Returns
840
+ -------
841
+
842
+ out: array
843
+ an array with the same data and data type as ``x`` and located on the
844
+ specified ``device``.
845
+
846
+ Notes
847
+ -----
848
+
849
+ For NumPy, this function effectively does nothing since the only supported
850
+ device is the CPU. For CuPy, this method supports CuPy CUDA
851
+ :external+cupy:class:`Device <cupy.cuda.Device>` and
852
+ :external+cupy:class:`Stream <cupy.cuda.Stream>` objects. For PyTorch,
853
+ this is the same as :external+torch:meth:`x.to(device) <torch.Tensor.to>`
854
+ (the ``stream`` argument is not supported in PyTorch).
855
+
856
+ See Also
857
+ --------
858
+
859
+ device : Hardware device the array data resides on.
860
+
861
+ """
862
+ if is_numpy_array(x):
863
+ if stream is not None:
864
+ raise ValueError("The stream argument to to_device() is not supported")
865
+ if device == "cpu":
866
+ return x
867
+ raise ValueError(f"Unsupported device {device!r}")
868
+ elif is_cupy_array(x):
869
+ # cupy does not yet have to_device
870
+ return _cupy_to_device(x, device, stream=stream)
871
+ elif is_torch_array(x):
872
+ return _torch_to_device(x, device, stream=stream) # pyright: ignore[reportArgumentType]
873
+ elif is_dask_array(x):
874
+ if stream is not None:
875
+ raise ValueError("The stream argument to to_device() is not supported")
876
+ # TODO: What if our array is on the GPU already?
877
+ if device == "cpu":
878
+ return x
879
+ raise ValueError(f"Unsupported device {device!r}")
880
+ elif is_jax_array(x):
881
+ if not hasattr(x, "__array_namespace__"):
882
+ # In JAX v0.4.31 and older, this import adds to_device method to x...
883
+ import jax.experimental.array_api # noqa: F401 # pyright: ignore
884
+
885
+ # ... but only on eager JAX. It won't work inside jax.jit.
886
+ if not hasattr(x, "to_device"):
887
+ return x
888
+ return x.to_device(device, stream=stream)
889
+ elif is_pydata_sparse_array(x) and device == _device(x):
890
+ # Perform trivial check to return the same array if
891
+ # device is same instead of err-ing.
892
+ return x
893
+ return x.to_device(device, stream=stream) # pyright: ignore
894
+
895
+
896
+ @overload
897
+ def size(x: HasShape[Collection[SupportsIndex]]) -> int: ...
898
+ @overload
899
+ def size(x: HasShape[Collection[None]]) -> None: ...
900
+ @overload
901
+ def size(x: HasShape[Collection[SupportsIndex | None]]) -> int | None: ...
902
+ def size(x: HasShape[Collection[SupportsIndex | None]]) -> int | None:
903
+ """
904
+ Return the total number of elements of x.
905
+
906
+ This is equivalent to `x.size` according to the `standard
907
+ <https://data-apis.org/array-api/latest/API_specification/generated/array_api.array.size.html>`__.
908
+
909
+ This helper is included because PyTorch defines `size` in an
910
+ :external+torch:meth:`incompatible way <torch.Tensor.size>`.
911
+ It also fixes dask.array's behaviour which returns nan for unknown sizes, whereas
912
+ the standard requires None.
913
+ """
914
+ # Lazy API compliant arrays, such as ndonnx, can contain None in their shape
915
+ if None in x.shape:
916
+ return None
917
+ out = math.prod(cast("Collection[SupportsIndex]", x.shape))
918
+ # dask.array.Array.shape can contain NaN
919
+ return None if math.isnan(out) else out
920
+
921
+
922
+ @lru_cache(100)
923
+ def _is_writeable_cls(cls: type) -> bool | None:
924
+ if (
925
+ _issubclass_fast(cls, "numpy", "generic")
926
+ or _issubclass_fast(cls, "jax", "Array")
927
+ or _issubclass_fast(cls, "sparse", "SparseArray")
928
+ ):
929
+ return False
930
+ if _is_array_api_cls(cls):
931
+ return True
932
+ return None
933
+
934
+
935
+ def is_writeable_array(x: object) -> bool:
936
+ """
937
+ Return False if ``x.__setitem__`` is expected to raise; True otherwise.
938
+ Return False if `x` is not an array API compatible object.
939
+
940
+ Warning
941
+ -------
942
+ As there is no standard way to check if an array is writeable without actually
943
+ writing to it, this function blindly returns True for all unknown array types.
944
+ """
945
+ cls = cast(Hashable, type(x))
946
+ if _issubclass_fast(cls, "numpy", "ndarray"):
947
+ return cast("npt.NDArray", x).flags.writeable
948
+ res = _is_writeable_cls(cls)
949
+ if res is not None:
950
+ return res
951
+ return hasattr(x, '__array_namespace__')
952
+
953
+
954
+ @lru_cache(100)
955
+ def _is_lazy_cls(cls: type) -> bool | None:
956
+ if (
957
+ _issubclass_fast(cls, "numpy", "ndarray")
958
+ or _issubclass_fast(cls, "numpy", "generic")
959
+ or _issubclass_fast(cls, "cupy", "ndarray")
960
+ or _issubclass_fast(cls, "torch", "Tensor")
961
+ or _issubclass_fast(cls, "sparse", "SparseArray")
962
+ ):
963
+ return False
964
+ if (
965
+ _issubclass_fast(cls, "jax", "Array")
966
+ or _issubclass_fast(cls, "dask.array", "Array")
967
+ or _issubclass_fast(cls, "ndonnx", "Array")
968
+ ):
969
+ return True
970
+ return None
971
+
972
+
973
+ def is_lazy_array(x: object) -> bool:
974
+ """Return True if x is potentially a future or it may be otherwise impossible or
975
+ expensive to eagerly read its contents, regardless of their size, e.g. by
976
+ calling ``bool(x)`` or ``float(x)``.
977
+
978
+ Return False otherwise; e.g. ``bool(x)`` etc. is guaranteed to succeed and to be
979
+ cheap as long as the array has the right dtype and size.
980
+
981
+ Note
982
+ ----
983
+ This function errs on the side of caution for array types that may or may not be
984
+ lazy, e.g. JAX arrays, by always returning True for them.
985
+ """
986
+ # **JAX note:** while it is possible to determine if you're inside or outside
987
+ # jax.jit by testing the subclass of a jax.Array object, as well as testing bool()
988
+ # as we do below for unknown arrays, this is not recommended by JAX best practices.
989
+
990
+ # **Dask note:** Dask eagerly computes the graph on __bool__, __float__, and so on.
991
+ # This behaviour, while impossible to change without breaking backwards
992
+ # compatibility, is highly detrimental to performance as the whole graph will end
993
+ # up being computed multiple times.
994
+
995
+ # Note: skipping reclassification of JAX zero gradient arrays, as one will
996
+ # exclusively get them once they leave a jax.grad JIT context.
997
+ cls = cast(Hashable, type(x))
998
+ res = _is_lazy_cls(cls)
999
+ if res is not None:
1000
+ return res
1001
+
1002
+ if not hasattr(x, "__array_namespace__"):
1003
+ return False
1004
+
1005
+ # Unknown Array API compatible object. Note that this test may have dire consequences
1006
+ # in terms of performance, e.g. for a lazy object that eagerly computes the graph
1007
+ # on __bool__ (dask is one such example, which however is special-cased above).
1008
+
1009
+ # Select a single point of the array
1010
+ s = size(cast("HasShape[Collection[SupportsIndex | None]]", x))
1011
+ if s is None:
1012
+ return True
1013
+ xp = array_namespace(x)
1014
+ if s > 1:
1015
+ x = xp.reshape(x, (-1,))[0]
1016
+ # Cast to dtype=bool and deal with size 0 arrays
1017
+ x = xp.any(x)
1018
+
1019
+ try:
1020
+ bool(x)
1021
+ return False
1022
+ # The Array API standard dictactes that __bool__ should raise TypeError if the
1023
+ # output cannot be defined.
1024
+ # Here we allow for it to raise arbitrary exceptions, e.g. like Dask does.
1025
+ except Exception:
1026
+ return True
1027
+
1028
+
1029
+ __all__ = [
1030
+ "array_namespace",
1031
+ "device",
1032
+ "get_namespace",
1033
+ "is_array_api_obj",
1034
+ "is_array_api_strict_namespace",
1035
+ "is_cupy_array",
1036
+ "is_cupy_namespace",
1037
+ "is_dask_array",
1038
+ "is_dask_namespace",
1039
+ "is_jax_array",
1040
+ "is_jax_namespace",
1041
+ "is_numpy_array",
1042
+ "is_numpy_namespace",
1043
+ "is_torch_array",
1044
+ "is_torch_namespace",
1045
+ "is_ndonnx_array",
1046
+ "is_ndonnx_namespace",
1047
+ "is_pydata_sparse_array",
1048
+ "is_pydata_sparse_namespace",
1049
+ "is_writeable_array",
1050
+ "is_lazy_array",
1051
+ "size",
1052
+ "to_device",
1053
+ ]
1054
+
1055
+ _all_ignore = ['lru_cache', 'sys', 'math', 'inspect', 'warnings']
1056
+
1057
+ def __dir__() -> list[str]:
1058
+ return __all__