FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Add unit tests and type-safe overloads for CUDA kernel functions by saschiwy · Pull Request #800 · taskflow/taskflow · GitHub

Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension .cu  (3) .dox  (1) .hpp  (2) .txt  (1) All 4 file types selected
Viewed files
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Unified
Split
Hide whitespace
Diff view
Unified
Split
Hide whitespace
11 changes: 11 additions & 0 deletions doxygen/releases/release-4.2.0.dox
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,17 @@ If your project does not support C++20, please use %Taskflow v3.11.0, which can
+ fixed missing tf::TaskParamsLike concept usage in async tasking
+ replaced generic template type in dependent async tasks with std::input_iterator

@subsection release-4-2-0_gpu_tasking GPU Tasking

@li added type-safe overloads of tf::cudaGraphBase::kernel and tf::cudaGraphExecBase::kernel
that accept a typed @c __global__ function pointer (@c void(*)(Params...)) and use
@c tf::detail::kernelArgCast to enforce correct argument types at compile time
@li scalar arguments use @c static_cast — width mismatches (e.g. @c int where @c size_t
is expected) are resolved correctly; truly incompatible types become compile errors
@li pointer arguments use @c reinterpret_cast — handles @c T*→const T* and typedef aliases
without requiring explicit casts at the call site
@li the existing generic overload (lambda/functor) is unchanged

@subsection release-4-2-0_utilities Utilities

@section release-4-2-0_bug_fixes Bug Fixes
Expand Down
113 changes: 113 additions & 0 deletions taskflow/cuda/cuda_graph.hpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -516,6 +516,40 @@ class cudaGraphDeleter {
};


namespace detail {

/**
@brief casts a kernel call-site argument to the exact type expected by the kernel parameter

Dispatches between two cast strategies based on whether the target parameter
type is a pointer. The dispatch is resolved entirely at compile time via
`if constexpr` — no runtime branching is generated.

- Pointer types: reinterpret_cast — handles T*→const T*, typedef aliases (e.g. boolean_T*
vs __nv_bool*), and other ABI-compatible but nominally distinct pointer pairs.
- Non-pointer (scalar) types: static_cast — widens int→size_t correctly and makes truly
incompatible conversions (non-convertible structs, integer→pointer) a compile error.

When TParam == TArg (identical types), static_cast<T>(t) is a no-op: the compiler emits no
additional instructions compared to passing the argument directly. The function is therefore
a zero-cost abstraction — it adds compile-time type safety without any runtime overhead.

@tparam TParam the exact type the kernel parameter expects (deduced from the function pointer)
@tparam TArg the type of the argument at the call site (deduced from the caller)
@param arg the argument value from the call site
@return the argument converted to TParam
*/
template<typename TParam, typename TArg>
constexpr TParam kernelArgCast(TArg&& arg) {
if constexpr (std::is_pointer_v<TParam>) {
return reinterpret_cast<TParam>(arg);
} else {
return static_cast<TParam>(std::forward<TArg>(arg));
}
}

} // namespace detail

/**
@class cudaGraphBase

Expand Down Expand Up @@ -636,6 +670,47 @@ class cudaGraphBase : public std::unique_ptr<std::remove_pointer_t<cudaGraph_t>,
template <typename F, typename... ArgsT>
cudaTask kernel(dim3 g, dim3 b, size_t s, F f, ArgsT... args);

/**
@brief creates a kernel task with compile-time type checking of arguments

This overload accepts a typed @c __global__ function pointer instead of a generic callable.
The compiler deduces the kernel's exact parameter types @c Params from the function pointer
type and casts each argument accordingly via @c tf::detail::kernelArgCast:
- Scalar arguments: @c static_cast — width mismatches (e.g. @c int where @c size_t is
required) produce a correctly widened value; genuinely incompatible types become compile
errors.
- Pointer arguments: @c reinterpret_cast — handles @c T* → @c const T* and typedef aliases
that are ABI-compatible but nominally distinct.

The number of arguments @c ArgsT must exactly match the number of kernel parameters @c Params;
a mismatch is a compile error.

The existing @c kernel(dim3,dim3,size_t,F,ArgsT...) overload is retained for lambdas and
functors where a typed function pointer is not available.

@tparam Params parameter types deduced from the typed function pointer @c f
@tparam ArgsT types of the arguments at the call site

@param g grid dimensions
@param b block dimensions
@param s shared memory size in bytes
@param f pointer to the @c __global__ kernel function
@param args arguments to forward to the kernel

@return a tf::cudaTask handle for the newly created kernel node

@code{.cpp}
// kernel declaration
__global__ void scale(float* data, size_t n, float factor);

tf::cudaGraph cg;
// typed overload: int literal for size_t parameter is widened correctly
auto task = cg.kernel({8,1,1}, {128,1,1}, 0, scale, d_data, (size_t)N, 2.0f);
@endcode
*/
template <typename... Params, typename... ArgsT>
cudaTask kernel(dim3 g, dim3 b, size_t s, void(*f)(Params...), ArgsT&&... args);

