diff --git a/.basedpyright/baseline.json b/.basedpyright/baseline.json index 081ac5ef3..4be4ff4c0 100644 --- a/.basedpyright/baseline.json +++ b/.basedpyright/baseline.json @@ -20876,22 +20876,6 @@ "endColumn": 47, "lineCount": 1 } - }, - { - "code": "reportUnknownMemberType", - "range": { - "startColumn": 18, - "endColumn": 30, - "lineCount": 1 - } - }, - { - "code": "reportUnknownArgumentType", - "range": { - "startColumn": 18, - "endColumn": 30, - "lineCount": 1 - } } ], "./loopy/library/function.py": [ @@ -55770,6 +55754,24 @@ } } ], + "./loopy/transform/loop_fusion.py": [ + { + "code": "reportUnknownMemberType", + "range": { + "startColumn": 22, + "endColumn": 34, + "lineCount": 1 + } + }, + { + "code": "reportUnknownArgumentType", + "range": { + "startColumn": 22, + "endColumn": 34, + "lineCount": 1 + } + } + ], "./loopy/transform/pack_and_unpack_args.py": [ { "code": "reportUnknownArgumentType", diff --git a/loopy/kernel/tools.py b/loopy/kernel/tools.py index 63db4b634..2402dfe92 100644 --- a/loopy/kernel/tools.py +++ b/loopy/kernel/tools.py @@ -24,13 +24,9 @@ THE SOFTWARE. """ -import dataclasses import itertools import logging -import operator import sys -from collections.abc import Set as AbstractSet -from functools import reduce from sys import intern from typing import ( TYPE_CHECKING, @@ -43,9 +39,7 @@ from typing_extensions import override import namedisl as nisl -import pymbolic.primitives as p -from pymbolic import Expression -from pytools import fset_union, memoize_on_first_arg, natsorted, set_union +from pytools import fset_union, memoize_on_first_arg, natsorted from loopy.diagnostic import LoopyError, warn_with_kernel from loopy.kernel import LoopKernel @@ -65,12 +59,20 @@ TUnitOrKernelT, for_each_kernel, ) +from loopy.typing import not_none if TYPE_CHECKING: - from collections.abc import Collection, Iterable, Mapping, Sequence + from collections.abc import ( + Collection, + Iterable, + Mapping, + Sequence, + Set as AbstractSet, + ) - from pymbolic import ArithmeticExpression + import pymbolic.primitives as p + from pymbolic import ArithmeticExpression, Expression from pytools.tag import Tag from loopy.types import ToLoopyTypeConvertible @@ -2155,67 +2157,28 @@ def get_hw_axis_base_for_codegen(kernel: LoopKernel, iname: str) -> nisl.Aff: # {{{ get access map from an instruction -@dataclasses.dataclass -class _IndexCollector(CombineMapper[AbstractSet[tuple[Expression, ...]], []]): - var: str - - def __post_init__(self) -> None: - super().__init__() - - @override - def combine(self, - values: Iterable[AbstractSet[tuple[Expression, ...]]] - ) -> AbstractSet[tuple[Expression, ...]]: - return set_union(values) - - @override - def map_subscript(self, expr: p.Subscript) -> AbstractSet[tuple[Expression, ...]]: - assert isinstance(expr.aggregate, p.Variable) - if expr.aggregate.name == self.var: - return (super().map_subscript(expr) | frozenset([expr.index_tuple])) - else: - return super().map_subscript(expr) - - @override - def map_algebraic_leaf( - self, expr: p.AlgebraicLeaf, - ) -> frozenset[tuple[Expression, ...]]: - return frozenset() - - @override - def map_constant( - self, expr: object - ) -> frozenset[tuple[Expression, ...]]: - return frozenset() - - -def _union_amaps(amaps: Sequence[nisl.Map]): - return reduce(operator.or_, amaps[1:], amaps[0]) - - -def get_insn_access_map(kernel: LoopKernel, insn_id: str, var: str): +def get_insn_access_maps( + kernel: LoopKernel, insn_id: str, var: str) -> list[nisl.Map]: from loopy.match import Id - from loopy.symbolic import get_access_map from loopy.transform.subst import expand_subst + kernel = expand_subst(kernel, within=Id(insn_id)) + insn = kernel.id_to_insn[insn_id] + insn_inames = kernel.insn_inames(insn) - kernel = expand_subst(kernel, within=Id(insn_id)) - indices = tuple( - _IndexCollector(var)( - (insn.expression, insn.assignees, tuple(insn.predicates)) - ) - ) + from loopy.symbolic import BatchedAccessMapMapper + bamm = BatchedAccessMapMapper(kernel, [var]) + bamm((insn.expression, insn.assignees, tuple(insn.predicates)), insn_inames) - amaps = [ - get_access_map( - kernel.get_inames_domain(insn.within_inames), - idx, kernel.assumptions - ) - for idx in indices - ] + if var in bamm.bad_subscripts: + from loopy.diagnostic import UnableToDetermineAccessRangeError + raise UnableToDetermineAccessRangeError( + f"cannot determine access range for '{var}' in '{insn_id}'") - return _union_amaps(amaps) + return [ + not_none(amap) + for amap in bamm.access_maps[var].values()] # }}} diff --git a/loopy/symbolic.py b/loopy/symbolic.py index 1dd729188..f1f3cd7fa 100644 --- a/loopy/symbolic.py +++ b/loopy/symbolic.py @@ -2935,6 +2935,11 @@ def __init__(self, self._var_names = set(var_names) super().__init__() + @memoize_method + def _get_inames_domain_strict(self, inames: AbstractSet[str]) -> nisl.Set: + return self.kernel.get_inames_domain(inames).project_out_except( + inames, dim_type=DimType.out) + def get_access_range(self, var_name: str) -> nisl.Set | None: loops_to_amaps = self.access_maps[var_name] if not loops_to_amaps: @@ -2947,7 +2952,7 @@ def get_access_range(self, var_name: str) -> nisl.Set | None: @override def map_subscript(self, expr: p.Subscript, /, inames: AbstractSet[str]) -> None: - domain = self.kernel.get_inames_domain(inames) + domain = self._get_inames_domain_strict(inames) super().map_subscript(expr, inames) assert isinstance(expr.aggregate, p.Variable) diff --git a/loopy/transform/loop_fusion.py b/loopy/transform/loop_fusion.py index 3504dd0a0..99dbf033f 100644 --- a/loopy/transform/loop_fusion.py +++ b/loopy/transform/loop_fusion.py @@ -391,25 +391,39 @@ def _compute_isinfusible_via_access_map( """ from loopy.diagnostic import UnableToDetermineAccessRangeError - from loopy.kernel.tools import get_insn_access_map + from loopy.kernel.tools import get_insn_access_maps try: - amap_pred = get_insn_access_map(kernel, insn_pred, var) - amap_succ = get_insn_access_map(kernel, insn_succ, var) + amaps_pred = get_insn_access_maps(kernel, insn_pred, var) + amaps_succ = get_insn_access_maps(kernel, insn_succ, var) except UnableToDetermineAccessRangeError: # either predecessors or successors has a non-affine access i.e. # fallback to the safer option => infusible return True + amaps_pred = [ + amap.project_out_except( + outer_inames | {candidate_pred}, dim_type=DimType.in_) + for amap in amaps_pred] + amaps_succ = [ + amap.project_out_except( + outer_inames | {candidate_succ}, dim_type=DimType.in_) + for amap in amaps_succ] + + # amaps should have the same space after projecting out the inner loops, so they + # can safely be unioned + def union_amaps(amaps: Sequence[nisl.Map]) -> nisl.Map: + import operator + from functools import reduce + return reduce(operator.or_, amaps[1:], amaps[0]) + + amap_pred = union_amaps(amaps_pred) + amap_succ = union_amaps(amaps_succ) + storage_dim_names = get_access_map_storage_names(amap_pred) assert set(storage_dim_names) == amap_pred.space.out_names assert set(storage_dim_names) == amap_succ.space.out_names - amap_pred = amap_pred.project_out_except( - outer_inames | {candidate_pred}, dim_type=DimType.in_) - amap_succ = amap_succ.project_out_except( - outer_inames | {candidate_succ}, dim_type=DimType.in_) - amap_pred = amap_pred.move_dims(outer_inames, DimType.param) amap_succ = amap_succ.move_dims(outer_inames, DimType.param) diff --git a/test/test_loop_fusion.py b/test/test_loop_fusion.py index 2a1dd46ff..86162b8f3 100644 --- a/test/test_loop_fusion.py +++ b/test/test_loop_fusion.py @@ -497,6 +497,39 @@ def test_reduction_loop_fusion_with_multiple_redn_in_same_insn( lp.auto_test_vs_ref(ref_t_unit, ctx, t_unit.with_kernel(knl)) +def test_loop_fusion_with_inner_reduction(ctx_factory: cl.CtxFactory): + ctx = ctx_factory() + + t_unit = lp.make_kernel( + ["{[i0, j0]: 0 <= i0, j0 < 10}", + "{[i1]: 0 <= i1 < 10}", + # Intentionally keeping j1 separate from i1 to test for regression. See + # https://github.com/inducer/loopy/pull/1009 for details. + "{[j1]: 0 <= j1 < 10}"], + """ + a[i0, j0] = j0 * 1.0 {id=insn1} + out[i1] = sum(j1, a[i1, j1]) {id=insn2} + """, + ) + ref_t_unit = t_unit + + knl = t_unit.default_entrypoint + + fused_chunks = lp.get_kennedy_unweighted_fusion_candidates( + knl, frozenset(["i0", "i1"]) + ) + knl = lp.rename_inames_in_batch(knl, fused_chunks) + + assert ( + len( + knl.id_to_insn["insn1"].within_inames + & knl.id_to_insn["insn2"].within_inames + ) == 1 + ) + + lp.auto_test_vs_ref(ref_t_unit, ctx, t_unit.with_kernel(knl)) + + if __name__ == "__main__": if len(sys.argv) > 1: exec(sys.argv[1])