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
|
@@ -0,0 +1,1374 @@
|
|
|
1
|
+
# -*- coding: utf-8 -*-
|
|
2
|
+
################################################################################
|
|
3
|
+
# immlib/test/util/test_numeric.py
|
|
4
|
+
#
|
|
5
|
+
# Tests of the numeric module in immlib: i.e., tests for the code in the
|
|
6
|
+
# immlib.util._numeric module.
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
# Dependencies #################################################################
|
|
10
|
+
|
|
11
|
+
from unittest import TestCase
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
# Tests ########################################################################
|
|
15
|
+
|
|
16
|
+
class TestUtilNumeric(TestCase):
|
|
17
|
+
"""Tests the immlib.util._numeric module."""
|
|
18
|
+
|
|
19
|
+
# Numeric Types ############################################################
|
|
20
|
+
def test_is_numberdata(self):
|
|
21
|
+
from immlib import is_numberdata
|
|
22
|
+
import torch, numpy as np
|
|
23
|
+
# is_numberdata returns True for numbers and False for non-numbers.
|
|
24
|
+
self.assertTrue(is_numberdata(0))
|
|
25
|
+
self.assertTrue(is_numberdata(5))
|
|
26
|
+
self.assertTrue(is_numberdata(10.0))
|
|
27
|
+
self.assertTrue(is_numberdata(-2.0 + 9.0j))
|
|
28
|
+
self.assertTrue(is_numberdata(True))
|
|
29
|
+
self.assertFalse(is_numberdata('abc'))
|
|
30
|
+
self.assertFalse(is_numberdata('10'))
|
|
31
|
+
self.assertFalse(is_numberdata(None))
|
|
32
|
+
# Arrays and tensors are checked for their dtype to be such that their
|
|
33
|
+
# elements are numbers.
|
|
34
|
+
self.assertTrue(is_numberdata(np.array(5)))
|
|
35
|
+
self.assertTrue(is_numberdata(np.array(10.0 + 2.0j)))
|
|
36
|
+
self.assertTrue(is_numberdata(torch.tensor(5)))
|
|
37
|
+
self.assertTrue(is_numberdata(torch.tensor(10.0 + 2.0j)))
|
|
38
|
+
self.assertTrue(is_numberdata(torch.tensor([1,2,3])))
|
|
39
|
+
self.assertTrue(is_numberdata(np.array([[-12.0]])))
|
|
40
|
+
self.assertFalse(is_numberdata(np.array(['abc'])))
|
|
41
|
+
def test_is_booldata(self):
|
|
42
|
+
from immlib import is_booldata
|
|
43
|
+
import torch, numpy as np
|
|
44
|
+
# is_booldata returns True for booleans and False for non-booleans.
|
|
45
|
+
self.assertFalse(is_booldata(0))
|
|
46
|
+
self.assertFalse(is_booldata(5))
|
|
47
|
+
self.assertFalse(is_booldata(10.0))
|
|
48
|
+
self.assertFalse(is_booldata(-2.0 + 9.0j))
|
|
49
|
+
self.assertTrue(is_booldata(True))
|
|
50
|
+
self.assertFalse(is_booldata('abc'))
|
|
51
|
+
self.assertFalse(is_booldata('10'))
|
|
52
|
+
self.assertFalse(is_booldata(None))
|
|
53
|
+
# Arrays and tensors are checked for their dtype to be such that their
|
|
54
|
+
# elements are numbers.
|
|
55
|
+
self.assertTrue(is_booldata(np.array(True)))
|
|
56
|
+
self.assertFalse(is_booldata(np.array(10.0 + 2.0j)))
|
|
57
|
+
self.assertFalse(is_booldata(torch.tensor(5)))
|
|
58
|
+
self.assertFalse(is_booldata(torch.tensor(10.0 + 2.0j)))
|
|
59
|
+
self.assertTrue(is_booldata(torch.tensor([True,False,False])))
|
|
60
|
+
self.assertFalse(is_booldata(np.array([[12.0]])))
|
|
61
|
+
def test_is_intdata(self):
|
|
62
|
+
from immlib import is_intdata
|
|
63
|
+
import torch, numpy as np
|
|
64
|
+
# is_intdata returns True for integers and False for non-integers.
|
|
65
|
+
self.assertTrue(is_intdata(0))
|
|
66
|
+
self.assertTrue(is_intdata(5))
|
|
67
|
+
self.assertFalse(is_intdata(10.0))
|
|
68
|
+
self.assertFalse(is_intdata(-2.0 + 9.0j))
|
|
69
|
+
self.assertTrue(is_intdata(True))
|
|
70
|
+
self.assertFalse(is_intdata('abc'))
|
|
71
|
+
self.assertFalse(is_intdata('10'))
|
|
72
|
+
self.assertFalse(is_intdata(None))
|
|
73
|
+
# Arrays and tensors are checked for their dtype to be such that their
|
|
74
|
+
# elements are numbers.
|
|
75
|
+
self.assertTrue(is_intdata(np.array(5)))
|
|
76
|
+
self.assertFalse(is_intdata(np.array(10.0 + 2.0j)))
|
|
77
|
+
self.assertTrue(is_intdata(torch.tensor(5)))
|
|
78
|
+
self.assertFalse(is_intdata(torch.tensor(10.0 + 2.0j)))
|
|
79
|
+
self.assertTrue(is_intdata(torch.tensor([1,2,3])))
|
|
80
|
+
self.assertFalse(is_intdata(np.array([[12.0]])))
|
|
81
|
+
def test_is_realdata(self):
|
|
82
|
+
from immlib import is_realdata
|
|
83
|
+
import torch, numpy as np
|
|
84
|
+
# is_real returns True for reals and False for non-reals.
|
|
85
|
+
self.assertTrue(is_realdata(0))
|
|
86
|
+
self.assertTrue(is_realdata(5))
|
|
87
|
+
self.assertTrue(is_realdata(10.0))
|
|
88
|
+
self.assertTrue(is_realdata(True))
|
|
89
|
+
self.assertFalse(is_realdata(-2.0 + 9.0j))
|
|
90
|
+
self.assertFalse(is_realdata('abc'))
|
|
91
|
+
self.assertFalse(is_realdata('10'))
|
|
92
|
+
self.assertFalse(is_realdata(None))
|
|
93
|
+
# Arrays and tensors are checked for their dtype to be such that their
|
|
94
|
+
# elements are numbers.
|
|
95
|
+
self.assertTrue(is_realdata(np.array(5)))
|
|
96
|
+
self.assertFalse(is_realdata(np.array(10.0 + 2.0j)))
|
|
97
|
+
self.assertTrue(is_realdata(torch.tensor(5)))
|
|
98
|
+
self.assertFalse(is_realdata(torch.tensor(10.0 + 2.0j)))
|
|
99
|
+
self.assertTrue(is_realdata(torch.tensor([1,2,3])))
|
|
100
|
+
self.assertTrue(is_realdata(np.array([[12.0]])))
|
|
101
|
+
def test_is_complexdata(self):
|
|
102
|
+
from immlib import is_complexdata
|
|
103
|
+
import torch, numpy as np
|
|
104
|
+
# is_complexdata returns True for complexs and False for non-complexs.
|
|
105
|
+
self.assertTrue(is_complexdata(0))
|
|
106
|
+
self.assertTrue(is_complexdata(5))
|
|
107
|
+
self.assertTrue(is_complexdata(10.0))
|
|
108
|
+
self.assertTrue(is_complexdata(True))
|
|
109
|
+
self.assertTrue(is_complexdata(-2.0 + 9.0j))
|
|
110
|
+
self.assertFalse(is_complexdata('abc'))
|
|
111
|
+
self.assertFalse(is_complexdata('10'))
|
|
112
|
+
self.assertFalse(is_complexdata(None))
|
|
113
|
+
# Arrays and tensors are checked for their dtype to be such that their
|
|
114
|
+
# elements are numbers.
|
|
115
|
+
self.assertTrue(is_complexdata(np.array(5)))
|
|
116
|
+
self.assertTrue(is_complexdata(np.array(10.0 + 2.0j)))
|
|
117
|
+
self.assertTrue(is_complexdata(torch.tensor(5)))
|
|
118
|
+
self.assertTrue(is_complexdata(torch.tensor(10.0 + 2.0j)))
|
|
119
|
+
self.assertTrue(is_complexdata(torch.tensor([1,2,3])))
|
|
120
|
+
self.assertTrue(is_complexdata(np.array([[12.0]])))
|
|
121
|
+
def test_is_number(self):
|
|
122
|
+
from immlib import is_number
|
|
123
|
+
import torch, numpy as np
|
|
124
|
+
# is_number returns True for numbers and False for non-numbers.
|
|
125
|
+
self.assertTrue(is_number(0))
|
|
126
|
+
self.assertTrue(is_number(5))
|
|
127
|
+
self.assertTrue(is_number(10.0))
|
|
128
|
+
self.assertTrue(is_number(-2.0 + 9.0j))
|
|
129
|
+
self.assertTrue(is_number(True))
|
|
130
|
+
self.assertFalse(is_number('abc'))
|
|
131
|
+
self.assertFalse(is_number('10'))
|
|
132
|
+
self.assertFalse(is_number(None))
|
|
133
|
+
# Scalar arrays and tensors are counted as scalars.
|
|
134
|
+
self.assertTrue(is_number(np.array(5)))
|
|
135
|
+
self.assertTrue(is_number(np.array(10.0 + 2.0j)))
|
|
136
|
+
self.assertTrue(is_number(torch.tensor(5)))
|
|
137
|
+
self.assertTrue(is_number(torch.tensor(10.0 + 2.0j)))
|
|
138
|
+
# Arrays and tensors are *not* checked for their dtype to be such--that
|
|
139
|
+
# is instead performed by is_numberdata.
|
|
140
|
+
self.assertFalse(is_number(torch.tensor([1,2,3])))
|
|
141
|
+
self.assertFalse(is_number(np.array([[-12.0]])))
|
|
142
|
+
# Specific dtypes can also be tested for.
|
|
143
|
+
self.assertTrue(is_number(0, dtype=int))
|
|
144
|
+
self.assertTrue(is_number(5, dtype=int))
|
|
145
|
+
self.assertTrue(is_number(10.0, dtype=float))
|
|
146
|
+
self.assertTrue(is_number(-2.0 + 9.0j, dtype=complex))
|
|
147
|
+
self.assertTrue(is_number(True, dtype=bool))
|
|
148
|
+
self.assertFalse(is_number(0, dtype=bool))
|
|
149
|
+
self.assertTrue(is_number(5, dtype=float))
|
|
150
|
+
self.assertTrue(is_number(10.0, dtype=complex))
|
|
151
|
+
self.assertFalse(is_number(-2.0 + 9.0j, dtype=int))
|
|
152
|
+
self.assertTrue(is_number(True, dtype=float))
|
|
153
|
+
# is_number returns True if the argument is a scalar number; otherwise
|
|
154
|
+
# it returns false.
|
|
155
|
+
self.assertTrue(is_number(10))
|
|
156
|
+
self.assertTrue(is_number(10.0))
|
|
157
|
+
self.assertTrue(is_number(10.0 + 20.5j))
|
|
158
|
+
self.assertTrue(is_number(True))
|
|
159
|
+
self.assertTrue(is_number(np.array(10)))
|
|
160
|
+
self.assertTrue(is_number(torch.tensor(10)))
|
|
161
|
+
self.assertFalse(is_number([10]))
|
|
162
|
+
self.assertFalse(is_number([[10]]))
|
|
163
|
+
self.assertFalse(is_number([[[10]]]))
|
|
164
|
+
self.assertFalse(is_number(np.array([10])))
|
|
165
|
+
self.assertFalse(is_number(np.array([[10]])))
|
|
166
|
+
self.assertFalse(is_number(np.array([[[10]]])))
|
|
167
|
+
self.assertFalse(is_number(torch.tensor([10])))
|
|
168
|
+
self.assertFalse(is_number(torch.tensor([[10]])))
|
|
169
|
+
self.assertFalse(is_number(torch.tensor([[[10]]])))
|
|
170
|
+
self.assertFalse(is_number('10'))
|
|
171
|
+
self.assertFalse(is_number({'a':10}))
|
|
172
|
+
self.assertFalse(is_number([1,2,3]))
|
|
173
|
+
# Make sure it throws errors when appropriate:
|
|
174
|
+
with self.assertRaises(ValueError):
|
|
175
|
+
is_number(0, dtype=str)
|
|
176
|
+
def test_is_bool(self):
|
|
177
|
+
from immlib import is_bool
|
|
178
|
+
import torch, numpy as np
|
|
179
|
+
# is_bool returns True for bools and False for non-bools.
|
|
180
|
+
self.assertFalse(is_bool(0))
|
|
181
|
+
self.assertFalse(is_bool(5))
|
|
182
|
+
self.assertFalse(is_bool(10.0))
|
|
183
|
+
self.assertFalse(is_bool(-2.0 + 9.0j))
|
|
184
|
+
self.assertTrue(is_bool(True))
|
|
185
|
+
self.assertFalse(is_bool('abc'))
|
|
186
|
+
self.assertFalse(is_bool('10'))
|
|
187
|
+
self.assertFalse(is_bool(None))
|
|
188
|
+
# Scalar arrays and tensors are allowed.
|
|
189
|
+
self.assertTrue(is_bool(np.array(False)))
|
|
190
|
+
self.assertTrue(is_bool(torch.tensor(True)))
|
|
191
|
+
self.assertFalse(is_bool(np.array(10.0 + 2.0j)))
|
|
192
|
+
self.assertFalse(is_bool(torch.tensor(10.0 + 2.0j)))
|
|
193
|
+
# is_bool does not return True for collections (use is_intdata
|
|
194
|
+
# instead).
|
|
195
|
+
self.assertFalse(is_bool(torch.tensor([1,2,3])))
|
|
196
|
+
self.assertFalse(is_bool(np.array([[12.0]])))
|
|
197
|
+
def test_is_integer(self):
|
|
198
|
+
from immlib import is_integer
|
|
199
|
+
import torch, numpy as np
|
|
200
|
+
# is_integer returns True for integers and False for non-integers.
|
|
201
|
+
self.assertTrue(is_integer(0))
|
|
202
|
+
self.assertTrue(is_integer(5))
|
|
203
|
+
self.assertFalse(is_integer(10.0))
|
|
204
|
+
self.assertFalse(is_integer(-2.0 + 9.0j))
|
|
205
|
+
self.assertTrue(is_integer(True))
|
|
206
|
+
self.assertFalse(is_integer('abc'))
|
|
207
|
+
self.assertFalse(is_integer('10'))
|
|
208
|
+
self.assertFalse(is_integer(None))
|
|
209
|
+
# Scalar arrays and tensors are allowed.
|
|
210
|
+
self.assertTrue(is_integer(np.array(5)))
|
|
211
|
+
self.assertTrue(is_integer(torch.tensor(5)))
|
|
212
|
+
self.assertFalse(is_integer(np.array(10.0 + 2.0j)))
|
|
213
|
+
self.assertFalse(is_integer(torch.tensor(10.0 + 2.0j)))
|
|
214
|
+
# is_integer does not return True for collections (use is_intdata
|
|
215
|
+
# instead).
|
|
216
|
+
self.assertFalse(is_integer(torch.tensor([1,2,3])))
|
|
217
|
+
self.assertFalse(is_integer(np.array([[12.0]])))
|
|
218
|
+
def test_is_real(self):
|
|
219
|
+
from immlib import is_real
|
|
220
|
+
import torch, numpy as np
|
|
221
|
+
# is_real returns True for reals and False for non-reals.
|
|
222
|
+
self.assertTrue(is_real(0))
|
|
223
|
+
self.assertTrue(is_real(5))
|
|
224
|
+
self.assertTrue(is_real(10.0))
|
|
225
|
+
self.assertTrue(is_real(True))
|
|
226
|
+
self.assertFalse(is_real(-2.0 + 9.0j))
|
|
227
|
+
self.assertFalse(is_real('abc'))
|
|
228
|
+
self.assertFalse(is_real('10'))
|
|
229
|
+
self.assertFalse(is_real(None))
|
|
230
|
+
# Scalar arrays and tensors are allowed.
|
|
231
|
+
self.assertTrue(is_real(np.array(5)))
|
|
232
|
+
self.assertTrue(is_real(torch.tensor(5)))
|
|
233
|
+
self.assertFalse(is_real(np.array(10.0 + 2.0j)))
|
|
234
|
+
self.assertFalse(is_real(torch.tensor(10.0 + 2.0j)))
|
|
235
|
+
# Arrays and tensors are checked for their dtype to be such that their
|
|
236
|
+
# elements are numbers.
|
|
237
|
+
self.assertFalse(is_real(torch.tensor([1,2,3])))
|
|
238
|
+
self.assertFalse(is_real(np.array([[12.0]])))
|
|
239
|
+
def test_is_complex(self):
|
|
240
|
+
from immlib import is_complex
|
|
241
|
+
import torch, numpy as np
|
|
242
|
+
# is_complex returns True for complexs and False for non-complexs.
|
|
243
|
+
self.assertTrue(is_complex(0))
|
|
244
|
+
self.assertTrue(is_complex(5))
|
|
245
|
+
self.assertTrue(is_complex(10.0))
|
|
246
|
+
self.assertTrue(is_complex(True))
|
|
247
|
+
self.assertTrue(is_complex(-2.0 + 9.0j))
|
|
248
|
+
self.assertFalse(is_complex('abc'))
|
|
249
|
+
self.assertFalse(is_complex('10'))
|
|
250
|
+
self.assertFalse(is_complex(None))
|
|
251
|
+
# Scalar arrays/tensors are allowed.
|
|
252
|
+
self.assertTrue(is_complex(np.array(5)))
|
|
253
|
+
self.assertTrue(is_complex(np.array(10.0 + 2.0j)))
|
|
254
|
+
self.assertTrue(is_complex(torch.tensor(5)))
|
|
255
|
+
self.assertTrue(is_complex(torch.tensor(10.0 + 2.0j)))
|
|
256
|
+
# Arrays and tensors are only checked if they are scalars.
|
|
257
|
+
self.assertFalse(is_complex(torch.tensor([1,2,3])))
|
|
258
|
+
self.assertFalse(is_complex(np.array([[12.0]])))
|
|
259
|
+
|
|
260
|
+
# NumPy Utilities ##########################################################
|
|
261
|
+
def test_is_numpydtype(self):
|
|
262
|
+
from immlib.util import is_numpydtype
|
|
263
|
+
import torch, numpy as np
|
|
264
|
+
# is_numpydtype returns true for dtypes and dtypes alone.
|
|
265
|
+
self.assertTrue(is_numpydtype(np.dtype('int')))
|
|
266
|
+
self.assertTrue(is_numpydtype(np.dtype(float)))
|
|
267
|
+
self.assertTrue(is_numpydtype(np.dtype(np.bool_)))
|
|
268
|
+
self.assertFalse(is_numpydtype('int'))
|
|
269
|
+
self.assertFalse(is_numpydtype(float))
|
|
270
|
+
self.assertFalse(is_numpydtype(np.bool_))
|
|
271
|
+
self.assertFalse(is_numpydtype(torch.float))
|
|
272
|
+
def test_like_numpydtype(self):
|
|
273
|
+
from immlib.util import like_numpydtype
|
|
274
|
+
import torch, numpy as np
|
|
275
|
+
# Anything that can be converted into a numpy dtype object is considered
|
|
276
|
+
# to be like a numpy dtype.
|
|
277
|
+
self.assertTrue(like_numpydtype('int'))
|
|
278
|
+
self.assertTrue(like_numpydtype(float))
|
|
279
|
+
self.assertTrue(like_numpydtype(np.bool_))
|
|
280
|
+
self.assertFalse(like_numpydtype('abc'))
|
|
281
|
+
self.assertFalse(like_numpydtype(10))
|
|
282
|
+
self.assertFalse(like_numpydtype(...))
|
|
283
|
+
# Note that None can be converted to a numpy dtype (float64).
|
|
284
|
+
self.assertTrue(like_numpydtype(None))
|
|
285
|
+
# numpy dtypes themselves are like numpy dtypes, as are torch dtypes.
|
|
286
|
+
self.assertTrue(like_numpydtype(np.dtype(int)))
|
|
287
|
+
self.assertTrue(like_numpydtype(torch.float))
|
|
288
|
+
def test_to_numpydtype(self):
|
|
289
|
+
from immlib.util import to_numpydtype
|
|
290
|
+
import torch, numpy as np
|
|
291
|
+
# Converting a numpy dtype into a dtype results in the identical dtype.
|
|
292
|
+
dt = np.dtype(int)
|
|
293
|
+
self.assertIs(dt, to_numpydtype(dt))
|
|
294
|
+
# Torch dtypes can be converted into a numpy dtype.
|
|
295
|
+
self.assertEqual(np.dtype(np.float64), to_numpydtype(torch.float64))
|
|
296
|
+
self.assertEqual(np.dtype('int32'), to_numpydtype(torch.int32))
|
|
297
|
+
# Ordinary tags can be converted into dtypes as well.
|
|
298
|
+
self.assertEqual(np.dtype(np.float64), to_numpydtype(np.float64))
|
|
299
|
+
self.assertEqual(np.dtype('int32'), to_numpydtype('int32'))
|
|
300
|
+
def test_sparray_utils(self):
|
|
301
|
+
import scipy.sparse as sps, numpy as np, torch, pint
|
|
302
|
+
from immlib.util import (
|
|
303
|
+
sparse_find, sparse_data, sparse_indices, sparse_layout,
|
|
304
|
+
sparse_haslayout, sparse_tolayout, quant)
|
|
305
|
+
from immlib import (units, quant)
|
|
306
|
+
sparr = sps.csr_array(
|
|
307
|
+
([1.0, 0.5, 0.5, 0.2, 0.1],
|
|
308
|
+
([0, 0, 4, 5, 9], [4, 9, 4, 1, 8])),
|
|
309
|
+
shape=(10, 10),
|
|
310
|
+
dtype=float)
|
|
311
|
+
sptns = torch.sparse_coo_tensor(
|
|
312
|
+
torch.tensor([[0, 0, 4, 5, 9], [4, 9, 4, 1, 8]]),
|
|
313
|
+
torch.tensor([1.0, 0.5, 0.5, 0.2, 0.1]),
|
|
314
|
+
size=(10, 10),
|
|
315
|
+
dtype=float)
|
|
316
|
+
sptns = sptns.coalesce()
|
|
317
|
+
# We can get a sparse layout from names or objects.
|
|
318
|
+
for k in ('coo', 'csr', 'csc', 'bsr', 'bsc', 'dok', 'lil', 'dia'):
|
|
319
|
+
sl = sparse_layout(k)
|
|
320
|
+
self.assertEqual(k, sl.name)
|
|
321
|
+
self.assertIs(sl, sparse_layout(sl))
|
|
322
|
+
# We can also look up things by numpy array or torch tensor type.
|
|
323
|
+
self.assertEqual('csr', sparse_layout(sparr).name)
|
|
324
|
+
self.assertEqual('coo', sparse_layout(sptns).name)
|
|
325
|
+
# sparse_layout works with quantities.
|
|
326
|
+
self.assertEqual('csr', sparse_layout(quant(sparr, units.mm)).name)
|
|
327
|
+
self.assertEqual('coo', sparse_layout(quant(sptns, units.mm)).name)
|
|
328
|
+
# Unrecognized objects and names produce None:
|
|
329
|
+
self.assertIs(None, sparse_layout('???'))
|
|
330
|
+
self.assertIs(None, sparse_layout(object()))
|
|
331
|
+
# We can convert between layouts:
|
|
332
|
+
cooarr = sparse_tolayout(sparr, 'coo')
|
|
333
|
+
self.assertEqual(cooarr.format, 'coo')
|
|
334
|
+
self.assertTrue(np.array_equal(sparr.todense(), cooarr.todense()))
|
|
335
|
+
csrtns = sparse_tolayout(sptns, 'csr')
|
|
336
|
+
self.assertEqual(csrtns.layout, torch.sparse_csr)
|
|
337
|
+
self.assertTrue(torch.equal(sptns.to_dense(), csrtns.to_dense()))
|
|
338
|
+
q_sparr = quant(sparr, units.mm)
|
|
339
|
+
self.assertIsInstance(q_sparr, pint.Quantity)
|
|
340
|
+
q_cooarr = sparse_tolayout(q_sparr, 'coo')
|
|
341
|
+
self.assertIsInstance(q_cooarr, pint.Quantity)
|
|
342
|
+
self.assertTrue(np.array_equal(sparr.todense(), q_cooarr.m.todense()))
|
|
343
|
+
f_sparr = sparr.copy()
|
|
344
|
+
f_sparr.data.setflags(write=False)
|
|
345
|
+
f_cooarr = sparse_tolayout(f_sparr, 'coo')
|
|
346
|
+
self.assertTrue(np.array_equal(sparr.todense(), f_cooarr.todense()))
|
|
347
|
+
self.assertFalse(f_cooarr.data.flags['WRITEABLE'])
|
|
348
|
+
with self.assertRaises(TypeError):
|
|
349
|
+
sparse_tolayout(object(), 'coo')
|
|
350
|
+
# Check if a sparse object has a particular layout:
|
|
351
|
+
self.assertTrue(sparse_haslayout(sparr, 'csr'))
|
|
352
|
+
self.assertTrue(sparse_haslayout(cooarr, 'coo'))
|
|
353
|
+
self.assertFalse(sparse_haslayout(sparr, 'coo'))
|
|
354
|
+
self.assertFalse(sparse_haslayout(cooarr, 'csr'))
|
|
355
|
+
self.assertTrue(sparse_haslayout(csrtns, 'csr'))
|
|
356
|
+
self.assertTrue(sparse_haslayout(sptns, 'coo'))
|
|
357
|
+
self.assertFalse(sparse_haslayout(csrtns, 'coo'))
|
|
358
|
+
self.assertFalse(sparse_haslayout(sptns, 'csr'))
|
|
359
|
+
self.assertTrue(isinstance(q_sparr, pint.Quantity))
|
|
360
|
+
self.assertTrue(sps.issparse(q_sparr.m))
|
|
361
|
+
self.assertTrue(sparse_haslayout(q_sparr, 'csr'))
|
|
362
|
+
self.assertTrue(sparse_haslayout(q_cooarr, 'coo'))
|
|
363
|
+
self.assertFalse(sparse_haslayout(object(), 'csr'))
|
|
364
|
+
with self.assertRaises(ValueError):
|
|
365
|
+
sparse_haslayout(sptns, '???')
|
|
366
|
+
with self.assertRaises(ValueError):
|
|
367
|
+
sparse_haslayout(sptns, object())
|
|
368
|
+
# We can extract the various bits of data also:
|
|
369
|
+
self.assertTrue(
|
|
370
|
+
all(map(np.array_equal, sparse_find(sparr), sps.find(sparr))))
|
|
371
|
+
self.assertTrue(
|
|
372
|
+
np.array_equal(sparse_data(sparr), sparr.data))
|
|
373
|
+
ii = np.stack(sparse_find(sparr)[:-1])
|
|
374
|
+
self.assertTrue(
|
|
375
|
+
np.array_equal(sparse_indices(sparr), ii))
|
|
376
|
+
sfnd = sparse_find(sptns)
|
|
377
|
+
ii = sptns.indices()
|
|
378
|
+
vv = sptns.values()
|
|
379
|
+
tfnd = tuple(ii) + (vv,)
|
|
380
|
+
self.assertTrue(all(map(torch.equal, sfnd, tfnd)))
|
|
381
|
+
self.assertTrue(
|
|
382
|
+
torch.equal(vv, sparse_data(sptns)))
|
|
383
|
+
self.assertTrue(
|
|
384
|
+
torch.equal(ii, sparse_indices(sptns)))
|
|
385
|
+
# These also work for quantities:
|
|
386
|
+
q_sparr = quant(sparr, units.mm)
|
|
387
|
+
self.assertIsInstance(q_sparr, pint.Quantity)
|
|
388
|
+
self.assertTrue(
|
|
389
|
+
all(np.array_equal(u.m if isinstance(u, pint.Quantity) else u, v)
|
|
390
|
+
for (u,v) in zip(sparse_find(q_sparr), sps.find(sparr))))
|
|
391
|
+
q_spdat = sparse_data(q_sparr)
|
|
392
|
+
self.assertIsInstance(q_spdat, pint.Quantity)
|
|
393
|
+
self.assertTrue(
|
|
394
|
+
np.array_equal(q_spdat.m, q_sparr.data))
|
|
395
|
+
ii = np.stack(sparse_find(sparr)[:-1])
|
|
396
|
+
self.assertTrue(
|
|
397
|
+
np.array_equal(sparse_indices(q_sparr), ii))
|
|
398
|
+
# They throw reasonable errors too:
|
|
399
|
+
with self.assertRaises(TypeError):
|
|
400
|
+
sparse_find('')
|
|
401
|
+
with self.assertRaises(TypeError):
|
|
402
|
+
sparse_indices('')
|
|
403
|
+
with self.assertRaises(TypeError):
|
|
404
|
+
sparse_data('')
|
|
405
|
+
def test_is_array(self):
|
|
406
|
+
from immlib import (is_array, quant)
|
|
407
|
+
from numpy import (array, linspace, dot)
|
|
408
|
+
from scipy.sparse import csr_matrix
|
|
409
|
+
import torch, numpy as np
|
|
410
|
+
# By default, is_array() returns True for numpy arrays, scipy sparse
|
|
411
|
+
# matrices, and quantities of these.
|
|
412
|
+
arr = linspace(0, 1, 25)
|
|
413
|
+
mtx = dot(linspace(0, 1, 10)[:,None], linspace(1, 2, 10)[None,:])
|
|
414
|
+
sp_mtx = csr_matrix(
|
|
415
|
+
([1.0, 0.5, 0.5, 0.2, 0.1],
|
|
416
|
+
([0, 0, 4, 5, 9], [4, 9, 4, 1, 8])),
|
|
417
|
+
shape=(10, 10),
|
|
418
|
+
dtype=float)
|
|
419
|
+
q_arr = quant(arr, 'mm')
|
|
420
|
+
q_mtx = quant(arr, 'seconds')
|
|
421
|
+
q_sp_mtx = quant(sp_mtx, 'kg')
|
|
422
|
+
self.assertTrue(is_array(arr))
|
|
423
|
+
self.assertTrue(is_array(mtx))
|
|
424
|
+
self.assertTrue(is_array(sp_mtx))
|
|
425
|
+
self.assertTrue(is_array(q_arr))
|
|
426
|
+
self.assertTrue(is_array(q_mtx))
|
|
427
|
+
self.assertTrue(is_array(q_sp_mtx))
|
|
428
|
+
# Things like lists, numbers, and torch tensors are not arrays.
|
|
429
|
+
self.assertFalse(is_array('abc'))
|
|
430
|
+
self.assertFalse(is_array(10))
|
|
431
|
+
self.assertFalse(is_array([12.0, 0.5, 3.2]))
|
|
432
|
+
self.assertFalse(is_array(torch.tensor([1.0, 2.0, 3.0])))
|
|
433
|
+
self.assertFalse(is_array(quant(torch.tensor([1.0, 2.0, 3.0]), 'mm')))
|
|
434
|
+
# We can use the dtype argument to restrict what we consider an array by
|
|
435
|
+
# its dtype. The dtype of the is_array argument must be a sub-dtype of
|
|
436
|
+
# the dtype parameter.
|
|
437
|
+
self.assertTrue(is_array(arr, dtype=np.number))
|
|
438
|
+
self.assertTrue(is_array(arr, dtype=arr.dtype))
|
|
439
|
+
self.assertFalse(is_array(arr, dtype=np.str_))
|
|
440
|
+
# If a tuple is passed for the dtype, the dtype must match one of the
|
|
441
|
+
# tuple's members exactly.
|
|
442
|
+
self.assertTrue(is_array(mtx, dtype=(mtx.dtype,)))
|
|
443
|
+
self.assertTrue(is_array(mtx, dtype=(mtx.dtype,np.dtype(int),np.str_)))
|
|
444
|
+
self.assertFalse(is_array(mtx, dtype=(np.dtype(int),np.str_)))
|
|
445
|
+
self.assertFalse(is_array(np.array([], dtype=np.int32),
|
|
446
|
+
dtype=(np.int64,)))
|
|
447
|
+
# torch dtypes can be interpreted into numpy dtypes.
|
|
448
|
+
self.assertTrue(is_array(mtx, dtype=torch.as_tensor(mtx).dtype))
|
|
449
|
+
# We can use the ndim argument to restrict the number of dimensions that
|
|
450
|
+
# an array can have in order to be considered a matching array.
|
|
451
|
+
# Typically, this is just the number of dimensions.
|
|
452
|
+
self.assertTrue(is_array(arr, ndim=1))
|
|
453
|
+
self.assertTrue(is_array(mtx, ndim=2))
|
|
454
|
+
self.assertFalse(is_array(arr, ndim=2))
|
|
455
|
+
self.assertFalse(is_array(mtx, ndim=1))
|
|
456
|
+
# Alternately, a tuple may be given, in which case any of the dimension
|
|
457
|
+
# counts in the tuple are accepted.
|
|
458
|
+
self.assertTrue(is_array(mtx, ndim=(1,2)))
|
|
459
|
+
self.assertTrue(is_array(arr, ndim=(1,2)))
|
|
460
|
+
self.assertFalse(is_array(mtx, ndim=(1,3)))
|
|
461
|
+
self.assertFalse(is_array(arr, ndim=(0,2)))
|
|
462
|
+
# Scalar arrays have 0 dimensions.
|
|
463
|
+
self.assertTrue(is_array(array(0), ndim=0))
|
|
464
|
+
# The shape option is a more specific version of the ndim parameter. It
|
|
465
|
+
# lets you specify what kind of shape is required of the array. The most
|
|
466
|
+
# straightforward usage is to require a specific shape.
|
|
467
|
+
self.assertTrue(is_array(arr, shape=(25,)))
|
|
468
|
+
self.assertTrue(is_array(arr, shape=25))
|
|
469
|
+
self.assertTrue(is_array(mtx, shape=(10,10)))
|
|
470
|
+
self.assertFalse(is_array(arr, shape=(25,25)))
|
|
471
|
+
self.assertFalse(is_array(mtx, shape=(10,)))
|
|
472
|
+
self.assertTrue(is_array(np.array(100), shape=()))
|
|
473
|
+
# A -1 value that appears in the shape option represents any size along
|
|
474
|
+
# that dimension (a wildcard). Any number of -1s can be included.
|
|
475
|
+
self.assertTrue(is_array(arr, shape=(-1,)))
|
|
476
|
+
self.assertTrue(is_array(mtx, shape=(-1,10)))
|
|
477
|
+
self.assertTrue(is_array(mtx, shape=(10,-1)))
|
|
478
|
+
self.assertTrue(is_array(mtx, shape=(-1,-1)))
|
|
479
|
+
self.assertFalse(is_array(mtx, shape=(1,-1)))
|
|
480
|
+
# No more than 1 ellipsis may be included in the shape to indicate that
|
|
481
|
+
# any number of dimensions, with any sizes, can appear in place of the
|
|
482
|
+
# ellipsis.
|
|
483
|
+
self.assertTrue(is_array(arr, shape=(...,25)))
|
|
484
|
+
self.assertTrue(is_array(arr, shape=(25,...)))
|
|
485
|
+
self.assertFalse(is_array(arr, shape=(25,...,25)))
|
|
486
|
+
self.assertTrue(is_array(mtx, shape=(...,10)))
|
|
487
|
+
self.assertTrue(is_array(mtx, shape=(10,...)))
|
|
488
|
+
self.assertTrue(is_array(mtx, shape=(10,...,10)))
|
|
489
|
+
self.assertTrue(is_array(mtx, shape=(10,10,...)))
|
|
490
|
+
self.assertTrue(is_array(mtx, shape=(...,10,10)))
|
|
491
|
+
self.assertFalse(is_array(mtx, shape=(10,...,10,10)))
|
|
492
|
+
self.assertFalse(is_array(mtx, shape=(10,10,...,10)))
|
|
493
|
+
self.assertTrue(is_array(np.zeros((1,2,3,4,5)), shape=(1,...,4,5)))
|
|
494
|
+
# The numel option allows one to specify the number of elements that an
|
|
495
|
+
# object must have. This does not care about dimensionality.
|
|
496
|
+
self.assertTrue(is_array(arr, numel=25))
|
|
497
|
+
self.assertTrue(is_array(arr, numel=(25,26))) # Is numel 25 or 26?
|
|
498
|
+
self.assertFalse(is_array(arr, numel=26))
|
|
499
|
+
self.assertFalse(is_array(arr, numel=(24,26)))
|
|
500
|
+
self.assertTrue(is_array(np.array(0), numel=1))
|
|
501
|
+
self.assertTrue(is_array(np.array([0]), numel=1))
|
|
502
|
+
self.assertTrue(is_array(np.array([[0]]), numel=1))
|
|
503
|
+
# The frozen option can be used to test whether an array is frozen
|
|
504
|
+
# or not. This is judged by the array's 'WRITEABLE' flag.
|
|
505
|
+
self.assertFalse(is_array(arr, frozen=True))
|
|
506
|
+
self.assertTrue(is_array(arr, frozen=False))
|
|
507
|
+
self.assertFalse(is_array(mtx, frozen=True))
|
|
508
|
+
self.assertTrue(is_array(mtx, frozen=False))
|
|
509
|
+
with self.assertRaises(ValueError):
|
|
510
|
+
is_array(mtx, frozen='fail')
|
|
511
|
+
farr = arr.copy()
|
|
512
|
+
farr.setflags(write=False)
|
|
513
|
+
self.assertFalse(is_array(farr, frozen=False))
|
|
514
|
+
self.assertTrue(is_array(farr, frozen=True))
|
|
515
|
+
# If we change the flags of these arrays, they become frozen.
|
|
516
|
+
arr.setflags(write=False)
|
|
517
|
+
mtx.setflags(write=False)
|
|
518
|
+
self.assertTrue(is_array(arr, frozen=True))
|
|
519
|
+
self.assertFalse(is_array(arr, frozen=False))
|
|
520
|
+
self.assertTrue(is_array(mtx, frozen=True))
|
|
521
|
+
self.assertFalse(is_array(mtx, frozen=False))
|
|
522
|
+
# The sparse option can test whether an object is a sparse matrix or
|
|
523
|
+
# not. By default sparse is None, meaning that it doesn't matter whether
|
|
524
|
+
# an object is sparse, but sometimes you want to check for strict
|
|
525
|
+
# numpy arrays only.
|
|
526
|
+
self.assertTrue(is_array(arr, sparse=False))
|
|
527
|
+
self.assertTrue(is_array(mtx, sparse=False))
|
|
528
|
+
self.assertFalse(is_array(sp_mtx, sparse=False))
|
|
529
|
+
self.assertFalse(is_array(arr, sparse=True))
|
|
530
|
+
self.assertFalse(is_array(mtx, sparse=True))
|
|
531
|
+
self.assertTrue(is_array(sp_mtx, sparse=True))
|
|
532
|
+
# You can also require a kind of sparse matrix.
|
|
533
|
+
self.assertTrue(is_array(sp_mtx, sparse='csr'))
|
|
534
|
+
self.assertFalse(is_array(sp_mtx, sparse='csc'))
|
|
535
|
+
with self.assertRaises(ValueError):
|
|
536
|
+
is_array(sp_mtx, sparse='???')
|
|
537
|
+
with self.assertRaises(ValueError):
|
|
538
|
+
is_array(sp_mtx, sparse=object())
|
|
539
|
+
# Sparse and frozen can be tested together.
|
|
540
|
+
self.assertTrue(is_array(arr, sparse=False, frozen=True))
|
|
541
|
+
self.assertFalse(is_array(arr, sparse=False, frozen=False))
|
|
542
|
+
self.assertFalse(is_array(arr, sparse=True, frozen=False))
|
|
543
|
+
self.assertFalse(is_array(arr, sparse=True, frozen=True))
|
|
544
|
+
self.assertFalse(is_array(sp_mtx, sparse=True, frozen=True))
|
|
545
|
+
self.assertTrue(is_array(sp_mtx, sparse=True, frozen=False))
|
|
546
|
+
self.assertFalse(is_array(sp_mtx, sparse=False, frozen=True))
|
|
547
|
+
self.assertFalse(is_array(sp_mtx, sparse=False, frozen=False))
|
|
548
|
+
sp_mtx.data.setflags(write=False)
|
|
549
|
+
self.assertTrue(is_array(sp_mtx, sparse=True, frozen=True))
|
|
550
|
+
self.assertFalse(is_array(sp_mtx, sparse=True, frozen=False))
|
|
551
|
+
self.assertFalse(is_array(sp_mtx, sparse=False, frozen=True))
|
|
552
|
+
self.assertFalse(is_array(sp_mtx, sparse=False, frozen=False))
|
|
553
|
+
# The quant option can be used to control whether the object must or
|
|
554
|
+
# must not be a quantity.
|
|
555
|
+
self.assertTrue(is_array(arr, quant=False))
|
|
556
|
+
self.assertTrue(is_array(mtx, quant=False))
|
|
557
|
+
self.assertFalse(is_array(arr, quant=True))
|
|
558
|
+
self.assertFalse(is_array(mtx, quant=True))
|
|
559
|
+
self.assertTrue(is_array(q_arr, quant=True))
|
|
560
|
+
self.assertTrue(is_array(q_mtx, quant=True))
|
|
561
|
+
self.assertFalse(is_array(q_arr, quant=False))
|
|
562
|
+
self.assertFalse(is_array(q_mtx, quant=False))
|
|
563
|
+
# The units option can be used to require that either an object have
|
|
564
|
+
# no units (or is not a quantity) or that it have specific units.
|
|
565
|
+
self.assertTrue(is_array(arr, unit=None))
|
|
566
|
+
self.assertTrue(is_array(mtx, unit=None))
|
|
567
|
+
self.assertFalse(is_array(arr, unit='mm'))
|
|
568
|
+
self.assertFalse(is_array(mtx, unit='s'))
|
|
569
|
+
self.assertFalse(is_array(q_arr, unit=None))
|
|
570
|
+
self.assertFalse(is_array(q_mtx, unit=None))
|
|
571
|
+
self.assertTrue(is_array(q_arr, unit='mm'))
|
|
572
|
+
self.assertTrue(is_array(q_mtx, unit='s'))
|
|
573
|
+
self.assertFalse(is_array(q_arr, unit='s'))
|
|
574
|
+
self.assertFalse(is_array(q_mtx, unit='mm'))
|
|
575
|
+
# We can also specify the units registry (Ellipsis means immlib.units).
|
|
576
|
+
self.assertFalse(is_array(q_arr, unit='s', ureg=Ellipsis))
|
|
577
|
+
def test_to_array(self):
|
|
578
|
+
from immlib import (to_array, quant, is_quant, units, frozenarray)
|
|
579
|
+
from numpy import (array, linspace, dot)
|
|
580
|
+
from scipy.sparse import (csr_matrix, issparse)
|
|
581
|
+
import torch, numpy as np, pint
|
|
582
|
+
# We'll use a few objects throughout our tests, which we setup now.
|
|
583
|
+
arr = linspace(0, 1, 25)
|
|
584
|
+
tns = torch.linspace(0, 1, 25)
|
|
585
|
+
mtx = dot(linspace(0, 1, 10)[:,None], linspace(0, 2, 10)[None,:])
|
|
586
|
+
sp_mtx = csr_matrix(
|
|
587
|
+
([1.0, 0.5, 0.5, 0.2, 0.1],
|
|
588
|
+
([0, 0, 4, 5, 9], [4, 9, 4, 1, 8])),
|
|
589
|
+
shape=(10, 10),
|
|
590
|
+
dtype=float)
|
|
591
|
+
sp_tns = torch.sparse_coo_tensor(
|
|
592
|
+
torch.tensor([[0, 0, 4, 5, 9], [4, 9, 4, 1, 8]]),
|
|
593
|
+
torch.tensor([1.0, 0.5, 0.5, 0.2, 0.1]),
|
|
594
|
+
size=(10, 10),
|
|
595
|
+
dtype=float)
|
|
596
|
+
f_arr = frozenarray(arr)
|
|
597
|
+
f_sp_mtx = frozenarray(sp_mtx)
|
|
598
|
+
q_arr = quant(arr, 'mm')
|
|
599
|
+
q_mtx = quant(arr, 'seconds')
|
|
600
|
+
q_sp_mtx = quant(sp_mtx, 'kg')
|
|
601
|
+
# For an object that is already a numpy array, any call that doesn't
|
|
602
|
+
# request a copy and that doesn't change its parameters will return the
|
|
603
|
+
# identical object.
|
|
604
|
+
self.assertIs(arr, to_array(arr))
|
|
605
|
+
self.assertIs(arr, to_array(arr, sparse=False, frozen=False))
|
|
606
|
+
self.assertIs(arr, to_array(arr, quant=False))
|
|
607
|
+
self.assertIs(f_arr, to_array(f_arr))
|
|
608
|
+
self.assertIs(f_arr, to_array(f_arr, sparse=False, frozen=True))
|
|
609
|
+
self.assertIs(f_arr, to_array(f_arr, quant=False))
|
|
610
|
+
# to_array can be used to convert from tensors into arrays; the detach
|
|
611
|
+
# parameter lets us automatically detach tensors from the gradient
|
|
612
|
+
# system (this is the default) or raise an error if that would be
|
|
613
|
+
# required (detach=False).
|
|
614
|
+
self.assertIsInstance(to_array(tns), np.ndarray)
|
|
615
|
+
gradtns = tns.clone().requires_grad_(True)
|
|
616
|
+
self.assertIsInstance(to_array(gradtns), np.ndarray)
|
|
617
|
+
self.assertTrue(np.isclose(to_array(tns), arr).all())
|
|
618
|
+
self.assertTrue(np.isclose(to_array(gradtns), arr).all())
|
|
619
|
+
with self.assertRaises(ValueError):
|
|
620
|
+
to_array(gradtns, detach=False)
|
|
621
|
+
# Sparse arrays/tensors should also convert fine.
|
|
622
|
+
dn_tns = sp_tns.to_dense()
|
|
623
|
+
x = to_array(sp_tns)
|
|
624
|
+
self.assertTrue(issparse(x))
|
|
625
|
+
self.assertEqual(x.format, 'coo')
|
|
626
|
+
x = to_array(dn_tns, sparse='lil')
|
|
627
|
+
self.assertTrue(issparse(x))
|
|
628
|
+
self.assertEqual(x.format, 'lil')
|
|
629
|
+
x = to_array(dn_tns, sparse=torch.sparse_csr)
|
|
630
|
+
self.assertTrue(issparse(x))
|
|
631
|
+
self.assertEqual(x.format, 'csr')
|
|
632
|
+
x = to_array(dn_tns, sparse=False)
|
|
633
|
+
self.assertFalse(issparse(x))
|
|
634
|
+
self.assertTrue(
|
|
635
|
+
np.array_equal(dn_tns.numpy(), x))
|
|
636
|
+
x = to_array(sp_tns, sparse=False)
|
|
637
|
+
self.assertTrue(
|
|
638
|
+
np.array_equal(dn_tns.numpy(), x))
|
|
639
|
+
with self.assertRaises(ValueError):
|
|
640
|
+
to_array(dn_tns, sparse=object())
|
|
641
|
+
with self.assertRaises(ValueError):
|
|
642
|
+
to_array(dn_tns, sparse='???')
|
|
643
|
+
x = to_array(sp_mtx, copy=False, dtype=complex)
|
|
644
|
+
self.assertTrue(np.issubdtype(x.dtype, complex))
|
|
645
|
+
self.assertTrue(
|
|
646
|
+
np.all(np.isclose(x.todense().real, sp_mtx.todense())))
|
|
647
|
+
self.assertTrue(
|
|
648
|
+
np.all(np.abs(x.todense().imag) < 1e-9))
|
|
649
|
+
sp_tns = sp_tns.coalesce()
|
|
650
|
+
x = to_array(sp_tns, copy=False)
|
|
651
|
+
self.assertTrue(
|
|
652
|
+
np.shares_memory(x.data, sp_tns.values().detach().numpy()))
|
|
653
|
+
x = to_array(sp_tns, copy=True)
|
|
654
|
+
self.assertFalse(
|
|
655
|
+
np.shares_memory(x.data, sp_tns.values().detach().numpy()))
|
|
656
|
+
# If we change the parameters of the returned array, we will get
|
|
657
|
+
# different (but typically equal) objects back.
|
|
658
|
+
self.assertIsNot(arr, to_array(arr, frozen=True))
|
|
659
|
+
self.assertIsNot(f_arr, to_array(f_arr, frozen=False))
|
|
660
|
+
self.assertTrue(np.array_equal(arr, to_array(arr, frozen=True)))
|
|
661
|
+
self.assertTrue(np.array_equal(f_arr, to_array(f_arr, frozen=False)))
|
|
662
|
+
# We can also request that a copy be made like with np.array.
|
|
663
|
+
self.assertIsNot(arr, to_array(arr, copy=True))
|
|
664
|
+
self.assertTrue(np.array_equal(arr, to_array(arr, copy=True)))
|
|
665
|
+
# The sparse flag can be used to convert to/from a sparse array.
|
|
666
|
+
self.assertIsInstance(to_array(sp_mtx, sparse=False), np.ndarray)
|
|
667
|
+
self.assertTrue(np.array_equal(to_array(sp_mtx, sparse=False),
|
|
668
|
+
sp_mtx.todense()))
|
|
669
|
+
self.assertTrue(issparse(to_array(mtx, sparse=True)))
|
|
670
|
+
self.assertTrue(np.array_equal(to_array(mtx, sparse=True).todense(),
|
|
671
|
+
mtx))
|
|
672
|
+
# The frozen flag ensures that the return value does or does not have
|
|
673
|
+
# the writeable flag set.
|
|
674
|
+
self.assertFalse(to_array(mtx, frozen=True).flags['WRITEABLE'])
|
|
675
|
+
self.assertTrue(np.array_equal(to_array(mtx, frozen=True), mtx))
|
|
676
|
+
self.assertIsNot(to_array(mtx, frozen=True), mtx)
|
|
677
|
+
fsp_mtx = to_array(sp_mtx, frozen=True)
|
|
678
|
+
self.assertTrue(np.array_equal(fsp_mtx.todense(), sp_mtx.todense()))
|
|
679
|
+
self.assertFalse(fsp_mtx.data.flags['WRITEABLE'])
|
|
680
|
+
tfsp_mtx = to_array(fsp_mtx, frozen=False)
|
|
681
|
+
self.assertTrue(np.array_equal(tfsp_mtx.todense(), sp_mtx.todense()))
|
|
682
|
+
self.assertTrue(tfsp_mtx.data.flags['WRITEABLE'])
|
|
683
|
+
self.assertFalse(fsp_mtx.data.flags['WRITEABLE'])
|
|
684
|
+
with self.assertRaises(ValueError):
|
|
685
|
+
to_array(sp_mtx, frozen=object())
|
|
686
|
+
# The quant argument can be used to enforce the return of quantities or
|
|
687
|
+
# non-quantities, but you can't force a quantity without a unit:
|
|
688
|
+
with self.assertRaises(ValueError):
|
|
689
|
+
arr = to_array(arr, quant=True)
|
|
690
|
+
# The unit parameter can be used to specify what unit to use.
|
|
691
|
+
self.assertTrue(
|
|
692
|
+
np.array_equal(q_arr.m, to_array(arr, quant=True, unit='mm').m))
|
|
693
|
+
self.assertTrue(
|
|
694
|
+
np.all(
|
|
695
|
+
np.isclose(
|
|
696
|
+
to_array(q_arr, quant=True, unit='m').m,
|
|
697
|
+
to_array(arr, quant=True, unit='mm').m_as('m'))))
|
|
698
|
+
self.assertTrue(
|
|
699
|
+
np.all(
|
|
700
|
+
np.isclose(
|
|
701
|
+
to_array(q_arr, quant=True, unit='m').m,
|
|
702
|
+
to_array(q_arr, quant=True, unit='mm').m_as('m'))))
|
|
703
|
+
self.assertTrue(
|
|
704
|
+
np.all(
|
|
705
|
+
np.isclose(
|
|
706
|
+
to_array(arr, quant=False, unit='mm'),
|
|
707
|
+
to_array(q_arr, quant=False, unit='m')*1000)))
|
|
708
|
+
# We can also use quant=False and a unit to extract the array in a
|
|
709
|
+
# with a certain unit (like the mag function).
|
|
710
|
+
e_arr = to_array(q_arr, quant=False, unit=...)
|
|
711
|
+
self.assertIsInstance(e_arr, np.ndarray)
|
|
712
|
+
self.assertTrue(np.all(np.isclose(e_arr, arr)))
|
|
713
|
+
e_arr = to_array(q_arr, quant=False, unit='m')
|
|
714
|
+
self.assertIsInstance(e_arr, np.ndarray)
|
|
715
|
+
self.assertTrue(np.all(np.isclose(e_arr, arr/1000)))
|
|
716
|
+
self.assertTrue(
|
|
717
|
+
np.array_equal(q_arr.m, to_array(arr, quant=True, unit='mm').m))
|
|
718
|
+
with self.assertRaises(ValueError):
|
|
719
|
+
to_array(arr, quant=True, unit=Ellipsis)
|
|
720
|
+
with self.assertRaises(ValueError):
|
|
721
|
+
to_array(arr, quant=True, unit=None)
|
|
722
|
+
with self.assertRaises(ValueError):
|
|
723
|
+
to_array(q_arr, quant=True, unit=None)
|
|
724
|
+
with self.assertRaises(ValueError):
|
|
725
|
+
to_array(arr, quant=object())
|
|
726
|
+
# We can also specify the units registry (Ellipsis means immlib.units).
|
|
727
|
+
self.assertTrue(
|
|
728
|
+
np.all(
|
|
729
|
+
np.isclose(
|
|
730
|
+
to_array(q_arr, unit='m', ureg=Ellipsis).m,
|
|
731
|
+
q_arr.m / 1000.0)))
|
|
732
|
+
# We can also use unit to extract a specific unit from a quantity.
|
|
733
|
+
self.assertEqual(1000, to_array(quant(1, units.meter), unit='mm').m)
|
|
734
|
+
# However, a non-quantity is always assumed to already have the units
|
|
735
|
+
# requested, so converting it to a particular unit (but not converting
|
|
736
|
+
# it to a quantity) results in the same object.
|
|
737
|
+
self.assertIs(to_array(arr, quant=False, unit='mm'), arr)
|
|
738
|
+
# If we simply request an array with a unit, without specifying that it
|
|
739
|
+
# not be a quantity, we get a quantity back.
|
|
740
|
+
self.assertIsInstance(to_array(arr, unit='mm'), pint.Quantity)
|
|
741
|
+
# An error is raised if you try to request no units for a quantity.
|
|
742
|
+
with self.assertRaises(ValueError):
|
|
743
|
+
to_array(arr, quant=True, unit=None)
|
|
744
|
+
|
|
745
|
+
# PyTorch Utilities ########################################################
|
|
746
|
+
def test_is_torchdtype(self):
|
|
747
|
+
from immlib.util import is_torchdtype
|
|
748
|
+
import torch, numpy as np
|
|
749
|
+
# is_torchdtype returns true for torch's dtypes and its dtypes alone.
|
|
750
|
+
self.assertTrue(is_torchdtype(torch.int))
|
|
751
|
+
self.assertTrue(is_torchdtype(torch.float))
|
|
752
|
+
self.assertTrue(is_torchdtype(torch.bool))
|
|
753
|
+
self.assertFalse(is_torchdtype('int'))
|
|
754
|
+
self.assertFalse(is_torchdtype(float))
|
|
755
|
+
self.assertFalse(is_torchdtype(np.bool_))
|
|
756
|
+
def test_like_torchdtype(self):
|
|
757
|
+
from immlib.util import like_torchdtype
|
|
758
|
+
import torch, numpy as np
|
|
759
|
+
# Anything that can be converted into a torch dtype object is considered
|
|
760
|
+
# to be like a torch dtype.
|
|
761
|
+
self.assertTrue(like_torchdtype('int'))
|
|
762
|
+
self.assertTrue(like_torchdtype(float))
|
|
763
|
+
self.assertTrue(like_torchdtype(np.bool_))
|
|
764
|
+
self.assertFalse(like_torchdtype('abc'))
|
|
765
|
+
self.assertFalse(like_torchdtype(10))
|
|
766
|
+
self.assertFalse(like_torchdtype(...))
|
|
767
|
+
# Note that None can be converted to a torch dtype (float64).
|
|
768
|
+
self.assertTrue(like_torchdtype(None))
|
|
769
|
+
# torch dtypes themselves are like torch dtypes.
|
|
770
|
+
self.assertTrue(like_torchdtype(torch.float))
|
|
771
|
+
def test_to_torchdtype(self):
|
|
772
|
+
from immlib.util import to_torchdtype
|
|
773
|
+
import torch, numpy as np
|
|
774
|
+
# Converting a numpy dtype into a dtype results in the identical dtype.
|
|
775
|
+
dt = torch.int
|
|
776
|
+
self.assertIs(dt, to_torchdtype(dt))
|
|
777
|
+
# Numpy dtypes can be converted into a torch dtype.
|
|
778
|
+
self.assertEqual(torch.float64, to_torchdtype(np.dtype('float64')))
|
|
779
|
+
self.assertEqual(torch.int32, to_torchdtype(np.int32))
|
|
780
|
+
def test_is_tensor(self):
|
|
781
|
+
from immlib import (is_tensor, quant)
|
|
782
|
+
from scipy.sparse import csr_matrix, csr_array
|
|
783
|
+
import torch, numpy as np
|
|
784
|
+
# By default, is_tensor() returns True for PyTorch tensors and
|
|
785
|
+
# quantities whose magnitudes are PyTorch tensors.
|
|
786
|
+
arr = torch.linspace(0, 1, 25)
|
|
787
|
+
mtx = torch.mm(torch.linspace(0, 1, 10)[:,None],
|
|
788
|
+
torch.linspace(1, 2, 10)[None,:])
|
|
789
|
+
sp_mtx = torch.sparse_coo_tensor(
|
|
790
|
+
torch.tensor([[0, 0, 4, 5, 9],
|
|
791
|
+
[4, 9, 4, 1, 8]]),
|
|
792
|
+
torch.tensor([1, 0.5, 0.5, 0.2, 0.1]),
|
|
793
|
+
size=(10, 10),
|
|
794
|
+
dtype=float)
|
|
795
|
+
q_arr = quant(arr, 'mm')
|
|
796
|
+
q_mtx = quant(arr, 'seconds')
|
|
797
|
+
q_sp_mtx = quant(sp_mtx, 'kg')
|
|
798
|
+
self.assertTrue(is_tensor(arr))
|
|
799
|
+
self.assertTrue(is_tensor(mtx))
|
|
800
|
+
self.assertTrue(is_tensor(sp_mtx))
|
|
801
|
+
self.assertTrue(is_tensor(q_arr))
|
|
802
|
+
self.assertTrue(is_tensor(q_mtx))
|
|
803
|
+
self.assertTrue(is_tensor(q_sp_mtx))
|
|
804
|
+
# Things like lists, numbers, and numpy arrays are not tensors.
|
|
805
|
+
self.assertFalse(is_tensor('abc'))
|
|
806
|
+
self.assertFalse(is_tensor(10))
|
|
807
|
+
self.assertFalse(is_tensor([12.0, 0.5, 3.2]))
|
|
808
|
+
self.assertFalse(is_tensor(np.array([1.0, 2.0, 3.0])))
|
|
809
|
+
self.assertFalse(is_tensor(quant(np.array([1.0, 2.0, 3.0]), 'mm')))
|
|
810
|
+
# We can use the dtype argument to restrict what we consider an array by
|
|
811
|
+
# its dtype. The dtype of the is_array argument must be a sub-dtype of
|
|
812
|
+
# the dtype parameter.
|
|
813
|
+
self.assertTrue(is_tensor(arr, dtype=arr.dtype))
|
|
814
|
+
self.assertFalse(is_tensor(arr, dtype=torch.int))
|
|
815
|
+
# If a tuple is passed for the dtype, the dtype must match one of the
|
|
816
|
+
# tuple's elements.
|
|
817
|
+
self.assertTrue(is_tensor(mtx, dtype=(mtx.dtype,)))
|
|
818
|
+
self.assertTrue(is_tensor(mtx, dtype=(mtx.dtype, torch.int)))
|
|
819
|
+
self.assertFalse(is_tensor(mtx, dtype=(torch.int, torch.bool)))
|
|
820
|
+
self.assertFalse(is_tensor(torch.tensor([], dtype=torch.int32),
|
|
821
|
+
dtype=torch.int64))
|
|
822
|
+
# torch dtypes can be interpreted into PyTorch dtypes.
|
|
823
|
+
self.assertTrue(is_tensor(mtx, dtype=mtx.numpy().dtype))
|
|
824
|
+
# We can use the ndim argument to restrict the number of dimensions that
|
|
825
|
+
# an array can have in order to be considered a matching tensor.
|
|
826
|
+
# Typically, this is just the number of dimensions.
|
|
827
|
+
self.assertTrue(is_tensor(arr, ndim=1))
|
|
828
|
+
self.assertTrue(is_tensor(mtx, ndim=2))
|
|
829
|
+
self.assertFalse(is_tensor(arr, ndim=2))
|
|
830
|
+
self.assertFalse(is_tensor(mtx, ndim=1))
|
|
831
|
+
# Alternately, a tuple may be given, in which case any of the dimension
|
|
832
|
+
# counts in the tuple are accepted.
|
|
833
|
+
self.assertTrue(is_tensor(mtx, ndim=(1,2)))
|
|
834
|
+
self.assertTrue(is_tensor(arr, ndim=(1,2)))
|
|
835
|
+
self.assertFalse(is_tensor(mtx, ndim=(1,3)))
|
|
836
|
+
self.assertFalse(is_tensor(arr, ndim=(0,2)))
|
|
837
|
+
# Scalar tensors have 0 dimensions.
|
|
838
|
+
self.assertTrue(is_tensor(torch.tensor(0), ndim=0))
|
|
839
|
+
# The shape option is a more specific version of the ndim parameter. It
|
|
840
|
+
# lets you specify what kind of shape is required of the tensor. The
|
|
841
|
+
# most straightforward usage is to require a specific shape.
|
|
842
|
+
self.assertTrue(is_tensor(arr, shape=(25,)))
|
|
843
|
+
self.assertTrue(is_tensor(mtx, shape=(10,10)))
|
|
844
|
+
self.assertFalse(is_tensor(arr, shape=(25,25)))
|
|
845
|
+
self.assertFalse(is_tensor(mtx, shape=(10,)))
|
|
846
|
+
# A -1 value that appears in the shape option represents any size along
|
|
847
|
+
# that dimension (a wildcard). Any number of -1s can be included.
|
|
848
|
+
self.assertTrue(is_tensor(arr, shape=(-1,)))
|
|
849
|
+
self.assertTrue(is_tensor(mtx, shape=(-1,10)))
|
|
850
|
+
self.assertTrue(is_tensor(mtx, shape=(10,-1)))
|
|
851
|
+
self.assertTrue(is_tensor(mtx, shape=(-1,-1)))
|
|
852
|
+
self.assertFalse(is_tensor(mtx, shape=(1,-1)))
|
|
853
|
+
# No more than 1 ellipsis may be included in the shape to indicate that
|
|
854
|
+
# any number of dimensions, with any sizes, can appear in place of the
|
|
855
|
+
# ellipsis.
|
|
856
|
+
self.assertTrue(is_tensor(arr, shape=(...,25)))
|
|
857
|
+
self.assertTrue(is_tensor(arr, shape=(25,...)))
|
|
858
|
+
self.assertFalse(is_tensor(arr, shape=(25,...,25)))
|
|
859
|
+
self.assertTrue(is_tensor(mtx, shape=(...,10)))
|
|
860
|
+
self.assertTrue(is_tensor(mtx, shape=(10,...)))
|
|
861
|
+
self.assertTrue(is_tensor(mtx, shape=(10,...,10)))
|
|
862
|
+
self.assertTrue(is_tensor(mtx, shape=(10,10,...)))
|
|
863
|
+
self.assertTrue(is_tensor(mtx, shape=(...,10,10)))
|
|
864
|
+
self.assertFalse(is_tensor(mtx, shape=(10,...,10,10)))
|
|
865
|
+
self.assertFalse(is_tensor(mtx, shape=(10,10,...,10)))
|
|
866
|
+
self.assertTrue(is_tensor(torch.zeros((1,2,3,4,5)), shape=(1,...,4,5)))
|
|
867
|
+
# The numel option allows one to specify the number of elements that an
|
|
868
|
+
# object must have. This does not care about dimensionality.
|
|
869
|
+
self.assertTrue(is_tensor(arr, numel=25))
|
|
870
|
+
self.assertFalse(is_tensor(arr, numel=26))
|
|
871
|
+
self.assertTrue(is_tensor(torch.tensor(0), numel=1))
|
|
872
|
+
self.assertTrue(is_tensor(torch.tensor([0]), numel=1))
|
|
873
|
+
self.assertTrue(is_tensor(torch.tensor([[0]]), numel=1))
|
|
874
|
+
# The sparse option can test whether an object is a sparse tensor or
|
|
875
|
+
# not. By default sparse is None, meaning that it doesn't matter whether
|
|
876
|
+
# an object is sparse, but sometimes you want to check for strict
|
|
877
|
+
# sparsity requirements.
|
|
878
|
+
self.assertTrue(is_tensor(arr, sparse=False))
|
|
879
|
+
self.assertTrue(is_tensor(mtx, sparse=False))
|
|
880
|
+
self.assertFalse(is_tensor(sp_mtx, sparse=False))
|
|
881
|
+
self.assertFalse(is_tensor(arr, sparse=True))
|
|
882
|
+
self.assertFalse(is_tensor(mtx, sparse=True))
|
|
883
|
+
self.assertTrue(is_tensor(sp_mtx, sparse=True))
|
|
884
|
+
# You can also require a kind of sparse matrix.
|
|
885
|
+
self.assertTrue(is_tensor(sp_mtx, sparse='coo'))
|
|
886
|
+
self.assertFalse(is_tensor(sp_mtx, sparse='csc'))
|
|
887
|
+
with self.assertRaises(ValueError):
|
|
888
|
+
is_tensor(sp_mtx, sparse='???')
|
|
889
|
+
with self.assertRaises(ValueError):
|
|
890
|
+
is_tensor(sp_mtx, sparse=object())
|
|
891
|
+
# The quant option can be used to control whether the object must or
|
|
892
|
+
# must not be a quantity.
|
|
893
|
+
self.assertTrue(is_tensor(arr, quant=False))
|
|
894
|
+
self.assertTrue(is_tensor(mtx, quant=False))
|
|
895
|
+
self.assertFalse(is_tensor(arr, quant=True))
|
|
896
|
+
self.assertFalse(is_tensor(mtx, quant=True))
|
|
897
|
+
self.assertTrue(is_tensor(q_arr, quant=True))
|
|
898
|
+
self.assertTrue(is_tensor(q_mtx, quant=True))
|
|
899
|
+
self.assertFalse(is_tensor(q_arr, quant=False))
|
|
900
|
+
self.assertFalse(is_tensor(q_mtx, quant=False))
|
|
901
|
+
# The units option can be used to require that either an object have
|
|
902
|
+
# no units (or is not a quantity) or that it have specific units.
|
|
903
|
+
self.assertTrue(is_tensor(arr, unit=None))
|
|
904
|
+
self.assertTrue(is_tensor(mtx, unit=None))
|
|
905
|
+
self.assertFalse(is_tensor(arr, unit='mm'))
|
|
906
|
+
self.assertFalse(is_tensor(mtx, unit='s'))
|
|
907
|
+
self.assertFalse(is_tensor(q_arr, unit=None))
|
|
908
|
+
self.assertFalse(is_tensor(q_mtx, unit=None))
|
|
909
|
+
self.assertTrue(is_tensor(q_arr, unit='mm'))
|
|
910
|
+
self.assertTrue(is_tensor(q_mtx, unit='s'))
|
|
911
|
+
self.assertFalse(is_tensor(q_arr, unit='s'))
|
|
912
|
+
self.assertFalse(is_tensor(q_mtx, unit='mm'))
|
|
913
|
+
# The units option can be used to require that either an object have
|
|
914
|
+
# no units (or is not a quantity) or that it have specific units.
|
|
915
|
+
self.assertTrue(is_tensor(arr, unit=None))
|
|
916
|
+
self.assertTrue(is_tensor(mtx, unit=None))
|
|
917
|
+
self.assertFalse(is_tensor(arr, unit='mm'))
|
|
918
|
+
self.assertFalse(is_tensor(mtx, unit='s'))
|
|
919
|
+
self.assertFalse(is_tensor(q_arr, unit=None))
|
|
920
|
+
self.assertFalse(is_tensor(q_mtx, unit=None))
|
|
921
|
+
self.assertTrue(is_tensor(q_arr, unit='mm'))
|
|
922
|
+
self.assertTrue(is_tensor(q_mtx, unit='s'))
|
|
923
|
+
self.assertFalse(is_tensor(q_arr, unit='s'))
|
|
924
|
+
self.assertFalse(is_tensor(q_mtx, unit='mm'))
|
|
925
|
+
# We can also specify the units registry (Ellipsis means immlib.units).
|
|
926
|
+
self.assertFalse(is_tensor(q_arr, unit='s', ureg=Ellipsis))
|
|
927
|
+
# We can also test on torch data like device and requires_grad:
|
|
928
|
+
self.assertTrue(is_tensor(arr, device='cpu'))
|
|
929
|
+
self.assertFalse(is_tensor(arr, device='cuda'))
|
|
930
|
+
self.assertTrue(is_tensor(arr, requires_grad=False))
|
|
931
|
+
self.assertFalse(is_tensor(arr, requires_grad=True))
|
|
932
|
+
gradtns = arr.clone().requires_grad_(True)
|
|
933
|
+
self.assertFalse(is_tensor(gradtns, requires_grad=False))
|
|
934
|
+
self.assertTrue(is_tensor(gradtns, requires_grad=True))
|
|
935
|
+
def test_to_tensor(self):
|
|
936
|
+
from immlib import (to_tensor, quant, is_quant, units)
|
|
937
|
+
from immlib.util._numeric import torch__is_sparse
|
|
938
|
+
from numpy import (linspace, dot)
|
|
939
|
+
from scipy.sparse import (csr_array, issparse)
|
|
940
|
+
import torch, numpy as np, pint
|
|
941
|
+
# We'll use a few objects throughout our tests, which we setup now.
|
|
942
|
+
arr = linspace(0, 1, 25)
|
|
943
|
+
tns = torch.linspace(0, 1, 25)
|
|
944
|
+
sp_arr = csr_array(
|
|
945
|
+
([1.0, 0.5, 0.5, 0.2, 0.1],
|
|
946
|
+
([0, 0, 4, 5, 9], [4, 9, 4, 1, 8])),
|
|
947
|
+
shape=(10, 10),
|
|
948
|
+
dtype=float)
|
|
949
|
+
sp_tns = torch.sparse_coo_tensor(
|
|
950
|
+
torch.tensor([[0, 0, 4, 5, 9], [4, 9, 4, 1, 8]]),
|
|
951
|
+
torch.tensor([1.0, 0.5, 0.5, 0.2, 0.1]),
|
|
952
|
+
size=(10, 10),
|
|
953
|
+
dtype=float)
|
|
954
|
+
q_arr = quant(arr, 'mm')
|
|
955
|
+
q_tns = quant(tns, 'mm')
|
|
956
|
+
q_sp_tns = quant(sp_tns, 'kg')
|
|
957
|
+
q_sp_arr = quant(sp_arr, 'lb')
|
|
958
|
+
# For an object that is already a numpy tensor, any call that doesn't
|
|
959
|
+
# request a copy and that doesn't change its parameters will return the
|
|
960
|
+
# identical object.
|
|
961
|
+
self.assertIs(tns, to_tensor(tns))
|
|
962
|
+
self.assertIs(tns, to_tensor(tns, quant=False))
|
|
963
|
+
# to_tensor can be used to convert from arrayss into tensors
|
|
964
|
+
self.assertIsInstance(to_tensor(arr), torch.Tensor)
|
|
965
|
+
# Sparse tensors/tensors should also convert fine.
|
|
966
|
+
dn_tns = sp_tns.to_dense()
|
|
967
|
+
dn_arr = sp_arr.todense()
|
|
968
|
+
x = to_tensor(sp_arr)
|
|
969
|
+
self.assertTrue(torch__is_sparse(x))
|
|
970
|
+
self.assertEqual(x.layout, torch.sparse_csr)
|
|
971
|
+
x = to_tensor(dn_tns, sparse='coo')
|
|
972
|
+
self.assertTrue(torch__is_sparse(x))
|
|
973
|
+
self.assertEqual(x.layout, torch.sparse_coo)
|
|
974
|
+
self.assertIsInstance(x, torch.Tensor)
|
|
975
|
+
x = to_tensor(dn_tns, sparse=torch.sparse_csr)
|
|
976
|
+
self.assertTrue(torch__is_sparse(x))
|
|
977
|
+
self.assertEqual(x.layout, torch.sparse_csr)
|
|
978
|
+
self.assertIsInstance(x, torch.Tensor)
|
|
979
|
+
x = to_tensor(dn_tns, sparse=False)
|
|
980
|
+
self.assertFalse(torch__is_sparse(x))
|
|
981
|
+
self.assertTrue(torch.equal(dn_tns, x))
|
|
982
|
+
self.assertIsInstance(x, torch.Tensor)
|
|
983
|
+
x = to_tensor(sp_tns, sparse=False)
|
|
984
|
+
self.assertTrue(torch.equal(dn_tns, x))
|
|
985
|
+
self.assertIsInstance(x, torch.Tensor)
|
|
986
|
+
x = to_tensor(sp_arr, sparse=False)
|
|
987
|
+
self.assertTrue(torch.all(torch.isclose(dn_tns, x)))
|
|
988
|
+
self.assertIsInstance(x, torch.Tensor)
|
|
989
|
+
with self.assertRaises(ValueError):
|
|
990
|
+
to_tensor(dn_tns, sparse=object())
|
|
991
|
+
with self.assertRaises(ValueError):
|
|
992
|
+
to_tensor(dn_tns, sparse='???')
|
|
993
|
+
with self.assertRaises(ValueError):
|
|
994
|
+
x = to_tensor(sp_tns, copy=False, dtype=complex)
|
|
995
|
+
x = to_tensor(sp_tns, copy=None, dtype=complex)
|
|
996
|
+
self.assertEqual(x.dtype, torch.complex128)
|
|
997
|
+
self.assertTrue(
|
|
998
|
+
torch.all(torch.isclose(x.to_dense().real, sp_tns.to_dense())))
|
|
999
|
+
self.assertTrue(
|
|
1000
|
+
torch.all(torch.abs(x.to_dense().imag) < 1e-9))
|
|
1001
|
+
# We can also request that a copy be made like with np.array.
|
|
1002
|
+
self.assertIsNot(arr, to_tensor(arr, copy=True).numpy())
|
|
1003
|
+
self.assertTrue(torch.equal(tns, to_tensor(tns, copy=True)))
|
|
1004
|
+
self.assertTrue(
|
|
1005
|
+
np.shares_memory(arr.data, to_tensor(arr, copy=False).numpy().data))
|
|
1006
|
+
# The sparse flag can be used to convert to/from a sparse tensor.
|
|
1007
|
+
self.assertIsInstance(to_tensor(sp_tns, sparse=False), torch.Tensor)
|
|
1008
|
+
self.assertEqual(to_tensor(sp_tns, sparse=False).layout, torch.strided)
|
|
1009
|
+
self.assertTrue(
|
|
1010
|
+
torch.equal(to_tensor(sp_tns, sparse=False), sp_tns.to_dense()))
|
|
1011
|
+
self.assertTrue(torch__is_sparse(to_tensor(tns, sparse=True)))
|
|
1012
|
+
self.assertTrue(
|
|
1013
|
+
torch.equal(to_tensor(tns, sparse=True).to_dense(), tns))
|
|
1014
|
+
# The quant argument can be used to enforce the return of quantities or
|
|
1015
|
+
# non-quantities, but you can't force a quantity without a unit:
|
|
1016
|
+
with self.assertRaises(ValueError):
|
|
1017
|
+
arr = to_tensor(arr, quant=True, unit=None)
|
|
1018
|
+
# The unit parameter can be used to specify what unit to use.
|
|
1019
|
+
self.assertTrue(
|
|
1020
|
+
torch.equal(q_tns.m, to_tensor(tns, quant=True, unit='mm').m))
|
|
1021
|
+
self.assertTrue(
|
|
1022
|
+
torch.all(
|
|
1023
|
+
torch.isclose(
|
|
1024
|
+
to_tensor(q_arr, quant=True, unit='m').m,
|
|
1025
|
+
to_tensor(arr, quant=True, unit='mm').m_as('m'))))
|
|
1026
|
+
self.assertTrue(
|
|
1027
|
+
torch.all(
|
|
1028
|
+
torch.isclose(
|
|
1029
|
+
to_tensor(q_arr, quant=True, unit='m').m,
|
|
1030
|
+
to_tensor(q_arr, quant=True, unit='mm').m_as('m'))))
|
|
1031
|
+
self.assertTrue(
|
|
1032
|
+
torch.all(
|
|
1033
|
+
torch.isclose(
|
|
1034
|
+
to_tensor(arr, quant=False, unit='mm'),
|
|
1035
|
+
to_tensor(q_arr, quant=False, unit='m')*1000)))
|
|
1036
|
+
# We can also use quant=False and a unit to extract the tensor with a
|
|
1037
|
+
# certain unit (like the mag function).
|
|
1038
|
+
e_tns = to_tensor(q_tns, quant=False, unit=...)
|
|
1039
|
+
self.assertIsInstance(e_tns, torch.Tensor)
|
|
1040
|
+
self.assertTrue(torch.all(torch.isclose(e_tns, tns)))
|
|
1041
|
+
e_tns = to_tensor(q_tns, quant=False, unit='m')
|
|
1042
|
+
self.assertIsInstance(e_tns, torch.Tensor)
|
|
1043
|
+
self.assertTrue(torch.all(torch.isclose(e_tns, tns/1000)))
|
|
1044
|
+
self.assertTrue(
|
|
1045
|
+
torch.equal(q_tns.m, to_tensor(tns, quant=True, unit='mm').m))
|
|
1046
|
+
with self.assertRaises(ValueError):
|
|
1047
|
+
to_tensor(tns, quant=True, unit=Ellipsis)
|
|
1048
|
+
with self.assertRaises(ValueError):
|
|
1049
|
+
to_tensor(tns, quant=True, unit=None)
|
|
1050
|
+
with self.assertRaises(ValueError):
|
|
1051
|
+
to_tensor(q_tns, quant=True, unit=None)
|
|
1052
|
+
with self.assertRaises(ValueError):
|
|
1053
|
+
to_tensor(tns, quant=object())
|
|
1054
|
+
# We can also specify the units registry (Ellipsis means immlib.units).
|
|
1055
|
+
self.assertTrue(
|
|
1056
|
+
torch.all(
|
|
1057
|
+
torch.isclose(
|
|
1058
|
+
to_tensor(q_tns, unit='m', ureg=Ellipsis).m,
|
|
1059
|
+
q_tns.m / 1000.0)))
|
|
1060
|
+
# We can also use unit to extract a specific unit from a quantity.
|
|
1061
|
+
self.assertEqual(1000, to_tensor(quant(1, units.meter), unit='mm').m)
|
|
1062
|
+
# However, a non-quantity is always assumed to already have the units
|
|
1063
|
+
# requested, so converting it to a particular unit (but not converting
|
|
1064
|
+
# it to a quantity) results in the same object.
|
|
1065
|
+
self.assertIs(to_tensor(tns, quant=False, unit='mm'), tns)
|
|
1066
|
+
# If we simply request an tensor with a unit, without specifying that it
|
|
1067
|
+
# not be a quantity, we get a quantity back.
|
|
1068
|
+
self.assertIsInstance(to_tensor(tns, unit='mm'), pint.Quantity)
|
|
1069
|
+
# An error is raised if you try to request no units for a quantity.
|
|
1070
|
+
with self.assertRaises(ValueError):
|
|
1071
|
+
to_tensor(tns, quant=True, unit=None)
|
|
1072
|
+
# If we change the parameters of the returned array, we will get
|
|
1073
|
+
# different (but typically equal) objects back.
|
|
1074
|
+
self.assertTrue(torch.equal(tns, to_tensor(tns, requires_grad=True)))
|
|
1075
|
+
|
|
1076
|
+
# PyTorch and Numpy Helper Functions #######################################
|
|
1077
|
+
def test_is_numeric(self):
|
|
1078
|
+
from immlib import is_numeric
|
|
1079
|
+
import torch, numpy as np
|
|
1080
|
+
from scipy.sparse import csr_matrix
|
|
1081
|
+
# The is_numeric function is just a wrapper around is_array and
|
|
1082
|
+
# is_tensor that calls one or the other depending on whether the object
|
|
1083
|
+
# requested is a tensor or not. I.e., it passes all arguments through
|
|
1084
|
+
# and merely switches on the type.
|
|
1085
|
+
sp_a = csr_matrix(([0.5, 1.0], ([0,1], [3,2])), shape=(5,5))
|
|
1086
|
+
sp_t = torch.sparse_coo_tensor(torch.tensor([[0,1],[3,2]]),
|
|
1087
|
+
torch.tensor([0.5, 1]),
|
|
1088
|
+
(5,5))
|
|
1089
|
+
a = sp_a.todense()
|
|
1090
|
+
t = sp_t.to_dense()
|
|
1091
|
+
self.assertTrue(is_numeric(a))
|
|
1092
|
+
self.assertTrue(is_numeric(t))
|
|
1093
|
+
self.assertTrue(is_numeric(sp_a))
|
|
1094
|
+
self.assertTrue(is_numeric(sp_t))
|
|
1095
|
+
self.assertFalse(is_numeric('abc'))
|
|
1096
|
+
self.assertFalse(is_numeric([1,2,3]))
|
|
1097
|
+
def test_to_numeric(self):
|
|
1098
|
+
from immlib import to_numeric
|
|
1099
|
+
import torch, numpy as np
|
|
1100
|
+
from scipy.sparse import csr_matrix
|
|
1101
|
+
# The is_numeric function is just a wrapper around to_array and
|
|
1102
|
+
# to_tensor that calls one or the other depending on whether the object
|
|
1103
|
+
# requested is a tensor or not. I.e., it passes all arguments through
|
|
1104
|
+
# and merely switches on the type.
|
|
1105
|
+
sp_a = csr_matrix(([0.5, 1.0], ([0,1], [3,2])), shape=(5,5))
|
|
1106
|
+
sp_t = torch.sparse_coo_tensor(torch.tensor([[0,1],[3,2]]),
|
|
1107
|
+
torch.tensor([0.5, 1]),
|
|
1108
|
+
(5,5))
|
|
1109
|
+
a = np.array(sp_a.todense())
|
|
1110
|
+
t = sp_t.to_dense()
|
|
1111
|
+
self.assertIs(a, to_numeric(a))
|
|
1112
|
+
self.assertIs(t, to_numeric(t))
|
|
1113
|
+
self.assertIs(sp_a, to_numeric(sp_a))
|
|
1114
|
+
self.assertIs(sp_t, to_numeric(sp_t))
|
|
1115
|
+
self.assertIsInstance(to_numeric([1,2,3]), np.ndarray)
|
|
1116
|
+
def test_is_sparse(self):
|
|
1117
|
+
from immlib import is_sparse
|
|
1118
|
+
import torch, numpy as np
|
|
1119
|
+
from scipy.sparse import csr_matrix
|
|
1120
|
+
# is_sparse returns True for any sparse array and False for anything
|
|
1121
|
+
# other than a sparse array.
|
|
1122
|
+
sp_a = csr_matrix(([0.5, 1.0], ([0,1], [3,2])), shape=(5,5))
|
|
1123
|
+
sp_t = torch.sparse_coo_tensor(torch.tensor([[0,1],[3,2]]),
|
|
1124
|
+
torch.tensor([0.5, 1]),
|
|
1125
|
+
(5,5))
|
|
1126
|
+
self.assertTrue(is_sparse(sp_a))
|
|
1127
|
+
self.assertTrue(is_sparse(sp_t))
|
|
1128
|
+
self.assertFalse(is_sparse(sp_a.todense()))
|
|
1129
|
+
self.assertFalse(is_sparse(sp_t.to_dense()))
|
|
1130
|
+
def test_to_sparse(self):
|
|
1131
|
+
from immlib import to_sparse
|
|
1132
|
+
import torch, numpy as np
|
|
1133
|
+
from scipy.sparse import issparse
|
|
1134
|
+
# to_sparse supports the arguments of to_array and to_tensor (because it
|
|
1135
|
+
# simply calls through to these functions), but it always returns a
|
|
1136
|
+
# sparse object.
|
|
1137
|
+
m = np.array([[1.0, 0, 0, 0], [0, 0, 0, 0],
|
|
1138
|
+
[0, 1.0, 0, 0], [0, 0, 0, 1.0]])
|
|
1139
|
+
t = torch.tensor(m)
|
|
1140
|
+
self.assertTrue(issparse(to_sparse(m)))
|
|
1141
|
+
self.assertTrue(to_sparse(t).is_sparse)
|
|
1142
|
+
def test_is_dense(self):
|
|
1143
|
+
from immlib import is_dense
|
|
1144
|
+
import torch, numpy as np
|
|
1145
|
+
from scipy.sparse import csr_array
|
|
1146
|
+
# is_dense returns True for any dense array and False for anything
|
|
1147
|
+
# other than a dense array.
|
|
1148
|
+
sp_a = csr_array(([0.5, 1.0], ([0,1], [3,2])), shape=(5,5))
|
|
1149
|
+
sp_t = torch.sparse_coo_tensor(
|
|
1150
|
+
torch.tensor([[0,1],[3,2]]),
|
|
1151
|
+
torch.tensor([0.5, 1]),
|
|
1152
|
+
(5,5))
|
|
1153
|
+
self.assertTrue(is_dense(sp_a.todense()))
|
|
1154
|
+
x = sp_t.to_dense()
|
|
1155
|
+
q = is_dense(sp_t.to_dense())
|
|
1156
|
+
self.assertTrue(q)
|
|
1157
|
+
self.assertFalse(is_dense(sp_a))
|
|
1158
|
+
self.assertFalse(is_dense(sp_t))
|
|
1159
|
+
def test_to_dense(self):
|
|
1160
|
+
from immlib import to_dense
|
|
1161
|
+
import torch, numpy as np
|
|
1162
|
+
from scipy.sparse import (issparse, csr_matrix)
|
|
1163
|
+
# to_dense supports the arguments of to_array and to_tensor (because it
|
|
1164
|
+
# simply calls through to these functions), but it always returns a
|
|
1165
|
+
# dense object.
|
|
1166
|
+
sp_a = csr_matrix(([0.5, 1.0], ([0,1], [3,2])), shape=(5,5))
|
|
1167
|
+
sp_t = torch.sparse_coo_tensor(torch.tensor([[0,1],[3,2]]),
|
|
1168
|
+
torch.tensor([0.5, 1]),
|
|
1169
|
+
(5,5))
|
|
1170
|
+
self.assertFalse(issparse(to_dense(sp_a)))
|
|
1171
|
+
self.assertFalse(to_dense(sp_t).is_sparse)
|
|
1172
|
+
def test_like_number(self):
|
|
1173
|
+
from immlib import like_number
|
|
1174
|
+
import torch, numpy as np
|
|
1175
|
+
# like_number returns True if the argument is a scalar number or if it
|
|
1176
|
+
# is convertible into a scalar number by the to_scalar function. Such
|
|
1177
|
+
# values include numbers and any numeric numpy array or PyTorch tensor
|
|
1178
|
+
# that has exactly one value.
|
|
1179
|
+
self.assertTrue(like_number(10))
|
|
1180
|
+
self.assertTrue(like_number(10.0))
|
|
1181
|
+
self.assertTrue(like_number(10.0 + 20.5j))
|
|
1182
|
+
self.assertTrue(like_number(True))
|
|
1183
|
+
self.assertTrue(like_number(np.array(10)))
|
|
1184
|
+
self.assertTrue(like_number(torch.tensor(10)))
|
|
1185
|
+
self.assertTrue(like_number([10]))
|
|
1186
|
+
self.assertTrue(like_number([[10]]))
|
|
1187
|
+
self.assertTrue(like_number([[[10]]]))
|
|
1188
|
+
self.assertTrue(like_number(np.array([10])))
|
|
1189
|
+
self.assertTrue(like_number(np.array([[10]])))
|
|
1190
|
+
self.assertTrue(like_number(np.array([[[10]]])))
|
|
1191
|
+
self.assertTrue(like_number(torch.tensor([10])))
|
|
1192
|
+
self.assertTrue(like_number(torch.tensor([[10]])))
|
|
1193
|
+
self.assertTrue(like_number(torch.tensor([[[10]]])))
|
|
1194
|
+
self.assertFalse(like_number('10'))
|
|
1195
|
+
self.assertFalse(like_number({'a':10}))
|
|
1196
|
+
self.assertFalse(like_number([1,2,3]))
|
|
1197
|
+
# ragged arrays are not like numbers:
|
|
1198
|
+
self.assertFalse(like_number([[1,2,3],[2,3]]))
|
|
1199
|
+
def test_to_number(self):
|
|
1200
|
+
from immlib import to_number
|
|
1201
|
+
import torch, numpy as np
|
|
1202
|
+
# to_number returns a scalar version of the given argument assuming that
|
|
1203
|
+
# the argument is like a scalar (see like_number).
|
|
1204
|
+
self.assertEqual(to_number(10), 10)
|
|
1205
|
+
self.assertEqual(to_number(10.0), 10.0)
|
|
1206
|
+
self.assertEqual(to_number(10.0 + 20.5j), 10.0 + 20.5j)
|
|
1207
|
+
self.assertEqual(to_number(True), True)
|
|
1208
|
+
self.assertEqual(to_number(np.array(10)), 10)
|
|
1209
|
+
self.assertEqual(to_number(torch.tensor(10)), 10)
|
|
1210
|
+
self.assertEqual(to_number([10]), 10)
|
|
1211
|
+
self.assertEqual(to_number([[10]]), 10)
|
|
1212
|
+
self.assertEqual(to_number([[[10]]]), 10)
|
|
1213
|
+
self.assertEqual(to_number(np.array([10])), 10)
|
|
1214
|
+
self.assertEqual(to_number(np.array([[10]])), 10)
|
|
1215
|
+
self.assertEqual(to_number(np.array([[[10]]])), 10)
|
|
1216
|
+
self.assertEqual(to_number(torch.tensor([10])), 10)
|
|
1217
|
+
self.assertEqual(to_number(torch.tensor([[10]])), 10)
|
|
1218
|
+
self.assertEqual(to_number(torch.tensor([[[10]]])), 10)
|
|
1219
|
+
with self.assertRaises(TypeError): to_number('10')
|
|
1220
|
+
with self.assertRaises(TypeError): to_number({'a':10})
|
|
1221
|
+
with self.assertRaises(TypeError): to_number([1,2,3])
|
|
1222
|
+
|
|
1223
|
+
# The numapi Decorator #####################################################
|
|
1224
|
+
def test_numapi(self):
|
|
1225
|
+
from immlib.util import numapi
|
|
1226
|
+
import numpy as np, torch
|
|
1227
|
+
# Basic test:
|
|
1228
|
+
@numapi
|
|
1229
|
+
def l2_distance(pt1, pt2):
|
|
1230
|
+
"Calculates the L2 distance between two points."
|
|
1231
|
+
pass
|
|
1232
|
+
@l2_distance.array
|
|
1233
|
+
def _(pt1, pt2):
|
|
1234
|
+
return np.sqrt(np.sum((pt1 - pt2)**2, axis=0))
|
|
1235
|
+
@l2_distance.tensor
|
|
1236
|
+
def _(pt1, pt2):
|
|
1237
|
+
return torch.sqrt(torch.sum((pt1 - pt2)**2, axis=0))
|
|
1238
|
+
tns = l2_distance(torch.tensor([0,0]), [0,1])
|
|
1239
|
+
self.assertIsInstance(tns, torch.Tensor)
|
|
1240
|
+
self.assertEqual(tns, 1.0)
|
|
1241
|
+
tns = l2_distance([0,0], torch.tensor([0,1]))
|
|
1242
|
+
self.assertIsInstance(tns, torch.Tensor)
|
|
1243
|
+
self.assertEqual(tns, 1.0)
|
|
1244
|
+
arr = l2_distance([0,0], [0,1])
|
|
1245
|
+
self.assertIsInstance(arr, float)
|
|
1246
|
+
self.assertEqual(arr, 1.0)
|
|
1247
|
+
|
|
1248
|
+
# The tensor_args, array_args, and numeric_args decorators ################
|
|
1249
|
+
def test_tensor_args(self):
|
|
1250
|
+
from immlib.util import tensor_args
|
|
1251
|
+
import numpy as np, torch
|
|
1252
|
+
# Without any arguments, it should just auto-tensorify the args.
|
|
1253
|
+
@tensor_args
|
|
1254
|
+
def test1(a, b, c='test'):
|
|
1255
|
+
return (torch.is_tensor(a), torch.is_tensor(b), torch.is_tensor(c))
|
|
1256
|
+
# By default it shouldn't convert things like strings or dicts into
|
|
1257
|
+
# tensors.
|
|
1258
|
+
(a, b, c) = test1(10, {'a':12, 'b':13})
|
|
1259
|
+
self.assertTrue(a)
|
|
1260
|
+
self.assertFalse(b)
|
|
1261
|
+
self.assertFalse(c)
|
|
1262
|
+
# But lists, numbers, and compatible arrays should get converted.
|
|
1263
|
+
(a, b, c) = test1(5.5, [10.1, 12.7], c=np.linspace(0,1,5))
|
|
1264
|
+
self.assertTrue(a)
|
|
1265
|
+
self.assertTrue(b)
|
|
1266
|
+
self.assertTrue(c)
|
|
1267
|
+
# Tensors should be passed through.
|
|
1268
|
+
(a, b, c) = test1(5.5, [10.1, 12.7], c=torch.linspace(0,1,5))
|
|
1269
|
+
self.assertTrue(a)
|
|
1270
|
+
self.assertTrue(b)
|
|
1271
|
+
self.assertTrue(c)
|
|
1272
|
+
# With the keep_arrays argument set to True, non-tensor arguments get
|
|
1273
|
+
# converted back to arrays if all the arguments were non-tensors.
|
|
1274
|
+
@tensor_args(keep_arrays=True)
|
|
1275
|
+
def test2(a, b, c='test'):
|
|
1276
|
+
return (torch.sqrt(a**2 + b**2), c)
|
|
1277
|
+
(x, y) = test2(10.0, [11.1, 12.2])
|
|
1278
|
+
self.assertFalse(torch.is_tensor(x))
|
|
1279
|
+
self.assertFalse(torch.is_tensor(y))
|
|
1280
|
+
# If there were any tensors, the results should remain as tensors.
|
|
1281
|
+
(x, y) = test2(torch.tensor(10.0), [11.1, 12.2])
|
|
1282
|
+
self.assertTrue(torch.is_tensor(x))
|
|
1283
|
+
self.assertFalse(torch.is_tensor(y)) # Still a string 'test' here.
|
|
1284
|
+
# With named arguments in the tensor_args arguments, only those args
|
|
1285
|
+
# are converted or considered.
|
|
1286
|
+
@tensor_args('a', keep_arrays=True)
|
|
1287
|
+
def test3(a, b, c='test'):
|
|
1288
|
+
return (torch.sqrt(a**2 + b**2), c)
|
|
1289
|
+
with self.assertRaises(TypeError):
|
|
1290
|
+
(x, y) = test3(torch.tensor(10.0), [11.1, 12.2])
|
|
1291
|
+
with self.assertRaises(TypeError):
|
|
1292
|
+
(x, y) = test3(10.0, [11.1, 12.2])
|
|
1293
|
+
(x, y) = test3(10.0, torch.tensor([11.1, 12.2]))
|
|
1294
|
+
# Because it skips parameter b, it doesn't consider this example to be
|
|
1295
|
+
# a case where the tensors should be maintained (keep_arrays indicates
|
|
1296
|
+
# that if any of the converted parameters were tensors it should not
|
|
1297
|
+
# convert results into arrays, but in this case parameter b isn't one
|
|
1298
|
+
# of the named parameters, so all that the conversion algorithm sees is
|
|
1299
|
+
# that parameter a isn't a tensor).
|
|
1300
|
+
self.assertFalse(torch.is_tensor(x))
|
|
1301
|
+
self.assertTrue(isinstance(y, str))
|
|
1302
|
+
(x, y) = test3(torch.tensor(10.0), torch.tensor([11.1, 12.2]))
|
|
1303
|
+
self.assertTrue(torch.is_tensor(x))
|
|
1304
|
+
self.assertTrue(isinstance(y, str))
|
|
1305
|
+
def test_array_args(self):
|
|
1306
|
+
from immlib.util import array_args
|
|
1307
|
+
import numpy as np, torch
|
|
1308
|
+
# Without any arguments, it should just auto-array all the args.
|
|
1309
|
+
@array_args
|
|
1310
|
+
def test1(a, b, c='test'):
|
|
1311
|
+
return (type(a), type(b), type(c))
|
|
1312
|
+
# Even strings and dictionaries get converted into arrays.
|
|
1313
|
+
(a, b, c) = test1(10, {'a':12, 'b':13})
|
|
1314
|
+
self.assertIs(a, np.ndarray)
|
|
1315
|
+
self.assertIs(a, np.ndarray)
|
|
1316
|
+
self.assertIs(a, np.ndarray)
|
|
1317
|
+
# Even tensors should be converted down when possible.
|
|
1318
|
+
(a, b, c) = test1(5.5, [10.1, 12.7], c=torch.linspace(0,1,5))
|
|
1319
|
+
self.assertIs(a, np.ndarray)
|
|
1320
|
+
self.assertIs(b, np.ndarray)
|
|
1321
|
+
self.assertIs(c, np.ndarray)
|
|
1322
|
+
# Tensors should be passed through.
|
|
1323
|
+
(a, b, c) = test1(5.5, [10.1, 12.7], c=torch.linspace(0,1,5))
|
|
1324
|
+
self.assertTrue(a)
|
|
1325
|
+
self.assertTrue(b)
|
|
1326
|
+
self.assertTrue(c)
|
|
1327
|
+
# With named arguments in the tensor_args arguments, only those args
|
|
1328
|
+
# are converted or considered.
|
|
1329
|
+
@array_args('a')
|
|
1330
|
+
def test3(a, b, c='test'):
|
|
1331
|
+
if not isinstance(a, np.ndarray):
|
|
1332
|
+
raise TypeError()
|
|
1333
|
+
if torch.is_tensor(b):
|
|
1334
|
+
raise TypeError()
|
|
1335
|
+
return (np.sqrt(a**2 + b**2), c)
|
|
1336
|
+
with self.assertRaises(TypeError):
|
|
1337
|
+
(x, y) = test3(10.0, torch.tensor([11.1, 12.2]))
|
|
1338
|
+
with self.assertRaises(TypeError):
|
|
1339
|
+
# Fails because you can't run [11.1, 12.2]**2.
|
|
1340
|
+
(x, y) = test3(10.0, [11.1, 12.2])
|
|
1341
|
+
(x, y) = test3(torch.tensor(10.0), np.array([11.1, 12.2]))
|
|
1342
|
+
self.assertFalse(torch.is_tensor(x))
|
|
1343
|
+
self.assertTrue(isinstance(y, str))
|
|
1344
|
+
def test_numeric_args(self):
|
|
1345
|
+
from immlib.util import numeric_args
|
|
1346
|
+
import numpy as np, torch
|
|
1347
|
+
# Without any arguments, it should just auto-array all the args.
|
|
1348
|
+
@numeric_args
|
|
1349
|
+
def test1(a, b, c='test'):
|
|
1350
|
+
return (type(a), type(b), type(c))
|
|
1351
|
+
# Even strings and dictionaries get converted into arrays.
|
|
1352
|
+
(a, b, c) = test1(10, {'a':12, 'b':13})
|
|
1353
|
+
self.assertIs(a, np.ndarray)
|
|
1354
|
+
self.assertIs(a, np.ndarray)
|
|
1355
|
+
self.assertIs(a, np.ndarray)
|
|
1356
|
+
# If there are tensors, then all args should be converted into tensors.
|
|
1357
|
+
(a, b, c) = test1(5.5, [10.1, 12.7], c=torch.linspace(0,1,5))
|
|
1358
|
+
self.assertIs(a, torch.Tensor)
|
|
1359
|
+
self.assertIs(b, torch.Tensor)
|
|
1360
|
+
self.assertIs(c, torch.Tensor)
|
|
1361
|
+
# With named arguments in the tensor_args arguments, only those args
|
|
1362
|
+
# are converted or considered.
|
|
1363
|
+
@numeric_args('a')
|
|
1364
|
+
def test3(a, b, c='test'):
|
|
1365
|
+
return (np.sqrt(a**2 + b**2), c)
|
|
1366
|
+
(x, y) = test3(10.0, torch.tensor([11.1, 12.2]))
|
|
1367
|
+
self.assertIsInstance(x, torch.Tensor)
|
|
1368
|
+
self.assertIsInstance(y, str)
|
|
1369
|
+
with self.assertRaises(TypeError):
|
|
1370
|
+
# Fails because you can't run [11.1, 12.2]**2.
|
|
1371
|
+
(x, y) = test3(10.0, [11.1, 12.2])
|
|
1372
|
+
(x, y) = test3(torch.tensor(10.0), np.array([11.1, 12.2]))
|
|
1373
|
+
self.assertTrue(torch.is_tensor(x))
|
|
1374
|
+
self.assertTrue(isinstance(y, str))
|