/**
@brief creates a memset task that fills untyped data with a byte value

Expand Down Expand Up @@ -1031,6 +1106,44 @@ cudaTask cudaGraphBase<Creator, Deleter>::kernel(
return cudaTask(this->get(), node);
}

// Function: kernel (typed overload — compile-time argument type checking)
template <typename Creator, typename Deleter>
template <typename... Params, typename... ArgsT>
cudaTask cudaGraphBase<Creator, Deleter>::kernel(
dim3 g, dim3 b, size_t s, void(*f)(Params...), ArgsT&&... args
) {
static_assert(sizeof...(Params) == sizeof...(ArgsT),
"kernel: argument count does not match kernel parameter count");

// Cast every argument to the exact type the kernel parameter expects.
// Storing in a tuple keeps the cast values alive until cudaGraphAddKernelNode returns.
auto castedArgs = std::make_tuple(
detail::kernelArgCast<Params>(std::forward<ArgsT>(args))...
);

cudaGraphNode_t node;
cudaKernelNodeParams p;

// Build the void* array from addresses of the tuple elements.
void* arguments[sizeof...(Params)];
[&]<std::size_t... Is>(std::index_sequence<Is...>) {
((arguments[Is] = static_cast<void*>(&std::get<Is>(castedArgs))), ...);
}(std::make_index_sequence<sizeof...(Params)>{});

p.func = reinterpret_cast<void*>(f);
p.gridDim = g;
p.blockDim = b;
p.sharedMemBytes = s;
p.kernelParams = arguments;
p.extra = nullptr;

TF_CHECK_CUDA(
cudaGraphAddKernelNode(&node, this->get(), nullptr, 0, &p),
"failed to create a kernel task"
);
return cudaTask(this->get(), node);
}

// Function: zero
template <typename Creator, typename Deleter>
template <typename T, std::enable_if_t<
Expand Down
67 changes: 66 additions & 1 deletion taskflow/cuda/cuda_graph_exec.hpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -146,7 +146,39 @@ class cudaGraphExecBase : public std::unique_ptr<std::remove_pointer_t<cudaGraph
void kernel(
cudaTask task, dim3 g, dim3 b, size_t shm, F f, ArgsT... args
);


/**
@brief updates parameters of a kernel task with compile-time type checking of arguments

This is the update-path counterpart of the type-safe @c cudaGraphBase::kernel() overload.
It calls @c cudaGraphExecKernelNodeSetParams on the node referenced by @c task.
All type-checking and casting rules are identical to the creation overload.

The kernel function name must NOT change between creation and update (CUDA restriction).

@tparam Params parameter types deduced from the typed function pointer @c f
@tparam ArgsT types of the arguments at the call site

@param task the cudaTask handle that references the kernel node to update
@param g new grid dimensions
@param b new block dimensions
@param s new shared memory size in bytes
@param f pointer to the (same) @c __global__ kernel function
@param args new argument values

@code{.cpp}
__global__ void scale(float* data, size_t n, float factor);

tf::cudaGraph cg;
auto task = cg.kernel({8,1,1}, {128,1,1}, 0, scale, d_data, (size_t)N, 1.0f);
tf::cudaGraphExec exec(cg);
// update the factor — type safety ensured at compile time
exec.kernel(task, {8,1,1}, {128,1,1}, 0, scale, d_data, (size_t)N, 2.0f);
@endcode
*/
template <typename... Params, typename... ArgsT>
void kernel(cudaTask task, dim3 g, dim3 b, size_t s, void(*f)(Params...), ArgsT&&... args);

