feectools 0.1.0__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 (98) hide show
  1. feectools/__init__.py +0 -0
  2. feectools/accelerate/__init__.py +0 -0
  3. feectools/accelerate/accelerate.py +220 -0
  4. feectools/accelerate/compile_psydac.mk +52 -0
  5. feectools/api/__init__.py +0 -0
  6. feectools/api/essential_bc.py +122 -0
  7. feectools/api/fem_bilinear_form.py +2226 -0
  8. feectools/api/fem_common.py +286 -0
  9. feectools/api/fem_sum_form.py +123 -0
  10. feectools/api/settings.py +82 -0
  11. feectools/core/__init__.py +11 -0
  12. feectools/core/bsplines.py +1107 -0
  13. feectools/core/bsplines_kernels.py +1349 -0
  14. feectools/core/field_evaluation_kernels.py +5015 -0
  15. feectools/core/tests/__init__.py +0 -0
  16. feectools/core/tests/test_bsplines.py +263 -0
  17. feectools/core/tests/test_bsplines_kernel.py +40 -0
  18. feectools/core/tests/test_bsplines_pyccel.py +752 -0
  19. feectools/ddm/__init__.py +3 -0
  20. feectools/ddm/basic.py +78 -0
  21. feectools/ddm/blocking_data_exchanger.py +348 -0
  22. feectools/ddm/cart.py +1835 -0
  23. feectools/ddm/interface_data_exchanger.py +122 -0
  24. feectools/ddm/mpi.py +109 -0
  25. feectools/ddm/nonblocking_data_exchanger.py +331 -0
  26. feectools/ddm/partition.py +207 -0
  27. feectools/ddm/petsc.py +112 -0
  28. feectools/ddm/tests/__init__.py +0 -0
  29. feectools/ddm/tests/test_cart_1d.py +138 -0
  30. feectools/ddm/tests/test_cart_2d.py +164 -0
  31. feectools/ddm/tests/test_cart_3d.py +158 -0
  32. feectools/ddm/tests/test_multicart_2d.py +173 -0
  33. feectools/ddm/tests/test_partition.py +124 -0
  34. feectools/ddm/utilities.py +24 -0
  35. feectools/feec/__init__.py +0 -0
  36. feectools/feec/derivatives.py +780 -0
  37. feectools/feec/dof_kernels.py +210 -0
  38. feectools/feec/global_geometric_projectors.py +1073 -0
  39. feectools/feec/hodge.py +148 -0
  40. feectools/fem/__init__.py +0 -0
  41. feectools/fem/basic.py +465 -0
  42. feectools/fem/grid.py +181 -0
  43. feectools/fem/partitioning.py +344 -0
  44. feectools/fem/projectors.py +160 -0
  45. feectools/fem/splines.py +559 -0
  46. feectools/fem/tensor.py +1393 -0
  47. feectools/fem/tests/__init__.py +0 -0
  48. feectools/fem/tests/analytical_profiles_1d.py +100 -0
  49. feectools/fem/tests/analytical_profiles_base.py +34 -0
  50. feectools/fem/tests/splines_error_bounds.py +155 -0
  51. feectools/fem/tests/test_spline_histopolation.py +120 -0
  52. feectools/fem/tests/test_spline_interpolation.py +182 -0
  53. feectools/fem/tests/test_splines.py +184 -0
  54. feectools/fem/tests/test_splines_par.py +46 -0
  55. feectools/fem/tests/test_vector_spaces.py +150 -0
  56. feectools/fem/tests/utilities.py +47 -0
  57. feectools/fem/vector.py +729 -0
  58. feectools/linalg/__init__.py +0 -0
  59. feectools/linalg/basic.py +1386 -0
  60. feectools/linalg/block.py +1451 -0
  61. feectools/linalg/direct_solvers.py +201 -0
  62. feectools/linalg/fft.py +258 -0
  63. feectools/linalg/kernels/__init__.py +0 -0
  64. feectools/linalg/kernels/axpy_kernels.py +57 -0
  65. feectools/linalg/kernels/inner_kernels.py +100 -0
  66. feectools/linalg/kernels/matvec_kernels.py +206 -0
  67. feectools/linalg/kernels/stencil2IJV_kernels.py +227 -0
  68. feectools/linalg/kernels/stencil2coo_kernels.py +179 -0
  69. feectools/linalg/kernels/transpose_kernels.py +263 -0
  70. feectools/linalg/kron.py +911 -0
  71. feectools/linalg/solvers.py +1914 -0
  72. feectools/linalg/sparse.py +114 -0
  73. feectools/linalg/stencil.py +2923 -0
  74. feectools/linalg/stencil_dot_kernels.py +317 -0
  75. feectools/linalg/stencil_transpose_kernels.py +372 -0
  76. feectools/linalg/tests/__init__.py +0 -0
  77. feectools/linalg/tests/test_block.py +1588 -0
  78. feectools/linalg/tests/test_fft.py +106 -0
  79. feectools/linalg/tests/test_kron_stencil_matrix.py +114 -0
  80. feectools/linalg/tests/test_linalg.py +1065 -0
  81. feectools/linalg/tests/test_matrix_free.py +128 -0
  82. feectools/linalg/tests/test_solvers.py +213 -0
  83. feectools/linalg/tests/test_stencil_interface_matrix.py +379 -0
  84. feectools/linalg/tests/test_stencil_vector.py +1036 -0
  85. feectools/linalg/tests/test_stencil_vector_space.py +440 -0
  86. feectools/linalg/topetsc.py +522 -0
  87. feectools/linalg/utilities.py +200 -0
  88. feectools/utilities/__init__.py +0 -0
  89. feectools/utilities/quadratures.py +113 -0
  90. feectools/utilities/utils.py +166 -0
  91. feectools/version.py +1 -0
  92. feectools-0.1.0.dist-info/METADATA +66 -0
  93. feectools-0.1.0.dist-info/RECORD +98 -0
  94. feectools-0.1.0.dist-info/WHEEL +5 -0
  95. feectools-0.1.0.dist-info/entry_points.txt +3 -0
  96. feectools-0.1.0.dist-info/licenses/AUTHORS +22 -0
  97. feectools-0.1.0.dist-info/licenses/LICENSE +21 -0
  98. feectools-0.1.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,158 @@
