77# is strictly prohibited.
88
99
10- import warnings
1110import ctypes
11+ import warnings
1212
1313import pytest
1414
@@ -173,15 +173,18 @@ def test_saxpy_arguments(get_saxpy_kernel):
173173 arg_info = krn .arguments_info
174174 n_args = len (arg_info )
175175 assert n_args == krn .num_arguments
176+
176177 class ExpectedStruct (ctypes .Structure ):
177178 _fields_ = [
178- ('a' , ctypes .c_float ),
179- ('x' , ctypes .POINTER (ctypes .c_float )),
180- ('y' , ctypes .POINTER (ctypes .c_float )),
181- (' out' , ctypes .POINTER (ctypes .c_float )),
182- ('N' , ctypes .c_size_t )
179+ ("a" , ctypes .c_float ),
180+ ("x" , ctypes .POINTER (ctypes .c_float )),
181+ ("y" , ctypes .POINTER (ctypes .c_float )),
182+ (" out" , ctypes .POINTER (ctypes .c_float )),
183+ ("N" , ctypes .c_size_t ),
183184 ]
184- offsets , sizes = zip (* arg_info )
185+
186+ offsets = [p .offset for p in arg_info ]
187+ sizes = [p .size for p in arg_info ]
185188 members = [getattr (ExpectedStruct , name ) for name , _ in ExpectedStruct ._fields_ ]
186189 expected_offsets = tuple (m .offset for m in members )
187190 assert all (actual == expected for actual , expected in zip (offsets , expected_offsets ))
@@ -197,15 +200,16 @@ def test_num_arguments(init_cuda, nargs, c_type_name, c_type):
197200 prog = Program (src , code_type = "c++" )
198201 mod = prog .compile (
199202 "cubin" ,
200- name_expressions = (f"foo{ nargs } " , ),
203+ name_expressions = (f"foo{ nargs } " ,),
201204 )
202205 krn = mod .get_kernel (f"foo{ nargs } " )
203206 assert krn .num_arguments == nargs
204207
205208 class ExpectedStruct (ctypes .Structure ):
206- _fields_ = [
207- (f'arg_{ i } ' , c_type ) for i in range (nargs )
208- ]
209+ _fields_ = [(f"arg_{ i } " , c_type ) for i in range (nargs )]
210+
209211 members = tuple (getattr (ExpectedStruct , f"arg_{ i } " ) for i in range (nargs ))
210- expected = [(m .offset , m .size ) for m in members ]
211- assert krn .arguments_info == expected
212+
213+ arg_info = krn .arguments_info
214+ assert all ([actual .offset == expected .offset for actual , expected in zip (arg_info , members )])
215+ assert all ([actual .size == expected .size for actual , expected in zip (arg_info , members )])
0 commit comments