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.
- immlib/__init__.py +131 -0
- immlib/_init.py +108 -0
- immlib/_version.py +235 -0
- immlib/doc/__init__.py +38 -0
- immlib/doc/_core.py +311 -0
- immlib/iolib/__init__.py +29 -0
- immlib/iolib/_core.py +720 -0
- immlib/pathlib/__init__.py +69 -0
- immlib/pathlib/_cache.py +152 -0
- immlib/pathlib/_core.py +869 -0
- immlib/pathlib/_osf.py +538 -0
- immlib/test/__init__.py +16 -0
- immlib/test/__main__.py +10 -0
- immlib/test/doc/__init__.py +6 -0
- immlib/test/doc/test_core.py +91 -0
- immlib/test/iolib/__init__.py +7 -0
- immlib/test/iolib/test_core.py +81 -0
- immlib/test/pathlib/__init__.py +11 -0
- immlib/test/pathlib/test_core.py +146 -0
- immlib/test/pathlib/test_osf.py +54 -0
- immlib/test/types/__init__.py +5 -0
- immlib/test/types/test_core.py +110 -0
- immlib/test/util/__init__.py +11 -0
- immlib/test/util/test_core.py +681 -0
- immlib/test/util/test_numeric.py +1374 -0
- immlib/test/util/test_quantity.py +218 -0
- immlib/test/util/test_url.py +51 -0
- immlib/test/workflow/__init__.py +9 -0
- immlib/test/workflow/test_core.py +418 -0
- immlib/test/workflow/test_plantype.py +248 -0
- immlib/types/__init__.py +29 -0
- immlib/types/_core.py +333 -0
- immlib/util/__init__.py +283 -0
- immlib/util/_core.py +2524 -0
- immlib/util/_numeric.py +2651 -0
- immlib/util/_quantity.py +523 -0
- immlib/util/_url.py +114 -0
- immlib/workflow/__init__.py +48 -0
- immlib/workflow/_core.py +1635 -0
- immlib/workflow/_plantype.py +334 -0
- immlib-1.0.0.dev2.dist-info/METADATA +76 -0
- immlib-1.0.0.dev2.dist-info/RECORD +45 -0
- immlib-1.0.0.dev2.dist-info/WHEEL +5 -0
- immlib-1.0.0.dev2.dist-info/licenses/LICENSE +21 -0
- immlib-1.0.0.dev2.dist-info/top_level.txt +1 -0
immlib/util/_numeric.py
ADDED
|
@@ -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)
|