immlib 1.0.0.dev2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (45) hide show
  1. immlib/__init__.py +131 -0
  2. immlib/_init.py +108 -0
  3. immlib/_version.py +235 -0
  4. immlib/doc/__init__.py +38 -0
  5. immlib/doc/_core.py +311 -0
  6. immlib/iolib/__init__.py +29 -0
  7. immlib/iolib/_core.py +720 -0
  8. immlib/pathlib/__init__.py +69 -0
  9. immlib/pathlib/_cache.py +152 -0
  10. immlib/pathlib/_core.py +869 -0
  11. immlib/pathlib/_osf.py +538 -0
  12. immlib/test/__init__.py +16 -0
  13. immlib/test/__main__.py +10 -0
  14. immlib/test/doc/__init__.py +6 -0
  15. immlib/test/doc/test_core.py +91 -0
  16. immlib/test/iolib/__init__.py +7 -0
  17. immlib/test/iolib/test_core.py +81 -0
  18. immlib/test/pathlib/__init__.py +11 -0
  19. immlib/test/pathlib/test_core.py +146 -0
  20. immlib/test/pathlib/test_osf.py +54 -0
  21. immlib/test/types/__init__.py +5 -0
  22. immlib/test/types/test_core.py +110 -0
  23. immlib/test/util/__init__.py +11 -0
  24. immlib/test/util/test_core.py +681 -0
  25. immlib/test/util/test_numeric.py +1374 -0
  26. immlib/test/util/test_quantity.py +218 -0
  27. immlib/test/util/test_url.py +51 -0
  28. immlib/test/workflow/__init__.py +9 -0
  29. immlib/test/workflow/test_core.py +418 -0
  30. immlib/test/workflow/test_plantype.py +248 -0
  31. immlib/types/__init__.py +29 -0
  32. immlib/types/_core.py +333 -0
  33. immlib/util/__init__.py +283 -0
  34. immlib/util/_core.py +2524 -0
  35. immlib/util/_numeric.py +2651 -0
  36. immlib/util/_quantity.py +523 -0
  37. immlib/util/_url.py +114 -0
  38. immlib/workflow/__init__.py +48 -0
  39. immlib/workflow/_core.py +1635 -0
  40. immlib/workflow/_plantype.py +334 -0
  41. immlib-1.0.0.dev2.dist-info/METADATA +76 -0
  42. immlib-1.0.0.dev2.dist-info/RECORD +45 -0
  43. immlib-1.0.0.dev2.dist-info/WHEEL +5 -0
  44. immlib-1.0.0.dev2.dist-info/licenses/LICENSE +21 -0
  45. immlib-1.0.0.dev2.dist-info/top_level.txt +1 -0
