array-split 0.6.5__zip

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.
@@ -0,0 +1,1517 @@
1
+ """
2
+ ========================================
3
+ The :mod:`array_split.split_test` Module
4
+ ========================================
5
+
6
+ .. currentmodule:: array_split.split_test
7
+
8
+ Module defining :mod:`array_split.split` unit-tests.
9
+ Execute as::
10
+
11
+ python -m array_split.split_tests
12
+
13
+
14
+
15
+ Classes
16
+ =======
17
+
18
+ .. autosummary::
19
+ :toctree: generated/
20
+ :template: autosummary/inherits_TestCase_class.rst
21
+
22
+ SplitTest - :obj:`unittest.TestCase` for :mod:`array_split.split` functions.
23
+
24
+
25
+ """
26
+ from __future__ import absolute_import
27
+ import array_split as _array_split
28
+ import numpy as _np
29
+
30
+ from .license import license as _license, copyright as _copyright, version as _version
31
+ from . import unittest as _unittest
32
+ from . import logging as _logging
33
+
34
+ from .split import ShapeSplitter, array_split, shape_split
35
+ from .split import calculate_num_slices_per_axis, shape_factors
36
+ from .split import calculate_tile_shape_for_max_bytes, pad_with_object, convert_halo_to_array_form
37
+ from .split import ARRAY_BOUNDS, NO_BOUNDS
38
+
39
+ __author__ = "Shane J. Latham"
40
+ __license__ = _license()
41
+ __copyright__ = _copyright()
42
+ __version__ = _version()
43
+
44
+
45
+ class SplitTest(_unittest.TestCase):
46
+
47
+ """
48
+ Tests for :mod:`array_split.split` module.
49
+ """
50
+
51
+ #: Class attribute for :obj:`logging.Logger` logging.
52
+ logger = _logging.getLogger(__name__ + ".SplitTest")
53
+
54
+ def test_properties(self):
55
+ """
56
+ Test :attr:`array_split.split.ARRAY_BOUNDS`
57
+ and :attr:`array_split.split.NO_BOUNDS`.
58
+ """
59
+ self.assertNotEqual(None, _array_split.split.ARRAY_BOUNDS)
60
+ self.assertNotEqual(None, _array_split.split.NO_BOUNDS)
61
+
62
+ def test_pad_with_object(self):
63
+ """
64
+ Tests :func:`array_split.split.pad_with_object`.
65
+ """
66
+ lst = pad_with_object([1, 3, 4, ], 5, obj=1)
67
+ self.assertSequenceEqual([1, 3, 4, 1, 1], lst)
68
+
69
+ self.assertRaises(ValueError, pad_with_object, [1, 2, 3, 4], 3)
70
+
71
+ def test_convert_halo_to_array_form(self):
72
+ """
73
+ Tests :func:`array_split.split.convert_halo_to_array_form`.
74
+ """
75
+ self.assertRaises(
76
+ ValueError,
77
+ convert_halo_to_array_form,
78
+ halo=[0, 2, 4],
79
+ ndim=2
80
+ )
81
+ self.assertRaises(
82
+ ValueError,
83
+ convert_halo_to_array_form,
84
+ halo=[0, 2, 4],
85
+ ndim=4
86
+ )
87
+ self.assertTrue(
88
+ _np.all(
89
+ convert_halo_to_array_form(1, 4)
90
+ ==
91
+ [[1, 1], [1, 1], [1, 1], [1, 1]]
92
+ )
93
+ )
94
+
95
+ self.assertTrue(
96
+ _np.all(
97
+ convert_halo_to_array_form([1, 2, 3, 4], 4)
98
+ ==
99
+ [[1, 1], [2, 2], [3, 3], [4, 4]]
100
+ )
101
+ )
102
+
103
+ self.assertTrue(
104
+ _np.all(
105
+ convert_halo_to_array_form([[1, 2], [3, 4], [5, 6], [7, 8]], 4)
106
+ ==
107
+ [[1, 2], [3, 4], [5, 6], [7, 8]]
108
+ )
109
+ )
110
+
111
+ def test_shape_factors(self):
112
+ """
113
+ Tests for :func:`array_split.split.shape_factors`.
114
+ """
115
+ f = shape_factors(4, 2)
116
+ self.assertTrue(_np.all(f == 2))
117
+
118
+ f = shape_factors(4, 1)
119
+ self.assertTrue(_np.all(f == 4))
120
+
121
+ f = shape_factors(5, 2)
122
+ self.assertTrue(_np.all(f == [1, 5]))
123
+
124
+ f = shape_factors(6, 2)
125
+ self.assertTrue(_np.all(f == [2, 3]))
126
+
127
+ f = shape_factors(6, 3)
128
+ self.assertTrue(_np.all(f == [1, 2, 3]))
129
+
130
+ def test_calculate_num_slices_per_axis(self):
131
+ """
132
+ Tests for :func:`array_split.split.calculate_num_slices_per_axis`.
133
+ """
134
+
135
+ self.assertRaises(
136
+ ValueError,
137
+ calculate_num_slices_per_axis,
138
+ [0, 1, 0],
139
+ 15,
140
+ [1, 0, 1024]
141
+ )
142
+
143
+ spa = calculate_num_slices_per_axis([0, ], 5)
144
+ self.assertEqual(1, len(spa))
145
+ self.assertTrue(_np.all(spa == 5))
146
+
147
+ spa = calculate_num_slices_per_axis([2, 0], 4)
148
+ self.assertEqual(2, len(spa))
149
+ self.assertTrue(_np.all(spa == 2))
150
+
151
+ spa = calculate_num_slices_per_axis([0, 2], 4)
152
+ self.assertEqual(2, len(spa))
153
+ self.assertTrue(_np.all(spa == 2))
154
+
155
+ spa = calculate_num_slices_per_axis([0, 0], 4)
156
+ self.assertEqual(2, len(spa))
157
+ self.assertTrue(_np.all(spa == 2))
158
+
159
+ spa = calculate_num_slices_per_axis([0, 0], 16)
160
+ self.assertEqual(2, len(spa))
161
+ self.assertTrue(_np.all(spa == 4))
162
+
163
+ spa = calculate_num_slices_per_axis([0, 0, 0], 8)
164
+ self.assertEqual(3, len(spa))
165
+ self.assertTrue(_np.all(spa == 2))
166
+
167
+ spa = calculate_num_slices_per_axis([0, 1, 0], 8)
168
+ self.assertEqual(3, len(spa))
169
+ self.assertTrue(_np.all(spa == [4, 1, 2]))
170
+
171
+ spa = calculate_num_slices_per_axis([0, 1, 0], 17)
172
+ self.assertEqual(3, len(spa))
173
+ self.assertTrue(_np.all(spa == [17, 1, 1]))
174
+
175
+ spa = calculate_num_slices_per_axis([0, 1, 0], 15, [1, _np.inf, _np.inf])
176
+ self.assertEqual(3, len(spa))
177
+ self.assertTrue(_np.all(spa == [1, 1, 15]))
178
+
179
+ spa = calculate_num_slices_per_axis([0, 1, 0], 16, [1, _np.inf, _np.inf])
180
+ self.assertEqual(3, len(spa))
181
+ self.assertTrue(_np.all(spa == [1, 1, 16]))
182
+
183
+ spa = calculate_num_slices_per_axis([0, 0, 0], 64, [1, 2, _np.inf])
184
+ self.assertEqual(3, len(spa))
185
+ self.assertTrue(_np.all(spa == [1, 2, 32]))
186
+
187
+ spa = calculate_num_slices_per_axis([0, 0, 0], 64, [_np.inf, 1, 2])
188
+ self.assertEqual(3, len(spa))
189
+ self.assertSequenceEqual([32, 1, 2], spa.tolist())
190
+
191
+ spa = calculate_num_slices_per_axis([0, 0, 0], 27, [_np.inf, 2, 2])
192
+ self.assertEqual(3, len(spa))
193
+ self.assertSequenceEqual([27, 1, 1], spa.tolist())
194
+
195
+ def test_calculate_tile_shape_for_max_bytes_1d(self):
196
+ """
197
+ Test case for :func:`array_split.split.calculate_tile_shape_for_max_bytes`,
198
+ where :samp:`array_shape` parameter is 1D, i.e. of the form :samp:`(N,)`.
199
+ """
200
+ self.assertRaises(
201
+ ValueError,
202
+ calculate_tile_shape_for_max_bytes,
203
+ array_shape=(512, 1024, 1024),
204
+ array_itemsize=1,
205
+ max_tile_bytes=2**20,
206
+ sub_tile_shape=(1024, 128, 128)
207
+ )
208
+ tile_shape = \
209
+ calculate_tile_shape_for_max_bytes(
210
+ array_shape=(512,),
211
+ array_itemsize=1,
212
+ max_tile_bytes=1024
213
+ )
214
+ self.assertSequenceEqual((512,), tile_shape)
215
+
216
+ tile_shape = \
217
+ calculate_tile_shape_for_max_bytes(
218
+ array_shape=(512,),
219
+ array_itemsize=1,
220
+ max_tile_bytes=1024,
221
+ sub_tile_shape=[64, ]
222
+ )
223
+ self.assertSequenceEqual((512,), tile_shape)
224
+
225
+ tile_shape = \
226
+ calculate_tile_shape_for_max_bytes(
227
+ array_shape=(512,),
228
+ array_itemsize=1,
229
+ max_tile_bytes=1024,
230
+ sub_tile_shape=[26, ]
231
+ )
232
+ self.assertSequenceEqual((260,), tile_shape)
233
+
234
+ tile_shape = \
235
+ calculate_tile_shape_for_max_bytes(
236
+ array_shape=(512,),
237
+ array_itemsize=1,
238
+ max_tile_bytes=512
239
+ )
240
+ self.assertSequenceEqual((512,), tile_shape)
241
+
242
+ tile_shape = \
243
+ calculate_tile_shape_for_max_bytes(
244
+ array_shape=(512,),
245
+ array_itemsize=2,
246
+ max_tile_bytes=512,
247
+ )
248
+ self.assertSequenceEqual((256,), tile_shape)
249
+
250
+ tile_shape = \
251
+ calculate_tile_shape_for_max_bytes(
252
+ array_shape=(512,),
253
+ array_itemsize=2,
254
+ max_tile_bytes=512,
255
+ sub_tile_shape=[32, ]
256
+ )
257
+ self.assertSequenceEqual((256,), tile_shape)
258
+
259
+ tile_shape = \
260
+ calculate_tile_shape_for_max_bytes(
261
+ array_shape=(512,),
262
+ array_itemsize=2,
263
+ max_tile_bytes=512,
264
+ sub_tile_shape=[60, ]
265
+ )
266
+ self.assertSequenceEqual((180,), tile_shape)
267
+
268
+ tile_shape = \
269
+ calculate_tile_shape_for_max_bytes(
270
+ array_shape=(512,),
271
+ array_itemsize=1,
272
+ max_tile_bytes=512,
273
+ halo=1
274
+ )
275
+ self.assertSequenceEqual((256,), tile_shape)
276
+
277
+ tile_shape = \
278
+ calculate_tile_shape_for_max_bytes(
279
+ array_shape=(512,),
280
+ array_itemsize=1,
281
+ max_tile_bytes=514,
282
+ halo=1
283
+ )
284
+ self.assertSequenceEqual((512,), tile_shape)
285
+
286
+ tile_shape = \
287
+ calculate_tile_shape_for_max_bytes(
288
+ array_shape=(512,),
289
+ array_itemsize=2,
290
+ max_tile_bytes=511
291
+ )
292
+ self.assertSequenceEqual((171,), tile_shape)
293
+
294
+ def test_calculate_tile_shape_for_max_bytes_2d(self):
295
+ """
296
+ Test case for :func:`array_split.split.calculate_tile_shape_for_max_bytes`,
297
+ where :samp:`array_shape` parameter is 2D, i.e. of the form :samp:`(H,W)`.
298
+ """
299
+ tile_shape = \
300
+ calculate_tile_shape_for_max_bytes(
301
+ array_shape=(512, 512),
302
+ array_itemsize=1,
303
+ max_tile_bytes=512 ** 2
304
+ )
305
+ self.assertSequenceEqual((512, 512), tile_shape.tolist())
306
+
307
+ tile_shape = \
308
+ calculate_tile_shape_for_max_bytes(
309
+ array_shape=(512, 512),
310
+ array_itemsize=1,
311
+ max_tile_bytes=512 ** 2 - 1
312
+ )
313
+ self.assertSequenceEqual((256, 512), tile_shape.tolist())
314
+
315
+ tile_shape = \
316
+ calculate_tile_shape_for_max_bytes(
317
+ array_shape=(513, 512),
318
+ array_itemsize=1,
319
+ max_tile_bytes=512 ** 2 - 1
320
+ )
321
+ self.assertSequenceEqual((257, 512), tile_shape.tolist())
322
+
323
+ tile_shape = \
324
+ calculate_tile_shape_for_max_bytes(
325
+ array_shape=(512, 512),
326
+ array_itemsize=1,
327
+ max_tile_bytes=512 ** 2 // 2
328
+ )
329
+ self.assertSequenceEqual((256, 512), tile_shape.tolist())
330
+
331
+ tile_shape = \
332
+ calculate_tile_shape_for_max_bytes(
333
+ array_shape=(512, 512),
334
+ array_itemsize=2,
335
+ max_tile_bytes=512 ** 2 // 2
336
+ )
337
+ self.assertSequenceEqual((128, 512), tile_shape.tolist())
338
+
339
+ tile_shape = \
340
+ calculate_tile_shape_for_max_bytes(
341
+ array_shape=(512, 512),
342
+ array_itemsize=2,
343
+ max_tile_bytes=512 ** 2 // 2,
344
+ sub_tile_shape=(32, 64)
345
+ )
346
+ self.assertSequenceEqual((128, 512), tile_shape.tolist())
347
+
348
+ tile_shape = \
349
+ calculate_tile_shape_for_max_bytes(
350
+ array_shape=(512, 512),
351
+ array_itemsize=1,
352
+ max_tile_bytes=512 ** 2 // 2,
353
+ sub_tile_shape=(30, 64)
354
+ )
355
+ self.assertSequenceEqual((180, 512), tile_shape.tolist())
356
+
357
+ tile_shape = \
358
+ calculate_tile_shape_for_max_bytes(
359
+ array_shape=(512, 512),
360
+ array_itemsize=2,
361
+ max_tile_bytes=512 ** 2 // 2,
362
+ sub_tile_shape=(30, 64)
363
+ )
364
+ self.assertSequenceEqual((120, 512), tile_shape.tolist())
365
+
366
+ tile_shape = \
367
+ calculate_tile_shape_for_max_bytes(
368
+ array_shape=(512, 1024),
369
+ array_itemsize=1,
370
+ max_tile_bytes=512 ** 2,
371
+ sub_tile_shape=(30, 60)
372
+ )
373
+ self.assertSequenceEqual((180, 540), tile_shape.tolist())
374
+
375
+ tile_shape = \
376
+ calculate_tile_shape_for_max_bytes(
377
+ array_shape=(512, 32),
378
+ array_itemsize=1,
379
+ max_tile_bytes=2 * 32 * 32,
380
+ sub_tile_shape=(32, 32)
381
+ )
382
+ self.assertSequenceEqual((64, 32), tile_shape.tolist())
383
+
384
+ tile_shape = \
385
+ calculate_tile_shape_for_max_bytes(
386
+ array_shape=(32, 512),
387
+ array_itemsize=1,
388
+ max_tile_bytes=2 * 32 * 32,
389
+ sub_tile_shape=(32, 32)
390
+ )
391
+ self.assertSequenceEqual((32, 64), tile_shape.tolist())
392
+
393
+ def test_multiple_parameter_groups_error(self):
394
+ """
395
+ Test for case for inconsistent parameter group arguments.
396
+ """
397
+
398
+ splitter = ShapeSplitter((100, ), axis=(3,), max_tile_bytes=30)
399
+ self.assertRaises(
400
+ ValueError,
401
+ splitter.calculate_split
402
+ )
403
+
404
+ splitter = ShapeSplitter((100, ))
405
+ self.assertRaises(
406
+ ValueError,
407
+ splitter.calculate_split
408
+ )
409
+
410
+ self.assertRaises(
411
+ ValueError,
412
+ splitter.set_split_extents_by_indices_per_axis
413
+ )
414
+
415
+ def test_array_split(self):
416
+ """
417
+ Test for case for :func:`array_split.split.array_split`.
418
+ """
419
+ x = _np.arange(9.0)
420
+ self.assertArraySplitEqual(
421
+ _np.array_split(x, 3),
422
+ array_split(x, 3)
423
+ )
424
+ self.assertArraySplitEqual(
425
+ _np.array_split(x, 4),
426
+ array_split(x, 4)
427
+ )
428
+ idx = [2, 3, 5, ]
429
+ self.assertArraySplitEqual(
430
+ _np.array_split(x, idx),
431
+ array_split(x, idx)
432
+ )
433
+
434
+ x = _np.arange(32)
435
+ x = x.reshape((4, 8))
436
+ self.logger.info("_np.array_split(x, 3, axis=0) = \n%s", _np.array_split(x, 3, axis=0))
437
+ self.logger.info(
438
+ "array_split.split.array_split(x, 3, axis=0) = \n%s", array_split(x, 3, axis=0)
439
+ )
440
+ self.assertArraySplitEqual(
441
+ _np.array_split(x, 3, axis=0),
442
+ array_split(x, 3, axis=0)
443
+ )
444
+
445
+ self.logger.info("_np.array_split(x, 3, axis=1) = \n%s", _np.array_split(x, 3, axis=1))
446
+ self.logger.info(
447
+ "array_split.split.array_split(x, 3, axis=1) = \n%s", array_split(x, 3, axis=1)
448
+ )
449
+ self.assertArraySplitEqual(
450
+ _np.array_split(x, 3, axis=1),
451
+ array_split(x, 3, axis=1)
452
+ )
453
+
454
+ self.logger.info("_np.array_split(x, 8, axis=0) = \n%s", _np.array_split(x, 8, axis=0))
455
+ self.assertArraySplitEqual(
456
+ _np.array_split(x, 8, axis=0),
457
+ array_split(x, 8, axis=0)
458
+ )
459
+
460
+ x = _np.arange(0, 64)
461
+ x = x.reshape((4, 16))
462
+ self.assertArraySplitEqual(
463
+ _np.array_split(x, [3, 8, 12], axis=1),
464
+ array_split(x, [3, 8, 12], axis=1)
465
+ )
466
+
467
+ x = _np.arange(0, 512, dtype="int16")
468
+ self.assertArraySplitEqual(
469
+ [_np.arange(0, 256), _np.arange(256, 512)],
470
+ array_split(x, max_tile_bytes=512)
471
+ )
472
+
473
+ def test_split_by_per_axis_indices(self):
474
+ """
475
+ Test for case for splitting by specified
476
+ indices. For example::
477
+
478
+ ShapeSplitter(array_shape=(10, 4), indices_or_sections=[[2, 6, 8], ]).calculate_split()
479
+
480
+ """
481
+ splitter = ShapeSplitter((10, 4), [[2, 6, 8], [1, ], [8, 4]])
482
+ self.assertRaises(
483
+ ValueError,
484
+ splitter.calculate_split
485
+ )
486
+ splitter = ShapeSplitter((10, 4), [[2, 6, 8], ])
487
+ split = splitter.calculate_split()
488
+ self.logger.info("split.shape = %s", split.shape)
489
+ self.logger.info("split =\n%s", split)
490
+ self.assertTrue(_np.all(_np.array(split.shape) == [4, 1]))
491
+ self.assertEqual(slice(0, 2), split[0, 0][0]) # axis 0 slice
492
+ self.assertEqual(slice(2, 6), split[1, 0][0]) # axis 0 slice
493
+ self.assertEqual(slice(6, 8), split[2, 0][0]) # axis 0 slice
494
+ self.assertEqual(slice(8, 10), split[3, 0][0]) # axis 0 slice
495
+ self.assertEqual(slice(0, 4), split[0, 0][1]) # axis 1 slice
496
+ self.assertEqual(slice(0, 4), split[1, 0][1]) # axis 1 slice
497
+ self.assertEqual(slice(0, 4), split[2, 0][1]) # axis 1 slice
498
+ self.assertEqual(slice(0, 4), split[3, 0][1]) # axis 1 slice
499
+
500
+ split1 = splitter.calculate_split_by_indices_per_axis()
501
+ self.assertTrue(_np.all(split == split1))
502
+
503
+ splitter = ShapeSplitter((10, 13), [None, [2, 5, 8], ])
504
+ split = splitter.calculate_split()
505
+ self.logger.info("split.shape = %s", split.shape)
506
+ self.logger.info("split =\n%s", split)
507
+ self.assertTrue(_np.all(_np.array(split.shape) == [1, 4]))
508
+ self.assertEqual(slice(0, 10), split[0, 0][0]) # axis 0 slice
509
+ self.assertEqual(slice(0, 10), split[0, 1][0]) # axis 0 slice
510
+ self.assertEqual(slice(0, 10), split[0, 2][0]) # axis 0 slice
511
+ self.assertEqual(slice(0, 10), split[0, 3][0]) # axis 0 slice
512
+ self.assertEqual(slice(0, 2), split[0, 0][1]) # axis 1 slice
513
+ self.assertEqual(slice(2, 5), split[0, 1][1]) # axis 1 slice
514
+ self.assertEqual(slice(5, 8), split[0, 2][1]) # axis 1 slice
515
+ self.assertEqual(slice(8, 13), split[0, 3][1]) # axis 1 slice
516
+
517
+ splitter = ShapeSplitter((10, 4), [[2, 6], [2, ]])
518
+ split = splitter.calculate_split()
519
+ self.logger.info("split.shape = %s", split.shape)
520
+ self.logger.info("split =\n%s", split)
521
+ self.assertTrue(_np.all(_np.array(split.shape) == [3, 2]))
522
+ self.assertEqual(slice(0, 2), split[0, 0][0]) # axis 0 slice
523
+ self.assertEqual(slice(2, 6), split[1, 0][0]) # axis 0 slice
524
+ self.assertEqual(slice(6, 10), split[2, 0][0]) # axis 0 slice
525
+ self.assertEqual(slice(0, 2), split[0, 1][0]) # axis 0 slice
526
+ self.assertEqual(slice(2, 6), split[1, 1][0]) # axis 0 slice
527
+ self.assertEqual(slice(6, 10), split[2, 1][0]) # axis 0 slice
528
+
529
+ self.assertEqual(slice(0, 2), split[0, 0][1]) # axis 1 slice
530
+ self.assertEqual(slice(0, 2), split[1, 0][1]) # axis 1 slice
531
+ self.assertEqual(slice(0, 2), split[2, 0][1]) # axis 1 slice
532
+ self.assertEqual(slice(2, 4), split[0, 1][1]) # axis 1 slice
533
+ self.assertEqual(slice(2, 4), split[1, 1][1]) # axis 1 slice
534
+ self.assertEqual(slice(2, 4), split[2, 1][1]) # axis 1 slice
535
+
536
+ splitter = ShapeSplitter((10,), [[2, 6, 8], ])
537
+ split = splitter.calculate_split()
538
+ self.logger.info("split.shape = %s", split.shape)
539
+ self.logger.info("split =\n%s", split)
540
+ self.assertTrue(_np.all(_np.array(split.shape) == [4, ]))
541
+ self.assertEqual(slice(0, 2), split[0][0]) # axis 0 slice
542
+ self.assertEqual(slice(2, 6), split[1][0]) # axis 0 slice
543
+ self.assertEqual(slice(6, 8), split[2][0]) # axis 0 slice
544
+ self.assertEqual(slice(8, 10), split[3][0]) # axis 0 slice
545
+
546
+ def test_split_by_num_slices_1d(self):
547
+ """
548
+ Test for case for splitting by number of
549
+ slice elements. For example::
550
+
551
+ ShapeSplitter(array_shape=(10, ), indices_or_sections=3).calculate_split()
552
+ ShapeSplitter(array_shape=(10, ), axis=[2, ]).calculate_split()
553
+
554
+ """
555
+
556
+ self.assertRaises(
557
+ ValueError,
558
+ ShapeSplitter,
559
+ (10,),
560
+ axis=[3, 4]
561
+ )
562
+
563
+ splitter = ShapeSplitter((10,), axis=[0, ])
564
+ self.assertRaises(
565
+ ValueError,
566
+ splitter.set_split_extents_by_split_size
567
+ )
568
+
569
+ splitter = ShapeSplitter((10,), axis=[2, ])
570
+ splitter.split_num_slices_per_axis = [2, 2]
571
+ self.assertRaises(
572
+ ValueError,
573
+ splitter.check_consistent_parameter_dimensions
574
+ )
575
+
576
+ splitter = ShapeSplitter((10,), 3)
577
+ split = splitter.calculate_split()
578
+ self.logger.info("split.shape = %s", split.shape)
579
+ self.logger.info("split =\n%s", split)
580
+ self.assertTrue(_np.all(_np.array(split.shape) == [3, ]))
581
+ self.assertEqual(slice(0, 4), split[0][0]) # axis 0 slice
582
+ self.assertEqual(slice(4, 7), split[1][0]) # axis 0 slice
583
+ self.assertEqual(slice(7, 10), split[2][0]) # axis 0 slice
584
+
585
+ split1 = splitter.calculate_split_by_split_size()
586
+ self.assertTrue(_np.all(split == split1))
587
+
588
+ splitter = ShapeSplitter((10,), axis=[3, ])
589
+ split = splitter.calculate_split()
590
+ self.logger.info("split.shape = %s", split.shape)
591
+ self.logger.info("split =\n%s", split)
592
+ self.assertTrue(_np.all(_np.array(split.shape) == [3, ]))
593
+ self.assertEqual(slice(0, 4), split[0][0]) # axis 0 slice
594
+ self.assertEqual(slice(4, 7), split[1][0]) # axis 0 slice
595
+ self.assertEqual(slice(7, 10), split[2][0]) # axis 0 slice
596
+
597
+ splitter = ShapeSplitter((10,), 3, axis=[3, ])
598
+ split = splitter.calculate_split()
599
+ self.logger.info("split.shape = %s", split.shape)
600
+ self.logger.info("split =\n%s", split)
601
+ self.assertTrue(_np.all(_np.array(split.shape) == [3, ]))
602
+ self.assertEqual(slice(0, 4), split[0][0]) # axis 0 slice
603
+ self.assertEqual(slice(4, 7), split[1][0]) # axis 0 slice
604
+ self.assertEqual(slice(7, 10), split[2][0]) # axis 0 slice
605
+
606
+ splitter = ShapeSplitter((10,), 3, axis=[0, ])
607
+ split = splitter.calculate_split()
608
+ self.logger.info("split.shape = %s", split.shape)
609
+ self.logger.info("split =\n%s", split)
610
+ self.assertTrue(_np.all(_np.array(split.shape) == [3, ]))
611
+ self.assertEqual(slice(0, 4), split[0][0]) # axis 0 slice
612
+ self.assertEqual(slice(4, 7), split[1][0]) # axis 0 slice
613
+ self.assertEqual(slice(7, 10), split[2][0]) # axis 0 slice
614
+
615
+ splitter = ShapeSplitter((10,), 2, axis=[0, ])
616
+ split = splitter.calculate_split()
617
+ self.logger.info("split.shape = %s", split.shape)
618
+ self.logger.info("split =\n%s", split)
619
+ self.assertTrue(_np.all(_np.array(split.shape) == [2, ]))
620
+ self.assertEqual(slice(0, 5), split[0][0]) # axis 0 slice
621
+ self.assertEqual(slice(5, 10), split[1][0]) # axis 0 slice
622
+
623
+ def test_split_by_num_slices_2d_non_0_axis_elems(self):
624
+ """
625
+ Test for case for splitting by number of
626
+ slice elements. For example::
627
+
628
+ ShapeSplitter(array_shape=(10, 13), axis=[2, 3]).calculate_split()
629
+
630
+ """
631
+
632
+ splitter = ShapeSplitter((10, 13), axis=[2, 2])
633
+ split = splitter.calculate_split()
634
+ self.logger.info("split.shape = %s", split.shape)
635
+ self.logger.info("split =\n%s", split)
636
+ self.assertTrue(_np.all(_np.array(split.shape) == [2, 2]))
637
+ self.assertEqual(slice(0, 5), split[0, 0][0]) # axis 0 slice
638
+ self.assertEqual(slice(0, 5), split[0, 1][0]) # axis 0 slice
639
+ self.assertEqual(slice(5, 10), split[1, 0][0]) # axis 0 slice
640
+ self.assertEqual(slice(5, 10), split[1, 1][0]) # axis 0 slice
641
+ self.assertEqual(slice(0, 7), split[0, 0][1]) # axis 1 slice
642
+ self.assertEqual(slice(7, 13), split[0, 1][1]) # axis 1 slice
643
+ self.assertEqual(slice(0, 7), split[1, 0][1]) # axis 1 slice
644
+ self.assertEqual(slice(7, 13), split[1, 1][1]) # axis 1 slice
645
+
646
+ splitter = ShapeSplitter((10, 13), 4, axis=[2, 2])
647
+ split = splitter.calculate_split()
648
+ self.logger.info("split.shape = %s", split.shape)
649
+ self.logger.info("split =\n%s", split)
650
+ self.assertTrue(_np.all(_np.array(split.shape) == [2, 2]))
651
+ self.assertEqual(slice(0, 5), split[0, 0][0]) # axis 0 slice
652
+ self.assertEqual(slice(0, 5), split[0, 1][0]) # axis 0 slice
653
+ self.assertEqual(slice(5, 10), split[1, 0][0]) # axis 0 slice
654
+ self.assertEqual(slice(5, 10), split[1, 1][0]) # axis 0 slice
655
+ self.assertEqual(slice(0, 7), split[0, 0][1]) # axis 1 slice
656
+ self.assertEqual(slice(7, 13), split[0, 1][1]) # axis 1 slice
657
+ self.assertEqual(slice(0, 7), split[1, 0][1]) # axis 1 slice
658
+ self.assertEqual(slice(7, 13), split[1, 1][1]) # axis 1 slice
659
+
660
+ def test_split_by_num_slices_2d_0_axis_elems(self):
661
+ """
662
+ Test for case for splitting by number of
663
+ slice elements. For example::
664
+
665
+ ShapeSplitter(array_shape=(10, 13), 4, axis=[2, 0]).calculate_split()
666
+
667
+ """
668
+
669
+ splitter = ShapeSplitter((10, 13), 4, axis=[1, 0])
670
+ split = splitter.calculate_split()
671
+ self.logger.info("split.shape = %s", split.shape)
672
+ self.logger.info("split =\n%s", split)
673
+ self.assertTrue(_np.all(_np.array(split.shape) == [1, 4]))
674
+ self.assertEqual(slice(0, 10), split[0, 0][0]) # axis 0 slice
675
+ self.assertEqual(slice(0, 10), split[0, 1][0]) # axis 0 slice
676
+ self.assertEqual(slice(0, 10), split[0, 2][0]) # axis 0 slice
677
+ self.assertEqual(slice(0, 10), split[0, 3][0]) # axis 0 slice
678
+ self.assertEqual(slice(0, 4), split[0, 0][1]) # axis 1 slice
679
+ self.assertEqual(slice(4, 7), split[0, 1][1]) # axis 1 slice
680
+ self.assertEqual(slice(7, 10), split[0, 2][1]) # axis 1 slice
681
+ self.assertEqual(slice(10, 13), split[0, 3][1]) # axis 1 slice
682
+
683
+ splitter = ShapeSplitter((10, 13), 4, axis=[0, 2])
684
+ split = splitter.calculate_split()
685
+ self.logger.info("split.shape = %s", split.shape)
686
+ self.logger.info("split =\n%s", split)
687
+ self.assertTrue(_np.all(_np.array(split.shape) == [2, 2]))
688
+ self.assertEqual(slice(0, 5), split[0, 0][0]) # axis 0 slice
689
+ self.assertEqual(slice(0, 5), split[0, 1][0]) # axis 0 slice
690
+ self.assertEqual(slice(5, 10), split[1, 0][0]) # axis 0 slice
691
+ self.assertEqual(slice(5, 10), split[1, 1][0]) # axis 0 slice
692
+ self.assertEqual(slice(0, 7), split[0, 0][1]) # axis 1 slice
693
+ self.assertEqual(slice(7, 13), split[0, 1][1]) # axis 1 slice
694
+ self.assertEqual(slice(0, 7), split[1, 0][1]) # axis 1 slice
695
+ self.assertEqual(slice(7, 13), split[1, 1][1]) # axis 1 slice
696
+
697
+ splitter = ShapeSplitter((10, 13), 4, axis=[2, 0])
698
+ split = splitter.calculate_split()
699
+ self.logger.info("split.shape = %s", split.shape)
700
+ self.logger.info("split =\n%s", split)
701
+ self.assertTrue(_np.all(_np.array(split.shape) == [2, 2]))
702
+ self.assertEqual(slice(0, 5), split[0, 0][0]) # axis 0 slice
703
+ self.assertEqual(slice(0, 5), split[0, 1][0]) # axis 0 slice
704
+ self.assertEqual(slice(5, 10), split[1, 0][0]) # axis 0 slice
705
+ self.assertEqual(slice(5, 10), split[1, 1][0]) # axis 0 slice
706
+ self.assertEqual(slice(0, 7), split[0, 0][1]) # axis 1 slice
707
+ self.assertEqual(slice(7, 13), split[0, 1][1]) # axis 1 slice
708
+ self.assertEqual(slice(0, 7), split[1, 0][1]) # axis 1 slice
709
+ self.assertEqual(slice(7, 13), split[1, 1][1]) # axis 1 slice
710
+
711
+ splitter = ShapeSplitter((10, 13), 4, axis=[0, 0])
712
+ split = splitter.calculate_split()
713
+ self.logger.info("split.shape = %s", split.shape)
714
+ self.logger.info("split =\n%s", split)
715
+ self.assertTrue(_np.all(_np.array(split.shape) == [2, 2]))
716
+ self.assertEqual(slice(0, 5), split[0, 0][0]) # axis 0 slice
717
+ self.assertEqual(slice(0, 5), split[0, 1][0]) # axis 0 slice
718
+ self.assertEqual(slice(5, 10), split[1, 0][0]) # axis 0 slice
719
+ self.assertEqual(slice(5, 10), split[1, 1][0]) # axis 0 slice
720
+ self.assertEqual(slice(0, 7), split[0, 0][1]) # axis 1 slice
721
+ self.assertEqual(slice(7, 13), split[0, 1][1]) # axis 1 slice
722
+ self.assertEqual(slice(0, 7), split[1, 0][1]) # axis 1 slice
723
+ self.assertEqual(slice(7, 13), split[1, 1][1]) # axis 1 slice
724
+
725
+ def test_calculate_split_by_tile_shape_1d(self):
726
+ """
727
+ Test for case for splitting by explicit tile shape. For example::
728
+
729
+ ShapeSplitter(array_shape=(10, ), tile_shape=(3, )).calculate_split()
730
+
731
+ """
732
+
733
+ splitter = ShapeSplitter((10, ), tile_shape=(3, 3))
734
+ self.assertRaises(
735
+ ValueError,
736
+ splitter.calculate_split
737
+ )
738
+
739
+ splitter = ShapeSplitter((100, ), max_tile_bytes=25, sub_tile_shape=(5, 5))
740
+ self.assertRaises(
741
+ ValueError,
742
+ splitter.calculate_split
743
+ )
744
+
745
+ splitter = \
746
+ ShapeSplitter((100, ), max_tile_bytes=25, sub_tile_shape=(5, ), max_tile_shape=(10, 10))
747
+ self.assertRaises(
748
+ ValueError,
749
+ splitter.calculate_split
750
+ )
751
+
752
+ splitter = \
753
+ ShapeSplitter((100, ), max_tile_shape=(10, ))
754
+ self.assertRaises(
755
+ ValueError,
756
+ splitter.calculate_split
757
+ )
758
+
759
+ splitter = \
760
+ ShapeSplitter((100, ), sub_tile_shape=(10, ))
761
+ self.assertRaises(
762
+ ValueError,
763
+ splitter.calculate_split
764
+ )
765
+
766
+ splitter = ShapeSplitter((10, ), tile_shape=(3,))
767
+ split = splitter.calculate_split()
768
+ self.logger.info("split.shape = %s", split.shape)
769
+ self.logger.info("split =\n%s", split)
770
+ self.assertSequenceEqual((4,), split.shape)
771
+ self.assertSequenceEqual(
772
+ [(slice(0, 3),), (slice(3, 6),), (slice(6, 9),), (slice(9, 10),)],
773
+ split.tolist()
774
+ )
775
+
776
+ split1 = splitter.calculate_split_by_tile_shape()
777
+ self.assertTrue(_np.all(split == split1))
778
+
779
+ splitter = ShapeSplitter((10, ), tile_shape=(4,))
780
+ split = splitter.calculate_split()
781
+ self.logger.info("split.shape = %s", split.shape)
782
+ self.logger.info("split =\n%s", split)
783
+ self.assertSequenceEqual((3,), split.shape)
784
+ self.assertSequenceEqual(
785
+ [(slice(0, 4),), (slice(4, 8),), (slice(8, 10),)],
786
+ split.tolist()
787
+ )
788
+
789
+ splitter = ShapeSplitter((10, ), tile_shape=(5,))
790
+ split = splitter.calculate_split()
791
+ self.logger.info("split.shape = %s", split.shape)
792
+ self.logger.info("split =\n%s", split)
793
+ self.assertSequenceEqual((2,), split.shape)
794
+ self.assertSequenceEqual(
795
+ [(slice(0, 5),), (slice(5, 10),)],
796
+ split.tolist()
797
+ )
798
+
799
+ splitter = ShapeSplitter((10, ), tile_shape=(10,))
800
+ split = splitter.calculate_split()
801
+ self.logger.info("split.shape = %s", split.shape)
802
+ self.logger.info("split =\n%s", split)
803
+ self.assertSequenceEqual((1,), split.shape)
804
+ self.assertSequenceEqual(
805
+ [(slice(0, 10),)],
806
+ split.tolist()
807
+ )
808
+
809
+ splitter = ShapeSplitter((10, ), tile_shape=(11,))
810
+ split = splitter.calculate_split()
811
+ self.logger.info("split.shape = %s", split.shape)
812
+ self.logger.info("split =\n%s", split)
813
+ self.assertSequenceEqual((1,), split.shape)
814
+ self.assertSequenceEqual(
815
+ [(slice(0, 10),)],
816
+ split.tolist()
817
+ )
818
+
819
+ def test_calculate_split_by_tile_shape_2d(self):
820
+ """
821
+ Test for case for splitting by explicit tile shape. For example::
822
+
823
+ ShapeSplitter(array_shape=(10, 17), tile_shape=(3, 8)).calculate_split()
824
+
825
+ """
826
+
827
+ splitter = ShapeSplitter((10, 17), tile_shape=(3, 8))
828
+ split = splitter.calculate_split()
829
+ self.logger.info("split.shape = %s", split.shape)
830
+ self.logger.info("split =\n%s", split)
831
+ self.assertSequenceEqual((4, 3), split.shape)
832
+ self.assertSequenceEqual(
833
+ shape_split(splitter.array_shape, [[3, 6, 9], [8, 16]]).flatten().tolist(),
834
+ split.flatten().tolist()
835
+ )
836
+
837
+ splitter = ShapeSplitter((10, 17), tile_shape=(2, 9))
838
+ split = splitter.calculate_split()
839
+ self.logger.info("split.shape = %s", split.shape)
840
+ self.logger.info("split =\n%s", split)
841
+ self.assertSequenceEqual((5, 2), split.shape)
842
+ self.assertSequenceEqual(
843
+ shape_split(splitter.array_shape, [[2, 4, 6, 8], [9, ]]).flatten().tolist(),
844
+ split.flatten().tolist()
845
+ )
846
+
847
+ def test_calculate_split_by_tile_max_bytes_1d(self):
848
+ """
849
+ Test for case for splitting with maximum number of tile bytes constraint.
850
+ For example::
851
+
852
+ ShapeSplitter((512, ), max_tile_bytes=256, array_itemsize=1).calculate_split()
853
+
854
+ """
855
+
856
+ splitter = ShapeSplitter((512, ), max_tile_bytes=256, array_itemsize=1)
857
+ split = splitter.calculate_split()
858
+ self.logger.info("split.shape = %s", split.shape)
859
+ self.logger.info("split =\n%s", split)
860
+ self.assertSequenceEqual((2,), split.shape)
861
+ self.assertSequenceEqual(
862
+ [(slice(0, 256),), (slice(256, 512),)],
863
+ split.tolist()
864
+ )
865
+
866
+ splitter = ShapeSplitter((512, ), max_tile_bytes=256, array_itemsize=2)
867
+ split = splitter.calculate_split()
868
+ self.logger.info("split.shape = %s", split.shape)
869
+ self.logger.info("split =\n%s", split)
870
+ self.assertSequenceEqual((4,), split.shape)
871
+ self.assertSequenceEqual(
872
+ [(slice(0, 128),), (slice(128, 256),), (slice(256, 384),), (slice(384, 512),)],
873
+ split.tolist()
874
+ )
875
+
876
+ split1 = splitter.calculate_split_by_tile_max_bytes()
877
+ self.assertTrue(_np.all(split == split1))
878
+
879
+ splitter = ShapeSplitter((512, ), max_tile_bytes=511, array_itemsize=2)
880
+ split = splitter.calculate_split()
881
+ self.logger.info("split.shape = %s", split.shape)
882
+ self.logger.info("split =\n%s", split)
883
+ self.assertSequenceEqual((3,), split.shape)
884
+ self.assertSequenceEqual(
885
+ [(slice(0, 171),), (slice(171, 342),), (slice(342, 512),)],
886
+ split.tolist()
887
+ )
888
+
889
+ splitter = ShapeSplitter((512, ), max_tile_bytes=256, array_itemsize=1, halo=1)
890
+ split = splitter.calculate_split()
891
+ self.logger.info("split.shape = %s", split.shape)
892
+ self.logger.info("split =\n%s", split)
893
+ self.assertSequenceEqual((3,), split.shape)
894
+ self.assertSequenceEqual(
895
+ [(slice(0, 172),), (slice(170, 343),), (slice(341, 512),)],
896
+ split.tolist()
897
+ )
898
+
899
+ splitter = \
900
+ ShapeSplitter((512, ), max_tile_bytes=256, array_itemsize=1, max_tile_shape=(128,))
901
+ split = splitter.calculate_split()
902
+ self.logger.info("split.shape = %s", split.shape)
903
+ self.logger.info("split =\n%s", split)
904
+ self.assertSequenceEqual((4,), split.shape)
905
+ self.assertSequenceEqual(
906
+ [(slice(0, 128),), (slice(128, 256),), (slice(256, 384),), (slice(384, 512),)],
907
+ split.tolist()
908
+ )
909
+
910
+ splitter = \
911
+ ShapeSplitter((512, ), max_tile_bytes=256, array_itemsize=1, sub_tile_shape=(130,))
912
+ split = splitter.calculate_split()
913
+ self.logger.info("split.shape = %s", split.shape)
914
+ self.logger.info("split =\n%s", split)
915
+ self.assertSequenceEqual((4,), split.shape)
916
+ self.assertSequenceEqual(
917
+ [(slice(0, 130),), (slice(130, 260),), (slice(260, 390),), (slice(390, 512),)],
918
+ split.tolist()
919
+ )
920
+
921
+ def test_calculate_split_with_array_start_1d(self):
922
+ """
923
+ Test for case for splitting with explicit array start multi-index. For example::
924
+
925
+ shape_split((10,), 2, array_start=(32,))
926
+
927
+ """
928
+
929
+ self.assertRaises(
930
+ ValueError,
931
+ shape_split,
932
+ (10,),
933
+ 2,
934
+ array_start=(2, 3)
935
+ )
936
+ split = shape_split((10,), 2, array_start=(0,))
937
+ self.assertSequenceEqual(
938
+ [(slice(0, 5),), (slice(5, 10),)],
939
+ split.tolist()
940
+ )
941
+
942
+ split = shape_split((10,), 2, array_start=(32,))
943
+ self.assertSequenceEqual(
944
+ [(slice(32, 37),), (slice(37, 42),)],
945
+ split.tolist()
946
+ )
947
+
948
+ def test_calculate_split_with_array_start_2d(self):
949
+ """
950
+ Test for case for splitting with explicit array start multi-index. For example::
951
+
952
+ shape_split((10, 12), axis=(2, 2), array_start=(32, 16))
953
+
954
+ """
955
+
956
+ split = shape_split((10, 12), axis=(2, 2), array_start=(0, 0))
957
+ self.assertSequenceEqual(
958
+ [
959
+ [(slice(0, 5), slice(0, 6)), (slice(0, 5), slice(6, 12))],
960
+ [(slice(5, 10), slice(0, 6)), (slice(5, 10), slice(6, 12))]
961
+ ],
962
+ split.tolist()
963
+ )
964
+
965
+ split = shape_split((10, 12), axis=(2, 2), array_start=(32, 16))
966
+ self.assertSequenceEqual(
967
+ [
968
+ [(slice(32, 37), slice(16, 22)), (slice(32, 37), slice(22, 28))],
969
+ [(slice(37, 42), slice(16, 22)), (slice(37, 42), slice(22, 28))]
970
+ ],
971
+ split.tolist()
972
+ )
973
+
974
+ def test_calculate_split_with_halo_1d(self):
975
+ """
976
+ Test for case for splitting with explicit halo. For example::
977
+
978
+ shape_split((10,), 3, halo=[(1, 2), ])
979
+
980
+ """
981
+
982
+ split = shape_split((10,), 3, halo=(0,))
983
+ self.assertSequenceEqual(
984
+ [(slice(0, 4),), (slice(4, 7),), (slice(7, 10),)],
985
+ split.tolist()
986
+ )
987
+
988
+ split = shape_split((10,), 3, halo=(0, 0))
989
+ self.assertSequenceEqual(
990
+ [(slice(0, 4),), (slice(4, 7),), (slice(7, 10),)],
991
+ split.tolist()
992
+ )
993
+
994
+ split = shape_split((10,), 3, halo=(1, 0))
995
+ self.assertSequenceEqual(
996
+ [(slice(0, 4),), (slice(3, 7),), (slice(6, 10),)],
997
+ split.tolist()
998
+ )
999
+
1000
+ split = shape_split((10,), 3, halo=(0, 1))
1001
+ self.assertSequenceEqual(
1002
+ [(slice(0, 5),), (slice(4, 8),), (slice(7, 10),)],
1003
+ split.tolist()
1004
+ )
1005
+
1006
+ split = shape_split((10,), 3, halo=(1, 1))
1007
+ self.assertSequenceEqual(
1008
+ [(slice(0, 5),), (slice(3, 8),), (slice(6, 10),)],
1009
+ split.tolist()
1010
+ )
1011
+
1012
+ split = shape_split((10,), 3, halo=[(1, 2), ])
1013
+ self.assertSequenceEqual(
1014
+ [(slice(0, 6),), (slice(3, 9),), (slice(6, 10),)],
1015
+ split.tolist()
1016
+ )
1017
+
1018
+ split = shape_split((10,), 3, halo=1)
1019
+ self.assertSequenceEqual(
1020
+ [(slice(0, 5),), (slice(3, 8),), (slice(6, 10),)],
1021
+ split.tolist()
1022
+ )
1023
+
1024
+ split = shape_split((10,), 3, halo=1, tile_bounds_policy=ARRAY_BOUNDS)
1025
+ self.assertSequenceEqual(
1026
+ [(slice(0, 5),), (slice(3, 8),), (slice(6, 10),)],
1027
+ split.tolist()
1028
+ )
1029
+
1030
+ split = shape_split((10,), 3, halo=1, tile_bounds_policy=None)
1031
+ self.assertSequenceEqual(
1032
+ [(slice(0, 5),), (slice(3, 8),), (slice(6, 10),)],
1033
+ split.tolist()
1034
+ )
1035
+
1036
+ split = shape_split((10,), 3, halo=1, tile_bounds_policy=NO_BOUNDS)
1037
+ self.assertSequenceEqual(
1038
+ [(slice(-1, 5),), (slice(3, 8),), (slice(6, 11),)],
1039
+ split.tolist()
1040
+ )
1041
+
1042
+ split = shape_split((10,), 3, halo=((2, 3),), tile_bounds_policy=NO_BOUNDS)
1043
+ self.assertSequenceEqual(
1044
+ [(slice(-2, 7),), (slice(2, 10),), (slice(5, 13),)],
1045
+ split.tolist()
1046
+ )
1047
+
1048
+ split = shape_split((10,), 3, halo=(2, 3), tile_bounds_policy=NO_BOUNDS)
1049
+ self.assertSequenceEqual(
1050
+ [(slice(-2, 7),), (slice(2, 10),), (slice(5, 13),)],
1051
+ split.tolist()
1052
+ )
1053
+
1054
+ def test_calculate_split_with_halo_2d(self):
1055
+ """
1056
+ Test for case for splitting with explicit halo. For example::
1057
+
1058
+ shape_split(
1059
+ (15, 13),
1060
+ axis=[3, 3],
1061
+ halo=[[1, 2], [2, 3]],
1062
+ tile_bounds_policy=ARRAY_BOUNDS
1063
+ )
1064
+
1065
+ """
1066
+
1067
+ self.assertRaises(
1068
+ ValueError,
1069
+ shape_split,
1070
+ (15, 13),
1071
+ axis=[3, 3],
1072
+ halo=[0, 1, 2]
1073
+ )
1074
+
1075
+ self.assertRaises(
1076
+ ValueError,
1077
+ shape_split,
1078
+ (15, 13),
1079
+ axis=[3, 3],
1080
+ halo=[0, 1],
1081
+ tile_bounds_policy="bogus"
1082
+ )
1083
+
1084
+ split = shape_split((15, 13), axis=[3, 3], halo=0)
1085
+ self.assertSequenceEqual(
1086
+ [
1087
+ [
1088
+ (slice(0, 5), slice(0, 5)),
1089
+ (slice(0, 5), slice(5, 9)),
1090
+ (slice(0, 5), slice(9, 13))
1091
+ ],
1092
+ [
1093
+ (slice(5, 10), slice(0, 5)),
1094
+ (slice(5, 10), slice(5, 9)),
1095
+ (slice(5, 10), slice(9, 13))
1096
+ ],
1097
+ [
1098
+ (slice(10, 15), slice(0, 5)),
1099
+ (slice(10, 15), slice(5, 9)),
1100
+ (slice(10, 15), slice(9, 13))
1101
+ ],
1102
+ ],
1103
+ split.tolist()
1104
+ )
1105
+
1106
+ split = shape_split((15, 13), axis=[3, 3], halo=(0, 0))
1107
+ self.assertSequenceEqual(
1108
+ [
1109
+ [
1110
+ (slice(0, 5), slice(0, 5)),
1111
+ (slice(0, 5), slice(5, 9)),
1112
+ (slice(0, 5), slice(9, 13))
1113
+ ],
1114
+ [
1115
+ (slice(5, 10), slice(0, 5)),
1116
+ (slice(5, 10), slice(5, 9)),
1117
+ (slice(5, 10), slice(9, 13))
1118
+ ],
1119
+ [
1120
+ (slice(10, 15), slice(0, 5)),
1121
+ (slice(10, 15), slice(5, 9)),
1122
+ (slice(10, 15), slice(9, 13))
1123
+ ],
1124
+ ],
1125
+ split.tolist()
1126
+ )
1127
+
1128
+ split = shape_split((15, 13), axis=[3, 3], halo=[[0, 0], [0, 0]])
1129
+ self.assertSequenceEqual(
1130
+ [
1131
+ [
1132
+ (slice(0, 5), slice(0, 5)),
1133
+ (slice(0, 5), slice(5, 9)),
1134
+ (slice(0, 5), slice(9, 13))
1135
+ ],
1136
+ [
1137
+ (slice(5, 10), slice(0, 5)),
1138
+ (slice(5, 10), slice(5, 9)),
1139
+ (slice(5, 10), slice(9, 13))
1140
+ ],
1141
+ [
1142
+ (slice(10, 15), slice(0, 5)),
1143
+ (slice(10, 15), slice(5, 9)),
1144
+ (slice(10, 15), slice(9, 13))
1145
+ ],
1146
+ ],
1147
+ split.tolist()
1148
+ )
1149
+
1150
+ split = \
1151
+ shape_split(
1152
+ (15, 13),
1153
+ axis=[3, 3],
1154
+ halo=[[0, 0], [0, 0]],
1155
+ tile_bounds_policy=ARRAY_BOUNDS
1156
+ )
1157
+ self.assertSequenceEqual(
1158
+ [
1159
+ [
1160
+ (slice(0, 5), slice(0, 5)),
1161
+ (slice(0, 5), slice(5, 9)),
1162
+ (slice(0, 5), slice(9, 13))
1163
+ ],
1164
+ [
1165
+ (slice(5, 10), slice(0, 5)),
1166
+ (slice(5, 10), slice(5, 9)),
1167
+ (slice(5, 10), slice(9, 13))
1168
+ ],
1169
+ [
1170
+ (slice(10, 15), slice(0, 5)),
1171
+ (slice(10, 15), slice(5, 9)),
1172
+ (slice(10, 15), slice(9, 13))
1173
+ ],
1174
+ ],
1175
+ split.tolist()
1176
+ )
1177
+
1178
+ split = \
1179
+ shape_split(
1180
+ (15, 13),
1181
+ axis=[3, 3],
1182
+ halo=[[0, 0], [0, 0]],
1183
+ tile_bounds_policy=NO_BOUNDS
1184
+ )
1185
+ self.assertSequenceEqual(
1186
+ [
1187
+ [
1188
+ (slice(0, 5), slice(0, 5)),
1189
+ (slice(0, 5), slice(5, 9)),
1190
+ (slice(0, 5), slice(9, 13))
1191
+ ],
1192
+ [
1193
+ (slice(5, 10), slice(0, 5)),
1194
+ (slice(5, 10), slice(5, 9)),
1195
+ (slice(5, 10), slice(9, 13))
1196
+ ],
1197
+ [
1198
+ (slice(10, 15), slice(0, 5)),
1199
+ (slice(10, 15), slice(5, 9)),
1200
+ (slice(10, 15), slice(9, 13))
1201
+ ],
1202
+ ],
1203
+ split.tolist()
1204
+ )
1205
+
1206
+ split = \
1207
+ shape_split(
1208
+ (15, 13),
1209
+ axis=[3, 3],
1210
+ halo=1,
1211
+ tile_bounds_policy=ARRAY_BOUNDS
1212
+ )
1213
+ self.assertSequenceEqual(
1214
+ [
1215
+ [
1216
+ (slice(0, 6), slice(0, 6)),
1217
+ (slice(0, 6), slice(4, 10)),
1218
+ (slice(0, 6), slice(8, 13))
1219
+ ],
1220
+ [
1221
+ (slice(4, 11), slice(0, 6)),
1222
+ (slice(4, 11), slice(4, 10)),
1223
+ (slice(4, 11), slice(8, 13))
1224
+ ],
1225
+ [
1226
+ (slice(9, 15), slice(0, 6)),
1227
+ (slice(9, 15), slice(4, 10)),
1228
+ (slice(9, 15), slice(8, 13))
1229
+ ],
1230
+ ],
1231
+ split.tolist()
1232
+ )
1233
+
1234
+ split = \
1235
+ shape_split(
1236
+ (15, 13),
1237
+ axis=[3, 3],
1238
+ halo=(2, 3),
1239
+ tile_bounds_policy=ARRAY_BOUNDS
1240
+ )
1241
+ self.assertSequenceEqual(
1242
+ [
1243
+ [
1244
+ (slice(0, 7), slice(0, 8)),
1245
+ (slice(0, 7), slice(2, 12)),
1246
+ (slice(0, 7), slice(6, 13))
1247
+ ],
1248
+ [
1249
+ (slice(3, 12), slice(0, 8)),
1250
+ (slice(3, 12), slice(2, 12)),
1251
+ (slice(3, 12), slice(6, 13))
1252
+ ],
1253
+ [
1254
+ (slice(8, 15), slice(0, 8)),
1255
+ (slice(8, 15), slice(2, 12)),
1256
+ (slice(8, 15), slice(6, 13))
1257
+ ],
1258
+ ],
1259
+ split.tolist()
1260
+ )
1261
+
1262
+ split = \
1263
+ shape_split(
1264
+ (15, 13),
1265
+ axis=[3, 3],
1266
+ halo=[[1, 2], [2, 3]],
1267
+ tile_bounds_policy=ARRAY_BOUNDS
1268
+ )
1269
+ self.assertSequenceEqual(
1270
+ [
1271
+ [
1272
+ (slice(0, 7), slice(0, 8)),
1273
+ (slice(0, 7), slice(3, 12)),
1274
+ (slice(0, 7), slice(7, 13))
1275
+ ],
1276
+ [
1277
+ (slice(4, 12), slice(0, 8)),
1278
+ (slice(4, 12), slice(3, 12)),
1279
+ (slice(4, 12), slice(7, 13))
1280
+ ],
1281
+ [
1282
+ (slice(9, 15), slice(0, 8)),
1283
+ (slice(9, 15), slice(3, 12)),
1284
+ (slice(9, 15), slice(7, 13))
1285
+ ],
1286
+ ],
1287
+ split.tolist()
1288
+ )
1289
+
1290
+ # NO_BOUNDS
1291
+
1292
+ split = \
1293
+ shape_split(
1294
+ (15, 13),
1295
+ axis=[3, 3],
1296
+ halo=1,
1297
+ tile_bounds_policy=NO_BOUNDS
1298
+ )
1299
+ self.assertSequenceEqual(
1300
+ [
1301
+ [
1302
+ (slice(-1, 6), slice(-1, 6)),
1303
+ (slice(-1, 6), slice(4, 10)),
1304
+ (slice(-1, 6), slice(8, 14))
1305
+ ],
1306
+ [
1307
+ (slice(4, 11), slice(-1, 6)),
1308
+ (slice(4, 11), slice(4, 10)),
1309
+ (slice(4, 11), slice(8, 14))
1310
+ ],
1311
+ [
1312
+ (slice(9, 16), slice(-1, 6)),
1313
+ (slice(9, 16), slice(4, 10)),
1314
+ (slice(9, 16), slice(8, 14))
1315
+ ],
1316
+ ],
1317
+ split.tolist()
1318
+ )
1319
+
1320
+ split = \
1321
+ shape_split(
1322
+ (15, 13),
1323
+ axis=[3, 3],
1324
+ halo=(2, 3),
1325
+ tile_bounds_policy=NO_BOUNDS
1326
+ )
1327
+ self.assertSequenceEqual(
1328
+ [
1329
+ [
1330
+ (slice(-2, 7), slice(-3, 8)),
1331
+ (slice(-2, 7), slice(2, 12)),
1332
+ (slice(-2, 7), slice(6, 16))
1333
+ ],
1334
+ [
1335
+ (slice(3, 12), slice(-3, 8)),
1336
+ (slice(3, 12), slice(2, 12)),
1337
+ (slice(3, 12), slice(6, 16))
1338
+ ],
1339
+ [
1340
+ (slice(8, 17), slice(-3, 8)),
1341
+ (slice(8, 17), slice(2, 12)),
1342
+ (slice(8, 17), slice(6, 16))
1343
+ ],
1344
+ ],
1345
+ split.tolist()
1346
+ )
1347
+
1348
+ split = \
1349
+ shape_split(
1350
+ (15, 13),
1351
+ axis=[3, 3],
1352
+ halo=[[1, 2], [2, 3]],
1353
+ tile_bounds_policy=NO_BOUNDS
1354
+ )
1355
+ self.assertSequenceEqual(
1356
+ [
1357
+ [
1358
+ (slice(-1, 7), slice(-2, 8)),
1359
+ (slice(-1, 7), slice(3, 12)),
1360
+ (slice(-1, 7), slice(7, 16))
1361
+ ],
1362
+ [
1363
+ (slice(4, 12), slice(-2, 8)),
1364
+ (slice(4, 12), slice(3, 12)),
1365
+ (slice(4, 12), slice(7, 16))
1366
+ ],
1367
+ [
1368
+ (slice(9, 17), slice(-2, 8)),
1369
+ (slice(9, 17), slice(3, 12)),
1370
+ (slice(9, 17), slice(7, 16))
1371
+ ],
1372
+ ],
1373
+ split.tolist()
1374
+ )
1375
+
1376
+ def test_calculate_split_with_halo_for_empty_tiles(self):
1377
+ """
1378
+ Tests :func:`array_split.shape_split` for case of
1379
+ empty tiles and non-zero halo to ensure halo elements
1380
+ are not added to empty tiles.
1381
+ """
1382
+ # Zero halo, empty tiles.
1383
+ split = shape_split((5, 12), axis=[8, 1], halo=0)
1384
+ self.assertSequenceEqual(
1385
+ (
1386
+ slice(4, 5, None),
1387
+ slice(0, 12, None)
1388
+ ),
1389
+ split[4, 0].tolist()
1390
+ )
1391
+ for i in range(5, 8):
1392
+ self.assertSequenceEqual(
1393
+ (
1394
+ slice(5, 5, None),
1395
+ slice(0, 12, None)
1396
+ ),
1397
+ split[i, 0].tolist()
1398
+ )
1399
+
1400
+ # Now ensure that empty tiles remain empty despite halo=1
1401
+ split = shape_split((5, 12), axis=[8, 1], halo=1)
1402
+ self.assertSequenceEqual(
1403
+ (
1404
+ slice(3, 5, None),
1405
+ slice(0, 12, None)
1406
+ ),
1407
+ split[4, 0].tolist()
1408
+ )
1409
+ for i in range(5, 8):
1410
+ self.assertSequenceEqual(
1411
+ (
1412
+ slice(5, 5, None),
1413
+ slice(0, 12, None)
1414
+ ),
1415
+ split[i, 0].tolist()
1416
+ )
1417
+
1418
+ split = shape_split((5, 12), axis=[8, 15], halo=1)
1419
+ self.assertSequenceEqual(
1420
+ (
1421
+ slice(3, 5, None),
1422
+ slice(0, 2, None)
1423
+ ),
1424
+ split[4, 0].tolist()
1425
+ )
1426
+ for i in range(5, 8):
1427
+ for j in range(0, 12):
1428
+ self.assertEqual(
1429
+ slice(5, 5, None),
1430
+ split[i, j].tolist()[0]
1431
+ )
1432
+ for j in range(12, 15):
1433
+ self.assertSequenceEqual(
1434
+ (
1435
+ slice(5, 5, None),
1436
+ slice(12, 12, None)
1437
+ ),
1438
+ split[i, j].tolist()
1439
+ )
1440
+ for i in range(0, 5):
1441
+ for j in range(12, 15):
1442
+ self.assertEqual(
1443
+ slice(12, 12, None),
1444
+ split[i, j].tolist()[1]
1445
+ )
1446
+
1447
+ def test_calculate_split_halos_from_extents(self):
1448
+ """
1449
+ Tests the :meth:`array_split.split.ShapeSplitter.calculate_split_halos_from_extents`
1450
+ method.
1451
+ """
1452
+
1453
+ # Tiles wider than halo width
1454
+ splitter = ShapeSplitter((15, 13), axis=[3, 3], halo=0)
1455
+ splt = splitter.calculate_split()
1456
+ splt_halos = splitter.calculate_split_halos_from_extents()
1457
+ self.assertSequenceEqual(splt.shape, splt_halos.shape)
1458
+ self.assertTrue(_np.all(_np.asarray(splt_halos.tolist()) == 0))
1459
+
1460
+ # Some tiles narrower than halo width
1461
+ splitter = ShapeSplitter((15, 13), axis=[3, 3], halo=5, tile_bounds_policy=ARRAY_BOUNDS)
1462
+ splt = splitter.calculate_split()
1463
+ splt_halos = splitter.calculate_split_halos_from_extents()
1464
+ self.assertSequenceEqual(splt.shape, splt_halos.shape)
1465
+ for i in range(3):
1466
+ self.assertSequenceEqual([0, 5], tuple(splt_halos[0, i][0]))
1467
+ self.assertSequenceEqual([5, 5], tuple(splt_halos[1, i][0]))
1468
+ self.assertSequenceEqual([5, 0], tuple(splt_halos[2, i][0]))
1469
+ self.assertSequenceEqual([0, 5], tuple(splt_halos[i, 0][1]))
1470
+ self.assertSequenceEqual([5, 4], tuple(splt_halos[i, 1][1]))
1471
+ self.assertSequenceEqual([5, 0], tuple(splt_halos[i, 2][1]))
1472
+
1473
+ splitter = ShapeSplitter((15, 13), axis=[3, 3], halo=0)
1474
+ splt = splitter.calculate_split()
1475
+ splt_halos = splitter.calculate_split_halos_from_extents()
1476
+ self.assertSequenceEqual(splt.shape, splt_halos.shape)
1477
+ self.assertTrue(_np.all(_np.asarray(splt_halos.tolist()) == 0))
1478
+
1479
+ # Tiles narrower than halo width
1480
+ splitter = ShapeSplitter((5, 13), axis=[5, 3], halo=5, tile_bounds_policy=ARRAY_BOUNDS)
1481
+ splt = splitter.calculate_split()
1482
+ splt_halos = splitter.calculate_split_halos_from_extents()
1483
+ self.assertSequenceEqual(splt.shape, splt_halos.shape)
1484
+ for i in range(3):
1485
+ self.assertSequenceEqual([0, 4], tuple(splt_halos[0, i][0]))
1486
+ self.assertSequenceEqual([1, 3], tuple(splt_halos[1, i][0]))
1487
+ self.assertSequenceEqual([2, 2], tuple(splt_halos[2, i][0]))
1488
+ self.assertSequenceEqual([3, 1], tuple(splt_halos[3, i][0]))
1489
+ self.assertSequenceEqual([4, 0], tuple(splt_halos[4, i][0]))
1490
+ for i in range(5):
1491
+ self.assertSequenceEqual([0, 5], tuple(splt_halos[i, 0][1]))
1492
+ self.assertSequenceEqual([5, 4], tuple(splt_halos[i, 1][1]))
1493
+ self.assertSequenceEqual([5, 0], tuple(splt_halos[i, 2][1]))
1494
+
1495
+ # Zero sized tiles
1496
+ # Tiles narrower than halo width
1497
+ splitter = ShapeSplitter((5, 13), axis=[7, 3], halo=5, tile_bounds_policy=ARRAY_BOUNDS)
1498
+ splt = splitter.calculate_split()
1499
+ splt_halos = splitter.calculate_split_halos_from_extents()
1500
+ self.assertSequenceEqual(splt.shape, splt_halos.shape)
1501
+ for i in range(3):
1502
+ self.assertSequenceEqual([0, 4], tuple(splt_halos[0, i][0]))
1503
+ self.assertSequenceEqual([1, 3], tuple(splt_halos[1, i][0]))
1504
+ self.assertSequenceEqual([2, 2], tuple(splt_halos[2, i][0]))
1505
+ self.assertSequenceEqual([3, 1], tuple(splt_halos[3, i][0]))
1506
+ self.assertSequenceEqual([4, 0], tuple(splt_halos[4, i][0]))
1507
+ self.assertSequenceEqual([0, 0], tuple(splt_halos[5, i][0]))
1508
+ self.assertSequenceEqual([0, 0], tuple(splt_halos[6, i][0]))
1509
+ for i in range(5):
1510
+ self.assertSequenceEqual([0, 5], tuple(splt_halos[i, 0][1]))
1511
+ self.assertSequenceEqual([5, 4], tuple(splt_halos[i, 1][1]))
1512
+ self.assertSequenceEqual([5, 0], tuple(splt_halos[i, 2][1]))
1513
+
1514
+
1515
+ __all__ = [s for s in dir() if not s.startswith('_')]
1516
+
1517
+ _unittest.main(__name__)