Skip to content

Commit 1ba576e

Browse files
authored
Merge pull request #674 from OP2/ksagiyam/composed_map
add ComposedMap
2 parents a0259a9 + fcf4250 commit 1ba576e

8 files changed

Lines changed: 369 additions & 23 deletions

File tree

‎pyop2/codegen/builder.py‎

Lines changed: 27 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
import numpy
77
from loopy.types import OpaqueType
88
from pyop2.global_kernel import (GlobalKernelArg, DatKernelArg, MixedDatKernelArg,
9-
MatKernelArg, MixedMatKernelArg, PermutedMapKernelArg)
9+
MatKernelArg, MixedMatKernelArg, PermutedMapKernelArg, ComposedMapKernelArg)
1010
from pyop2.codegen.representation import (Accumulate, Argument, Comparison,
1111
DummyInstruction, Extent, FixedIndex,
1212
FunctionCall, Index, Indexed,
@@ -154,6 +154,28 @@ def indexed_vector(self, n, shape, layer=None):
154154
return super().indexed_vector(n, shape, layer=layer, permute=permute)
155155

156156

157+
class CMap(Map):
158+
159+
def __init__(self, *maps_):
160+
# Copy over properties
161+
self.variable = maps_[0].variable
162+
self.unroll = maps_[0].unroll
163+
self.layer_bounds = maps_[0].layer_bounds
164+
self.interior_horizontal = maps_[0].interior_horizontal
165+
self.prefetch = {}
166+
self.values = maps_[0].values
167+
self.offset = maps_[0].offset
168+
self.maps_ = maps_
169+
170+
def indexed(self, multiindex, layer=None):
171+
n, i, f = multiindex
172+
n_ = n
173+
for map_ in reversed(self.maps_):
174+
if map_ is not self.maps_[0]:
175+
n_, (_, _) = map_.indexed(MultiIndex(n_, FixedIndex(0), Index()), layer=None)
176+
return self.maps_[0].indexed(MultiIndex(n_, i, f), layer=layer)
177+
178+
157179
class Pack(metaclass=ABCMeta):
158180

159181
def pick_loop_indices(self, loop_index, layer_index=None, entity_index=None):
@@ -835,6 +857,8 @@ def _add_map(self, map_, unroll=False):
835857
if isinstance(map_, PermutedMapKernelArg):
836858
imap = self._add_map(map_.base_map, unroll)
837859
map_ = PMap(imap, numpy.asarray(map_.permutation, dtype=IntType))
860+
elif isinstance(map_, ComposedMapKernelArg):
861+
map_ = CMap(*(self._add_map(m, unroll) for m in map_.base_maps))
838862
else:
839863
map_ = Map(interior_horizontal,
840864
(self.bottom_layer, self.top_layer),
@@ -878,7 +902,8 @@ def wrapper_args(self):
878902
# But we don't need to emit stuff for PMaps because they
879903
# are a Map (already seen + a permutation [encoded in the
880904
# indexing]).
881-
if not isinstance(map_, PMap):
905+
# CMaps do not have their own arguments, either.
906+
if not isinstance(map_, (PMap, CMap)):
882907
args.append(map_.values)
883908
return tuple(args)
884909

‎pyop2/global_kernel.py‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -61,6 +61,26 @@ def cache_key(self):
6161
return type(self), self.base_map.cache_key, tuple(self.permutation)
6262

6363

64+
@dataclass(eq=False, init=False)
65+
class ComposedMapKernelArg:
66+
"""Class representing a composed map input to the kernel.
67+
68+
:param base_maps: An arbitrary combination of :class:`MapKernelArg`s, :class:`PermutedMapKernelArg`s, and :class:`ComposedMapKernelArg`s.
69+
"""
70+
71+
def __init__(self, *base_maps):
72+
self.base_maps = base_maps
73+
74+
def __post_init__(self):
75+
for m in self.base_maps:
76+
if not isinstance(m, (MapKernelArg, PermutedMapKernelArg, ComposedMapKernelArg)):
77+
raise TypeError("base_maps must be a combination of MapKernelArgs, PermutedMapKernelArgs, and ComposedMapKernelArgs")
78+
79+
@property
80+
def cache_key(self):
81+
return type(self), tuple(m.cache_key for m in self.base_maps)
82+
83+
6484
@dataclass(frozen=True)
6585
class GlobalKernelArg:
6686
"""Class representing a :class:`pyop2.types.Global` being passed to the kernel.

‎pyop2/op2.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -41,7 +41,7 @@
4141

4242
from pyop2.types import (
4343
Set, ExtrudedSet, MixedSet, Subset, DataSet, MixedDataSet,
44-
Map, MixedMap, PermutedMap, Sparsity, Halo,
44+
Map, MixedMap, PermutedMap, ComposedMap, Sparsity, Halo,
4545
Global, GlobalDataSet,
4646
Dat, MixedDat, DatView, Mat
4747
)
@@ -64,7 +64,7 @@
6464
'MixedSet', 'Subset', 'DataSet', 'GlobalDataSet', 'MixedDataSet',
6565
'Halo', 'Dat', 'MixedDat', 'Mat', 'Global', 'Map', 'MixedMap',
6666
'Sparsity', 'parloop', 'Parloop', 'ParLoop', 'par_loop',
67-
'DatView', 'PermutedMap']
67+
'DatView', 'PermutedMap', 'ComposedMap']
6868

6969

7070
_initialised = False

‎pyop2/parloop.py‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
MatKernelArg, MixedMatKernelArg, GlobalKernel)
1717
from pyop2.local_kernel import LocalKernel, CStringLocalKernel, CoffeeLocalKernel, LoopyLocalKernel
1818
from pyop2.types import (Access, Global, Dat, DatView, MixedDat, Mat, Set,
19-
MixedSet, ExtrudedSet, Subset, Map, MixedMap)
19+
MixedSet, ExtrudedSet, Subset, Map, ComposedMap, MixedMap)
2020
from pyop2.utils import cached_property
2121

2222

@@ -25,7 +25,10 @@ class ParloopArg(abc.ABC):
2525
@staticmethod
2626
def check_map(m):
2727
if configuration["type_check"]:
28-
if m.iterset.total_size > 0 and len(m.values_with_halo) == 0:
28+
if isinstance(m, ComposedMap):
29+
for m_ in m.maps_:
30+
ParloopArg.check_map(m_)
31+
elif m.iterset.total_size > 0 and len(m.values_with_halo) == 0:
2932
raise MapValueError(f"{m} is not initialized")
3033

3134

‎pyop2/sparsity.pyx‎

Lines changed: 63 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,9 @@ cdef extern from "petsc.h":
5151
PETSC_INSERT_VALUES "INSERT_VALUES"
5252
int PetscCalloc1(size_t, void*)
5353
int PetscMalloc1(size_t, void*)
54+
int PetscMalloc2(size_t, void*, size_t, void*)
5455
int PetscFree(void*)
56+
int PetscFree2(void*,void*)
5557
int MatSetValuesBlockedLocal(PETSc.PetscMat, PetscInt, PetscInt*, PetscInt, PetscInt*,
5658
PetscScalar*, PetscInsertMode)
5759
int MatSetValuesLocal(PETSc.PetscMat, PetscInt, PetscInt*, PetscInt, PetscInt*,
@@ -193,7 +195,9 @@ def fill_with_zeros(PETSc.Mat mat not None, dims, maps, iteration_regions, set_d
193195
PetscScalar zero = 0.0
194196
PetscInt nrow, ncol
195197
PetscInt rarity, carity, tmp_rarity, tmp_carity
196-
PetscInt[:, ::1] rmap, cmap
198+
PetscInt[:, ::1] rmap, cmap, tempmap
199+
PetscInt **rcomposedmaps = NULL, **ccomposedmaps = NULL
200+
PetscInt nrcomposedmaps = 0, nccomposedmaps = 0, rset_entry, cset_entry
197201
PetscInt *rvals
198202
PetscInt *cvals
199203
PetscInt *roffset
@@ -213,23 +217,52 @@ def fill_with_zeros(PETSc.Mat mat not None, dims, maps, iteration_regions, set_d
213217
set_size = pair[0].iterset.size
214218
if set_size == 0:
215219
continue
216-
# Memoryviews require writeable buffers
217-
rflag = set_writeable(pair[0])
218-
cflag = set_writeable(pair[1])
219-
# Map values
220-
rmap = pair[0].values_with_halo
221-
cmap = pair[1].values_with_halo
220+
rflags = []
221+
cflags = []
222+
if isinstance(pair[0], op2.ComposedMap):
223+
m = pair[0].flattened_maps[0]
224+
rflags.append(set_writeable(m))
225+
rmap = m.values_with_halo
226+
nrcomposedmaps = len(pair[0].flattened_maps) - 1
227+
else:
228+
rflags.append(set_writeable(pair[0])) # Memoryviews require writeable buffers
229+
rmap = pair[0].values_with_halo # Map values
230+
if isinstance(pair[1], op2.ComposedMap):
231+
m = pair[1].flattened_maps[0]
232+
cflags.append(set_writeable(m))
233+
cmap = m.values_with_halo
234+
nccomposedmaps = len(pair[1].flattened_maps) - 1
235+
else:
236+
cflags.append(set_writeable(pair[1]))
237+
cmap = pair[1].values_with_halo
238+
# Handle ComposedMaps
239+
CHKERR(PetscMalloc2(nrcomposedmaps, &rcomposedmaps, nccomposedmaps, &ccomposedmaps))
240+
for i in range(nrcomposedmaps):
241+
m = pair[0].flattened_maps[1 + i]
242+
rflags.append(set_writeable(m))
243+
tempmap = m.values_with_halo
244+
rcomposedmaps[i] = &tempmap[0, 0]
245+
for i in range(nccomposedmaps):
246+
m = pair[1].flattened_maps[1 + i]
247+
cflags.append(set_writeable(m))
248+
tempmap = m.values_with_halo
249+
ccomposedmaps[i] = &tempmap[0, 0]
222250
# Arity of maps
223251
rarity = pair[0].arity
224252
carity = pair[1].arity
225-
226253
if not extruded:
227254
# The non-extruded case is easy, we just walk over the
228255
# rmap and cmap entries and set a block of values.
229256
CHKERR(PetscCalloc1(rarity*carity*rdim*cdim, &values))
230257
for set_entry in range(set_size):
231-
CHKERR(MatSetValuesBlockedLocal(mat.mat, rarity, &rmap[set_entry, 0],
232-
carity, &cmap[set_entry, 0],
258+
rset_entry = <PetscInt>set_entry
259+
cset_entry = <PetscInt>set_entry
260+
for i in range(nrcomposedmaps):
261+
rset_entry = rcomposedmaps[nrcomposedmaps - 1 - i][rset_entry]
262+
for i in range(nccomposedmaps):
263+
cset_entry = ccomposedmaps[nccomposedmaps - 1 - i][cset_entry]
264+
CHKERR(MatSetValuesBlockedLocal(mat.mat, rarity, &rmap[<int>rset_entry, 0],
265+
carity, &cmap[<int>cset_entry, 0],
233266
values, PETSC_INSERT_VALUES))
234267
else:
235268
# The extruded case needs a little more work.
@@ -268,6 +301,12 @@ def fill_with_zeros(PETSc.Mat mat not None, dims, maps, iteration_regions, set_d
268301
for i in range(carity):
269302
coffset[i] = pair[1].offset[i]
270303
for set_entry in range(set_size):
304+
rset_entry = <PetscInt>set_entry
305+
cset_entry = <PetscInt>set_entry
306+
for i in range(nrcomposedmaps):
307+
rset_entry = rcomposedmaps[nrcomposedmaps - 1 - i][rset_entry]
308+
for i in range(nccomposedmaps):
309+
cset_entry = ccomposedmaps[nccomposedmaps - 1 - i][cset_entry]
271310
if constant_layers:
272311
layer_start = layers[0, 0]
273312
layer_end = layers[0, 1] - 1
@@ -287,15 +326,15 @@ def fill_with_zeros(PETSc.Mat mat not None, dims, maps, iteration_regions, set_d
287326

288327
# In the case of tmp_rarity == rarity this is just:
289328
#
290-
# rvals[i] = rmap[set_entry, i] + layer_start * roffset[i]
329+
# rvals[i] = rmap[rset_entry, i] + layer_start * roffset[i]
291330
#
292331
# But this means less special casing.
293332
for i in range(tmp_rarity):
294-
rvals[i] = rmap[set_entry, i % rarity] + \
333+
rvals[i] = rmap[<int>rset_entry, i % rarity] + \
295334
(layer_start - layer_bottom + i // rarity) * roffset[i % rarity]
296335
# Ditto
297336
for i in range(tmp_carity):
298-
cvals[i] = cmap[set_entry, i % carity] + \
337+
cvals[i] = cmap[<int>cset_entry, i % carity] + \
299338
(layer_start - layer_bottom + i // carity) * coffset[i % carity]
300339
for layer in range(layer_start, layer_end):
301340
CHKERR(MatSetValuesBlockedLocal(mat.mat, tmp_rarity, rvals,
@@ -310,6 +349,15 @@ def fill_with_zeros(PETSc.Mat mat not None, dims, maps, iteration_regions, set_d
310349
CHKERR(PetscFree(cvals))
311350
CHKERR(PetscFree(roffset))
312351
CHKERR(PetscFree(coffset))
313-
restore_writeable(pair[0], rflag)
314-
restore_writeable(pair[1], cflag)
352+
CHKERR(PetscFree2(rcomposedmaps, ccomposedmaps))
353+
if isinstance(pair[0], op2.ComposedMap):
354+
for m, rflag in zip(pair[0].flattened_maps, rflags):
355+
restore_writeable(m, rflag)
356+
else:
357+
restore_writeable(pair[0], rflags[0])
358+
if isinstance(pair[1], op2.ComposedMap):
359+
for m, cflag in zip(pair[1].flattened_maps, cflags):
360+
restore_writeable(m, cflag)
361+
else:
362+
restore_writeable(pair[1], cflags[0])
315363
CHKERR(PetscFree(values))

‎pyop2/types/map.py‎

Lines changed: 94 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -149,6 +149,13 @@ def __le__(self, o):
149149
"""self<=o if o equals self or self._parent <= o."""
150150
return self == o
151151

152+
@utils.cached_property
153+
def flattened_maps(self):
154+
"""Return all component maps.
155+
156+
This is useful to flatten nested :class:`ComposedMap`s."""
157+
return (self, )
158+
152159

153160
class PermutedMap(Map):
154161
"""Composition of a standard :class:`Map` with a constant permutation.
@@ -173,6 +180,10 @@ class PermutedMap(Map):
173180
want two global-sized data structures.
174181
"""
175182
def __init__(self, map_, permutation):
183+
if not isinstance(map_, Map):
184+
raise TypeError("map_ must be a Map instance")
185+
if isinstance(map_, ComposedMap):
186+
raise NotImplementedError("PermutedMap of ComposedMap not implemented: simply permute before composing")
176187
self.map_ = map_
177188
self.permutation = np.asarray(permutation, dtype=Map.dtype)
178189
assert (np.unique(permutation) == np.arange(map_.arity, dtype=Map.dtype)).all()
@@ -192,6 +203,85 @@ def __getattr__(self, name):
192203
return getattr(self.map_, name)
193204

194205

206+
class ComposedMap(Map):
207+
"""Composition of :class:`Map`s, :class:`PermutedMap`s, and/or :class:`ComposedMap`s.
208+
209+
:arg maps_: The maps to compose.
210+
211+
Where normally staging to element data is performed as
212+
213+
.. code-block::
214+
215+
local[i] = global[map[i]]
216+
217+
With a :class:`ComposedMap` we instead get
218+
219+
.. code-block::
220+
221+
local[i] = global[maps_[0][maps_[1][maps_[2][...[i]]]]]
222+
223+
This might be useful if the map you want can be represented by
224+
a composition of existing maps.
225+
"""
226+
def __init__(self, *maps_, name=None):
227+
if not all(isinstance(m, Map) for m in maps_):
228+
raise TypeError("All maps must be Map instances")
229+
for tomap, frommap in zip(maps_[:-1], maps_[1:]):
230+
if tomap.iterset is not frommap.toset:
231+
raise ex.MapTypeError("tomap.iterset must match frommap.toset")
232+
if tomap.comm is not frommap.comm:
233+
raise ex.MapTypeError("All maps needs to share a communicator")
234+
if frommap.arity != 1:
235+
raise ex.MapTypeError("frommap.arity must be 1")
236+
self._iterset = maps_[-1].iterset
237+
self._toset = maps_[0].toset
238+
self.comm = self._toset.comm
239+
self._arity = maps_[0].arity
240+
# Don't call super().__init__() to avoid calling verify_reshape()
241+
self._values = None
242+
self.shape = (self._iterset.total_size, self._arity)
243+
self._name = name or "cmap_#x%x" % id(self)
244+
self._offset = maps_[0]._offset
245+
# A cache for objects built on top of this map
246+
self._cache = {}
247+
self.maps_ = tuple(maps_)
248+
249+
@utils.cached_property
250+
def _kernel_args_(self):
251+
return tuple(itertools.chain(*[m._kernel_args_ for m in self.maps_]))
252+
253+
@utils.cached_property
254+
def _wrapper_cache_key_(self):
255+
return tuple(m._wrapper_cache_key_ for m in self.maps_)
256+
257+
@utils.cached_property
258+
def _global_kernel_arg(self):
259+
from pyop2.global_kernel import ComposedMapKernelArg
260+
261+
return ComposedMapKernelArg(*(m._global_kernel_arg for m in self.maps_))
262+
263+
@utils.cached_property
264+
def values(self):
265+
raise RuntimeError("ComposedMap does not store values directly")
266+
267+
@utils.cached_property
268+
def values_with_halo(self):
269+
raise RuntimeError("ComposedMap does not store values directly")
270+
271+
def __str__(self):
272+
return "OP2 ComposedMap of Maps: [%s]" % ",".join([str(m) for m in self.maps_])
273+
274+
def __repr__(self):
275+
return "ComposedMap(%s)" % ",".join([repr(m) for m in self.maps_])
276+
277+
def __le__(self, o):
278+
raise NotImplementedError("__le__ not implemented for ComposedMap")
279+
280+
@utils.cached_property
281+
def flattened_maps(self):
282+
return tuple(itertools.chain(*(m.flattened_maps for m in self.maps_)))
283+
284+
195285
class MixedMap(Map, caching.ObjectCached):
196286
r"""A container for a bag of :class:`Map`\s."""
197287

@@ -315,3 +405,7 @@ def __str__(self):
315405

316406
def __repr__(self):
317407
return "MixedMap(%r)" % (self._maps,)
408+
409+
@utils.cached_property
410+
def flattened_maps(self):
411+
raise NotImplementedError("flattend_maps should not be necessary for MixedMap")

‎pyop2/types/mat.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
from pyop2.types.access import Access
1919
from pyop2.types.data_carrier import DataCarrier
2020
from pyop2.types.dataset import DataSet, GlobalDataSet, MixedDataSet
21-
from pyop2.types.map import Map
21+
from pyop2.types.map import Map, ComposedMap
2222
from pyop2.types.set import MixedSet, Set, Subset
2323

2424

@@ -165,7 +165,7 @@ def _process_args(cls, dsets, maps, *, iteration_regions=None, name=None, nest=N
165165
if not isinstance(m, Map):
166166
raise ex.MapTypeError(
167167
"All maps must be of type map, not type %r" % type(m))
168-
if len(m.values_with_halo) == 0 and m.iterset.total_size > 0:
168+
if not isinstance(m, ComposedMap) and len(m.values_with_halo) == 0 and m.iterset.total_size > 0:
169169
raise ex.MapValueError(
170170
"Unpopulated map values when trying to build sparsity.")
171171
# Make sure that the "to" Set of each map in a pair is the set of

0 commit comments

Comments
 (0)