immlib 1.0.0.dev2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (45) hide show
  1. immlib/__init__.py +131 -0
  2. immlib/_init.py +108 -0
  3. immlib/_version.py +235 -0
  4. immlib/doc/__init__.py +38 -0
  5. immlib/doc/_core.py +311 -0
  6. immlib/iolib/__init__.py +29 -0
  7. immlib/iolib/_core.py +720 -0
  8. immlib/pathlib/__init__.py +69 -0
  9. immlib/pathlib/_cache.py +152 -0
  10. immlib/pathlib/_core.py +869 -0
  11. immlib/pathlib/_osf.py +538 -0
  12. immlib/test/__init__.py +16 -0
  13. immlib/test/__main__.py +10 -0
  14. immlib/test/doc/__init__.py +6 -0
  15. immlib/test/doc/test_core.py +91 -0
  16. immlib/test/iolib/__init__.py +7 -0
  17. immlib/test/iolib/test_core.py +81 -0
  18. immlib/test/pathlib/__init__.py +11 -0
  19. immlib/test/pathlib/test_core.py +146 -0
  20. immlib/test/pathlib/test_osf.py +54 -0
  21. immlib/test/types/__init__.py +5 -0
  22. immlib/test/types/test_core.py +110 -0
  23. immlib/test/util/__init__.py +11 -0
  24. immlib/test/util/test_core.py +681 -0
  25. immlib/test/util/test_numeric.py +1374 -0
  26. immlib/test/util/test_quantity.py +218 -0
  27. immlib/test/util/test_url.py +51 -0
  28. immlib/test/workflow/__init__.py +9 -0
  29. immlib/test/workflow/test_core.py +418 -0
  30. immlib/test/workflow/test_plantype.py +248 -0
  31. immlib/types/__init__.py +29 -0
  32. immlib/types/_core.py +333 -0
  33. immlib/util/__init__.py +283 -0
  34. immlib/util/_core.py +2524 -0
  35. immlib/util/_numeric.py +2651 -0
  36. immlib/util/_quantity.py +523 -0
  37. immlib/util/_url.py +114 -0
  38. immlib/workflow/__init__.py +48 -0
  39. immlib/workflow/_core.py +1635 -0
  40. immlib/workflow/_plantype.py +334 -0
  41. immlib-1.0.0.dev2.dist-info/METADATA +76 -0
  42. immlib-1.0.0.dev2.dist-info/RECORD +45 -0
  43. immlib-1.0.0.dev2.dist-info/WHEEL +5 -0
  44. immlib-1.0.0.dev2.dist-info/licenses/LICENSE +21 -0
  45. immlib-1.0.0.dev2.dist-info/top_level.txt +1 -0
@@ -0,0 +1,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))