Skip to content

Commit c0c7779

Browse files
committed
TMP SPLIT REVIEW ME
1 parent beed5ed commit c0c7779

4 files changed

Lines changed: 55 additions & 16 deletions

File tree

loopy/check.py

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -862,10 +862,23 @@ def map_call(self, expr: p.Call, domain: nisl.Set, insn_id: str):
862862

863863
kw_to_pos, _ = get_kw_pos_association(subkernel)
864864

865-
arg_id_to_arg = self.kernel.id_to_insn[insn_id].arg_id_to_arg()
865+
insn = self.kernel.id_to_insn[insn_id]
866+
if isinstance(insn, CallInstruction):
867+
arg_id_to_arg = insn.arg_id_to_arg()
866868

867-
kwargs = {k: _mark_variables_from_caller(arg_id_to_arg[kw_to_pos[k]])
868-
for k in subkernel.get_unwritten_value_args()}
869+
kwargs = {k: _mark_variables_from_caller(
870+
cast("ArithmeticExpression",
871+
arg_id_to_arg[kw_to_pos[k]]))
872+
for k in subkernel.get_unwritten_value_args()}
873+
else:
874+
# The call appears as a plain expression (e.g. on the right
875+
# hand side of an assignment); its arguments are found in
876+
# *expr* itself.
877+
kwargs = {k: _mark_variables_from_caller(
878+
cast("ArithmeticExpression",
879+
expr.parameters[kw_to_pos[k]]))
880+
for k in subkernel.get_unwritten_value_args()
881+
if kw_to_pos[k] < len(expr.parameters)}
869882

870883
kw_space = nisl.Space.from_names(
871884
out=[], param=[*get_dependencies(tuple(kwargs.values())),

loopy/kernel/creation.py

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2500,8 +2500,26 @@ def make_function(
25002500
assumptions_space = nisl.Space.from_names(param=dom0.space.param_names)
25012501
assumptions = nisl.Set.universe(assumptions_space)
25022502
elif isinstance(assumptions, str):
2503+
# Value args are also legitimate assumption parameters, even when
2504+
# they do not appear in any domain (e.g. when they only occur in
2505+
# an array shape). For string declarations the eventual argument
2506+
# type is not known at this point, so declare all of their names.
2507+
import re
2508+
2509+
param_names = set(get_outer_params(parsed_domains))
2510+
for arg in kernel_args:
2511+
if isinstance(arg, str):
2512+
# kernel_data strings are split on commas up front, so
2513+
# fragments (e.g. " 2]" of "x: float64[n, 2]") may end up
2514+
# here; only valid names are useful.
2515+
arg = arg.strip()
2516+
if re.fullmatch(r"[a-zA-Z_][a-zA-Z0-9_]*", arg):
2517+
param_names.add(arg)
2518+
elif isinstance(arg, ValueArg):
2519+
param_names.add(arg.name)
2520+
25032521
assumptions_set_str = "[%s] -> { : %s}" \
2504-
% (",".join(s for s in get_outer_params(parsed_domains)),
2522+
% (",".join(sorted(param_names)),
25052523
assumptions)
25062524
assumptions = nisl.make_set(assumptions_set_str)
25072525
else:

loopy/transform/callable.py

Lines changed: 16 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -204,9 +204,10 @@ def substitute_into_domain(
204204
param_name: str,
205205
expr: ArithmeticExpression,
206206
allowed_param_dims: Collection[str],
207-
):
207+
) -> nisl.Set:
208208
"""
209-
:arg allowed_deps: A :class:`list` of :class:`str` that are
209+
:arg allowed_param_dims: Names that may be introduced as new parameters
210+
of *domain* (i.e. parameters of the caller).
210211
"""
211212
import pymbolic.primitives as prim
212213

@@ -217,24 +218,29 @@ def substitute_into_domain(
217218

218219
# {{{ rename 'param_name' to avoid namespace pollution with allowed_param_dims
219220

220-
domain = domain.rename_dims([(
221-
param_name, UniqueNameGenerator(set(allowed_param_dims))(param_name))])
221+
renamed_param_name = UniqueNameGenerator(set(allowed_param_dims))(param_name)
222+
domain = domain.rename_dims([(param_name, renamed_param_name)])
222223

223224
# }}}
224225

225-
domain.add_dims(DimType.param, [
226-
dep for dep in get_dependencies(expr)
227-
if dep in allowed_param_dims
228-
])
226+
new_param_names: list[str] = []
227+
for dep in get_dependencies(expr):
228+
if dep in allowed_param_dims:
229+
new_param_names.append(dep)
230+
else:
231+
raise ValueError("Augmenting caller's domain "
232+
f"with '{dep}' is not allowed.")
233+
234+
domain = domain.add_dims(DimType.param, new_param_names)
229235

230236
set_ = isl_set_from_expr(domain.var_pw_affs,
231-
prim.Comparison(prim.Variable(param_name),
237+
prim.Comparison(prim.Variable(renamed_param_name),
232238
"==",
233239
expr))
234240

235241
domain = domain & set_
236242

237-
return domain.project_out([param_name])
243+
return domain.project_out([renamed_param_name])
238244

239245

240246
def get_valid_domain_param_names(knl):

test/test_reduction.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -147,7 +147,8 @@ def test_multi_nested_dependent_reduction():
147147
lp.ValueArg("ntgts", np.int32),
148148
lp.ValueArg("nboxes", np.int32),
149149
],
150-
assumptions="ntgts>=1",
150+
# 'a' is indexed by 'itgt', so the bounds check needs 'ntgts <= n'
151+
assumptions="ntgts>=1 and ntgts <= n",
151152
target=lp.PyOpenCLTarget())
152153

153154
print(lp.generate_code_v2(knl).device_code())
@@ -179,7 +180,8 @@ def test_recursive_nested_dependent_reduction():
179180
lp.ValueArg("ntgts", np.int32),
180181
lp.ValueArg("nboxes", np.int32),
181182
],
182-
assumptions="ntgts>=1",
183+
# 'a' is indexed by 'itgt', so the bounds check needs 'ntgts <= n'
184+
assumptions="ntgts>=1 and ntgts <= n",
183185
target=lp.PyOpenCLTarget())
184186

185187
print(lp.generate_code_v2(knl).device_code())

0 commit comments

Comments
 (0)