99from devito .exceptions import CompilationError
1010from devito .ir import (
1111 Backward , Cluster , Forward , GuardBound , GuardFactor , InitArray , Interval ,
12- IntervalGroup , IterationSpace , Properties , Queue , Vector , detect_halo_writes ,
13- lower_exprs , vmax , vmin
12+ IterationSpace , Properties , Queue , Vector , detect_halo_writes , lower_exprs , vmax , vmin
1413)
1514from devito .logger import warning
1615from devito .passes .clusters .utils import is_memcpy
@@ -496,61 +495,61 @@ def expand_halo_transfers(clusters, mapper):
496495 writes. For example, `usave` in `Eq(usave, u)` must eventually receive `u`'s
497496 populated HALO if a preceding `Eq` writes into `u`'s HALO.
498497 """
499- buffered = {f for f , _ in mapper }
500- if not buffered :
498+ if not mapper :
501499 return clusters
502500
501+ # Get HALO writes along the buffered dimensions
503502 bdims = set ()
504503 for b in mapper .values ():
505504 bdims .update (d for d in b .dimensions if not isinstance (d , BufferDimension ))
506- key = lambda d : d in bdims
507505
508506 halo_writes = set ()
509507 for c in clusters :
510- for w in detect_halo_writes (c , key ):
508+ for w in detect_halo_writes (c , bdims . __contains__ ):
511509 halo_writes .add (w .function )
510+ if not halo_writes :
511+ return clusters
512512
513+ # Expand the IterationSpace over the necessary amount of HALO; in doing so,
514+ # check the expanded footprint of every access, including shifted reads.
515+ # Writes must be pointwise so that the whole destination halo is filled
516+ buffered = {f for f , _ in mapper }
513517 processed = []
514518 for c in clusters :
515- scope = c .scope
516- targets = set (scope .writes ) & buffered
517- if c .is_wild or not targets or not halo_writes .intersection (scope .reads ):
519+ writes = c .scope .writes_tensor
520+
521+ if c .is_wild or \
522+ writes .isdisjoint (buffered ) or \
523+ halo_writes .isdisjoint (c .scope .reads ):
518524 processed .append (c )
519525 continue
520526
521- if scope . writes_tensor != targets :
527+ if not writes <= buffered :
522528 raise CompilationError (
523529 "Cannot expand a mixed Cluster over the halo while buffering"
524530 )
525531
526532 ispace = c .ispace
527- for f in targets :
533+ for f in writes :
528534 ispace = _include_halo (ispace , f )
529535
530- # Check the expanded footprint of every access, including shifted reads.
531- # Writes must be pointwise so that the whole destination halo is filled
532- for a in scope .accesses :
536+ for a in c .scope .accesses :
533537 f = a .function
534- if not f .is_AbstractFunction :
535- continue
536-
537- for d in f .dimensions :
538- if not key (d ):
539- continue
540- if d not in ispace .dimensions :
541- raise CompilationError (
542- f"Cannot expand access to `{ f .name } ` over the halo"
543- )
544538
539+ for d in bdims .intersection (a .findices ):
545540 size = f ._size_nodomain [d ]
546- offset = simplify (a [d ] - d )
547- if not is_integer (offset ) or (a .is_write and offset != size .left ):
541+ offset = simplify (a [d ] - d - size .left )
542+
543+ if d not in ispace .dimensions or \
544+ not is_integer (offset ) or \
545+ (a .is_write and offset != 0 ):
548546 raise CompilationError (
549- f"Cannot expand non-pointwise access to `{ f .name } ` over the halo"
547+ f"Cannot expand access to `{ f .name } ` over the halo"
550548 )
551549
552550 i = ispace [d ]
553- if i .lower + offset < 0 or i .upper + offset > sum (size ):
551+ if i .lower + offset < - size .left or \
552+ i .upper + offset > size .right :
554553 raise CompilationError (
555554 f"Insufficient halo for `{ f .name } ` in buffered write"
556555 )
@@ -562,10 +561,10 @@ def expand_halo_transfers(clusters, mapper):
562561
563562def _include_halo (ispace , f ):
564563 """Extend `ispace` to include `f`'s HALO."""
565- ihalo = IntervalGroup ( [
564+ ihalo = [
566565 Interval (i .dim , - f ._size_halo [i .dim ].left , f ._size_halo [i .dim ].right , i .stamp )
567566 for i in ispace if i .dim in f .dimensions
568- ])
567+ ]
569568
570569 return IterationSpace .union (ispace , IterationSpace (ihalo ))
571570
0 commit comments