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
4 changes: 4 additions & 0 deletions src/cunumpy/_cuda_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -2421,6 +2421,10 @@ def __call__(
if shared_mem < 0:
raise ValueError(f"shared_mem must be non-negative, got {shared_mem}")
values = self.prepare_args(*args)
# RawKernel accepts size-one NumPy arrays for structs passed by value,
# but not structured NumPy scalars (np.void). Keep their packed bytes
# and alignment intact, including array-view pointers and strides.
values = tuple(np.asarray(v) if isinstance(v, np.void) else v for v in values)
if 0 in grid_shape:
return

Expand Down
23 changes: 23 additions & 0 deletions tests/unit/test_cuda_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -2306,6 +2306,7 @@ def __init__(self):

def __call__(self, grid, block, args, shared_mem=0):
self.launches.append((grid, block, shared_mem))
self.args = args


@pytest.fixture
Expand Down Expand Up @@ -2333,6 +2334,28 @@ def recorded(monkeypatch):
return kernel, raw


def test_launch_passes_structs_as_size_one_numpy_arrays(recorded, monkeypatch):
kernel, raw = recorded
packed = np.zeros((), dtype=[("data", np.uintp), ("shape", np.int64, (1,))])
packed["data"] = 0x1000
packed["shape"] = (7,)
scalar = np.float64(2.0)
device_array = FakeDeviceArray(np.float64, shape=(7,))
monkeypatch.setattr(
kernel, "prepare_args", lambda *args: (packed[()], scalar, device_array)
)

kernel(n_threads=1)

struct_arg, scalar_arg, pointer_arg = raw.args
assert isinstance(struct_arg, np.ndarray)
assert struct_arg.size == 1
assert struct_arg.dtype == packed.dtype
assert struct_arg.tobytes() == packed.tobytes()
assert scalar_arg is scalar
assert pointer_arg is device_array


def test_n_threads_from_first_array(recorded):
kernel, raw = recorded
kernel.n_threads_from = "first_array"
Expand Down
Loading