@@ -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
240246def get_valid_domain_param_names (knl ):
0 commit comments