1
+ # File test_cart_3d.py
2
+
3
+ from feectools.ddm.blocking_data_exchanger import BlockingCartDataExchanger
4
+ from feectools.ddm.nonblocking_data_exchanger import NonBlockingCartDataExchanger
5
+
6
+ #===============================================================================
7
+ # TEST CartDecomposition and CartDataExchanger in 3D
8
+ #===============================================================================
9
+ def run_cart_3d( data_exchanger_type, verbose=False ):
10
+
11
+ import numpy as np
12
+ from feectools.ddm.mpi import mpi as MPI
13
+ from feectools.ddm.cart import DomainDecomposition, CartDecomposition
14
+
15
+ #---------------------------------------------------------------------------
16
+ # INPUT PARAMETERS
17
+ #---------------------------------------------------------------------------
18
+
19
+
20
+ # Number of cells
21
+ nc1 = 135
22
+ nc2 = 77
23
+ nc3 = 98
24
+
25
+ # Padding ('thickness' of ghost region)
26
+ p1 = 3
27
+ p2 = 2
28
+ p3 = 5
29
+
30
+ # Periodicity
31
+ period1 = True
32
+ period2 = False
33
+ period3 = True
34
+
35
+ # Number of Points
36
+ n1 = nc1 + p1*(1-period1)
37
+ n2 = nc2 + p2*(1-period2)
38
+ n3 = nc3 + p3*(1-period3)
39
+ #---------------------------------------------------------------------------
40
+ # DOMAIN DECOMPOSITION
41
+ #---------------------------------------------------------------------------
42
+
43
+ # Parallel info
44
+ comm = MPI.COMM_WORLD
45
+ size = comm.Get_size()
46
+ rank = comm.Get_rank()
47
+
48
+ domain_decomposition = DomainDecomposition(ncells=[nc1,nc2,nc3], periods=[period1,period2,period3], comm=comm)
49
+
50
+ npts = [n1,n2,n3]
51
+ global_starts = [None]*3
52
+ global_ends = [None]*3
53
+ for axis in range(3):
54
+ es = domain_decomposition.global_element_starts[axis]
55
+ ee = domain_decomposition.global_element_ends [axis]
56
+
57
+ global_ends [axis] = (ee+1)-1
58
+ global_ends [axis][-1] = npts[axis]-1
59
+ global_starts[axis] = np.array([0] + (global_ends[axis][:-1]+1).tolist())
60
+
61
+ # Decomposition of Cartesian domain
62
+ cart = CartDecomposition(
63
+ domain_decomposition = domain_decomposition,
64
+ npts = [n1,n2,n3],
65
+ global_starts = global_starts,
66
+ global_ends = global_ends,
67
+ pads = [p1,p2,p3],
68
+ shifts = [1,1,1],
69
+ )
70
+
71
+ # Local 3D array with 3D vector data (extended domain)
72
+ shape = list( cart.shape ) + [3]
73
+ u = np.zeros( shape, dtype=int )
74
+
75
+ # Global indices of first and last elements of array
76
+ s1,s2,s3 = cart.starts
77
+ e1,e2,e3 = cart.ends
78
+
79
+ # Create object in charge of exchanging data between subdomains
80
+ synchronizer = data_exchanger_type( cart, u.dtype, coeff_shape=[3] )
81
+
82
+ # Print some info
83
+ if rank == 0:
84
+ print( "" )
85
+
86
+ for k in range(size):
87
+ if k == rank:
88
+ print( "Proc. # {}".format( rank ) )
89
+ print( "---------" )
90
+ print( ". s1:e1 = {:2d}:{:2d}".format( s1,e1 ) )
91
+ print( ". s2:e2 = {:2d}:{:2d}".format( s2,e2 ) )
92
+ print( ". s3:e3 = {:2d}:{:2d}".format( s3,e3 ) )
93
+ print( "", flush=True )
94
+ comm.Barrier()
95
+
96
+ #---------------------------------------------------------------------------
97
+ # TEST
98
+ #---------------------------------------------------------------------------
99
+
100
+ # Fill in true domain with u[i1_loc,i2_loc,i3_loc,:]=[i1_glob,i2_glob,i3_glob]
101
+ u[p1:-p1,p2:-p2,p3:-p3,:] = [[[(i1,i2,i3) for i3 in range(s3,e3+1)] \
102
+ for i2 in range(s2,e2+1)] \
103
+ for i1 in range(s1,e1+1)]
104
+
105
+ request = synchronizer.prepare_communications(u)
106
+ # Update ghost regions
107
+ synchronizer.start_update_ghost_regions( u, request )
108
+ synchronizer.end_update_ghost_regions( u, request )
109
+
110
+ #---------------------------------------------------------------------------
111
+ # CHECK RESULTS
112
+ #---------------------------------------------------------------------------
113
+
114
+ # Verify that ghost cells contain correct data (note periodic domain!)
115
+ val = lambda i1,i2,i3: (i1%n1,i2,i3%n3) if 0<=i2<n2 else (0,0,0)
116
+
117
+ uex = [[[val(i1,i2,i3) for i3 in range(s3-p3,e3+p3+1)] \
118
+ for i2 in range(s2-p2,e2+p2+1)] \
119
+ for i1 in range(s1-p1,e1+p1+1)]
120
+
121
+ success = (u == uex).all()
122
+
123
+ # MASTER only: collect information from all processes
124
+ success_global = comm.reduce( success, op=MPI.LAND, root=0 )
125
+
126
+ return locals()
127
+
128
+ #===============================================================================
129
+ # RUN TEST WITH PYTEST
130
+ #===============================================================================
131
+ import pytest
132
+
133
+ @pytest.mark.parametrize( 'data_exchanger_type', [BlockingCartDataExchanger, NonBlockingCartDataExchanger] )
134
+ @pytest.mark.parallel
135
+ def test_cart_3d(data_exchanger_type):
136
+
137
+ namespace = run_cart_3d(data_exchanger_type)
138
+
139
+ assert namespace['success']
140
+
141
+ #===============================================================================
142
+ # RUN TEST MANUALLY
143
+ #===============================================================================
144
+ if __name__=='__main__':
145
+
146
+ locals().update( run_cart_3d( BlockingCartDataExchanger, verbose=True ) )
147
+
148
+ # Print error messages (if any) in orderly fashion
149
+ for k in range(size):
150
+ if k == rank and not success:
151
+ print( "Rank {}: wrong ghost cell data!".format( rank ), flush=True )
152
+ comm.Barrier()
153
+
154
+ if rank == 0:
155
+ if success_global:
156
+ print( "PASSED", end='\n\n', flush=True )
157
+ else:
158
+ print( "FAILED", end='\n\n', flush=True )
@@ -0,0 +1,173 @@
1
+ # File test_multicart_2d.py
2
+
3
+ #===============================================================================
4
+ # TEST MultiCartDecomposition in 2D
5
+ #===============================================================================
6
+
7
+ #------------------------------------------------------------------------------
8
+ def get_minus_starts_ends(plus_starts, plus_ends, minus_npts, plus_npts, minus_axis, plus_axis,
9
+ minus_ext, plus_ext, minus_pads, plus_pads, minus_shifts, plus_shifts,
10
+ diff):
11
+ """
12
+ Compute the coefficients needed by the minus patch in a given interface.
13
+ """
14
+ starts = [max(0,s-m*p) for s,m,p in zip(plus_starts, minus_shifts, minus_pads)]
15
+ ends = [min(n,e+m*p) for e,n,m,p in zip(plus_ends, minus_npts, minus_shifts, minus_pads)]
16
+ starts[minus_axis] = 0 if minus_ext == -1 else ends[minus_axis]-minus_pads[minus_axis]
17
+ ends[minus_axis] = ends[minus_axis] if minus_ext == 1 else minus_pads[minus_axis]
18
+ return starts, ends
19
+
20
+ #------------------------------------------------------------------------------
21
+ def get_plus_starts_ends(minus_starts, minus_ends, minus_npts, plus_npts, minus_axis, plus_axis,
22
+ minus_ext, plus_ext, minus_pads, plus_pads, minus_shifts, plus_shifts,
23
+ diff):
24
+ """
25
+ Compute the coefficients needed by the plus patch in a given interface.
26
+ """
27
+ starts = [max(0,s-m*p) for s,m,p in zip(minus_starts, plus_shifts, plus_pads)]
28
+ ends = [min(n,e+m*p) for e,n,m,p in zip(minus_ends, plus_npts, plus_shifts, plus_pads)]
29
+ starts[plus_axis] = 0 if plus_ext == -1 else ends[plus_axis]-plus_pads[plus_axis]
30
+ ends[plus_axis] = ends[plus_axis] if plus_ext == 1 else plus_pads[plus_axis]
31
+ return starts, ends
32
+
33
+
34
+ #===============================================================================
35
+ # TEST MultiPatchDomainDecomposition and CartDataExchanger in 2D
36
+ #===============================================================================
37
+ def run_carts_2d():
38
+ import numpy as np
39
+
40
+ from feectools.ddm.mpi import mpi as MPI
41
+ from feectools.ddm.cart import MultiPatchDomainDecomposition, CartDecomposition, create_interfaces_cart
42
+ from feectools.ddm.blocking_data_exchanger import BlockingCartDataExchanger
43
+ from feectools.ddm.interface_data_exchanger import InterfaceCartDataExchanger
44
+
45
+ #---------------------------------------------------------------------------
46
+ # INPUT PARAMETERS
47
+ #---------------------------------------------------------------------------
48
+
49
+ # Number of patches
50
+ N = 2
51
+
52
+ # Number of cells
53
+ nc1,nc2 = 16,16
54
+ nc = [[nc1,nc2] for i in range(N)]
55
+
56
+ # Padding ('thickness' of ghost region)
57
+ p1,p2 = 2,2
58
+ p = [[p1,p2] for i in range(N)]
59
+
60
+ # Periodicity
61
+ P = [[False, False] for i in range(N)]
62
+
63
+ connectivity = {(i,i+1):( ((0,1),(0,-1)) if i%2 ==0 else ((0,-1),(0,1))) for i in range(N-1)}
64
+
65
+ #---------------------------------------------------------------------------
66
+ # DOMAIN DECOMPOSITION
67
+ #---------------------------------------------------------------------------
68
+
69
+ # Parallel info
70
+ comm = MPI.COMM_WORLD
71
+
72
+ domain_decomposition = MultiPatchDomainDecomposition(nc, P, comm=comm)
73
+
74
+ # Number of Points
75
+ n = [[ncij + pij*(1-periodij) for ncij,pij,periodij in zip(nci,pi,periodi)] for nci,pi,periodi in zip(nc,p,P)]
76
+
77
+ carts = []
78
+ for i in range(N):
79
+ global_starts = [None]*2
80
+ global_ends = [None]*2
81
+ for axis in range(2):
82
+ es = domain_decomposition.domains[i].global_element_starts[axis]
83
+ ee = domain_decomposition.domains[i].global_element_ends [axis]
84
+
85
+ global_ends [axis] = (ee+1)-1
86
+ global_ends [axis][-1] = n[i][axis]-1
87
+ global_starts[axis] = np.array([0] + (global_ends[axis][:-1]+1).tolist())
88
+
89
+ carts.append(CartDecomposition(
90
+ domain_decomposition = domain_decomposition.domains[i],
91
+ npts = n[i],
92
+ global_starts = global_starts,
93
+ global_ends = global_ends,
94
+ pads = p[i],
95
+ shifts = [1,1]))
96
+ carts = tuple(carts)
97
+
98
+ communication_info = (get_minus_starts_ends, get_plus_starts_ends)
99
+ interfaces_cart = create_interfaces_cart(domain_decomposition, carts, connectivity, communication_info=communication_info)
100
+
101
+ us = [None]*len(carts)
102
+ syn = [None]*len(carts)
103
+ syn_interface = {}
104
+ dtype = int
105
+
106
+ val = lambda k,i1,i2: k*n[k][0]*n[k][1]+i1*n[k][0]+i2 if (0<=i1<n[k][0] and 0<=i2<n[k][1]) else 0
107
+ for i,ci in enumerate(carts):
108
+ if not ci.is_comm_null:
109
+ s1,s2 = ci.starts
110
+ e1,e2 = ci.ends
111
+ m1,m2 = ci.shifts
112
+ us[i] = np.zeros( ci.shape, dtype=dtype )
113
+ us[i][m1*p1:-m1*p1,m2*p2:-m2*p2] = [[val(i,i1,i2)for i2 in range(s2,e2+1)] for i1 in range(s1,e1+1)]
114
+ synchronizer = BlockingCartDataExchanger( ci, us[i].dtype)
115
+ syn[i] = synchronizer
116
+
117
+ for i,j in connectivity:
118
+ if not interfaces_cart[i,j].is_comm_null:
119
+ if carts[i].is_comm_null:
120
+ shape = interfaces_cart[i,j].get_interface_communication_infos(interfaces_cart[i,j]._axis)['gbuf_recv_shape'][0]
121
+ us[i] = np.zeros(shape, dtype=dtype)
122
+
123
+ if carts[j].is_comm_null:
124
+ shape = interfaces_cart[i,j].get_interface_communication_infos(interfaces_cart[i,j]._axis)['gbuf_recv_shape'][0]
125
+ us[j] = np.zeros(shape, dtype=dtype)
126
+
127
+ syn_interface[i,j] = InterfaceCartDataExchanger(interfaces_cart[i,j], dtype)
128
+
129
+ for minus,plus in connectivity:
130
+ if not interfaces_cart[minus,plus].is_comm_null:
131
+ req = syn_interface[minus,plus].start_update_ghost_regions(us[minus], us[plus])
132
+ syn_interface[minus,plus].end_update_ghost_regions(req=req)
133
+
134
+ for i,ci in enumerate(carts):
135
+ if not ci.is_comm_null:
136
+ # Update ghost regions
137
+ syn[i].start_update_ghost_regions( us[i], None )
138
+ syn[i].end_update_ghost_regions( us[i], None )
139
+
140
+ for i,ci in enumerate(carts):
141
+ if not ci.is_comm_null:
142
+ s1,s2 = ci.starts
143
+ e1,e2 = ci.ends
144
+ m1,m2 = ci.shifts
145
+ uex = [[val(i,i1,i2) for i2 in range(s2-m2*p2,e2+m2*p2+1)] for i1 in range(s1-m1*p1,e1+m1*p1+1)]
146
+ success = (us[i] == uex).all()
147
+ assert success
148
+
149
+ # for minus,plus in connectivity:
150
+ # if not interfaces[minus,plus].is_comm_null:
151
+ # axis = interfaces[minus,plus].axis
152
+ # I = interfaces[minus,plus]
153
+
154
+ # if not carts[minus].is_comm_null:
155
+ # uex = [[val(plus,i1,i2)for i2 in range(*ranges[1])] for i1 in range(*ranges[0])]
156
+ # uex = np.pad(uex, [(m*p,m*p) for m,p in zip(carts[minus].shifts, carts[minus].pads)])
157
+ # u_ij = us[plus]
158
+ # elif not carts[plus].is_comm_null:
159
+ # uex = [[val(minus,i1,i2)for i2 in range(*ranges[1])] for i1 in range(*ranges[0])]
160
+ # uex = np.pad(uex, [(m*p,m*p) for m,p in zip(carts[plus].shifts, carts[plus].pads)])
161
+ # u_ij = us[minus]
162
+
163
+ # success = (u_ij == uex).all()
164
+ # assert success
165
+
166
+ #===============================================================================
167
+ # RUN TEST MANUALLY
168
+ #===============================================================================
169
+ if __name__=='__main__':
170
+
171
+ run_carts_2d()
172
+
173
+
@@ -0,0 +1,124 @@
1
+ import pytest
2
+
3
+ from feectools.ddm.partition import compute_dims
4
+
5
+ #==============================================================================
6
+ @pytest.mark.parametrize( 'mpi_size', [1,2,5,10] )
7
+
8
+ def test_partition_1d_uniform( mpi_size ):
9
+
10
+ # ...
11
+ # Should pass: all blocks are identical and have size=11
12
+ n1 = 11 * mpi_size
13
+ p1 = 3
14
+
15
+ dims, blocksizes = compute_dims( mpi_size, [n1,], [p1,] )
16
+
17
+ assert dims[0] == mpi_size
18
+ assert blocksizes[0] == 11
19
+
20
+ # ...
21
+ # Should fail: minimum block size is too large
22
+ n1 = 4 * mpi_size
23
+ p1 = 5
24
+
25
+ with pytest.raises( Exception ):
26
+ dims, blocksizes = compute_dims( mpi_size, [n1,], [p1,] )
27
+
28
+ #==============================================================================
29
+ @pytest.mark.parametrize( 'mpi_size', [1,2,5,10] )
30
+
31
+ def test_partition_1d_general( mpi_size ):
32
+
33
+ # ...
34
+ # Should pass, nominal block size is 11
35
+ n1 = 11 * mpi_size + int( mpi_size > 1 )
36
+ p1 = 4
37
+
38
+ dims, blocksizes = compute_dims( mpi_size, [n1,], [p1,] )
39
+
40
+ assert dims[0] == mpi_size
41
+ assert blocksizes[0] == 11
42
+
43
+ # ...
44
+ # Should fail: minimum block size is too large
45
+ n1 = 4 * mpi_size + int( mpi_size > 1 )
46
+ p1 = 5
47
+
48
+ with pytest.raises( Exception ):
49
+ dims, blocksizes = compute_dims( mpi_size, [n1,], [p1,] )
50
+
51
+ #==============================================================================
52
+ @pytest.mark.parametrize( 'mpi_size', [1,2,5,10] )
53
+ @pytest.mark.parametrize( 'mask', [[True, False], [False, True]] )
54
+ @pytest.mark.parametrize( 'npts', [[64, 64], [58, 64], [64, 31]] )
55
+
56
+ def test_partition_2d_dims_mask( mpi_size, npts, mask ):
57
+
58
+ # General partition: blocks are not all identical but closer to a cube
59
+ dims, blocksizes = compute_dims( mpi_size, npts, [3,3] )
60
+
61
+ # Mask dimensions
62
+ dims, blocksizes = compute_dims( mpi_size, npts, [3,3], mpi_dims_mask=mask )
63
+
64
+ # test
65
+ assert dims[0]*dims[1] == mpi_size
66
+ for bsize, n, use_dim in zip(blocksizes, npts, mask):
67
+ if not use_dim:
68
+ assert bsize == n
69
+ else:
70
+ assert bsize == n//mpi_size
71
+
72
+ #==============================================================================
73
+ def test_partition_3d():
74
+
75
+ npts = [64,128,50]
76
+ mpi_size = 100
77
+
78
+ # ...
79
+ # Uniform partition, yields small block size along 3rd dimension
80
+ dims, blocksizes = compute_dims( mpi_size, npts, try_uniform=True )
81
+
82
+ assert tuple( dims ) == (2, 2, 25)
83
+ assert tuple( blocksizes ) == (32, 64, 2)
84
+
85
+ # ...
86
+ # General partition: blocks are not all identical but closer to a cube
87
+ dims, blocksizes = compute_dims( mpi_size, npts, [3,3,3] )
88
+
89
+ assert tuple( dims ) == (5, 5, 4)
90
+ assert tuple( blocksizes ) == (12, 25, 12)
91
+
92
+ #==============================================================================
93
+ @pytest.mark.parametrize( 'mpi_size', [1,2,5,10] )
94
+ @pytest.mark.parametrize( 'mask', [[True, False, False],
95
+ [False, True, False],
96
+ [False, False, True],
97
+ [True, True, False],
98
+ [True, False, True],
99
+ [False, True, True]] )
100
+ @pytest.mark.parametrize( 'npts', [[32, 64, 128], [62, 59, 41]] )
101
+
102
+ def test_partition_3d_dims_mask( mpi_size, npts, mask ):
103
+
104
+ # General partition: blocks are not all identical but closer to a cube
105
+ dims, blocksizes = compute_dims( mpi_size, npts, [3,3,3] )
106
+
107
+ # Mask dimensions
108
+ dims, blocksizes = compute_dims( mpi_size, npts, [3,3,3], mpi_dims_mask=mask )
109
+
110
+ # test
111
+ assert dims[0]*dims[1]*dims[2] == mpi_size
112
+ for bsize, n, use_dim in zip(blocksizes, npts, mask):
113
+ if not use_dim:
114
+ assert bsize == n
115
+
116
+
117
+ if __name__ == '__main__':
118
+ # test_partition_2d_dims_mask(10, [64, 64], [True, False])
119
+ # test_partition_2d_dims_mask(10, [58, 64], [True, False])
120
+
121
+ test_partition_3d_dims_mask(10, [32, 64, 128], [True, False, True])
122
+
123
+
124
+
@@ -0,0 +1,24 @@
1
+ # coding: utf-8
2
+
3
+ from .cart import CartDecomposition, InterfaceCartDecomposition
4
+ from .blocking_data_exchanger import BlockingCartDataExchanger
5
+ from .nonblocking_data_exchanger import NonBlockingCartDataExchanger
6
+ from .interface_data_exchanger import InterfaceCartDataExchanger
7
+
8
+ __all__ = ('get_data_exchanger',)
9
+
10
+ def get_data_exchanger(cart, dtype, *, coeff_shape=(), assembly=False, axis=None, shape=None, blocking=True):
11
+
12
+ if isinstance(cart, InterfaceCartDecomposition):
13
+ return InterfaceCartDataExchanger(cart, dtype, coeff_shape=coeff_shape)
14
+
15
+ elif isinstance(cart, CartDecomposition):
16
+ if blocking:
17
+ return BlockingCartDataExchanger(cart, dtype, coeff_shape=coeff_shape, assembly=assembly, axis=axis, shape=shape)
18
+
19
+ else:
20
+ return NonBlockingCartDataExchanger(cart, dtype, coeff_shape=coeff_shape, assembly=assembly, axis=axis, shape=shape)
21
+ else:
22
+ raise TypeError('cart can only be of type CartDecomposition or InterfaceCartDecomposition')
23
+
24
+
File without changes