@@ -0,0 +1,2651 @@
1
+ # -*- coding: utf-8 -*-
2
+ ###############################################################################
3
+ # immlib/util/_numeric.py
4
+
5
+
6
+ # Dependencies ################################################################
7
+
8
+ import inspect
9
+ from functools import (partial, wraps, update_wrapper)
10
+ from collections import namedtuple
11
+
12
+ import pint
13
+ import numpy as np
14
+ import scipy as sp
15
+ import scipy.sparse as sps
16
+ from scipy.sparse import issparse as scipy__is_sparse
17
+ from pcollections import *
18
+
19
+ from ..doc import docwrap
20
+ from ._core import (
21
+ is_tuple, is_list, is_aseq, is_aset,
22
+ is_str, streq, strnorm,
23
+ frozenarray, freezearray,
24
+ unitregistry)
25
+
26
+
27
+
28
+ # PyTorch Configuration #######################################################
29
+
30
+ # If torch isn't imported or configured, that's fine, we just write our methods
31
+ # to generate errors. We want these errors to explain the problem, so we create
32
+ # our own error type, then have a wrapper for the functions that follow that
33
+ # automatically raise the error when torch isn't found.
34
+
35
+ class TorchNotFound(Exception):
36
+ """Exception raised when PyTorch is requested but is not installed."""
37
+ def __str__(self):
38
+ return (
39
+ "pytorch not found.\n\n"
40
+ "Immlib does not require pytorch, but it must be installed for\n"
41
+ "certain operations to work.\n\n"
42
+ "See https://pytorch.org/get-started/locally/ for help\n"
43
+ "installing pytorch.")
44
+ @staticmethod
45
+ def raise_self(*args, **kw):
46
+ """Raises a `TorchNotFound` error."""
47
+ raise TorchNotFound()
48
+ class FakeTorchPackage:
49
+ """A class that raises errors for the keras package if it cannot be loaded.
50
+ """
51
+ __slots__ = ('__version__')
52
+ def __new__(cls):
53
+ self = object.__new__(cls)
54
+ object.__setattr__(self, '__version__', '0.0.0')
55
+ return self
56
+ def __getattr__(self, k):
57
+ raise TorchNotFound()
58
+ @classmethod
59
+ def is_tensor(cls, arg):
60
+ return False
61
+ try:
62
+ import torch
63
+ torch_found = True
64
+ @docwrap('immlib.util.checktorch', indent=8)
65
+ def checktorch(f):
66
+ """Decorator, ensures that PyTorch functions throw an informative error
67
+ when PyTorch isn't found.
68
+
69
+ A function that is wrapped with the ``@checktorch`` decorator will
70
+ always throw a descriptive error message when PyTorch isn't found on
71
+ the system rather than raising a complex exception. Any ``immlib``
72
+ function that uses the ``torch`` library should use this decorator.
73
+
74
+ The ``torch`` library was found on this system, so ``checktorch(f)``
75
+ always returns ``f``.
76
+ """
77
+ return f
78
+ @docwrap('immlib.util.alttorch', indent=8)
79
+ def alttorch(f_alt):
80
+ """Decorator that runs an alternative function when PyTorch isn't
81
+ found on the system.
82
+
83
+ A function ``f`` that is wrapped with the ``@alttorch(f_alt)``
84
+ decorator will always run `f_alt` instead of ``f`` when called if
85
+ PyTorch is not found on the system and will always run ``f`` when
86
+ `PyTorch` is found.
87
+
88
+ The ``torch`` library was found on this system, so
89
+ ``alttorch(f)(f_alt)`` always returns ``f``.
90
+ """
91
+ return (lambda f: f)
92
+ except (ModuleNotFoundError, ImportError) as e:
93
+ torch = FakeTorchPackage()
94
+ torch_found = False
95
+ @docwrap('immlib.util.checktorch', indent=8)
96
+ def checktorch(f):
97
+ """Decorator that ensures that PyTorch functions throw an informative
98
+ error when PyTorch isn't found.
99
+
100
+ A function that is wrapped with the ``@checktorch`` decorator will
101
+ always throw a descriptive error message when PyTorch isn't found on
102
+ the system rather than raising a complex exception. Any ``immlib``
103
+ function that uses the ``torch`` library should use this decorator.
104
+
105
+ The ``torch`` library was not found on this system, so
106
+ ``checktorch(f)`` always returns a function with the same docstring as
107
+ `f` but which raises a ``TorchNotFound`` exception.
108
+ """
109
+ from functools import wraps
110
+ return wraps(f)(TorchNotFound.raise_self)
111
+ @docwrap('immlib.util.alttorch', indent=8)
112
+ def alttorch(f_alt):
113
+ """Decorator that runs an alternative function when PyTorch isn't
114
+ found on the system.
115
+
116
+ A function ``f`` that is wrapped with the ``@alttorch(f_alt)``
117
+ decorator will always run `f_alt` instead of ``f`` when called if
118
+ PyTorch is not found on the system and will always run ``f`` when
119
+ PyTorch is found.
120
+
121
+ The ``torch`` library was not found on this system, so
122
+ ``alttorch(f)(f_alt)`` always returns `f_alt`, or rather a version of
123
+ `f_alt` wrapped to ``f``.
124
+ """
125
+ from functools import wraps
126
+ return (lambda f: wraps(f)(f_alt))
127
+ _sparse_torch_types = pdict()
128
+ _sparse_torch_layouts = pdict()
129
+ # Get the torch version setup.
130
+ _torch_version = torch.__version__.split('.')
131
+ try:
132
+ _torch_version = (
133
+ int(_torch_version[0]),
134
+ int(_torch_version[1]),
135
+ int(_torch_version[2]))
136
+ except Exception:
137
+ _torch_version = (
138
+ int(_torch_version[0]),
139
+ int(_torch_version[1]),
140
+ _torch_version[2])
141
+
142
+
143
+ # Numerical Types #############################################################
144
+
145
+ from numpy import ndarray
146
+ def _is_numtype(obj, numtype, dtypes):
147
+ if isinstance(obj, numtype):
148
+ return True
149
+ elif isinstance(obj, ndarray):
150
+ return any(map(partial(np.issubdtype, obj.dtype), dtypes))
151
+ elif torch.is_tensor(obj):
152
+ return any(map(partial(np.issubdtype, obj.numpy().dtype), dtypes))
153
+ else:
154
+ return False
155
+ from numbers import Number
156
+ _number_dtypes = (np.number, np.bool_)
157
+ @docwrap('immlib.is_numberdata')
158
+ def is_numberdata(obj, /):
159
+ """Returns ``True`` if an object is a Python number, otherwise ``False``.
160
+
161
+ ``is_numberdata(obj)`` returns ``True`` if the given object ``obj`` is an
162
+ instance of the ``numbers.Number`` type or if it is an instance of a
163
+ numeric NumPy array or PyTorch tensor.
164
+
165
+ Except in special cases, ``is_numberdata(x)`` is equivalent to
166
+ ``is_complexdata(x)``.
167
+
168
+ ``is_numberdata`` is related to the function ``is_numeric``: if
169
+ ``is_numeric(x)`` is ``True`` then ``is_numberdata(x)`` is also
170
+ ``True``. However, ``is_numberdata(10)`` is ``True`` while
171
+ ``is_numeric(10)`` is not. ``is_numberdata`` is designed for determining
172
+ whether an object represents numbers, whereas ``is_numeric`` is designed
173
+ for querying the properties of NumPy arrays and PyTorch tensors such as
174
+ their shapes and data types.
175
+
176
+ Parameters
177
+ ----------
178
+ obj : object
179
+ The object whose quality as a ``Number`` object or numerical array or
180
+ tensor is to be assessed.
181
+
182
+ Returns
183
+ -------
184
+ boolean
185
+ ``True`` if `obj` is an instance of ``Number`` or is a numerical
186
+ array or tensor, otherwise ``False``.
187
+
188
+ See Also
189
+ --------
190
+ is_booldata, is_intdata, is_realdata, is_complexdata
191
+ """
192
+ return _is_numtype(obj, Number, _number_dtypes)
193
+ _bool_dtypes = (np.bool_,)
194
+ @docwrap('immlib.is_booldata')
195
+ def is_booldata(obj, /):
196
+ """Returns ``True`` if an object is a boolean, otherwise ``False``.
197
+
198
+ ``is_booldata(obj)`` returns ``True`` if the given object `obj` is an
199
+ instance of the ``bool`` type or if it is an instance of a boolean NumPy
200
+ array or PyTorch tensor.
201
+
202
+ Parameters
203
+ ----------
204
+ obj : object
205
+ The object whose quality as a ``bool`` object or boolean array or
206
+ tensor is to be assessed.
207
+
208
+ Returns
209
+ -------
210
+ boolean
211
+ ``True`` if `obj` is an instance of ``bool`` or is a boolean array or
212
+ tensor, otherwise ``False``.
213
+
214
+ See Also
215
+ --------
216
+ is_intdata, is_realdata, is_complexdata, is_numberdata
217
+ """
218
+ return _is_numtype(obj, bool, _bool_dtypes)
219
+ from numbers import Integral
220
+ _integer_dtypes = (np.integer, np.bool_)
221
+ @docwrap('immlib.is_intdata')
222
+ def is_intdata(obj, /):
223
+ """Returns ``True`` if an object is a Python integer, otherwise ``False``.
224
+
225
+ ``is_intdata(obj)`` returns ``True`` if the given object `obj` is an
226
+ instance of the ``numbers.Integral`` type or if it is an instance of a
227
+ numeric NumPy array or PyTorch tensor whose dtype is an integer type.
228
+
229
+ Parameters
230
+ ----------
231
+ obj : object
232
+ The object whose quality as a ``Integral`` object or integer-valued
233
+ array or tensor is to be assessed.
234
+
235
+ Returns
236
+ -------
237
+ boolean
238
+ ``True`` if `obj` is an instance of ``Integral`` or is an integer numpy
239
+ array, otherwise ``False``.
240
+
241
+ See Also
242
+ --------
243
+ is_booldata, is_realdata, is_complexdata, is_numberdata
244
+ """
245
+ return _is_numtype(obj, Integral, _integer_dtypes)
246
+ from numbers import Real
247
+ _real_dtypes = (np.floating, np.integer, np.bool_)
248
+ @docwrap('immlib.is_intdata')
249
+ def is_realdata(obj, /):
250
+ """Returns ``True`` if an object is a Python number, otherwise ``False``.
251
+
252
+ ``is_realdata(obj)`` returns ``True`` if the given object `obj` is an
253
+ instance of the ``numbers.Real`` type or of a real-valued NumPy array or
254
+ PyTorch tensor.
255
+
256
+ Parameters
257
+ ----------
258
+ obj : object
259
+ The object whose quality as a ``Real`` object or real-values NumPy
260
+ array ot PyTorch tensor is to be assessed.
261
+
262
+ Returns
263
+ -------
264
+ bool
265
+ ``True`` if `obj` is an instance of ``Real`` or is a real-valued array
266
+ or tensor, otherwise ``False``.
267
+
268
+ See Also
269
+ --------
270
+ is_booldata, is_intdata, is_complexdata, is_numberdata
271
+ """
272
+ return _is_numtype(obj, Real, _real_dtypes)
273
+ from numbers import Complex
274
+ _complex_dtypes = (np.number, np.bool_)
275
+ @docwrap('immlib.is_complexdata')
276
+ def is_complexdata(obj):
277
+ """Returns ``True`` if an object is a complex number, otherwise ``False``.
278
+
279
+ ``is_complexdata(obj)`` returns ``True`` if the given object `obj` is an
280
+ instance of the ``numbers.Complex`` type or an instance of a complex-valued
281
+ NumPy array or PyTorch tensor.
282
+
283
+ Parameters
284
+ ----------
285
+ obj : object
286
+ The object whose quality as a ``Complex`` object is to be assessed.
287
+
288
+ Returns
289
+ -------
290
+ boolean
291
+ ``True`` if `obj` is an instance of ``Complex``, otherwise ``False``.
292
+
293
+ See Also
294
+ --------
295
+ is_booldata, is_intdata, is_realdata, is_numberdata
296
+ """
297
+ return _is_numtype(obj, Complex, _complex_dtypes)
298
+
299
+
300
+ # Scalar Utilities ############################################################
301
+
302
+ def _is_scalar(obj, numtype):
303
+ if isinstance(obj, np.ndarray) or torch.is_tensor(obj):
304
+ if obj.shape != ():
305
+ return False
306
+ obj = obj.item()
307
+ return isinstance(obj, numtype)
308
+ @docwrap('immlib.is_number')
309
+ def is_number(obj, /, dtype=None):
310
+ """Determines whether the argument is a scalar number or not.
311
+
312
+ ``is_number(obj)`` returns ``True`` if `obj` is a scalar number and
313
+ ``False`` otherwise. The following are considered scalar numbers:
314
+
315
+ - Any instances of ``numbers.Number``,
316
+ - Any numpy array ``x`` whose shape is ``()`` such that ``x.item()`` is a
317
+ scalar.
318
+
319
+ Parameters
320
+ ----------
321
+ obj : object
322
+ The object whose quality as a scalar number is to be tested.
323
+ dtype : bool, int, float, complex, or None, optional
324
+ The type of the scalar. If this is ``None`` (the default)``, then the
325
+ type of the scalar must be a number but it needn't be any particular
326
+ number. Otherwise, it must match the given type.
327
+
328
+ Returns
329
+ -------
330
+ bool
331
+ ``True`` if `obj` is a scalar number value and ``False`` otherwise.
332
+
333
+ See Also
334
+ --------
335
+ like_number, is_numberdata, is_numeric, is_bool, is_integer, is_real,
336
+ is_complex
337
+ """
338
+ if dtype is None:
339
+ return _is_scalar(obj, Number)
340
+ elif dtype is bool:
341
+ return _is_scalar(obj, bool)
342
+ elif dtype is int:
343
+ return _is_scalar(obj, Integral)
344
+ elif dtype is float:
345
+ return _is_scalar(obj, Real)
346
+ elif dtype is complex:
347
+ return _is_scalar(obj, Complex)
348
+ else:
349
+ raise ValueError(f"invalid dtype: {dtype}")
350
+ @docwrap('immlib.is_bool')
351
+ def is_bool(obj, /):
352
+ """Determines whether the argument is a scalar boolean or not.
353
+
354
+ ``is_bool(obj)`` returns ``True`` if `obj` is a scalar boolean and
355
+ ``False`` otherwise.
356
+
357
+ See Also
358
+ --------
359
+ is_scalar, is_booldata
360
+ """
361
+ return _is_scalar(obj, bool)
362
+ @docwrap('immlib.is_integer')
363
+ def is_integer(obj, /):
364
+ """Determines whether the argument is a scalar integer or not.
365
+
366
+ ``is_integer(obj)`` returns ``True`` if `obj` is a scalar integer and
367
+ ``False`` otherwise. Note that booleans are considered integers.
368
+
369
+ See Also
370
+ --------
371
+ is_scalar, is_intdata
372
+ """
373
+ return _is_scalar(obj, Integral)
374
+ @docwrap('immlib.is_real')
375
+ def is_real(obj, /):
376
+ """Determines whether the argument is a scalar real number or not.
377
+
378
+ ``is_real(obj)`` returns ``True`` if `obj` is a scalar real number and
379
+ ``False`` otherwise. Note that booleans and integers are considered real
380
+ numbers.
381
+
382
+ See Also
383
+ --------
384
+ is_scalar, is_realdata
385
+ """
386
+ return _is_scalar(obj, Real)
387
+ @docwrap('immlib.is_complex')
388
+ def is_complex(obj, /):
389
+ """Determines whether the argument is a scalar complex number or not.
390
+
391
+ ``is_complex(obj)`` returns ``True`` if `obj` is a scalar complex number
392
+ and ``False`` otherwise. Note that booleans, integers, and real numbers are
393
+ all considered valid complex numbers.
394
+
395
+ See Also
396
+ --------
397
+ is_number, is_complexdata
398
+ """
399
+ return _is_scalar(obj, Complex)
400
+ @docwrap('immlib.like_number')
401
+ def like_number(obj, /):
402
+ """Determines whether the argument holds a scalar number value or not.
403
+
404
+ ``like_number(x)`` returns ``True`` if ``x`` is already a scalar number, if
405
+ ``x`` is a single-element numpy array or tensor, or if ``x`` is a sequence
406
+ or set that has only one numerical element; otherwise, it returns
407
+ ``False``.
408
+
409
+ If ``like_number(x)`` returns ``True``, then ``to_number(x)`` will always
410
+ return a valid Python number (i.e., an object of type ``numbers.Number``).
411
+
412
+ See Also
413
+ --------
414
+ is_number, to_number
415
+ """
416
+ if isinstance(obj, Number):
417
+ return True
418
+ if torch.is_tensor(obj):
419
+ return torch.numel(obj) == 1
420
+ if not isinstance(obj, np.ndarray):
421
+ try:
422
+ obj = np.asarray(obj)
423
+ except (TypeError, ValueError):
424
+ return False
425
+ return obj.size == 1 and is_numberdata(obj)
426
+ @docwrap('immlib.to_number')
427
+ def to_number(obj, /, unit=Ellipsis, *, ureg=None):
428
+ """Converts the argument into a simple Python number.
429
+
430
+ ``to_number(obj)`` returns a simple Python number representation of `obj`
431
+ (in other words, `obj` will be a subtype of Python's ``numbers.Number``
432
+ type). Any number, any NumPy array with only one element, and any PyTorch
433
+ tensor with only one element can be converted into a scalar. If `obj` is a
434
+ ``pint.Quantity`` then the return value is a quantity with the same unit as
435
+ `obj` and whose magnitude is ``to_number(obj.m)``.
436
+
437
+ Parameters
438
+ ----------
439
+ obj : object
440
+ The object that is to be converted into a scalar number.
441
+ unit : unit-like, bool, or None, optional
442
+ If `obj` is a ``pint.Quantity`` object, the `unit` parameter determines
443
+ how it is handled by ``to_number``. If `unit` is ``None`` and `obj` is
444
+ a quantity, then an error will be raised. If `unit` is a valid
445
+ ``pint.Unit`` or the an object that can be converted into a unit vis
446
+ that ``immlib.unit`` function. then an error is raised if `obj` is not
447
+ a quantity with alike units. If `unit` is ``Ellipsis`` (the default
448
+ value), then the behavior depends on whether `obj` is a quantity: if
449
+ `obj` is a quantity, the ``to_number`` function is run on its magnitude
450
+ and a quantity with the same unit is returned; if `obj` is not a
451
+ quantity, then the a non-quantity is returned.
452
+ ureg : pint.UnitRegistry, None, or Ellipsis, optional
453
+ The ``pint.UnitRegistry`` object to use for units. If `ureg` is
454
+ ``Ellipsis``, then ``immlib.units`` is used. If `ureg` is ``None`` (the
455
+ default), then the registry of `obj` is used if `obj` is a quantity,
456
+ and ``immlib.units`` is used if not.
457
+
458
+
459
+ Returns
460
+ -------
461
+ number or pint.Quantity
462
+ A scalar number that is an object whose class is a subtype of
463
+ ``numbers.Number`` or a ``pint.Quantity`` object whose magnitude is
464
+ such a number.
465
+
466
+ Raises
467
+ ------
468
+ TypeError
469
+ If the argument is not like a scalar number.
470
+ """
471
+ if ureg is Ellipsis:
472
+ from immlib import units as ureg
473
+ # If obj is a quantity, we handle things differently.
474
+ if isinstance(obj, pint.Quantity):
475
+ if ureg is None:
476
+ from ._quantity import unitregistry
477
+ ureg = unitregistry(obj)
478
+ if unit is None:
479
+ raise ValueError("to_number: unit is None but Quantity given")
480
+ q = ureg.Quantity(to_number(obj.m, unit=None), obj.u)
481
+ if unit is not Ellipsis:
482
+ from ._quantity import unit as to_unit
483
+ q = q.to(to_unit(unit))
484
+ return q
485
+ elif isinstance(obj, Number):
486
+ return obj
487
+ elif torch.is_tensor(obj):
488
+ if torch.numel(obj) == 1:
489
+ return obj.item()
490
+ else:
491
+ u = np.asarray(obj)
492
+ if u.size == 1 and is_numberdata(u):
493
+ return u.item()
494
+ raise TypeError(f"given object is not scalar-like: {obj}")
495
+
496
+
497
+ # Numerical Collection Suport #################################################
498
+
499
+ # Numerical collections include numpy arrays and torch tensors. These objects
500
+ # are handled similarly due to their overall functional similarity, and certain
501
+ # support functions are used for both.
502
+ def _numcoll_match(numcoll_shape, numcoll_dtype, ndim, shape, numel, dtype):
503
+ """Checks that the actual numcoll shape and the actual numcol dtype match
504
+ the requirements of the ndim, shape, and dtype parameters.
505
+ """
506
+ # Parse the shape int front and back requirements and whether middle values
507
+ # are allowed.
508
+ if shape is None:
509
+ (sh_pre, sh_mid, sh_suf) = ((), True, ())
510
+ elif shape == ():
511
+ (sh_pre, sh_mid, sh_suf) = ((), False, ())
512
+ elif np.shape(shape) == ():
513
+ (sh_pre, sh_mid, sh_suf) = ((shape,), False, ())
514
+ else:
515
+ # We add things to the prefix until we get to an ellipsis...
516
+ sh_pre = []
517
+ for d in shape:
518
+ if d is Ellipsis: break
519
+ sh_pre.append(d)
520
+ sh_pre = tuple(sh_pre)
521
+ # We might have finished with just that; otherwise, note the ellipsis
522
+ # and move on.
523
+ if len(sh_pre) == len(shape):
524
+ (sh_mid, sh_suf) = (False, ())
525
+ else:
526
+ sh_suf = []
527
+ for d in reversed(shape):
528
+ if d is Ellipsis: break
529
+ sh_suf.append(d)
530
+ sh_suf = tuple(sh_suf) # We leave this reversed!
531
+ sh_mid = len(sh_suf) + len(sh_pre) < len(shape)
532
+ assert len(sh_suf) + len(sh_pre) + int(sh_mid) == len(shape), \
533
+ "only one Ellipsis may be used in the shape filter"
534
+ # Parse ndim.
535
+ if not (is_tuple(ndim) or is_aset(ndim) or is_list(ndim) or ndim is None):
536
+ ndim = (ndim,)
537
+ # See if we match in terms of numel, ndim, and shape
538
+ sh = numcoll_shape
539
+ if ndim is not None and len(sh) not in ndim:
540
+ return False
541
+ if numel is not None:
542
+ n = np.prod(sh)
543
+ if is_tuple(numel):
544
+ if n not in numel:
545
+ return False
546
+ elif n != numel:
547
+ return False
548
+ ndim = len(sh)
549
+ if ndim < len(sh_pre) + len(sh_suf):
550
+ return False
551
+ (npre, nsuf) = (0,0)
552
+ for (s,p) in zip(sh, sh_pre):
553
+ if p != -1 and p != s: return False
554
+ npre += 1
555
+ for (s,p) in zip(reversed(sh), sh_suf):
556
+ if p != -1 and p != s: return False
557
+ nsuf += 1
558
+ # If there are extras in the middle and we don't allow them, we fail the
559
+ # match.
560
+ if not sh_mid and nsuf + npre != ndim:
561
+ return False
562
+ # See if we match the dtype.
563
+ if dtype is not None:
564
+ if is_numpydtype(numcoll_dtype):
565
+ if is_aseq(dtype) or is_aset(dtype):
566
+ dtype = [to_numpydtype(dt) for dt in dtype]
567
+ else:
568
+ # If we have been given a torch dtype, we convert it, but
569
+ # otherwise we let np.issubdtype do the converstion so that
570
+ # users can pass in things like np.integer meaningfully.
571
+ if is_torchdtype(dtype):
572
+ dtype = to_numpydtype(dtype)
573
+ if not np.issubdtype(numcoll_dtype, dtype):
574
+ return False
575
+ dtype = (numcoll_dtype,)
576
+ elif is_torchdtype(numcoll_dtype):
577
+ if is_aseq(dtype) or is_aset(dtype):
578
+ dtype = [to_torchdtype(dt) for dt in dtype]
579
+ else:
580
+ dtype = [to_torchdtype(dtype)]
581
+ if numcoll_dtype not in dtype:
582
+ return False
583
+ # We match everything!
584
+ return True
585
+
586
+
587
+ # Numpy Arrays ################################################################
588
+
589
+ # For testing whether numpy arrays or pytorch tensors have the appropriate
590
+ # dimensionality, shape, and dtype, we use some helper functions.
591
+ from numpy import dtype as numpy_dtype
592
+ @docwrap('immlib.util.is_numpydype')
593
+ def is_numpydtype(obj, /):
594
+ """Returns ``True`` for a ``numpy.dtype`` object and ``False`` otherwise.
595
+
596
+ ``is_numpydtype(obj)`` returns ``True`` if the given object `obj` is an
597
+ instance of the ``numpy.dtype`` class.
598
+
599
+ Parameters
600
+ ----------
601
+ obj : object
602
+ The object whose quality as a NumPy ``dtype`` object is to be assessed.
603
+
604
+ Returns
605
+ -------
606
+ boolean
607
+ ``True`` if `obj` is a valid ``numpy.dtype``, otherwise ``False``.
608
+ """
609
+ return isinstance(obj, numpy_dtype)
610
+ @docwrap('immlib.util.like_numydtype')
611
+ def like_numpydtype(obj, /):
612
+ """Returns ``True`` for any object that can be converted into a
613
+ ``numpy.dtype`` object.
614
+
615
+ ``like_numpydtype(obj)`` returns ``True`` if the given object `obj` is an
616
+ instance of the ``numpy.dtype`` class, is a string that can be used to
617
+ construct a ``numpy.dtype`` object, or is a ``torch.dtype`` object.
618
+
619
+ Parameters
620
+ ----------
621
+ obj : object
622
+ The object whose quality as a NumPy ``dtype`` object is to be assessed.
623
+
624
+ Returns
625
+ -------
626
+ bool
627
+ ``True`` if `obj` can be converted into a valid numpy ``dtype``,
628
+ otherwise ``False``.
629
+ """
630
+ if is_numpydtype(obj) or is_torchdtype(obj):
631
+ return True
632
+ else:
633
+ try:
634
+ return is_numpydtype(np.dtype(obj))
635
+ except TypeError:
636
+ return False
637
+ @docwrap('immlib.util.to_numpydtype')
638
+ def to_numpydtype(obj, /):
639
+ """Returns a ``numpy.dtype`` object equivalent to the given argument.
640
+
641
+ ``to_numpydtype(obj)`` attempts to coerce the given `obj` into a
642
+ ``numpy.dtype`` object. If `obj` is already a ``numpy.dtype`` object, then
643
+ `obj` itself is returned. If the object cannot be converted into a
644
+ ``numpy.dtype`` object, then an error is raised.
645
+
646
+ The following kinds of objects can be converted into a ``numpy.dtype`` (see
647
+ also ``like_numpydtype()``):
648
+ - ``numpy.dtype`` objects;
649
+ - ``torch.dtype`` objects;
650
+ - ``None`` (the default ``numpy.dtype``);
651
+ - strings that name ``numpy.dtype`` objects; or
652
+ - any object that can be passed to ``numpy.dtype()``, such as
653
+ ``numpy.int32``.
654
+
655
+ Parameters
656
+ ----------
657
+ obj : object
658
+ The object whose quality as a NumPy ``dtype`` object is to be assessed.
659
+
660
+ Returns
661
+ -------
662
+ numpy.dtype
663
+ The ``numpy.dtype`` object that is equivalent to the argument `obj`.
664
+
665
+ Raises
666
+ ------
667
+ TypeError
668
+ If the given argument `obj` cannot be converted into a ``numpy.dtype``
669
+ object.
670
+ """
671
+ if is_numpydtype(obj):
672
+ return obj
673
+ elif is_torchdtype(obj):
674
+ return torch.as_tensor((), dtype=obj).numpy().dtype
675
+ else:
676
+ return np.dtype(obj)
677
+ # Sparse Array/Tensor stuff.
678
+ def sparray_isfrozen(obj):
679
+ return not obj.data.flags['WRITEABLE']
680
+ def sparray_freeze(obj):
681
+ obj.data.setflags(write=False)
682
+ def sparray_frozen(obj):
683
+ if not obj.data.flags['WRITEABLE']:
684
+ return obj
685
+ obj = obj.copy()
686
+ obj.data.setflags(write=False)
687
+ return obj
688
+ def ndarray_isfrozen(obj):
689
+ return not obj.flags['WRITEABLE']
690
+ def ndarray_freeze(obj):
691
+ obj.setflags(write=False)
692
+ def ndarray_frozen(obj):
693
+ if not obj.flags['WRITEABLE']:
694
+ return obj
695
+ obj = obj.copy()
696
+ obj.setflags(write=False)
697
+ return obj
698
+ SparseLayout = namedtuple(
699
+ 'SparseLayout',
700
+ ('name',
701
+ 'scipy_type', 'scipy_matrix_type', 'scipy_tomethod',
702
+ 'torch_constructor', 'torch_layout', 'torch_tomethod'))
703
+ _sparse_layouts = tdict(
704
+ bsr=SparseLayout(
705
+ 'bsr',
706
+ sps.bsr_array, sps.bsr_matrix, 'tobsr',
707
+ 'sparse_bsr_tensor', 'sparse_bsr', 'to_sparse_bsr'),
708
+ bsc=SparseLayout(
709
+ 'bsc',
710
+ None, None, None,
711
+ 'sparse_bsc_tensor', 'sparse_bsc', 'to_sparse_bsc'),
712
+ coo=SparseLayout(
713
+ 'coo',
714
+ sps.coo_array, sps.coo_matrix, 'tocoo',
715
+ 'sparse_coo_tensor', 'sparse_coo', 'to_sparse_coo'),
716
+ csr=SparseLayout(
717
+ 'csr',
718
+ sps.csr_array, sps.csr_matrix, 'tocsr',
719
+ 'sparse_csr_tensor', 'sparse_csr', 'to_sparse_csr'),
720
+ csc=SparseLayout(
721
+ 'csc',
722
+ sps.csc_array, sps.csc_matrix, 'tocsc',
723
+ 'sparse_csc_tensor', 'sparse_csc', 'to_sparse_csc'),
724
+ dia=SparseLayout(
725
+ 'dia',
726
+ sps.dia_array, sps.dia_matrix, 'todia',
727
+ None, None, None),
728
+ dok=SparseLayout(
729
+ 'dok',
730
+ sps.dok_array, sps.dok_matrix, 'todok',
731
+ None, None, None),
732
+ lil=SparseLayout(
733
+ 'lil',
734
+ sps.lil_array, sps.lil_matrix, 'tolil',
735
+ None, None, None))
736
+ for (k,v) in tuple(_sparse_layouts.items()):
737
+ try:
738
+ con = getattr(torch, v.torch_constructor)
739
+ lay = getattr(torch, v.torch_layout)
740
+ cas = getattr(torch.Tensor, v.torch_tomethod)
741
+ _sparse_layouts[k] = SparseLayout(
742
+ v.name, v.scipy_type, v.scipy_matrix_type, v.scipy_tomethod,
743
+ con, lay, cas)
744
+ except Exception:
745
+ _sparse_layouts[k] = SparseLayout(
746
+ v.name, v.scipy_type, v.scipy_matrix_type, v.scipy_tomethod,
747
+ None, None, None)
748
+ _sparse_layouts = _sparse_layouts.persistent()
749
+ # Indices for going from type or layout to SparseLayout:
750
+ _sparse_index = tdict()
751
+ for (k,st) in _sparse_layouts.items():
752
+ if st.scipy_type is not None:
753
+ _sparse_index[st.scipy_type] = st
754
+ if st.scipy_matrix_type is not None:
755
+ _sparse_index[st.scipy_matrix_type] = st
756
+ if st.torch_layout is not None:
757
+ _sparse_index[st.torch_layout] = st
758
+ _sparse_index = _sparse_index.persistent()
759
+ _sparse_torch_layouts = frozenset(
760
+ st.torch_layout
761
+ for st in _sparse_layouts.values()
762
+ if st.torch_layout is not None)
763
+ def torch__is_sparse(obj):
764
+ if not torch.is_tensor(obj):
765
+ return False
766
+ return obj.layout in _sparse_torch_layouts
767
+ @docwrap('immlib.util.sparse_layout')
768
+ def sparse_layout(obj, /):
769
+ """Returns a tuple containing data about a sparse array layout.
770
+
771
+ ``sparse_layout(name)`` returns the ``SparseLayout`` tuple for the sparse
772
+ array layout with the given ``name``. The ``name`` must be one of the
773
+ following (see ``scipy.sparse`` for more information about layouts):
774
+ - ``'bsr'``
775
+ - ``'bsc'``
776
+ - ``'coo'``
777
+ - ``'csr'``
778
+ - ``'csc'``
779
+ - ``'dia'``
780
+ - ``'dok'``
781
+ - ``'lil'``
782
+
783
+ Alternatively, ``sparse_layout(obj)`` returns the sparse layout information
784
+ for the given sparse array or sparse tensor `obj`.
785
+
786
+ Whether the argument is a string or another type, the value ``None`` is
787
+ returned if the object does not correspond to a sparse layout.
788
+
789
+ The ``SparseLayout`` namedtuple that is returned has the following
790
+ elements:
791
+ - ``scipy_type``: the scipy type (e.g., ``scipy.sparse.csr_array``).
792
+ - ``scipy_matrix_type``: the matrix type (e.g.,
793
+ ``scipy.sparse.csr_matrix``).
794
+ - ``scipy_tomethod``: the scipy casting method name (e.g., ``'tocsr'``).
795
+ - ``torch_constructor``: The torch constructor function (e.g.,
796
+ ``torch.sparse_csr_tensor``).
797
+ - ``torch_layout``: The torch layout object (e.g., ``torch.sparse_csr``).
798
+ - ``torch_tomethod``: The name of the torch ``Tensor`` method for casting
799
+ (e.g., ``'to_sparse_csr'``).
800
+ """
801
+ if isinstance(obj, SparseLayout):
802
+ return obj
803
+ elif isinstance(obj, pint.Quantity):
804
+ return sparse_layout(obj.m)
805
+ elif isinstance(obj, str):
806
+ return _sparse_layouts.get(obj, None)
807
+ elif torch.is_tensor(obj):
808
+ if obj.layout in _sparse_torch_layouts:
809
+ return _sparse_index.get(obj.layout, None)
810
+ elif scipy__is_sparse(obj):
811
+ return _sparse_index.get(type(obj), None)
812
+ else:
813
+ return _sparse_index.get(obj, None)
814
+ @docwrap('immlib.util.sparse_haslayout')
815
+ def sparse_haslayout(arr, layout):
816
+ """Returns ``True`` if the given sparse array or tensor has the given
817
+ layout.
818
+
819
+ If the first argument ``arr`` is not a sparse array nor a sparse tensor,
820
+ then the return value is ``False`` (i.e., no, the object does not have the
821
+ given sparse array or sparse tensor layout).
822
+
823
+ The second argument ``layout`` may be any valid argument to the
824
+ ``sparse_layout`` function.
825
+ """
826
+ if isinstance(arr, pint.Quantity):
827
+ return sparse_haslayout(arr.m, layout)
828
+ dstlay = sparse_layout(layout)
829
+ if dstlay is None:
830
+ raise ValueError(f"invalid layout: {layout}")
831
+ if not (scipy__is_sparse(arr) or torch__is_sparse(arr)):
832
+ return False
833
+ srclay = sparse_layout(arr)
834
+ return srclay.name == dstlay.name
835
+ @docwrap('immlib.sparse_find')
836
+ def sparse_find(arr, /):
837
+ """Returns the indices and values of nonzero elements of a sparse object.
838
+
839
+ ``sparse_find(sp_array)`` is equivalent to ``scipy.sparse.find(sp_array)``
840
+ for a sparse array ``sp_array``.
841
+
842
+ ``sparse_find(sp_tensor)`` is equivalent to ``s.indices() + (s.values(),)``
843
+ for a sparse PyTorch tensor ``sp_tensor`` and a version of it that has been
844
+ coalesced, ``s = sp_tensor.coalesce()``. Note that the ``s.values()``
845
+ tensor is cloned and detached before being returned.
846
+
847
+ ``sparse_find(q)`` for a quantity ``q`` returns the equivalent of
848
+ ``sparse_find(q.m)`` except that the returned value array will have the
849
+ same magnitude as ``q``.
850
+
851
+ Raises
852
+ ------
853
+ TypeError
854
+ If `arr` is not a sparse array or sparse tensor.
855
+
856
+ See Also
857
+ --------
858
+ sparse_data, sparse_indices
859
+ """
860
+ if isinstance(arr, pint.Quantity):
861
+ from ._quantity import quant
862
+ f = sparse_find(arr.m)
863
+ return f[:-1] + (quant(f[-1], arr.u),)
864
+ elif scipy__is_sparse(arr):
865
+ return sps.find(arr)
866
+ elif torch__is_sparse(arr):
867
+ arr = arr.coalesce()
868
+ return tuple(arr.indices()) + (arr.values().clone().detach(),)
869
+ else:
870
+ raise TypeError(f"sparse_find requires a sparse array or sparse tensor")
871
+ @docwrap('immlib.util.sparse_indices')
872
+ def sparse_indices(arr, /):
873
+ """Returns the indices of the nonzero values in the given sparse object.
874
+
875
+ ``sparse_indices(arr)`` is roughly equivalent to the expression
876
+ ``stack(sparse_find(arr)[:-1])``---i.e., it returns a numpy array or a
877
+ pytorch tensor of the index matrix of the nonzero values in `arr`.
878
+
879
+ See Also
880
+ --------
881
+ sparse_find, sparse_data
882
+ """
883
+ if isinstance(arr, pint.Quantity):
884
+ return sparse_indices(arr.m)
885
+ elif scipy__is_sparse(arr):
886
+ return np.stack(sps.find(arr)[:-1])
887
+ elif torch__is_sparse(arr):
888
+ arr = arr.coalesce()
889
+ return arr.indices()
890
+ else:
891
+ raise TypeError(f"sparse_data requires a sparse array or sparse tensor")
892
+ @docwrap('immlib.util.sparse_data')
893
+ def sparse_data(arr, /):
894
+ """Returns the data vector for the given sparse array or sparse tensor.
895
+
896
+ ``sparse_data(arr)`` is equivalent to ``sparse_find(arr)[-1]``---i.e., it
897
+ returns a vector of non-zero values in the sparse array or sparse tensor
898
+ `arr`---with the exception that it returns the actual vector itself and
899
+ not a copy of the data vector. Changes to the return value of this function
900
+ will be reflected in `arr`.
901
+
902
+ See Also
903
+ --------
904
+ sparse_find, sparse_indices
905
+ """
906
+ if isinstance(arr, pint.Quantity):
907
+ from ._quantity import quant
908
+ return quant(sparse_data(arr.m), arr.u)
909
+ elif scipy__is_sparse(arr):
910
+ return arr.data
911
+ elif torch__is_sparse(arr):
912
+ arr = arr.coalesce()
913
+ return arr.values()
914
+ else:
915
+ raise TypeError(f"sparse_data requires a sparse array or sparse tensor")
916
+ @docwrap('immlib.util.sparse_tolayout')
917
+ def sparse_tolayout(obj, layout):
918
+ """Copies a sparse object into another sparse object with a given layout.
919
+
920
+ ``sparse_tolayout(sparr, layout)`` copies the given sparse SciPy array or
921
+ sparse PyTorch tensor ``sparr`` into an equivalent array or tensor that
922
+ uses the sparse layout given by the ``layout`` argument, which must be
923
+ compatible with the ``sparse_layout`` function. The backend (PyTorch or
924
+ SciPy) will not be changed.
925
+
926
+ See Also
927
+ --------
928
+ sparse_layout
929
+ """
930
+ if isinstance(obj, pint.Quantity):
931
+ from ._quantity import quant
932
+ arr = sparse_tolayout(obj.m, layout)
933
+ return obj if arr is obj.m else quant(arr, obj.u)
934
+ lay = sparse_layout(layout)
935
+ if scipy__is_sparse(obj):
936
+ method = getattr(obj, lay.scipy_tomethod)
937
+ if method is None:
938
+ raise ValueError(f"layout is invalid for scipy arrays: {layout}")
939
+ arr = method()
940
+ # If obj was frozen, duplicate that.
941
+ if arr is not obj and sparray_isfrozen(obj):
942
+ sparray_freeze(arr)
943
+ return arr
944
+ elif torch__is_sparse(obj):
945
+ obj = obj.coalesce()
946
+ mtd = lay.torch_tomethod
947
+ if mtd is None:
948
+ raise ValueError(f"layout is invalid for pytorch tensors: {layout}")
949
+ return mtd(obj)
950
+ else:
951
+ raise TypeError(
952
+ "sparse_tolayout requires a sparse scipy array or"
953
+ " a sparse pytorch tensor")
954
+ @docwrap('immlib.is_array')
955
+ def is_array(obj, /, *,
956
+ dtype=None, shape=None, ndim=None, numel=None, frozen=None,
957
+ sparse=None, quant=None, unit=Ellipsis, ureg=None):
958
+ """Returns ``True`` if an object is a ``numpy.ndarray`` object, otherwise
959
+ returns ``False``.
960
+
961
+ ``is_array(obj)`` returns ``True`` if the given object `obj` is an instance
962
+ of the ``numpy.ndarray`` class or is a ``scipy.sparse`` array, or if `obj`
963
+ is a ``pint.Quantity`` object whose magnitude is one of these. Additional
964
+ constraints may be placed on the object via the optional argments.
965
+
966
+ Note that to ``immlib``, both ``numpy.ndarray`` arrays and ``scipy.sparse``
967
+ arrays are considered "arrays". This behavior can be changed with the
968
+ ``sparse`` parameter.
969
+
970
+ Parameters
971
+ ----------
972
+ obj : object
973
+ The object whose quality as a NumPy array object is to be assessed.
974
+ dtype : dtype-like or None, optional
975
+ The NumPy `dtype` that is required of the `obj` in order to be
976
+ considered a valid ``ndarray``. The ``obj.dtype`` matches the given
977
+ `dtype` parameter if either `dtype` is ``None`` (the default) or if
978
+ ``obj.dtype`` is a sub-dtype of `dtype` according to
979
+ ``numpy.issubdtype``. Alternately, `dtype` can be a tuple, in which
980
+ case, `obj` is considered valid if its dtype is any of the dtypes in
981
+ `dtype`. Note that in the case of a tuple, the dtype of `obj` must
982
+ appear exactly in the tuple rather than be a subtype of one of the
983
+ objects in the tuple.
984
+ ndim : int, tuple or ints, or None, optional
985
+ The number of dimensions that the object must have in order to be
986
+ considered a valid numpy array. If ``None``, then any number of
987
+ dimensions is acceptable (this is the default). If this is an integer,
988
+ then the number of dimensions must be exactly that integer. If this is
989
+ a list or tuple of integers, then the dimensionality must be one of
990
+ these numbers.
991
+ shape : int, tuple of ints, or None, optional
992
+ If the ``shape`` parameter is not ``None``, then the given `obj` must
993
+ have a shape that matches the parameter value. The value `shape` must
994
+ be a tuple that is equal to the `obj`'s shape tuple with the following
995
+ additional rules: a ``-1`` value in the ``shape`` tuple will match any
996
+ value in the `obj`'s shape tuple, and a single ``Ellipsis`` may appear
997
+ in `shape`, which matches any number of values in the `obj`'s shape
998
+ tuple. The default value of ``None`` indicates that no restriction
999
+ should be applied to the `obj`'s shape.
1000
+ numel : int, tuple of ints, or None, optional
1001
+ If the `numel` parameter is not ``None``, then the given `obj` must
1002
+ have the same number of elements as given by `numel`. If `numel` is a
1003
+ tuple, then the number of elements in `obj` must be in the `numel`
1004
+ tuple. The number of elements is the product of its shape.
1005
+ frozen : bool or None, optional
1006
+ If ``None``, then no restrictions are placed on the ``'WRITEABLE'``
1007
+ flag of `obj`. If ``True``, then the data in `obj` must be read-only in
1008
+ order for `obj` to be considered a valid array. If ``False``, then the
1009
+ data in `obj` must not be read-only.
1010
+ sparse : boolean or False, optional
1011
+ If the `sparse`` parameter is ``None``, then no requirements are placed
1012
+ on the sparsity of `obj` for it to be considered a valid array. If
1013
+ `sparse` is ``True`` or ``False``, then `obj` must either be sparse or
1014
+ not be sparse, respectively, for `obj` to be considered valid. If
1015
+ ``sparse`` is a string, then it must be either ``'coo'``, ``'lil'``,
1016
+ ``'csr'``, or ``'csr'``, indicating the required sparse array
1017
+ type. Only ``scipy.sparse`` matrices are considered valid sparse
1018
+ arrays.
1019
+ quant : bool, optional
1020
+ Whether ``Quantity`` objects should be considered valid arrays or not.
1021
+ If ``quant=True`` then `obj` is considered a valid array only when
1022
+ ``obj`` is a quantity object with a ``numpy`` array as the
1023
+ magnitude. If ``False``, then `obj` must be a ``numpy`` array itself
1024
+ and not a ``Quantity`` to be considered valid. If ``None`` (the
1025
+ default), then either quantities or ``numpy`` arrays are considered
1026
+ valid arrays.
1027
+ unit : unit-like, Ellipsis, or None, optional
1028
+ A unit with which the object `obj`'s unit must be compatible in order
1029
+ for `obj` to be considered a valid array. An `obj` that is not a
1030
+ quantity is considered to have a unit of ``None``, which is not the
1031
+ same as being a quantity with a dimensionless unit. In other words,
1032
+ ``is_array(array, quant=None)`` will return ``True`` for a numpy array
1033
+ while ``is_array(arary, quant='dimensionless')`` will return
1034
+ ``False``. If ``unit=Ellipsis`` (the default), then the object's unit
1035
+ is ignored.
1036
+ ureg : pint.UnitRegistry, None, or Ellipsis, optional
1037
+ The ``pint.UnitRegistry`` object to use for units. If `ureg` is
1038
+ ``Ellipsis``, then ``immlib.units`` is used. If `ureg` is ``None`` (the
1039
+ default), then the registry of `obj` is used if `obj` is a quantity,
1040
+ and ``immlib.units`` is used if not.
1041
+
1042
+ Returns
1043
+ -------
1044
+ bool
1045
+ ``True`` if `obj` is a valid numpy array, otherwise ``False``.
1046
+
1047
+ See Also
1048
+ --------
1049
+ is_tensor, is_numeric
1050
+ """
1051
+ if ureg is Ellipsis:
1052
+ from immlib import units as ureg
1053
+ # If this is a quantity, just extract the magnitude.
1054
+ if isinstance(obj, pint.Quantity):
1055
+ if quant is False:
1056
+ return False
1057
+ if ureg is None:
1058
+ from ._quantity import unitregistry
1059
+ ureg = unitregistry(obj)
1060
+ u = obj.u
1061
+ obj = obj.m
1062
+ elif quant is True:
1063
+ return False
1064
+ else:
1065
+ if ureg is None:
1066
+ from immlib import units as ureg
1067
+ u = None
1068
+ # At this point we want to check if this is a valid numpy array or scipy
1069
+ # sparse matrix; however how we handle the answer to this question depends
1070
+ # on the sparse parameter.
1071
+ if sparse is True:
1072
+ if not scipy__is_sparse(obj):
1073
+ return False
1074
+ elif sparse is False:
1075
+ if not isinstance(obj, ndarray):
1076
+ return False
1077
+ elif sparse is None:
1078
+ # Set sparse to True/False:
1079
+ sparse = scipy__is_sparse(obj)
1080
+ # Also, it still has to be either a numpy array or a scipy sparse array.
1081
+ if not (isinstance(obj, ndarray) or sparse):
1082
+ return False
1083
+ else:
1084
+ if is_str(sparse):
1085
+ sparse = strnorm(sparse.strip(), case=True, unicode=False)
1086
+ layout = sparse_layout(sparse)
1087
+ if layout is None:
1088
+ raise ValueError(f"invalid sparse array type: {sparse}")
1089
+ else:
1090
+ layout = sparse_layout(sparse)
1091
+ if layout is None:
1092
+ tt = type(sparse)
1093
+ raise ValueError(f"invalid sparse parameter of type {tt}")
1094
+ ltypes = (layout.scipy_type, layout.scipy_matrix_type)
1095
+ if not isinstance(obj, ltypes):
1096
+ return False
1097
+ sparse = True
1098
+ # At this point, the sparse parameter has been checked against the object
1099
+ # and sparse is now True if the object is a sparse array and False if not.
1100
+ # Next, check that the object is read-only or not.
1101
+ if frozen is True:
1102
+ if sparse:
1103
+ if not sparray_isfrozen(obj):
1104
+ return False
1105
+ elif not ndarray_isfrozen(obj):
1106
+ return False
1107
+ elif frozen is False:
1108
+ if sparse:
1109
+ if sparray_isfrozen(obj):
1110
+ return False
1111
+ else:
1112
+ if ndarray_isfrozen(obj):
1113
+ return False
1114
+ elif frozen is None:
1115
+ frozen = sparray_isfrozen(obj) if sparse else ndarray_isfrozen(obj)
1116
+ else:
1117
+ raise ValueError(
1118
+ f"frozen option must be boolean or None; got type {type(frozen)}")
1119
+ # Next, check compatibility of the units.
1120
+ if unit is None:
1121
+ # We are required to not be a quantity.
1122
+ if u is not None:
1123
+ return False
1124
+ elif unit is not Ellipsis:
1125
+ from ._quantity import alike_units
1126
+ if not is_tuple(unit):
1127
+ unit = (unit,)
1128
+ if not any(map(partial(alike_units, u), unit)):
1129
+ return False
1130
+ # Check the match to the numeric collection last.
1131
+ if dtype is None and shape is None and ndim is None and numel is None:
1132
+ return True
1133
+ return _numcoll_match(obj.shape, obj.dtype, ndim, shape, numel, dtype)
1134
+ def to_array(obj, /, dtype=None, *,
1135
+ order=None, copy=False, sparse=None, frozen=None,
1136
+ quant=None, ureg=None, unit=Ellipsis, detach=True):
1137
+ """Reinterprets `obj` as a NumPy array or quantity with an array magnitude.
1138
+
1139
+ ``immlib.to_array`` is roughly equivalent to the ``numpy.asarray`` function
1140
+ with a few exceptions:
1141
+
1142
+ - ``to_array(obj)`` allows quantities for `obj` and, in such a case, will
1143
+ return a quantity whose magnitude has been reinterpreted as an array,
1144
+ though this behavior can be altered with the `quant` parameter;
1145
+ - ``to_array(obj)`` can extract the ``numpy`` array from ``torch`` tensor
1146
+ objects.
1147
+
1148
+ Parameters
1149
+ ----------
1150
+ obj : object
1151
+ The object that is to be reinterpreted as, or if necessary covnerted
1152
+ to, a NumPy array object.
1153
+ dtype : data-type, optional
1154
+ The `dtype` that is passed to ``numpy.asarray()``.
1155
+ order : {'C', 'F'}, optional
1156
+ The array order that is passed to ``numpy.asarray()``.
1157
+ copy : boolean, optional
1158
+ Whether to copy the data in `obj` or not. If ``False``, then `obj` is
1159
+ only copied if doing so is required by the optional parameters. If
1160
+ ``True``, then `obj` is always copied if possible.
1161
+ sparse : bool, 'csr', 'coo', or None, optional
1162
+ If ``None``, then the sparsity of `obj` is the same as the sparsity of
1163
+ the array that is returned. Otherwise, the return value will always be
1164
+ either a ``scipy.spase`` matrix (``sparse=True``) or a
1165
+ ``numpy.ndarray`` (``sparse=False``) based on the given value of
1166
+ `sparse`. The `sparse` parameter may also be set to ``'bsr'``,
1167
+ ``'coo'``, ``'csc'``, ``'csr'``, ``'dia'``, or ``'dok'`` to return
1168
+ specific sparse matrix types.
1169
+ frozen : bool or None, optional
1170
+ Whether the return value should be read-only or not. If ``None``, then
1171
+ no changes are made to the return value; if a new array is allocated in
1172
+ the ``to_array()`` function call, then it is returned in a writeable
1173
+ form. If ``frozen=True``, then the return value is always a read-only
1174
+ array; if `obj` is not already read-only, then a copy of `obj` is
1175
+ always returned in this case. If ``frozen=False``, then the
1176
+ return-value is never read-only.
1177
+ quant : bool or None, optional
1178
+ Whether the return value should be a ``Quantity`` object wrapping the
1179
+ array (``quant=True``) or the array itself (``quant=False``). If
1180
+ `quant` is ``None`` (the default) then the return value is a quantity
1181
+ if either `obj` is a quantity or an explicit `unit` parameter is given
1182
+ and is not a quantity if `obj` is not a quantity.
1183
+ ureg : pint.UnitRegistry or None, optional
1184
+ The ``pint.UnitRegistry`` object to use for units. If `ureg` is
1185
+ ``Ellipsis``, then ``immlib.units`` is used. If `ureg` is ``None`` (the
1186
+ default), then no specific coersion to a ``UnitRegistry`` is performed
1187
+ (i.e., the same quantity class is returned).
1188
+ unit : unit-like, bool, or Ellipsis, optional
1189
+ The unit that should be used in the return value. When the return value
1190
+ of this function is a ``Quantity`` (see the `quant` parameter), the
1191
+ returned quantity always has a unit matching the `unit` parameter; if
1192
+ the provided `obj` is not a quantity, then its unit is presumed to be
1193
+ that requested by `unit`. When the return value of this function is not
1194
+ a ``Quantity`` object and is instead is a NumPy array object, then when
1195
+ `obj` is not a quantity the `unit` parameter is ignored, and when `obj`
1196
+ is a quantity, its magnitude is returned after conversion into
1197
+ `unit`. The default value of `unit`, ``Ellipsis``, indicates that, if
1198
+ `obj` is a quantity, its unit should be used, and `unit` should be
1199
+ considered dimensionless otherwise.
1200
+ detach : bool, optional
1201
+ If the argument is a PyTorch tensor that requires gradient tracking,
1202
+ then it must be detached from the gradient tracking system before it
1203
+ can be turned into an array. If `detach` is ``True`` (the default),
1204
+ then this detachment is performed automatically. Otherwise, an error is
1205
+ raised if a tensor would need to be detached.
1206
+
1207
+ Returns
1208
+ -------
1209
+ numpy.ndarray or pint.Quantity
1210
+ Either a NumPy array equivalent to `obj` or a ``Quantity`` whose
1211
+ magnitude is a NumPy array equivalent to `obj`.
1212
+
1213
+ Raises
1214
+ ------
1215
+ ValueError
1216
+ If invalid parameter values are given or if the parameters conflict.
1217
+
1218
+ See Also
1219
+ --------
1220
+ to_tensor, to_numeric
1221
+ """
1222
+ if ureg is Ellipsis:
1223
+ from immlib import units as ureg
1224
+ # If obj is a quantity, we handle things differently.
1225
+ if isinstance(obj, pint.Quantity):
1226
+ q = obj
1227
+ obj = q.m
1228
+ if ureg is None:
1229
+ from ._quantity import unitregistry
1230
+ ureg = unitregistry(q)
1231
+ else:
1232
+ q = None
1233
+ if ureg is None:
1234
+ from immlib import units as ureg
1235
+ # Translate obj depending on whether it's a pytorch array / scipy sparse
1236
+ # matrix. We need to think about whether the output array is being
1237
+ # requested in sparse format. If so, we handle the conversion differently.
1238
+ obj_is_spsparse = scipy__is_sparse(obj)
1239
+ obj_is_tensor = not obj_is_spsparse and torch.is_tensor(obj)
1240
+ obj_is_sparse = obj_is_spsparse or torch__is_sparse(obj)
1241
+ # If this is a tensor and it requires grad, we can check whether we can
1242
+ # duplicate it now or not.
1243
+ if obj_is_tensor and obj.requires_grad:
1244
+ if detach:
1245
+ obj = obj.detach()
1246
+ else:
1247
+ raise ValueError(
1248
+ f"to_array: tensor requires grad but detach is non-true")
1249
+ newarr = False # True means we own the memory of arr; False means we don't.
1250
+ if sparse is not False and (sparse is not None or obj_is_sparse):
1251
+ # That condition is rough to parse; essentially, in this then-clause:
1252
+ # * the user isn't requesting a dense output explicitly, and
1253
+ # * the inputs tell us that we need a sparse output (because either
1254
+ # the user is requesting a sparse output explicitly or they have
1255
+ # requested no change in output sparsity and the input is sparse).
1256
+ if sparse is None or sparse is True:
1257
+ layout = sparse_layout(obj if obj_is_sparse else 'coo')
1258
+ elif isinstance(sparse, str):
1259
+ sparse = strnorm(sparse.strip(), case=True, unicode=False)
1260
+ layout = sparse_layout(sparse)
1261
+ if layout is None:
1262
+ raise ValueError(
1263
+ f"invalid scipy sparse array layout name: {sparse}")
1264
+ else:
1265
+ layout = sparse_layout(sparse)
1266
+ if layout is None:
1267
+ raise ValueError(
1268
+ f"invalid scipy sparse array layout type: {type(sparse)}")
1269
+ # We now have a layout that we are converting into.
1270
+ if obj_is_sparse:
1271
+ # We're creating a scipy sparse output from a sparse input.
1272
+ if obj_is_tensor:
1273
+ # We're creating a scipy sparse output from a sparse tensor.
1274
+ arr = obj.coalesce()
1275
+ if obj is not arr:
1276
+ newarr = True
1277
+ ii = arr.indices().detach().numpy()
1278
+ uu = arr.values().detach().numpy()
1279
+ if copy:
1280
+ vv = np.array(uu, dtype=dtype, order=order)
1281
+ else:
1282
+ vv = np.asarray(uu, dtype=dtype, order=order)
1283
+ newarr = newarr or (uu is not vv)
1284
+ arr = layout.scipy_type(
1285
+ (vv, tuple(ii)),
1286
+ shape=arr.shape,
1287
+ dtype=dtype)
1288
+ else:
1289
+ # We're creating a scipy sparse output from another scipy
1290
+ # sparse matrix. The scipy.sparse API (i.e., the find()
1291
+ # function) does not let us get access to the data itself, and
1292
+ # using obj.data along with the first two return values of
1293
+ # sps.find (the indices) can cause problems because find
1294
+ # doesn't always return the data in the same order as
1295
+ # obj.data. So we must make a copy of the data in this case
1296
+ # unless we know that the dtype and order have not been
1297
+ # changed; since order doesn't apply to 1d vectors we can
1298
+ # ignore it.
1299
+ if dtype is None or dtype == obj.dtype:
1300
+ arr = obj
1301
+ else:
1302
+ (rr,cc,uu) = sps.find(obj)
1303
+ vv = np.asarray(uu, dtype=dtype)
1304
+ if vv is uu:
1305
+ arr = obj
1306
+ else:
1307
+ arr = layout.scipy_type(
1308
+ (vv, (rr,cc)),
1309
+ shape=obj.shape,
1310
+ dtype=dtype)
1311
+ else:
1312
+ # We're creating a scipy sparse matrix from a dense matrix.
1313
+ arr = obj.numpy() if obj_is_tensor else obj
1314
+ # Make sure our dtype and order match.
1315
+ arr = np.asarray(arr, dtype=dtype, order=order)
1316
+ # We just call the appropriate constructor.
1317
+ arr = layout.scipy_type(arr)
1318
+ newarr = True
1319
+ # We mark sparse as True so that below we know that the output is
1320
+ # sparse.
1321
+ sparse = True
1322
+ else:
1323
+ # We are creating a dense array output.
1324
+ if obj_is_sparse:
1325
+ # We are creating a dense array output from a sparse input.
1326
+ if obj_is_tensor:
1327
+ # We are creating a dense array output from a sparse tensor
1328
+ # input.
1329
+ arr = obj.to_dense().numpy()
1330
+ else:
1331
+ # We are creating a dense array output from a scipy sparse
1332
+ # array input.
1333
+ arr = obj.todense()
1334
+ # In both of these cases, a copy has already been made.
1335
+ arr = np.asarray(arr, dtype=dtype, order=order)
1336
+ newarr = True
1337
+ else:
1338
+ # We are creating a dense array output from a dense input.
1339
+ if obj_is_tensor:
1340
+ # We are creating a dense array output from a dense tensor
1341
+ # input.
1342
+ arr = obj.numpy()
1343
+ else:
1344
+ arr = obj
1345
+ # Whether we call array() or asarray() depends on the copy
1346
+ # parameter.
1347
+ if copy:
1348
+ tmp = np.array(arr, dtype=dtype, order=order)
1349
+ else:
1350
+ tmp = np.asarray(arr, dtype=dtype, order=order)
1351
+ newarr = tmp is not arr
1352
+ arr = tmp
1353
+ # We mark sparse as False so that below we know that the output is
1354
+ # dense.
1355
+ sparse = False
1356
+ # If a read-only array is requested, we either return the object itself (if
1357
+ # it is already a read-only array), or we make a copy and make it
1358
+ # read-only.
1359
+ if frozen is True:
1360
+ if sparse:
1361
+ arr = sparray_frozen(arr)
1362
+ else:
1363
+ arr = ndarray_frozen(arr)
1364
+ elif frozen is False:
1365
+ if sparse:
1366
+ if sparray_isfrozen(arr):
1367
+ if not newarr:
1368
+ arr = arr.copy()
1369
+ arr.data.setflags(write=True)
1370
+ elif ndarray_isfrozen(arr):
1371
+ arr = np.array(arr)
1372
+ elif frozen is None:
1373
+ frozen = sparray_isfrozen(arr) if sparse else ndarray_isfrozen(arr)
1374
+ else:
1375
+ raise ValueError(f"bad parameter value for frozen: {frozen}")
1376
+ # Next, we switch on whether we are being asked to return a quantity or
1377
+ # not.
1378
+ if quant is None:
1379
+ quant = (q if unit is Ellipsis else unit) is not None
1380
+ if quant is True:
1381
+ if unit is None:
1382
+ raise ValueError(
1383
+ "to_array: cannot make a quantity (quant=True) without a unit"
1384
+ " (unit=None)")
1385
+ if q is None:
1386
+ if unit is Ellipsis:
1387
+ raise ValueError(
1388
+ "to_array(x): cannot make a quantity (quant=True) with the"
1389
+ " same unit as x (unit=...) when the x is not a quantity")
1390
+ return ureg.Quantity(arr, unit)
1391
+ else:
1392
+ from ._quantity import unitregistry
1393
+ if ureg is not unitregistry(q) or obj is not arr:
1394
+ q = ureg.Quantity(arr, q.u)
1395
+ if unit is not Ellipsis and ureg.Unit(unit) != q.u:
1396
+ return q.to(unit)
1397
+ else:
1398
+ return q
1399
+ elif quant is False:
1400
+ # Don't return a quantity, whatever the input argument.
1401
+ if unit is Ellipsis:
1402
+ # We return the current array/magnitude whatever its unit.
1403
+ return arr
1404
+ elif q is None:
1405
+ # We just pretend this was already in the given unit (i.e., ignore
1406
+ # unit).
1407
+ return arr
1408
+ elif unit is None:
1409
+ raise ValueError(
1410
+ "to_tensor: cannot extract unit None from quantity; to get the"
1411
+ " native unit, use unit=Ellipsis")
1412
+ else:
1413
+ if obj is not arr:
1414
+ q = ureg.Quantity(arr, q.u)
1415
+ # We convert to the given unit and return that.
1416
+ return q.m_as(unit)
1417
+ else:
1418
+ raise ValueError(
1419
+ f"to_array: quant must be boolean or None;"
1420
+ f" got object of type {type(quant)}")
1421
+
1422
+
1423
+ # PyTorch Tensors #############################################################
1424
+
1425
+ # At this point, either torch has been imported or it hasn't, but either way,
1426
+ # we can use @checktorch to make sure that errors are thrown when torch isn't
1427
+ # present. Otherwise, we can just write the functions assuming that torch is
1428
+ # imported.
1429
+ @docwrap('immlib.unit.is_torchdtype')
1430
+ @alttorch(lambda dt: False)
1431
+ def is_torchdtype(obj, /):
1432
+ """Returns ``True`` for a PyTroch ``dtype`` object and ``False`` otherwise.
1433
+
1434
+ ``is_torchdtype(obj)`` returns ``True`` if the given object `obj` is an
1435
+ instance of the ``torch.dtype`` class.
1436
+
1437
+ Parameters
1438
+ ----------
1439
+ obj : object
1440
+ The object whose quality as a PyTorch ``dtype`` object is to be
1441
+ assessed.
1442
+
1443
+ Returns
1444
+ -------
1445
+ bool
1446
+ ``True`` if `obj` is a valid ``torch.dtype``, otherwise ``False``.
1447
+ """
1448
+ return isinstance(obj, torch.dtype)
1449
+ @docwrap('immlib.unit.like_torchdtype')
1450
+ def like_torchdtype(obj, /):
1451
+ """Returns ``True`` for any object that can be converted into a
1452
+ ``torch.dtype``.
1453
+
1454
+ ``like_torchdtype(obj)`` returns ``True`` if the given object `obj` is an
1455
+ instance of the ``torch.dtype`` class, is a string that names a
1456
+ ``torch.dtype`` object, or is a ``numpy.dtype`` object that is compatible
1457
+ with PyTorch. Note that ``None`` is equivalent to ``torch``'s default
1458
+ ``dtype``.
1459
+
1460
+ Parameters
1461
+ ----------
1462
+ obj : object
1463
+ The object whose quality as a PyTorch ``dtype`` object is to be
1464
+ assessed.
1465
+
1466
+ Returns
1467
+ -------
1468
+ bool
1469
+ ``True`` if `obj` can be converted into a valid ``torch.dtype``,
1470
+ otherwise ``False``.
1471
+ """
1472
+ if is_torchdtype(obj):
1473
+ return True
1474
+ elif is_numpydtype(obj):
1475
+ try:
1476
+ return None is not torch.from_numpy(np.array((), dtype=obj))
1477
+ except TypeError:
1478
+ return False
1479
+ elif is_str(obj):
1480
+ try:
1481
+ return is_torchdtype(getattr(torch, obj))
1482
+ except AttributeError:
1483
+ return False
1484
+ elif obj is None:
1485
+ return True
1486
+ else:
1487
+ try:
1488
+ return None is not torch.from_numpy(np.array((), dtype=obj))
1489
+ except Exception:
1490
+ return False
1491
+ @docwrap('immlib.unit.to_torchdtype')
1492
+ @checktorch
1493
+ def to_torchdtype(obj, /):
1494
+ """Returns a ``torch.dtype`` object equivalent to the given argument `obj`.
1495
+
1496
+ ``to_torchdtype(obj)`` attempts to coerce the given `obj` into a
1497
+ ``torch.dtype`` object. If `obj` is already a ``torch.dtype`` object,
1498
+ then `obj` itself is returned. If the object cannot be converted into a
1499
+ ``torch.dtype`` object, then an error is raised.
1500
+
1501
+ The following kinds of objects can be converted into a ``torch.dtype`` (see
1502
+ also ``like_numpydtype()``):
1503
+ - ``torch.dtype`` objects;
1504
+ - ``numpy.dtype`` objects with compatible (numeric) types;
1505
+ - strings that name ``torch.dtype`` objects; or
1506
+ - any object that can be passed to ``numpy.dtype()``, such as
1507
+ ``numpy.int32``, that also creates a compatible (numeric) type.
1508
+
1509
+ Parameters
1510
+ ----------
1511
+ obj : object
1512
+ The object whose quality as a NumPy ``dtype`` object is to be assessed.
1513
+
1514
+ Returns
1515
+ -------
1516
+ numpy.dtype
1517
+ The ``numpy.dtype`` object that is equivalent to the argument `obj`.
1518
+
1519
+ Raises
1520
+ ------
1521
+ TypeError
1522
+ If the given argument `obj` cannot be converted into a ``numpy.dtype``
1523
+ object.
1524
+ """
1525
+ if is_torchdtype(obj):
1526
+ return obj
1527
+ else:
1528
+ return torch.as_tensor(np.array([], dtype=obj)).dtype
1529
+ def _is_never_tensor(obj,
1530
+ dtype=None, shape=None, ndim=None, numel=None,
1531
+ device=None, requires_grad=None,
1532
+ sparse=None, quant=None, unit=Ellipsis, ureg=None):
1533
+ return False
1534
+ @docwrap('immlib.is_tensor')
1535
+ @alttorch(_is_never_tensor)
1536
+ def is_tensor(obj, /, dtype=None, *,
1537
+ shape=None, ndim=None, numel=None,
1538
+ device=None, requires_grad=None,
1539
+ sparse=None, quant=None, unit=Ellipsis, ureg=None):
1540
+ """Returns ``True`` if the argument is a ``torch.tensor`` object, otherwise
1541
+ returns ``False``.
1542
+
1543
+ ``is_tensor(obj)`` returns ``True`` if the given object `obj` is an
1544
+ instance of the ``torch.Tensor`` class or is a ``pint.Quantity`` object
1545
+ whose magnitude is an instance of ``torch.Tensor``. Additional constraints
1546
+ may be placed on the object via the optional argments.
1547
+
1548
+ Parameters
1549
+ ----------
1550
+ obj : object
1551
+ The object whose quality as a PyTorch tensor object is to be assessed.
1552
+ dtype : dtype-like or None, optional
1553
+ The PyTorch `dtype` or a dtype-like object that is required to match
1554
+ that of the `obj` in order to be considered a valid tensor. The
1555
+ ``obj.dtype`` matches the given `dtype` parameter if either `dtype` is
1556
+ ``None`` (the default) or if ``obj.dtype`` is equal to the PyTorch
1557
+ equivalent ot `dtype`. Alternately, `dtype` can be a tuple, in which
1558
+ case, `obj` is considered valid if its dtype is any of the dtypes in
1559
+ `dtype`. ndim : int or tuple or ints or None, optional The number of
1560
+ dimensions that the object must have in order to be considered a valid
1561
+ tensor. If ``None``, then any number of dimensions is acceptable (this
1562
+ is the default). If this is an integer, then the number of dimensions
1563
+ must be exactly that integer. If this is a list or tuple of integers,
1564
+ then the dimensionality must be one of these numbers.
1565
+ shape : int, tuple of ints, None, optional
1566
+ If the `shape` parameter is not ``None``, then the given `obj` must
1567
+ have a shape shape that matches the parameter value. The value of
1568
+ `shape` must be a tuple that is equal to `obj`'s shape tuple with the
1569
+ following additional rules: a ``-1`` value in the `shape` tuple will
1570
+ match any value in the `obj`'s shape tuple, and a single ``Ellipsis``
1571
+ may appear in `shape`, which matches any number of values in the
1572
+ `obj`'s shape tuple. The default value of ``None`` indicates that no
1573
+ restriction should be applied to the `obj`'s shape.
1574
+ numel : int, tuple of ints, or None, optional
1575
+ If the `numel` parameter is not ``None``, then the given `obj` must
1576
+ have the same number of elements as given by `numel`. If `numel` is a
1577
+ tuple, then the number of elements in `obj` must be in the `numel`
1578
+ tuple. The number of elements is the product of its shape.
1579
+ device : device-name or None, optional
1580
+ If `device` is ``None``, then a tensor with any `device` field is
1581
+ considered valid; otherwise, the `device` parameter must equal
1582
+ ``obj.device`` for `obj` to be considered a valid tensor. The default
1583
+ value is ``None``.
1584
+ requires_grad : bool or None, optional
1585
+ If ``None``, then a tensor with any `requires_grad` field is considered
1586
+ valid; otherwise, the `requires_grad` parameter must equal
1587
+ ``obj.requires_grad`` for `obj` to be considered a valid tensor. The
1588
+ default value is ``None``.
1589
+ sparse : bool or None, optional
1590
+ If the `sparse` parameter is ``None``, then no requirements are placed
1591
+ on the sparsity of `obj` for it to be considered a valid tensor. If
1592
+ `sparse` is ``True`` or ``False``, then `obj` must either be sparse or
1593
+ not be sparse, respectively, for `obj` to be considered valid. If
1594
+ `sparse` is a string, then it must be either ``'coo'`` or ``'csr'``,
1595
+ indicating the required sparse array type.
1596
+ quant : bool, optional
1597
+ Whether ``pint.Quantity`` objects should be considered valid tensors or
1598
+ not. If `quant` is ``True`` then `obj` is considered a valid array
1599
+ only when `obj` is a quantity object with a ``torch`` tensor as the
1600
+ magnitude. If `quant` is ``False``, then `obj` must be a ``torch``
1601
+ tensor itself and not a ``Quantity`` to be considered valid. If `quant`
1602
+ is ``None`` (the default), then either quantities or ``torch`` tensors
1603
+ are considered valid.
1604
+ unit : unit-like, None, Ellipsis, optional
1605
+ A unit with which the object obj's unit must be compatible in order for
1606
+ `obj` to be considered a valid tensor. An `obj` that is not a quantity
1607
+ is considered to have a unit of ``None``. If ``unit=Ellipsis`` (the
1608
+ default), then the object's unit is ignored.
1609
+ ureg : UnitRegistry, None, or Ellipsis, optional
1610
+ The ``pint.UnitRegistry`` object to use for units. If `ureg` is
1611
+ ``Ellipsis``, then ``immlib.units`` is used. If `ureg` is ``None`` (the
1612
+ default), then the registry of `obj` is used if `obj` is a quantity,
1613
+ and ``immlib.units`` is used if not.
1614
+
1615
+ Returns
1616
+ -------
1617
+ boolean
1618
+ ``True`` if `obj` is a valid PyTorch tensor whose properties match the
1619
+ requirements spelled out by the optional parameters, otherwise
1620
+ ``False``.
1621
+
1622
+ See Also
1623
+ --------
1624
+ is_array, is_numeric
1625
+ """
1626
+ # If so, we can process the arguments.
1627
+ if ureg is Ellipsis:
1628
+ from immlib import units as ureg
1629
+ # If this is a quantity, just extract the magnitude.
1630
+ if isinstance(obj, pint.Quantity):
1631
+ if quant is False:
1632
+ return False
1633
+ if ureg is None:
1634
+ from ._quantity import unitregistry
1635
+ ureg = unitregistry(obj)
1636
+ u = obj.u
1637
+ obj = obj.m
1638
+ else:
1639
+ if quant is True:
1640
+ return False
1641
+ if ureg is None:
1642
+ from immlib import units as ureg
1643
+ u = None
1644
+ # Right away: is this a torch tensor or not?
1645
+ if not torch.is_tensor(obj):
1646
+ return False
1647
+ # Do we match the various torch field requirements?
1648
+ if device is not None:
1649
+ device = torch.device(device)
1650
+ if obj.device != device:
1651
+ return False
1652
+ if requires_grad is not None:
1653
+ if obj.requires_grad != requires_grad:
1654
+ return False
1655
+ # Do we match the sparsity requirement?
1656
+ if sparse is True:
1657
+ if not torch__is_sparse(obj):
1658
+ return False
1659
+ elif sparse is False:
1660
+ if torch__is_sparse(obj):
1661
+ return False
1662
+ elif sparse is not None:
1663
+ layout = sparse_layout(sparse)
1664
+ if layout is None:
1665
+ if isinstance(sparse, str):
1666
+ raise ValueError(f"invalid sparse option: {sparse}")
1667
+ else:
1668
+ raise ValueError(
1669
+ f"invalid sparse option of type {type(sparse)}")
1670
+ if not sparse_haslayout(obj, layout):
1671
+ return False
1672
+ # Next, check compatibility of the units.
1673
+ if unit is None:
1674
+ # We are required to not be a quantity.
1675
+ if u is not None:
1676
+ return False
1677
+ elif unit is not Ellipsis:
1678
+ from ._quantity import alike_units
1679
+ if not is_tuple(unit):
1680
+ unit = (unit,)
1681
+ if not any(alike_units(u, uu) for uu in unit):
1682
+ return False
1683
+ # Check the match to the numeric collection last.
1684
+ if dtype is None and shape is None and ndim is None and numel is None:
1685
+ return True
1686
+ return _numcoll_match(obj.shape, obj.dtype, ndim, shape, numel, dtype)
1687
+ @docwrap('immlib.to_tensor')
1688
+ def to_tensor(obj, /, dtype=None, *,
1689
+ device=None, requires_grad=None, copy=False,
1690
+ sparse=None, quant=None, ureg=None, unit=Ellipsis):
1691
+ """Reinterprets `obj` as a PyTorch tensor or as a ``pint`` quantity with
1692
+ a tensor magnitude.
1693
+
1694
+ ``immlib.to_tensor`` is roughly equivalent to the ``torch.as_tensor``
1695
+ function with a few exceptions:
1696
+
1697
+ - ``to_tensor(obj)`` allows quantities for `obj` and, in such a case,
1698
+ will return a quantity whose magnitude has been reinterpreted as a
1699
+ tensor, though this behavior can be altered with the `quant`
1700
+ parameter;
1701
+ - ``to_tensor(obj)`` can convet a SciPy sparse matrix into a sparse
1702
+ tensor.
1703
+
1704
+ Parameters
1705
+ ----------
1706
+ obj : object
1707
+ The object that is to be reinterpreted as or covnerted to, a PyTorch
1708
+ tensor object.
1709
+ dtype : dtype-like, optional
1710
+ The `dtype` that is passed to ``torch.as_tensor(obj)``.
1711
+ device : device name or None, optional
1712
+ The `device` parameter that is passed to ``torch.as_tensor(obj)``,
1713
+ ``None`` by default.
1714
+ requires_grad : bool or None, optional
1715
+ Whether the returned tensor should require gradient calculations or
1716
+ not. If ``None`` (the default), then the objecct `obj` is not changed
1717
+ from its current gradient settings, if `obj` is a tensor, and `obj` is
1718
+ not made to track its gradient if it is converted into a tensor. If the
1719
+ `requires_grad` parameter does not match the given tensor's
1720
+ `requires_grad` field, then a copy is always returned.
1721
+ copy : bool, optional
1722
+ Whether to copy the data in `obj` or not. If ``False``, then `obj` is
1723
+ only copied if doing so is required by the optional parameters. If
1724
+ ``True``, then `obj` is always copied if possible.
1725
+ sparse : bool, {'csr','csc','bsr','bsc','coo'}, or None, optional
1726
+ If ``None``, then the sparsity of `obj` is the same as the sparsity of
1727
+ the tensor that is returned. Otherwise, the return value will always be
1728
+ either a spase tensor (``sparse=True``) or a dense tensor
1729
+ (``sparse=False``) based on the given value of ``sparse``. The
1730
+ ``sparse`` parameter may also be set to the name of a sparse layout in
1731
+ order to convert the object into that layout. Possible sparse layouts
1732
+ include ``'coo'``, ``'csr'``, ``'csc'``, ``'bsr'``, and ``'bsc'``.
1733
+ quant : bool or None, optional
1734
+ Whether the return value should be a ``Quantity`` object wrapping the
1735
+ array (``quant=True``) or the tensor itself (``quant=False``). If
1736
+ `quant` is ``None`` (the default) then the return value is a quantity
1737
+ if either `obj` is a quantity or an explicit `unit` parameter is given
1738
+ and is not a quantity if `obj` is not a quantity.
1739
+ ureg : pint.UnitRegistry or None, optional
1740
+ The ``pint.UnitRegistry`` object to use for units. If `ureg` is
1741
+ ``Ellipsis``, then ``immlib.units`` is used. If `ureg` is ``None`` (the
1742
+ default), then no specific coersion to a ``pint.UnitRegistry`` is
1743
+ performed (i.e., the specific subclass of ``pint.Quantity`` used by
1744
+ `obj` is not changed).
1745
+ unit : unit-like, bool, None, or Ellipsis, optional
1746
+ The unit that should be used in the return value. When the return value
1747
+ of this function is a ``pint.Quantity`` (see the `quant` parameter),
1748
+ the returned quantity always has units matching the `unit` parameter;
1749
+ if the provided `obj` is not a quantity, then its unit is presumed to
1750
+ be that requested by `unit`. When the return value of this function is
1751
+ not a ``pint.Quantity`` object and is instead a PyTorch tensor object,
1752
+ then when `obj` is not a quantity the `unit` parameter is ignored, and
1753
+ when `obj` is a quantity, its magnitude is returned after conversion
1754
+ into `unit`. The default value of `unit`, ``Ellipsis``, indicates that,
1755
+ if `obj` is a quantity, its unit should be used, and `unit` should be
1756
+ considered dimensionless otherwise.
1757
+
1758
+ Returns
1759
+ -------
1760
+ torch.Tensor or pint.Quantity
1761
+ Either a PyTorch tensor equivalent to `obj` or a ``pint.Quantity``
1762
+ whose magnitude is a PyTorch tensor equivalent to `obj`.
1763
+
1764
+ Raises
1765
+ ------
1766
+ ValueError
1767
+ If invalid parameter values are given or if the parameters conflict.
1768
+
1769
+ See Also
1770
+ --------
1771
+ to_array, to_numeric
1772
+ """
1773
+ if ureg is Ellipsis:
1774
+ from immlib import units as ureg
1775
+ if dtype is not None:
1776
+ dtype = to_torchdtype(dtype)
1777
+ # If obj is a quantity, we handle things differently.
1778
+ if isinstance(obj, pint.Quantity):
1779
+ q = obj
1780
+ obj = q.m
1781
+ if ureg is None:
1782
+ from ._quantity import unitregistry
1783
+ ureg = unitregistry(q)
1784
+ else:
1785
+ q = None
1786
+ if ureg is None:
1787
+ from immlib import units as ureg
1788
+ # Translate obj depending on whether it's a pytorch tensor already or a
1789
+ # scipy sparse matrix.
1790
+ if torch.is_tensor(obj):
1791
+ if requires_grad is None:
1792
+ requires_grad = obj.requires_grad
1793
+ if device is None:
1794
+ device = obj.device
1795
+ if dtype is None:
1796
+ dtype = obj.dtype
1797
+ needs_copy = device != obj.device or dtype != obj.dtype
1798
+ prefs_copy = copy or requires_grad != obj.requires_grad
1799
+ if copy is False:
1800
+ if needs_copy:
1801
+ if device == obj.device:
1802
+ msg = "dtype change"
1803
+ elif dtype == obj.dtype:
1804
+ msg = "device change"
1805
+ else:
1806
+ msg = "device and dtype change"
1807
+ raise ValueError(
1808
+ f"copy=False requested, but copy required by {msg}")
1809
+ else:
1810
+ arr = obj
1811
+ elif needs_copy:
1812
+ arr = obj.to(dtype=dtype, device=device)
1813
+ elif copy:
1814
+ arr = obj.detach().clone()
1815
+ else:
1816
+ arr = obj
1817
+ if arr.requires_grad != requires_grad:
1818
+ arr = arr.requires_grad_(requires_grad)
1819
+ else:
1820
+ if requires_grad is None:
1821
+ requires_grad = False
1822
+ if scipy__is_sparse(obj):
1823
+ (rows, cols, vals) = sps.find(obj)
1824
+ # Process these into a PyTorch COO matrix.
1825
+ ii = torch.as_tensor(
1826
+ np.array([rows, cols], dtype=np.int_),
1827
+ dtype=torch.long,
1828
+ device=device)
1829
+ vals = torch.as_tensor(vals, dtype=dtype, device=device)
1830
+ if dtype is None:
1831
+ dtype = vals.dtype
1832
+ arr = torch.sparse_coo_tensor(
1833
+ ii, vals, obj.shape,
1834
+ dtype=dtype,
1835
+ device=device,
1836
+ requires_grad=requires_grad)
1837
+ # If possible, convert to the layout we want.
1838
+ arr = sparse_tolayout(arr, obj.format)
1839
+ elif (copy or requires_grad is True or
1840
+ (isinstance(obj, np.ndarray) and not obj.flags['WRITEABLE'])):
1841
+ arr = torch.tensor(
1842
+ obj,
1843
+ dtype=dtype,
1844
+ device=device,
1845
+ requires_grad=requires_grad)
1846
+ dtype = arr.dtype
1847
+ else:
1848
+ arr = torch.as_tensor(obj, dtype=dtype, device=device)
1849
+ dtype = arr.dtype
1850
+ # If there is an instruction regarding the output's sparsity, handle that
1851
+ # now.
1852
+ if sparse is True:
1853
+ # arr must be sparse (COO by default); make sure it is.
1854
+ if not torch__is_sparse(arr):
1855
+ arr = arr.to_sparse()
1856
+ elif sparse is False:
1857
+ # arr must not be a sparse array; make sure it isn't.
1858
+ if torch__is_sparse(arr):
1859
+ arr = arr.to_dense()
1860
+ elif sparse is not None:
1861
+ layout = sparse_layout(sparse)
1862
+ if layout is None:
1863
+ if isinstance(sparse, str):
1864
+ raise ValueError(f"invalid pytorch sparse layout: {sparse}")
1865
+ else:
1866
+ raise ValueError(
1867
+ f"to_tensor: invalid sparse option of type {type(sparse)}")
1868
+ if arr.layout != layout.torch_layout:
1869
+ arr = layout.torch_tomethod(arr)
1870
+ # Next, we switch on whether we are being asked to return a quantity or
1871
+ # not.
1872
+ if quant is None:
1873
+ quant = (q if unit is Ellipsis else unit) is not None
1874
+ if quant is True:
1875
+ if unit is None:
1876
+ raise ValueError(
1877
+ "to_tensor: cannot make a quantity (quant=True) without a unit"
1878
+ " (unit=None)")
1879
+ if q is None:
1880
+ if unit is Ellipsis:
1881
+ raise ValueError(
1882
+ "to_tensor(x): cannot make a quantity (quant=True) with"
1883
+ " the same unit as x (unit=Ellipsis) when x is not a"
1884
+ " quantity")
1885
+ return ureg.Quantity(arr, unit)
1886
+ else:
1887
+ from ._quantity import unitregistry
1888
+ if ureg is not unitregistry(q) or obj is not arr:
1889
+ q = ureg.Quantity(arr, q.u)
1890
+ if unit is not Ellipsis and ureg.Unit(unit) != q.u:
1891
+ return q.to(unit)
1892
+ else:
1893
+ return q
1894
+ elif quant is False:
1895
+ # Don't return a quantity, whatever the input argument.
1896
+ if unit is Ellipsis:
1897
+ # We return the current array/magnitude whatever its unit.
1898
+ return arr
1899
+ elif q is None:
1900
+ # We just pretend this was already in the given unit (i.e., ignore
1901
+ # unit).
1902
+ return arr
1903
+ elif unit is None:
1904
+ raise ValueError(
1905
+ "to_tensor: cannot extract unit None from quantity; to get the"
1906
+ " native unit, use unit=Ellipsis")
1907
+ else:
1908
+ if obj is not arr:
1909
+ q = ureg.Quantity(arr, q.u)
1910
+ # We convert to the given unit and return that.
1911
+ return q.m_as(unit)
1912
+ else:
1913
+ raise ValueError(
1914
+ f"to_tensor: quant must be boolean or None;"
1915
+ f" got object of type {type(quant)}")
1916
+
1917
+
1918
+ # General Numeric Collection Functions ########################################
1919
+
1920
+ @docwrap('immlib.is_numeric')
1921
+ def is_numeric(obj, /, dtype=None, *,
1922
+ shape=None, ndim=None, numel=None,
1923
+ sparse=None, quant=None, unit=Ellipsis, ureg=None):
1924
+ """Returns ``True`` if an object is a numerical collection type and
1925
+ ``False`` otherwise.
1926
+
1927
+ ``is_numeric(obj)`` returns ``True`` if the given object `obj` is an
1928
+ instance of the ``torch.Tensor`` class, the ``numpy.ndarray`` class, one
1929
+ one of the ``scipy.sparse`` array classes, or is a ``pint.Quantity``
1930
+ object whose magnitude is an instance of one of these types. Additional
1931
+ constraints may be placed on the object via the optional argments.
1932
+
1933
+ .. Note:: The ``is_numeric`` function is similar to the ``is_array`` and
1934
+ ``is_tensor`` functions butis agnostic about whether its argument is a
1935
+ PyTorch tensor, a NumPy array, or an object that can be converted into
1936
+ one of these types.
1937
+
1938
+ Parameters
1939
+ ----------
1940
+ obj : object
1941
+ The object whose quality as a numeric object is to be assessed.
1942
+ dtype : dtype-like or None, optional
1943
+ The NumPy or PyTorch `dtype` or dtype-like object that is required to
1944
+ match that of the `obj` in order to be considered valid. The
1945
+ ``obj.dtype`` matches the given `dtype` parameter if either `dtype` is
1946
+ ``None`` (the default) or if ``obj.dtype`` is equivalent to
1947
+ `dtype`. Alternately, `dtype` can be a tuple, in which case, `obj` is
1948
+ considered valid if its dtype is any of the dtypes in `dtype`.
1949
+ ndim : int, tuple or ints, or None, optional
1950
+ The number of dimensions that the object must have in order to be
1951
+ considered valid. If `ndim` is ``None``, then any number of dimensions
1952
+ is acceptable (this is the default). If it is an integer, then the
1953
+ number of dimensions must be exactly that integer. If this is a list or
1954
+ tuple of integers, then the dimensionality must be one of the listedn
1955
+ numbers.
1956
+ shape : int, tuple of ints, or None, optional
1957
+ If the `shape` parameter is not ``None``, then the given `obj` must
1958
+ have a shape that matches the parameter value. The value of `shape`
1959
+ must be a tuple that is equal to the shape of `obj` with the following
1960
+ additional rules: a ``-1`` value in the `shape` tuple will match any
1961
+ value in the shape of `obj`, and a single ``Ellipsis`` may appear in
1962
+ `shape`, which matches any number of values in the shape tuple of
1963
+ `obj`. The default value of ``None`` indicates that no restriction
1964
+ should be applied to the shape of `obj`.
1965
+ sparse : bool or False, optional
1966
+ If the ``sparse`` parameter is ``None``, then no requirements are
1967
+ placed on the sparsity of `obj` for it to be considered valid. If
1968
+ `sparse` is ``True`` or ``False``, then `obj` must either be sparse or
1969
+ not be sparse, respectively, for `obj` to be considered valid. If
1970
+ `sparse` is a string, then it must be a valid sparse array type that
1971
+ matches the type of `obj` for `obj` to be considered valid.
1972
+ numel : int, tuple of ints, or None, optional
1973
+ If the `numel` parameter is not ``None``, then the given `obj` must
1974
+ have the same number of elements as given by `numel`. If `numel` is a
1975
+ tuple, then the number of elements in `obj` must be in the `numel`
1976
+ tuple. The number of elements is the product of its shape.
1977
+ quant : bool, optional
1978
+ Whether ``pint.Quantity`` objects should be considered valid or not.
1979
+ If ``quant=True`` then `obj` is considered a valid numerical object
1980
+ only when `obj` is a quantity object with a valid numerical object as
1981
+ the magnitude. If ``quant=False``, then `obj` must be a numerical
1982
+ object itself and not a ``pint.Quantity`` to be considered valid. If
1983
+ ``quant=None`` (the default), then either quantities or numerical
1984
+ objects are considered valid.
1985
+ unit : unit-like or Ellipsis, optional
1986
+ A unit with which the unit of `obj` must be compatible in order for
1987
+ `obj` to be considered a valid numerical object. An `obj` that is not a
1988
+ quantity is considered to have dimensionless units. If
1989
+ ``unit=Ellipsis`` (the default), then the object's unit is ignored.
1990
+ ureg : UnitRegistry, None, or Ellipsis, optional
1991
+ The ``pint.UnitRegistry`` object to use for units. If `ureg` is
1992
+ ``Ellipsis``, then ``immlib.units`` is used. If `ureg` is ``None`` (the
1993
+ default), then the registry of `obj` is used if `obj` is a quantity,
1994
+ and ``immlib.units`` is used if not.
1995
+
1996
+ Returns
1997
+ -------
1998
+ bool
1999
+ ``True`` if `obj` is a valid numerical object, otherwise ``False``.
2000
+
2001
+ See Also
2002
+ --------
2003
+ is_array, is_tensor
2004
+ """
2005
+ if isinstance(obj, pint.Quantity):
2006
+ istns = torch.is_tensor(obj.m)
2007
+ else:
2008
+ istns = torch.is_tensor(obj)
2009
+ if istns:
2010
+ return is_tensor(obj,
2011
+ dtype=dtype, shape=shape, ndim=ndim, numel=numel,
2012
+ sparse=sparse, quant=quant, unit=unit, ureg=ureg)
2013
+ else:
2014
+ return is_array(obj,
2015
+ dtype=dtype, shape=shape, ndim=ndim, numel=numel,
2016
+ sparse=sparse, quant=quant, unit=unit, ureg=ureg)
2017
+ @docwrap('immlib.to_numeric')
2018
+ def to_numeric(obj, /, dtype=None, *,
2019
+ copy=False, sparse=None, quant=None, ureg=None, unit=Ellipsis):
2020
+ """Reinterprets `obj` as a numeric type or quantity with such a magnitude.
2021
+
2022
+ ``immlib.to_numeric`` is roughly equivalent to the ``torch.as_tensor`` or
2023
+ ``numpy.asarray`` function with a few exceptions:
2024
+
2025
+ - ``to_numeric(obj)`` allows quantities for `obj` and, in such a case,
2026
+ will return a quantity whose magnitude has been reinterpreted as a
2027
+ numeric object, though this behavior can be altered with the ``quant``
2028
+ parameter;
2029
+ - ``to_numeric(obj)`` correctly handles SciPy sparse matrices, NumPy
2030
+ arrays, and PyTorch tensors.
2031
+
2032
+ If the object `obj` passed to ``immlib.to_numeric(obj)`` is a PyTorch
2033
+ tensor or a quantity whose magnitude is a PyTorch tensor, then a PyTorch
2034
+ tensor or a quantity with a PyTorch tensor magnitude is
2035
+ returned. Otherwise, a NumPy array, SciPy sparse matrix, or quantity with a
2036
+ magnitude matching one of these types is returned.
2037
+
2038
+ Parameters
2039
+ ----------
2040
+ obj : object
2041
+ The object that is to be reinterpreted as, or if necessary covnerted
2042
+ to, a numeric object.
2043
+ dtype : dtype-like or None, optional
2044
+ The dtype that is passed to ``torch.as_tensor(obj)`` or
2045
+ ``np.asarray(obj)``.
2046
+ copy : bool, optional
2047
+ Whether to copy the data in `obj` or not. If ``False``, then `obj` is
2048
+ only copied if doing so is required by the optional parameters. If
2049
+ ``True``, then `obj` is always copied if possible.
2050
+ sparse : bool, {'csr','csc','bsr','bsc','coo'}, or None, optional
2051
+ If ``None``, then the sparsity of `obj` is the same as the sparsity of
2052
+ the object that is returned. Otherwise, the return value will always be
2053
+ either a spase object (``sparse=True``) or a dense object
2054
+ (``sparse=False``) based on the given value of `sparse`. The `sparse`
2055
+ parameter may also be set to ``'coo'``, ``'csr'``, or other sparse
2056
+ matrix names to return specific sparse layouts.
2057
+ quant : bool or None, optional
2058
+ Whether the return value should be a ``pint.Quantity`` object wrapping
2059
+ wrapping the object (``quant=True``) or the object itself
2060
+ (``quant=False``). If `quant` is ``None`` (the default) then the return
2061
+ value is a quantity if `obj` is a quantity and is not a quantity if
2062
+ `obj` is not a quantity.
2063
+ ureg : pint.UnitRegistry or None, optional
2064
+ The ``pint.UnitRegistry`` object to use for units. If `ureg` is
2065
+ ``Ellipsis``, then ``immlib.units`` is used. If `ureg` is ``None`` (the
2066
+ default), then no specific coersion to a ``pint.UnitRegistry`` is
2067
+ performed (i.e., the same quantity class is returned).
2068
+ unit : unit-like, bool, None, or Ellipsis, optional
2069
+ The unit that should be used in the return value. When the return value
2070
+ of this function is a ``pint.Quantity`` (see the `quant` parameter),
2071
+ the returned quantity always has a unit matching the `unit` parameter;
2072
+ if the provided `obj` is not a quantity, then its unit is presumed to
2073
+ be those requested by `unit`. When the return value of this function is
2074
+ not a ``pint.Quantity`` object and is instead a numeric object, then
2075
+ when `obj` is not a quantity the `unit` parameter is ignored, and when
2076
+ `obj` is a quantity, its magnitude is returned after conversion into
2077
+ `unit`. The default value of `unit`, ``Ellipsis``, indicates that, if
2078
+ `obj` is a quantity, its unit should be used, and `unit` should be
2079
+ considered to be ``None`` otherwise.
2080
+
2081
+ Returns
2082
+ -------
2083
+ NumPy array or PyTorch tensor or Quantity
2084
+ Either a NumPy array or PyTorch tensor equivalent to `obj` or a
2085
+ ``pint.Quantity`` whose magnitude is such an object.
2086
+
2087
+ Raises
2088
+ ------
2089
+ ValueError
2090
+ If invalid parameter values are given or if the parameters conflict.
2091
+
2092
+ See Also
2093
+ --------
2094
+ to_array, to_tensor
2095
+ """
2096
+ if torch.is_tensor(obj):
2097
+ return to_tensor(obj,
2098
+ dtype=dtype, sparse=sparse,
2099
+ quant=quant, unit=unit, ureg=ureg)
2100
+ else:
2101
+ return to_array(obj,
2102
+ dtype=dtype, sparse=sparse,
2103
+ quant=quant, unit=unit, ureg=ureg)
2104
+
2105
+
2106
+ # Sparse Matrices and Dense Collections #######################################
2107
+
2108
+ @docwrap('immlib.is_sparse')
2109
+ def is_sparse(obj, /, dtype=None, *,
2110
+ shape=None, ndim=None, numel=None,
2111
+ quant=None, ureg=None, unit=Ellipsis):
2112
+ """Returns ``True`` if an object is a sparse SciPy array or a sparse
2113
+ PyTorch tensor.
2114
+
2115
+ ``is_sparse(obj)`` returns ``True`` if the given object `obj` is an
2116
+ instance of one of the SciPy sparse array classes, is a sparse PyTorch
2117
+ tensor, or is a ``pint.Quantity`` whose magnintude is one of
2118
+ theese. Additional constraints may be placed on the object via the optional
2119
+ argments.
2120
+
2121
+ Parameters
2122
+ ----------
2123
+ obj : object
2124
+ The object whose quality as a sparse numerical object is to be
2125
+ assessed.
2126
+ %(immlib.is_numeric.parameters.dtype)s
2127
+ %(immlib.is_numeric.parameters.ndim)s
2128
+ %(immlib.is_numeric.parameters.shape)s
2129
+ %(immlib.is_numeric.parameters.numel)s
2130
+ %(immlib.is_numeric.parameters.quant)s
2131
+ %(immlib.is_numeric.parameters.ureg)s
2132
+ %(immlib.is_numeric.parameters.unit)s
2133
+ sparsetype : 'matrix', 'array', or None, optional
2134
+ The kind of sparse array to accept: either ``'matrix'`` for the scipy
2135
+ sparse matrix types (e.g., ``scipy.sparse.csr_matrix``) or ``'array'``
2136
+ for the sparse array types (e.g., ``scipy.sparse.csr_array``). If the
2137
+ value is ``None`` (the default) then either is accepted.
2138
+
2139
+ Returns
2140
+ -------
2141
+ bool
2142
+ ``True`` if `obj` is a valid sparse numerical object, otherwise
2143
+ ``False``.
2144
+ """
2145
+ return is_numeric(obj, sparse=True,
2146
+ dtype=dtype, shape=shape, ndim=ndim, numel=numel,
2147
+ quant=quant, ureg=ureg, unit=unit)
2148
+ @docwrap('immlib.to_sparse')
2149
+ def to_sparse(obj, /, dtype=None, *, quant=None, ureg=None, unit=Ellipsis):
2150
+ """Returns a sparse version of the numerical object `obj`.
2151
+
2152
+ ``to_sparse(obj)`` returns `obj` if it is already a PyTorch sparse tensor
2153
+ or a SciPy sparse matrix or a quantity whose magnitude is one of these.
2154
+ Otherwise, it converts `obj` into a sparse representation and returns
2155
+ this. Additional requirements on the output format of the return value can
2156
+ be added using the optional parameters.
2157
+
2158
+ Parameters
2159
+ ----------
2160
+ obj : object
2161
+ The object that is to be converted into a sparse representation.
2162
+ %(immlib.to_numeric.parameters.dtype)s
2163
+ %(immlib.to_numeric.parameters.quant)s
2164
+ %(immlib.to_numeric.parameters.ureg)s
2165
+ %(immlib.to_numeric.parameters.unit)s
2166
+
2167
+ Returns
2168
+ -------
2169
+ sparse tensor, sparse array, or quantity with a sparse magnitude
2170
+ A sparse version of the argument `obj`.
2171
+ """
2172
+ return to_numeric(obj, sparse=True,
2173
+ dtype=dtype, quant=quant,
2174
+ ureg=ureg, unit=unit)
2175
+ @docwrap('immlib.is_dense')
2176
+ def is_dense(obj, /, dtype=None, *,
2177
+ shape=None, ndim=None, numel=None,
2178
+ quant=None, ureg=None, unit=Ellipsis):
2179
+ """Returns ``True`` if an object is a dense NumPy array or PyTorch tensor.
2180
+
2181
+ ``is_dense(obj)`` returns ``True`` if the given object `obj` is an instance
2182
+ of one of the NumPy ``ndarray`` classes, is a dense PyTorch tensor, or is a
2183
+ quantity whose magnintude is one of theese. Additional constraints may be
2184
+ placed on the object via the optional argments.
2185
+
2186
+ Parameters
2187
+ ----------
2188
+ obj : object
2189
+ The object whose quality as a dense numerical object is to be assessed.
2190
+ %(immlib.is_numeric.parameters.dtype)s
2191
+ %(immlib.is_numeric.parameters.ndim)s
2192
+ %(immlib.is_numeric.parameters.shape)s
2193
+ %(immlib.is_numeric.parameters.numel)s
2194
+ %(immlib.is_numeric.parameters.quant)s
2195
+ %(immlib.is_numeric.parameters.ureg)s
2196
+ %(immlib.is_numeric.parameters.unit)s
2197
+
2198
+ Returns
2199
+ -------
2200
+ bool
2201
+ ``True`` if `obj` is a valid dense numerical object, otherwise
2202
+ ``False``.
2203
+ """
2204
+ return is_numeric(obj, sparse=False,
2205
+ dtype=dtype, shape=shape, ndim=ndim, numel=numel,
2206
+ quant=quant, ureg=ureg, unit=unit)
2207
+ @docwrap('immlib.to_dense')
2208
+ def to_dense(obj, /, dtype=None, *, quant=None, ureg=None, unit=Ellipsis):
2209
+ """Returns a dense version of the numerical object `obj`.
2210
+
2211
+ ``to_dense(obj)`` returns `obj` if it is already a PyTorch dense tensor or
2212
+ a NumPy ``ndarray`` or a quantity whose magnitude is one of these.
2213
+ Otherwise, it converts `obj` into a dense representation and returns
2214
+ this. Additional requirements on the output format of the return value can
2215
+ be added using the optional parameters.
2216
+
2217
+ Parameters
2218
+ ----------
2219
+ obj : object
2220
+ The object that is to be converted into a dense representation.
2221
+ %(immlib.to_numeric.parameters.dtype)s
2222
+ %(immlib.to_numeric.parameters.quant)s
2223
+ %(immlib.to_numeric.parameters.ureg)s
2224
+ %(immlib.to_numeric.parameters.unit)s
2225
+
2226
+ Returns
2227
+ -------
2228
+ dense tensor, dense array, or quantity with a dense magnitude
2229
+ A dense version of the argument `obj`.
2230
+ """
2231
+ return to_numeric(obj, sparse=False,
2232
+ dtype=dtype, quant=quant, ureg=ureg, unit=unit)
2233
+
2234
+
2235
+ # Numeric Decorators ##########################################################
2236
+
2237
+ @docwrap('immlib.numapi')
2238
+ class numapi:
2239
+ """An interface for defining functions that expect all arguments to be
2240
+ either numpy arrays or pytorch tensors.
2241
+
2242
+ A function decorated with ``@numapi`` is a placeholder for two
2243
+ subfunctions: one that is called when any of the arguments are pytorch
2244
+ tensors (all of whose arguments, when possible, are converted into pytorch
2245
+ tensors), and a version called otherwise, all of whose arguments are
2246
+ converted into numpy arrays when possible. The body of the decorated
2247
+ function is usually ``pass``, but, if desired, it can return either the
2248
+ pytorch or the numpy modules to indicate that a particular version of the
2249
+ function should be called (if necessary, tensors are converted into numpy
2250
+ arrays for this).
2251
+
2252
+ Once a function has been decorated with ``@numapi``, that function should
2253
+ be used to decorate two other functions. If, for example, the function
2254
+ ``f`` is decorated with ``@numapi``, then ``@f.array`` should be used to
2255
+ decorate the version of the function that accepts numpy arrays and
2256
+ ``@f.tensor`` should be used to decorate the version of the function that
2257
+ accepts pytorch tensors.
2258
+
2259
+ Examples
2260
+ --------
2261
+ >>> @numapi
2262
+ ... def l2_distance(pt1, pt2):
2263
+ ... "Calculates the L2 distance between two points."
2264
+ ... pass
2265
+
2266
+ >>> @l2_distance.array
2267
+ ... def _(pt1, pt2):
2268
+ ... return np.sqrt(np.sum((pt1 - pt2)**2, axis=0))
2269
+
2270
+ >>> @l2_distance.tensor
2271
+ ... def _(pt1, pt2):
2272
+ ... return torch.sqrt(torch.sum((pt1 - pt2)**2, axis=0))
2273
+
2274
+ >>> l2_distance(torch.tensor([0,0]), [0,1])
2275
+ tensor(1.)
2276
+
2277
+ >>> l2_distance([0,0], torch.tensor([0,1]))
2278
+ tensor(1.)
2279
+
2280
+ >>> l2_distance([0,0], [0,1])
2281
+ 1.0
2282
+ """
2283
+ # Static Methods ----------------------------------------------------------
2284
+ @staticmethod
2285
+ def _as_array(arg):
2286
+ argmod = type(arg).__module__
2287
+ if argmod == 'numpy' or argmod.startswith('numpy.'):
2288
+ return arg
2289
+ else:
2290
+ return to_array(arg)
2291
+ @staticmethod
2292
+ def _as_tensor(arg, device=None):
2293
+ if is_tensor(arg):
2294
+ return to_tensor(arg, device=device)
2295
+ argmod = type(arg).__module__
2296
+ if argmod == 'torch' or argmod.startswith('torch.'):
2297
+ return arg
2298
+ # Otherwise, try converting it to a tensor.
2299
+ try:
2300
+ return to_tensor(arg, device=device)
2301
+ except (TypeError, RuntimeError):
2302
+ pass
2303
+ # If all else fails, convert it to a numpy array.
2304
+ return numapi._as_array(arg)
2305
+ @staticmethod
2306
+ def _find_device(args, kwargs):
2307
+ dev = kwargs.get('device', None)
2308
+ if dev is not None:
2309
+ return torch.device(dev)
2310
+ tns = next(filter(is_tensor, args), None)
2311
+ if tns is None:
2312
+ tns = next(filter(is_tensor, kwargs.values()), None)
2313
+ if tns is None:
2314
+ return None
2315
+ return tns.device
2316
+ # Constructor -------------------------------------------------------------
2317
+ __slots__ = (
2318
+ 'base_func', 'wrap_func', 'array_func', 'tensor_func', 'signature')
2319
+ def __new__(cls, fn):
2320
+ self = object.__new__(cls)
2321
+ def wrap_fn(*args, **kwargs):
2322
+ return self(*args, **kwargs)
2323
+ wrap_fn.numapi = self
2324
+ wrap_fn.array = self.array
2325
+ wrap_fn.tensor = self.tensor
2326
+ self.base_func = fn
2327
+ self.array_func = None
2328
+ self.tensor_func = None
2329
+ self.wrap_func = wrap_fn
2330
+ self.signature = inspect.signature(fn)
2331
+ return wraps(fn)(wrap_fn)
2332
+ # Methods -----------------------------------------------------------------
2333
+ def array(self, f):
2334
+ self.array_func = wraps(self.base_func)(f)
2335
+ def tensor(self, f):
2336
+ self.tensor_func = wraps(self.base_func)(f)
2337
+ def _bind(self, args, kwargs):
2338
+ b = self.signature.bind(*args, **kwargs)
2339
+ b.apply_defaults()
2340
+ return (b.args, b.kwargs)
2341
+ def _call_torch(self, args, kwargs):
2342
+ if not self.tensor_func:
2343
+ raise RuntimeError(
2344
+ f"tensor function for {self.__name__} was not defined")
2345
+ dev = numapi._find_device(args, kwargs)
2346
+ (args, kwargs) = self._bind(args, kwargs)
2347
+ args = map(numapi._as_tensor, args)
2348
+ if kwargs:
2349
+ kwargs = {k: numapi._as_tensor(v) for (k,v) in kwargs.items()}
2350
+ return self.tensor_func(*args, **kwargs)
2351
+ def _call_numpy(self, args, kwargs):
2352
+ if not self.array_func:
2353
+ raise RuntimeError(
2354
+ f"array function for {self.__name__} was not defined")
2355
+ (args, kwargs) = self._bind(args, kwargs)
2356
+ args = map(numapi._as_array, args)
2357
+ if kwargs:
2358
+ kwargs = {k: numapi._as_array(v) for (k,v) in kwargs.items()}
2359
+ return self.array_func(*args, **kwargs)
2360
+ def __call__(self, *args, **kwargs):
2361
+ # First, call the original function
2362
+ rval = self.base_func(*args, **kwargs)
2363
+ if rval is not None:
2364
+ if isinstance(rval, tuple):
2365
+ rvallen = len(rval)
2366
+ if rvallen == 3:
2367
+ (rval, args, kwargs) = rval
2368
+ elif rvallen == 2:
2369
+ (rval, args) = rval
2370
+ elif rvallen == 1:
2371
+ rval = rval[0]
2372
+ else:
2373
+ raise ValueError(
2374
+ f"invalid value returned from numapi base_func: tuple"
2375
+ f" must have 1-3 values but got {rvallen}")
2376
+ if rval is np:
2377
+ return self._call_numpy(args, kwargs)
2378
+ elif rval is torch:
2379
+ return self._call_torch(args, kwargs)
2380
+ else:
2381
+ raise ValueError(
2382
+ f"invalid value returned from numapi base_func: {rval}")
2383
+ if torch_found:
2384
+ any_arg = any(map(is_tensor, args))
2385
+ any_inp = any_arg or any(map(is_tensor, kwargs.values()))
2386
+ if any_inp:
2387
+ return self._call_torch(args, kwargs)
2388
+ # Otherwise we use the array form.
2389
+ return self._call_numpy(args, kwargs)
2390
+
2391
+ # The tensor_args, array_args, and numeric_args decorators are very similar,
2392
+ # but tensor_args has some extra magic for handling the keep_arrays option, and
2393
+ # numeric_args also needs some magic for finding the device from the first
2394
+ # tensor argument.
2395
+ def _args_find_tensor(argvals, varargs, kwargs):
2396
+ try:
2397
+ return next(filter(is_tensor, argvals))
2398
+ except StopIteration:
2399
+ pass
2400
+ if varargs is not None:
2401
+ try:
2402
+ return next(filter(is_tensor, varargs))
2403
+ except StopIteration:
2404
+ pass
2405
+ if kwargs is not None:
2406
+ try:
2407
+ return next(filter(is_tensor, kwargs.values()))
2408
+ except StopIteration:
2409
+ pass
2410
+ return None
2411
+ def _args_try_tensor(val, name=None, binding=None, first_tensor=None):
2412
+ device = None if first_tensor is None else first_tensor.device
2413
+ try:
2414
+ tns = to_tensor(val, device=device)
2415
+ except (TypeError, RuntimeError):
2416
+ return val
2417
+ if name is not None and val is not tns:
2418
+ binding.arguments[name] = tns
2419
+ return tns
2420
+ def _args_try_array(val, name=None, binding=None, first_tensor=None):
2421
+ # (The first_tensor parameter is ignored, but included to be similar to the
2422
+ # _args_try_tensor function.)
2423
+ try:
2424
+ arr = to_array(val)
2425
+ except TypeError:
2426
+ return val
2427
+ if name is not None and val is not arr:
2428
+ binding.arguments[name] = arr
2429
+ return arr
2430
+ def _args_try_numeric(val, name=None, binding=None, first_tensor=None):
2431
+ try:
2432
+ if first_tensor is None:
2433
+ arr = to_array(val)
2434
+ else:
2435
+ arr = to_tensor(val, device=first_tensor.device)
2436
+ except TypeError:
2437
+ return val
2438
+ if name is not None and val is not arr:
2439
+ binding.arguments[name] = arr
2440
+ return arr
2441
+ def _args_dispatch(args_try_fn, fn,
2442
+ sig, sig_args, sig_vargs, sig_kwargs, keep_arrays,
2443
+ *args, **kwargs):
2444
+ binding = sig.bind(*args, **kwargs)
2445
+ binding.apply_defaults()
2446
+ vals = tuple(map(binding.arguments.__getitem__, sig_args))
2447
+ if sig_vargs:
2448
+ vargs = binding.arguments[sig_vargs]
2449
+ nvargs = len(vargs)
2450
+ else:
2451
+ vargs = None
2452
+ nvargs = 0
2453
+ if sig_kwargs:
2454
+ kwargs = binding.arguments[sig_kwargs]
2455
+ nkw = len(kw)
2456
+ else:
2457
+ kwargs = None
2458
+ nkw = 0
2459
+ if args_try_fn is _args_try_array:
2460
+ first_tensor = None
2461
+ else:
2462
+ first_tensor = _args_find_tensor(vals, vargs, kwargs)
2463
+ # Convert to the appropriate types (and update the values in the arguments
2464
+ # list if there is a change in any of the args):
2465
+ for (argname,val) in zip(sig_args, vals):
2466
+ args_try_fn(val, argname, binding, first_tensor)
2467
+ if sig_vargs:
2468
+ new_vargs = tuple(
2469
+ args_try_fn(val, first_tensor=first_tensor)
2470
+ for val in vals[nargs:nargs+nva])
2471
+ binding.arguments[sig_varargs] = new_varargs
2472
+ if sig_kwargs:
2473
+ for (k,val) in kw.items():
2474
+ cnv = args_try_fn(val, first_tensor=first_tensor)
2475
+ if cnv is not val:
2476
+ kw[k] = tns
2477
+ rval = fn(*binding.args, **binding.kwargs)
2478
+ if keep_arrays and first_tensor is None:
2479
+ if is_tuple(rval):
2480
+ rval = tuple(
2481
+ to_array(u, copy=False) if is_tensor(u) else u
2482
+ for u in rval)
2483
+ elif is_tensor(rval):
2484
+ rval = to_array(rval, copy=False)
2485
+ return rval
2486
+ def _promote_args_decorate(arglist, args_try_fn, keep_arrays, fn):
2487
+ "[Private] Dispatcher for tensor_args decorator."
2488
+ from ..workflow import calc
2489
+ sig = inspect.signature(fn)
2490
+ if arglist is None or len(arglist) == 0:
2491
+ # We convert all of the args.
2492
+ arglist = tuple(sig.parameters.keys())
2493
+ sig_args = []
2494
+ sig_kwargs = None
2495
+ sig_varargs = None
2496
+ params = sig.parameters
2497
+ for arg in arglist:
2498
+ p = params.get(arg)
2499
+ if p is None:
2500
+ raise ValueError(
2501
+ f"'{arg}' requested as tensor but not found in arguments")
2502
+ if p.kind is p.VAR_POSITIONAL:
2503
+ sig_varargs = p.name
2504
+ elif p.kind is p.VAR_KEYWORD:
2505
+ sig_kwargs = p.name
2506
+ else:
2507
+ sig_args.append(p.name)
2508
+ nargs = len(sig_args)
2509
+ dispatch = partial(
2510
+ _args_dispatch,
2511
+ args_try_fn, fn,
2512
+ sig, sig_args, sig_varargs, sig_kwargs, keep_arrays)
2513
+ return wraps(fn)(dispatch)
2514
+ @docwrap('immlib.tensor_args')
2515
+ def tensor_args(fn=None, /, *args, keep_arrays=False):
2516
+ """Converts arguments of the decorated function into PyTorch tensors.
2517
+
2518
+ The decorator ``@tensor_args``, when applied to a function, will convert
2519
+ all of that function's arguments into PyTorch tensors prior to invoking the
2520
+ function. ``tensor_args`` considers ``pint.Quantity`` objects whose
2521
+ magnitudes are tensors to be tensors and will convert arguments that are
2522
+ quantitites into new quantities with tensor magnitudes.
2523
+
2524
+ If a function is decorated with ``@tensor_args('arg1', 'arg2' ...)`` then
2525
+ only the arguments whose names are given (``arg1``, ``arg2``, ...) are
2526
+ converted into tensors.
2527
+
2528
+ When arguments are converted into PyTorch tensors, the first object in the
2529
+ argument list that is already a tensor is found and its device is used as
2530
+ the device for all converted objects. If no such object is found, then
2531
+ ``None`` is used for the device.
2532
+
2533
+ The optional argument `keep_arrays` (default: ``False``) can be set to
2534
+ ``True`` to indicate that the function should convert tensor return values
2535
+ back into NumPy arrays if none of the arguments to the function were
2536
+ originally tensors. This allows a function to be written using one
2537
+ numerical interface (PyTorch) but to work for either PyTorch tensors or
2538
+ NumPy arrays while returning values whose types match the input types.
2539
+
2540
+ """
2541
+ if fn is None:
2542
+ # A function is being decorated with `@tensor_args(keep_arrays=value)`
2543
+ # but not `@tensor_args` alone.
2544
+ return partial(
2545
+ _promote_args_decorate,
2546
+ args, _args_try_tensor, keep_arrays)
2547
+ elif is_str(fn):
2548
+ # A function is being decorated with `@tensor_args('arg1' ...)`.
2549
+ return partial(
2550
+ _promote_args_decorate,
2551
+ (fn,) + args, _args_try_tensor, keep_arrays)
2552
+ elif not callable(fn):
2553
+ # We weren't given a string or a valid function to decorate.
2554
+ raise TypeError(
2555
+ f"expected string or callable for first argument; got {type(fn)}")
2556
+ else:
2557
+ # Otherwise, we have a callable, and maybe a list of strings. If we
2558
+ # have a list of strings, we may as well use it, thus allowing the
2559
+ # tensor_args decorator to be used either as:
2560
+ # @tensor_args('a')
2561
+ # def fn(a, b): ...
2562
+ # or as
2563
+ # fn = tensor_args(lambda a,b: ..., 'a').
2564
+ return _promote_args_decorate(args, _args_try_tensor, keep_arrays, fn)
2565
+ @docwrap('immlib.array_args')
2566
+ def array_args(fn=None, /, *args):
2567
+ """Converts arguments of the decorated function into NumPy arrays.
2568
+
2569
+ The decorator ``@array_args``, when applied to a function, will convert all
2570
+ of that function's arguments into NumPy arrays prior to invoking the
2571
+ function. ``array_args`` considers ``pint.Quantity`` objects whose
2572
+ magnitudes are arrays to be arrays and will convert arguments that are
2573
+ quantitites whose magnitudes are not arrays into new quantities with array
2574
+ magnitudes.
2575
+
2576
+ If a function is decorated with ``@array_args('arg1', 'arg2' ...)`` then
2577
+ only the arguments whose names are given (``arg1``, ``arg2``, ...) are
2578
+ converted into arrays.
2579
+ """
2580
+ if fn is None:
2581
+ # A function is being decorated with `@array_args()` or `@array_args`
2582
+ # alone.
2583
+ return partial(
2584
+ _promote_args_decorate,
2585
+ args, _args_try_array, False)
2586
+ elif is_str(fn):
2587
+ # A function is being decorated with `@array_args('arg1' ...)`
2588
+ return partial(
2589
+ _promote_args_decorate,
2590
+ (fn,) + args, _args_try_array, False)
2591
+ elif not callable(fn):
2592
+ # We weren't given a string or a valid function to decorate.
2593
+ raise TypeError(
2594
+ f"expected string or callable for first argument; got {type(fn)}")
2595
+ else:
2596
+ # Otherwise, we have a callable, and maybe a list of strings. If we
2597
+ # have a list of strings, we may as well use it, thus allowing the
2598
+ # tensor_args decorator to be used either as:
2599
+ # @array_args('a')
2600
+ # def fn(a, b): ...
2601
+ # or as
2602
+ # fn = array_args(lambda a,b: ..., 'a').
2603
+ return _promote_args_decorate(args, _args_try_array, False, fn)
2604
+ @docwrap('immlib.numeric_args')
2605
+ def numeric_args(fn=None, /, *args):
2606
+ """Converts arguments of the decorated function into either NumPy arrays or
2607
+ PyTorch tensors.
2608
+
2609
+ The decorator ``@numeric_args``, when applied to a function, will convert
2610
+ all of that function's arguments into numeric collections--either PyTorch
2611
+ tensors or NumPy arrays--prior to invoking the function. Either all
2612
+ arguments are converted into either NumPy arrays or all arguments are
2613
+ converted into PyTorch tensors; the former only occurs when no PyTorch
2614
+ tensors occur in the argument list. ``numeric_args`` considers
2615
+ ``pint.Quantity`` objects whose magnitudes are numeric collections to be
2616
+ numeric collections and will convert arguments that are quantitites into
2617
+ new quantities with numeric magnitudes if necessary.
2618
+
2619
+ If a function is decorated with ``@numeric_args('arg1', 'arg2' ...)`` then
2620
+ only the arguments whose names are given (``arg1``, ``arg2``, ...) are
2621
+ converted into numeric collections.
2622
+
2623
+ When arguments are converted into PyTorch tensors, the first object in the
2624
+ argument list that is already a tensor is found and its device is used as
2625
+ the device for all converted objects. If no such object is found, then
2626
+ ``None`` is used for the device.
2627
+ """
2628
+ if fn is None:
2629
+ # A function is being decorated with `@numeric_args(keep_arrays=value)`
2630
+ # or `@numeric_args` alone.
2631
+ return partial(
2632
+ _promote_args_decorate,
2633
+ args, _args_try_numeric, False)
2634
+ elif is_str(fn):
2635
+ # A function is being decorated with `@array_args('arg1' ...)`
2636
+ return partial(
2637
+ _promote_args_decorate,
2638
+ (fn,) + args, _args_try_numeric, False)
2639
+ elif not callable(fn):
2640
+ # We weren't given a string or a valid function to decorate.
2641
+ raise TypeError(
2642
+ f"expected string or callable for first argument; got {type(fn)}")
2643
+ else:
2644
+ # Otherwise, we have a callable, and maybe a list of strings. If we
2645
+ # have a list of strings, we may as well use it, thus allowing the
2646
+ # numeric_args decorator to be used either as:
2647
+ # @numeric_args('a')
2648
+ # def fn(a, b): ...
2649
+ # or as
2650
+ # fn = numeric_args(lambda a,b: ..., 'a').
2651
+ return _promote_args_decorate(args, _args_try_numeric, False, fn)