|
15 | 15 | from devito.passes.clusters.cse import _cse |
16 | 16 | from devito.passes.clusters.utils import expose_tuning_knobs |
17 | 17 | from devito.symbolics import ( |
18 | | - Uxmapper, estimate_cost, retrieve_functions, reuse_if_untouched, search, sympy_dtype, |
19 | | - uxreplace |
| 18 | + DOUBLE, INT, Cast, Uxmapper, estimate_cost, retrieve_functions, reuse_if_untouched, |
| 19 | + search, sympy_dtype, uxreplace |
20 | 20 | ) |
21 | 21 | from devito.tools import ( |
22 | 22 | Reconstructable, Stamp, as_mapper, as_tuple, flatten, frozendict, generator, |
@@ -295,6 +295,10 @@ def _do_generate(self, exprs, exclude, cbk_search, cbk_compose=None): |
295 | 295 |
|
296 | 296 | class CireInvariants(CireTransformerLegacy, Queue): |
297 | 297 |
|
| 298 | + # Predicate on Cluster used to pick which ones this pass fires on. |
| 299 | + # Subclasses override to target a different kind of cluster. |
| 300 | + _cluster_filter = staticmethod(lambda c: c.is_dense) |
| 301 | + |
298 | 302 | def __init__(self, sregistry, options, platform): |
299 | 303 | super().__init__(sregistry, options, platform) |
300 | 304 |
|
@@ -324,7 +328,8 @@ def callback(self, clusters, prefix, xtracted=None): |
324 | 328 | key = lambda c: self._lookup_key(c, d) |
325 | 329 | processed = list(clusters) |
326 | 330 | for ak, group in as_mapper(clusters, key=key).items(): |
327 | | - g = [c for c in group if c.is_dense and c not in xtracted] |
| 331 | + g = [c for c in group |
| 332 | + if self._cluster_filter(c) and c not in xtracted] |
328 | 333 | if not g: |
329 | 334 | continue |
330 | 335 |
|
@@ -387,6 +392,73 @@ def _generate(self, cgroup, exclude): |
387 | 392 | yield self._do_generate(exprs, exclude, cbk_search) |
388 | 393 |
|
389 | 394 |
|
| 395 | +def _is_floor(e): |
| 396 | + return getattr(e, 'is_Function', False) and e.func.__name__ == 'floor' |
| 397 | + |
| 398 | + |
| 399 | +class CireInvariantsSparse(CireInvariants): |
| 400 | + |
| 401 | + """ |
| 402 | + Hoist sparse-point position temps `pos = (c - o)/h` and their |
| 403 | + `floor(pos)` out of the per-stencil-point inner loop into preamble |
| 404 | + Arrays computed once per source. The inner loop then reads |
| 405 | + tabulated values instead of recomputing `floor((c - o)/h)` for each |
| 406 | + `(rp_srcx, rp_srcy, rp_srcz)` combination. `Lift` then moves the |
| 407 | + preamble out of the time loop. |
| 408 | + """ |
| 409 | + |
| 410 | + _cluster_filter = staticmethod(lambda c: not c.is_dense) |
| 411 | + |
| 412 | + def _generate(self, cgroup, exclude): |
| 413 | + # Tabulate `pos` as fp64 and `INT(floor(pos))` as int32. The int32 |
| 414 | + # tab feeds both integer position lookups and (via DOUBLE cast) the |
| 415 | + # bare `floor(pos)` uses in `pos - floor(pos)`. |
| 416 | + counter = generator() |
| 417 | + make_f64 = lambda: Symbol(name=f'dummy{counter()}', dtype=np.float64) |
| 418 | + make_i32 = lambda: Symbol(name=f'dummy{counter()}', dtype=np.int32) |
| 419 | + |
| 420 | + mapper = Uxmapper() |
| 421 | + |
| 422 | + def _add(expr, make): |
| 423 | + if expr.is_commutative is False: |
| 424 | + return |
| 425 | + if {a.function for a in expr.free_symbols} & exclude: |
| 426 | + return |
| 427 | + mapper.add(expr, make, None) |
| 428 | + |
| 429 | + for e in cgroup.exprs: |
| 430 | + for f in search(e, _is_floor, 'all', 'bfs'): |
| 431 | + _add(f.args[0], make_f64) |
| 432 | + _add(INT(f), make_i32) |
| 433 | + |
| 434 | + yield mapper |
| 435 | + |
| 436 | + def _choose(self, aliases, cgroup, mapper): |
| 437 | + # Skip score-based filtering and fold bare `floor(pos)` onto the |
| 438 | + # int32 alias built for `INT(floor(pos))` via a `DOUBLE(...)` cast. |
| 439 | + exprs = cgroup.exprs |
| 440 | + |
| 441 | + aliases = AliasList(aliases) |
| 442 | + if not aliases: |
| 443 | + return exprs, aliases |
| 444 | + |
| 445 | + aliaseds = set(aliases.aliaseds) |
| 446 | + subs = {k: v for k, v in mapper.items() if v.free_symbols & aliaseds} |
| 447 | + |
| 448 | + for k, v in list(mapper.items()): |
| 449 | + if not isinstance(k, Cast): |
| 450 | + continue |
| 451 | + if not (isinstance(k.dtype, type) and |
| 452 | + issubclass(k.dtype, np.integer)): |
| 453 | + continue |
| 454 | + inner = k.base |
| 455 | + if _is_floor(inner) and (v.free_symbols & aliaseds): |
| 456 | + subs[inner] = DOUBLE(v) |
| 457 | + |
| 458 | + exprs = [uxreplace(e, subs) for e in exprs] |
| 459 | + return exprs, aliases |
| 460 | + |
| 461 | + |
390 | 462 | class CireDerivatives(CireTransformerLegacy): |
391 | 463 |
|
392 | 464 | def __init__(self, sregistry, options, platform): |
@@ -519,7 +591,8 @@ def _cbk_search2(self, expr, rank): |
519 | 591 | # Subpass mapper |
520 | 592 | modes = { |
521 | 593 | 'invariants': [CireInvariantsElementary, |
522 | | - CireInvariantsDivs], |
| 594 | + CireInvariantsDivs, |
| 595 | + CireInvariantsSparse], |
523 | 596 | 'eval-derivs': [CireEvalDerivatives], # NOTE: legacy pass |
524 | 597 | 'index-derivs': [CireIndexDerivatives], |
525 | 598 | } |
|
0 commit comments