Skip to content

Commit a8340f7

Browse files
committed
Argument passing/return: consider is_{input,output} instead of is_written
1 parent 7f57d1e commit a8340f7

3 files changed

Lines changed: 44 additions & 29 deletions

File tree

loopy/target/c/c_execution.py

Lines changed: 18 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -161,21 +161,28 @@ def generate_output_handler(
161161

162162
from loopy.kernel.data import KernelArgument
163163

164+
def is_output(idi):
165+
from loopy.kernel.array import ArrayBase
166+
if not issubclass(idi.arg_class, ArrayBase):
167+
return False
168+
169+
arg = kernel.impl_arg_to_arg[idi.name]
170+
return arg.is_output
171+
164172
if options.return_dict:
165173
gen("return None, {%s}"
166-
% ", ".join(f'"{arg.name}": {arg.name}'
167-
for arg in implemented_data_info
168-
if issubclass(arg.arg_class, KernelArgument)
169-
if arg.base_name in
170-
kernel.get_written_variables()))
174+
% ", ".join(f'"{idi.name}": {idi.name}'
175+
for idi in implemented_data_info
176+
if issubclass(idi.arg_class, KernelArgument)
177+
if is_output(idi)))
171178
else:
172-
out_args = [arg
173-
for arg in implemented_data_info
174-
if issubclass(arg.arg_class, KernelArgument)
175-
if arg.base_name in kernel.get_written_variables()]
176-
if out_args:
179+
out_idis = [idi
180+
for idi in implemented_data_info
181+
if issubclass(idi.arg_class, KernelArgument)
182+
if is_output(idi)]
183+
if out_idis:
177184
gen("return None, (%s,)"
178-
% ", ".join(arg.name for arg in out_args))
185+
% ", ".join(idi.name for idi in out_idis))
179186
else:
180187
gen("return None, ()")
181188

loopy/target/execution.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -414,7 +414,7 @@ def generate_arg_setup(
414414
if not options.no_numpy:
415415
self.handle_non_numpy_arg(gen, arg)
416416

417-
if not options.skip_arg_checks and not is_written:
417+
if not options.skip_arg_checks and kernel_arg.is_input:
418418
gen("if %s is None:" % arg.name)
419419
with Indentation(gen):
420420
gen("raise RuntimeError(\"input argument '%s' must "
@@ -441,7 +441,8 @@ def generate_arg_setup(
441441

442442
# {{{ allocate written arrays, if needed
443443

444-
if is_written and arg.arg_class in [lp.ArrayArg, lp.ConstantArg] \
444+
if kernel_arg.is_output \
445+
and arg.arg_class in [lp.ArrayArg, lp.ConstantArg] \
445446
and arg.shape is not None \
446447
and all(si is not None for si in arg.shape):
447448

loopy/target/pyopencl_execution.py

Lines changed: 23 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -202,40 +202,47 @@ def generate_output_handler(
202202

203203
from loopy.kernel.data import KernelArgument
204204

205+
def is_output(idi):
206+
from loopy.kernel.array import ArrayBase
207+
if not issubclass(idi.arg_class, ArrayBase):
208+
return False
209+
210+
arg = kernel.impl_arg_to_arg[idi.name]
211+
return arg.is_output
212+
205213
if not options.no_numpy:
206214
gen("if out_host is None and (_lpy_encountered_numpy "
207215
"and not _lpy_encountered_dev):")
208216
with Indentation(gen):
209217
gen("out_host = True")
210218

211-
for arg in implemented_data_info:
212-
if not issubclass(arg.arg_class, KernelArgument):
219+
for idi in implemented_data_info:
220+
if not issubclass(idi.arg_class, KernelArgument):
213221
continue
214222

215-
is_written = arg.base_name in kernel.get_written_variables()
216-
if is_written:
217-
np_name = "_lpy_%s_np_input" % arg.name
223+
if is_output(idi):
224+
np_name = "_lpy_%s_np_input" % idi.name
218225
gen("if out_host or %s is not None:" % np_name)
219226
with Indentation(gen):
220227
gen("%s = %s.get(queue=queue, ary=%s)"
221-
% (arg.name, arg.name, np_name))
228+
% (idi.name, idi.name, np_name))
222229

223230
gen("")
224231

225232
if options.return_dict:
226233
gen("return _lpy_evt, {%s}"
227-
% ", ".join(f'"{arg.name}": {arg.name}'
228-
for arg in implemented_data_info
229-
if issubclass(arg.arg_class, KernelArgument)
230-
if arg.base_name in kernel.get_written_variables()))
234+
% ", ".join(f'"{idi.name}": {idi.name}'
235+
for idi in implemented_data_info
236+
if issubclass(idi.arg_class, KernelArgument)
237+
if is_output(idi)))
231238
else:
232-
out_args = [arg
233-
for arg in implemented_data_info
234-
if issubclass(arg.arg_class, KernelArgument)
235-
if arg.base_name in kernel.get_written_variables()]
236-
if out_args:
239+
out_idis = [idi
240+
for idi in implemented_data_info
241+
if issubclass(idi.arg_class, KernelArgument)
242+
if is_output(idi)]
243+
if out_idis:
237244
gen("return _lpy_evt, (%s,)"
238-
% ", ".join(arg.name for arg in out_args))
245+
% ", ".join(idi.name for idi in out_idis))
239246
else:
240247
gen("return _lpy_evt, ()")
241248

0 commit comments

Comments
 (0)