/**
@brief updates parameters of a memset task

Expand Down Expand Up @@ -295,6 +327,39 @@ void cudaGraphExecBase<Creator, Deleter>::kernel(
);
}

// Function: update kernel parameters (typed overload — compile-time type-safe)
template <typename Creator, typename Deleter>
template <typename... Params, typename... ArgsT>
void cudaGraphExecBase<Creator, Deleter>::kernel(
cudaTask task, dim3 g, dim3 b, size_t s, void(*f)(Params...), ArgsT&&... args
) {
static_assert(sizeof...(Params) == sizeof...(ArgsT),
"kernel: argument count does not match kernel parameter count");

auto castedArgs = std::make_tuple(
detail::kernelArgCast<Params>(std::forward<ArgsT>(args))...
);

cudaKernelNodeParams p;

void* arguments[sizeof...(Params)];
[&]<std::size_t... Is>(std::index_sequence<Is...>) {
((arguments[Is] = static_cast<void*>(&std::get<Is>(castedArgs))), ...);
}(std::make_index_sequence<sizeof...(Params)>{});

p.func = reinterpret_cast<void*>(f);
p.gridDim = g;
p.blockDim = b;
p.sharedMemBytes = s;
p.kernelParams = arguments;
p.extra = nullptr;

TF_CHECK_CUDA(
cudaGraphExecKernelNodeSetParams(this->get(), task._native_node, &p),
"failed to update kernel parameters on ", task
);
}

// Function: update copy parameters
template <typename Creator, typename Deleter>
template <typename T, std::enable_if_t<!std::is_same_v<T, void>, void>*>
Expand Down
36 changes: 36 additions & 0 deletions unittests/cuda/CMakeLists.txt
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ list(APPEND TF_CUDA_UNITTESTS
test_cuda_kmeans
test_cuda_for_each
test_cuda_transform
test_cuda_type_safe_kernel
#test_cuda_reduce
#test_cuda_scan
#test_cuda_find
Expand Down Expand Up @@ -34,5 +35,40 @@ foreach(cudatest IN LISTS TF_CUDA_UNITTESTS)
doctest_discover_tests(${cudatest})
endforeach()

# ---------------------------------------------------------------------------
# Negative compile tests — these files MUST fail to compile.
# try_compile returns TRUE if compilation succeeds, FALSE if it fails.
# We expect failure, so we REQUIRE the result to be FALSE.
# ---------------------------------------------------------------------------

try_compile(
TF_COMPILE_FAIL_ARG_COUNT_RESULT
${CMAKE_BINARY_DIR}
SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/test_compile_fail_arg_count.cu
CMAKE_FLAGS
"-DINCLUDE_DIRECTORIES=${CMAKE_SOURCE_DIR}"
"-DCMAKE_CUDA_STANDARD=20"
OUTPUT_VARIABLE TF_COMPILE_FAIL_ARG_COUNT_OUTPUT
)
if(TF_COMPILE_FAIL_ARG_COUNT_RESULT)
message(FATAL_ERROR
"test_compile_fail_arg_count.cu should have failed to compile "
"(static_assert on wrong arg count), but it succeeded.\n"
"Output:\n${TF_COMPILE_FAIL_ARG_COUNT_OUTPUT}")
endif()

try_compile(
TF_COMPILE_FAIL_SCALAR_MISMATCH_RESULT
${CMAKE_BINARY_DIR}
SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/test_compile_fail_scalar_mismatch.cu
CMAKE_FLAGS
"-DINCLUDE_DIRECTORIES=${CMAKE_SOURCE_DIR}"
"-DCMAKE_CUDA_STANDARD=20"
OUTPUT_VARIABLE TF_COMPILE_FAIL_SCALAR_MISMATCH_OUTPUT
)
if(TF_COMPILE_FAIL_SCALAR_MISMATCH_RESULT)
message(FATAL_ERROR
"test_compile_fail_scalar_mismatch.cu should have failed to compile "
"(incompatible struct to scalar cast), but it succeeded.\n"
"Output:\n${TF_COMPILE_FAIL_SCALAR_MISMATCH_OUTPUT}")
endif()
12 changes: 12 additions & 0 deletions unittests/cuda/test_compile_fail_arg_count.cu
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
// This file must fail to compile due to the static_assert in the typed kernel() overload.
// Expected error: "kernel: argument count does not match kernel parameter count"
#include <taskflow/cuda/cudaflow.hpp>

__global__ void k_two_args(int* ptr, size_t N) { /* empty */ }

int main() {
tf::cudaGraph cg;
// Wrong: k_two_args expects 2 arguments, we pass 3
cg.kernel({1,1,1}, {1,1,1}, 0, k_two_args, (int*)nullptr, (size_t)0, 42);
return 0;
}
15 changes: 15 additions & 0 deletions unittests/cuda/test_compile_fail_scalar_mismatch.cu
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
// This file must fail to compile because static_cast<size_t>(incompatible_struct) is ill-formed.
// Expected error: cannot convert struct to size_t
#include <taskflow/cuda/cudaflow.hpp>

struct Incompatible { int x; int y; };

__global__ void k_scalar(size_t N) { /* empty */ }

int main() {
tf::cudaGraph cg;
Incompatible s{1, 2};
// Wrong: k_scalar expects size_t, we pass an incompatible struct
cg.kernel({1,1,1}, {1,1,1}, 0, k_scalar, s);
return 0;
}
Loading

Back | FazBrowse Home | New Git URL