Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 18 additions & 16 deletions .basedpyright/baseline.json
Original file line number Diff line number Diff line change
Expand Up @@ -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": [
Expand Down Expand Up @@ -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",
Expand Down
89 changes: 26 additions & 63 deletions loopy/kernel/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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]:
Comment thread
inducer marked this conversation as resolved.
Comment thread
inducer marked this conversation as resolved.
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()]

# }}}

Expand Down
7 changes: 6 additions & 1 deletion loopy/symbolic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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)
Expand Down
30 changes: 22 additions & 8 deletions loopy/transform/loop_fusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
33 changes: 33 additions & 0 deletions test/test_loop_fusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
Expand Down
Loading