@@ -311,7 +311,11 @@ def _optimize(self, clusters, descriptors):
311311 processed .append (c )
312312 continue
313313
314- key1 = lambda d : not d ._defines & v .dim ._defines # noqa: B023
314+ # Lift only along the buffer's own spatial Dimensions, not
315+ # unrelated ones (e.g. sparse points of an indirect read)
316+ key1 = lambda d : (not d ._defines & v .dim ._defines and # noqa: B023
317+ any (d ._defines & bd ._defines # noqa: B023
318+ for bd in v .bdims )) # noqa: B023
315319 dims = c .ispace .project (key1 ).itdims
316320 ispace = c .ispace .lift (dims , key0 ())
317321 processed .append (c .rebuild (ispace = ispace ))
@@ -331,8 +335,8 @@ def _make_task_groups(self, descriptors):
331335 return stamps
332336
333337 task_sets = (
334- [b for b , v in descriptors .items () if v .is_writeonly ],
335- [b for b , v in descriptors .items () if v .is_readonly ],
338+ [b for b , v in descriptors .items () if any ( vb .is_writeonly for vb in v ) ],
339+ [b for b , v in descriptors .items () if any ( vb .is_readonly for vb in v ) ],
336340 )
337341
338342 for tasks in task_sets :
@@ -519,30 +523,37 @@ def itdims(self):
519523 def ispace (self ):
520524 # The IterationSpace within which the buffer will be accessed
521525
522- # NOTE: The `key` is to avoid Clusters including `f` but not directly
523- # using it in an expression, such as HaloTouch Clusters
524- def key (c ):
525- bufferdim = any (i in c .ispace .dimensions for i in self .bdims )
526- xd_only = all (d ._defines & self .xd ._defines for d in c .ispace .dimensions )
527- return bufferdim or xd_only
528-
529526 ispaces = set ()
527+ indirect = set ()
530528 for c in self .clusters :
531- if not key (c ):
532- continue
533-
534- # Skip wild clusters (e.g. HaloTouch Clusters)
529+ # Skip wild clusters (e.g. HaloTouch Clusters), which include `f`
530+ # but do not directly use it in an expression
535531 if c .is_wild :
536532 continue
537533
538- # Iterations space and buffering dims
539- edims = [d for d in self .bdims if d not in c .ispace .dimensions ]
540- if not edims :
541- ispaces .add (c .ispace )
542- else :
543- # Add all missing buffering dimensions and reorder to
544- # avoid duplicates with different ordering
545- ispaces .add (c .ispace .insert (self .dim , edims ).reorder ())
534+ bufferdim = any (i in c .ispace .dimensions for i in self .bdims )
535+ xd_only = all (d ._defines & self .xd ._defines for d in c .ispace .dimensions )
536+
537+ if bufferdim or xd_only :
538+ # `c` iterates (at least some of) the buffer's spatial dims
539+ edims = [d for d in self .bdims if d not in c .ispace .dimensions ]
540+ if not edims :
541+ ispaces .add (c .ispace )
542+ else :
543+ # Add all missing buffering dimensions and reorder to
544+ # avoid duplicates with different ordering
545+ ispaces .add (c .ispace .insert (self .dim , edims ).reorder ())
546+ elif ((self .f in c .scope .reads or self .f in c .scope .writes ) and
547+ self .dim .root in c .ispace .dimensions ):
548+ # `c` accesses `f` indirectly (e.g. interpolation), so it doesn't
549+ # iterate `bdims`; span the buffer's own Dimensions instead
550+ tispace = c .ispace .project (lambda i : i ._defines & self .dim .root ._defines )
551+ indirect .add (tispace .insert (self .dim , list (self .bdims )))
552+
553+ # Indirect accessors define the ispace only for a read-only streamed
554+ # buffer, where nothing iterates the buffer's own Dimensions directly
555+ if not ispaces :
556+ ispaces = indirect
546557
547558 if len (ispaces ) > 1 :
548559 # Best effort to make buffering work in the presence of multiple
@@ -663,7 +674,11 @@ def last_idx(self):
663674 mapper = {}
664675 func = vmax if self .is_forward_buffering else vmin
665676 for c in self .lastwrite + self .firstread :
666- indices = extract_indices (self .f , self .dim , [c ])
677+ # Consider all Clusters sharing `c`'s guards, so the leading edge is
678+ # found even when `f` is accessed at different offsets across them
679+ # (e.g. `f` and `f.forward` in separate Eqs)
680+ group = [c1 for c1 in self .clusters if c1 .guards == c .guards ]
681+ indices = extract_indices (self .f , self .dim , group )
667682 idx = func (* [Vector (i ) for i in indices ])[0 ]
668683 mapper [c ] = idx
669684
0 commit comments