//===----------------------------------------------------------------------===// // // Part of CUDA Experimental in CUDA C++ Core Libraries, // under the Apache License v2.0 with LLVM Exceptions. // See https://llvm.org/LICENSE.txt for license information. // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception // SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. // //===----------------------------------------------------------------------===// #include // cub::detail::choose_offset_t #include // cub::detail::CudaDriverLauncherFactory #include #include // DeviceThreeWayPartition kernels #include // policy_selector #include #include #include #include #include #include #include #include #include // std::is_same_v #include #include #include "jit_templates/templates/input_iterator.h" #include "jit_templates/templates/operation.h" #include "jit_templates/templates/output_iterator.h" #include "jit_templates/traits.h" #include "util/context.h" #include "util/errors.h" #include "util/indirect_arg.h" #include "util/nvjitlink.h" #include "util/serialization.h" #include "util/types.h" #include #include #include #include #include #include struct device_three_way_partition_policy_selector; using OffsetT = ptrdiff_t; static_assert(std::is_same_v::type, OffsetT>, "OffsetT must be long"); // check we can map OffsetT to cuda::std::int64_t static_assert(std::is_signed_v); static_assert(sizeof(OffsetT) == sizeof(cuda::std::int64_t)); namespace three_way_partition { struct three_way_partition_kernel_source { cccl_device_three_way_partition_build_result_t& build; CUkernel ThreeWayPartitionInitKernel() const { return build.three_way_partition_init_kernel; } CUkernel ThreeWayPartitionKernel() const { return build.three_way_partition_kernel; } }; std::string get_three_way_partition_init_kernel_name(std::string_view num_selected_out_iterator_name) { constexpr std::string_view scan_tile_state_t = "cub::detail::three_way_partition::ScanTileStateT"; return std::format("cub::detail::three_way_partition::DeviceThreeWayPartitionInitKernel<{0}, {1}>", scan_tile_state_t, // 0 num_selected_out_iterator_name); // 1 } std::string get_three_way_partition_kernel_name( std::string_view d_in_iterator_name, std::string_view d_first_part_out_iterator_name, std::string_view d_second_part_out_iterator_name, std::string_view d_unselected_out_iterator_name, std::string_view d_num_selected_out_iterator_name, std::string_view select_first_part_op_name, std::string_view select_second_part_op_name) { std::string policy_selector_t; check(cccl_type_name_from_nvrtc(&policy_selector_t)); constexpr std::string_view scan_tile_state_t = "cub::detail::three_way_partition::ScanTileStateT"; std::string offset_t; check(cccl_type_name_from_nvrtc(&offset_t)); const std::string streaming_context_t = std::format("cub::detail::three_way_partition::streaming_context_t<{0}>", offset_t); return std::format( "cub::detail::three_way_partition::DeviceThreeWayPartitionKernel<{0}, {1}, {2}, {3}, {4}, {5}, {6}, {7}, {8}, {9}, " "{10}>", policy_selector_t, // 0 d_in_iterator_name, // 1 d_first_part_out_iterator_name, // 2 d_second_part_out_iterator_name, // 3 d_unselected_out_iterator_name, // 4 d_num_selected_out_iterator_name, // 5 scan_tile_state_t, // 6 select_first_part_op_name, // 7 select_second_part_op_name, // 8 "cub::detail::three_way_partition::per_partition_offset_t", // 9 streaming_context_t // 10 ); } } // namespace three_way_partition struct three_way_partition_input_iterator_tag; struct three_way_partition_first_part_output_iterator_tag; struct three_way_partition_second_part_output_iterator_tag; struct three_way_partition_unselected_output_iterator_tag; struct three_way_partition_num_selected_output_iterator_tag; struct three_way_partition_select_first_part_operation_tag; struct three_way_partition_select_second_part_operation_tag; CUresult cccl_device_three_way_partition_compile( cccl_device_three_way_partition_build_result_t* build_ptr, cccl_iterator_t d_in, cccl_iterator_t d_first_part_out, cccl_iterator_t d_second_part_out, cccl_iterator_t d_unselected_out, cccl_iterator_t d_num_selected_out, cccl_op_t select_first_part_op, cccl_op_t select_second_part_op, int cc_major, int cc_minor, const char* cub_path, const char* thrust_path, const char* libcudacxx_path, const char* ctk_path, cccl_build_config* config) try { const char* name = "device_three_way_partition"; const cuda::compute_capability cc{cc_major, cc_minor}; const auto [d_in_iterator_name, d_in_iterator_src] = get_specialization(template_id(), d_in); const auto [d_first_part_out_iterator_name, d_first_part_out_iterator_src] = get_specialization( template_id(), d_first_part_out, d_first_part_out.value_type); const auto [d_second_part_out_iterator_name, d_second_part_out_iterator_src] = get_specialization( template_id(), d_second_part_out, d_second_part_out.value_type); const auto [d_unselected_out_iterator_name, d_unselected_out_iterator_src] = get_specialization( template_id(), d_unselected_out, d_unselected_out.value_type); const auto [d_num_selected_out_iterator_name, d_num_selected_out_iterator_src] = get_specialization( template_id(), d_num_selected_out, d_num_selected_out.value_type); cccl_type_info selector_result_t{sizeof(bool), alignof(bool), cccl_type_enum::CCCL_BOOLEAN}; const auto [select_first_part_op_name, select_first_part_op_src] = get_specialization( template_id(), select_first_part_op, selector_result_t, d_in.value_type); const auto [select_second_part_op_name, select_second_part_op_src] = get_specialization( template_id(), select_second_part_op, selector_result_t, d_in.value_type); const auto offset_t = cccl_type_enum_to_name(cccl_type_enum::CCCL_INT64); const std::string key_t = cccl_type_enum_to_name(d_in.value_type.type); const auto policy_sel = cub::detail::three_way_partition::policy_selector{ cccl_type_enum_to_cub_type(d_in.value_type.type), static_cast(d_in.value_type.size), int{sizeof(OffsetT)}}; // TODO(bgruber): drop this if tuning policies become formattable std::stringstream policy_sel_str; policy_sel_str << policy_sel(cc); const auto policy_selector_expr = std::format( R"XXX(cub::detail::three_way_partition::policy_selector_from_types<{0}, {1}>)XXX", key_t, // 0 offset_t); // 1 std::string final_src = std::format( R"XXX( #include #include {0} struct __align__({2}) storage_t {{ char data[{1}]; }}; {3} {4} {5} {6} {7} {8} {9} using device_three_way_partition_policy_selector = {10}; using namespace cub; using namespace cub::detail; using namespace cub::detail::three_way_partition; static_assert( device_three_way_partition_policy_selector()(current_tuning_cc()) == {11}, "Host generated and JIT compiled policy mismatch"); )XXX", jit_template_header_contents, // 0 d_in.value_type.size, // 1 d_in.value_type.alignment, // 2 d_in_iterator_src, // 3 d_first_part_out_iterator_src, // 4 d_second_part_out_iterator_src, // 5 d_unselected_out_iterator_src, // 6 d_num_selected_out_iterator_src, // 7 select_first_part_op_src, // 8 select_second_part_op_src, // 9 policy_selector_expr, // 10 policy_sel_str.view()); // 11 #if false // CCCL_DEBUGGING_SWITCH fflush(stderr); printf("\nCODE4NVRTC BEGIN\n%sCODE4NVRTC END\n", final_src.c_str()); fflush(stdout); #endif std::string three_way_partition_init_kernel_name = three_way_partition::get_three_way_partition_init_kernel_name(d_num_selected_out_iterator_name); std::string three_way_partition_kernel_name = three_way_partition::get_three_way_partition_kernel_name( d_in_iterator_name, d_first_part_out_iterator_name, d_second_part_out_iterator_name, d_unselected_out_iterator_name, d_num_selected_out_iterator_name, select_first_part_op_name, select_second_part_op_name); std::string three_way_partition_init_kernel_lowered_name; std::string three_way_partition_kernel_lowered_name; const std::string arch = std::format("-arch=sm_{0}{1}", cc_major, cc_minor); std::vector args = { arch.c_str(), cub_path, thrust_path, libcudacxx_path, ctk_path, "-rdc=true", "-dlto", "-DCUB_DISABLE_CDP", "-std=c++20"}; cccl::detail::extend_args_with_build_config(args, config); if (is_custom_op(select_first_part_op) != is_custom_op(select_second_part_op)) { return CUDA_ERROR_INVALID_VALUE; } const bool kernel_only = is_custom_op(select_first_part_op) && is_custom_op(select_second_part_op); constexpr size_t num_lto_args = 2; const char* lopts[num_lto_args] = {"-lto", arch.c_str()}; // Collect all LTO-IRs to be linked (empty ops when kernel_only — ops have no code). nvrtc_linkable_list linkable_list; nvrtc_linkable_list_appender appender{linkable_list}; appender.append_operation(select_first_part_op); appender.append_operation(select_second_part_op); appender.add_iterator_definition(d_in); appender.add_iterator_definition(d_first_part_out); appender.add_iterator_definition(d_second_part_out); appender.add_iterator_definition(d_unselected_out); appender.add_iterator_definition(d_num_selected_out); auto post_build = begin_linking_nvrtc_program(kernel_only ? 0 : num_lto_args, kernel_only ? nullptr : lopts) ->add_program(nvrtc_translation_unit{final_src.c_str(), name}) ->add_expression({three_way_partition_init_kernel_name}) ->add_expression({three_way_partition_kernel_name}) ->compile_program({args.data(), args.size()}) ->get_name({three_way_partition_init_kernel_name, three_way_partition_init_kernel_lowered_name}) ->get_name({three_way_partition_kernel_name, three_way_partition_kernel_lowered_name}); struct free_deleter { void operator()(void* p) const { std::free(p); } }; static_assert(::cuda::is_trivially_copyable_v); const size_t policy_size = sizeof(policy_sel); std::unique_ptr policy_ptr(std::malloc(policy_size)); if (!policy_ptr) { return CUDA_ERROR_OUT_OF_MEMORY; } std::memcpy(policy_ptr.get(), &policy_sel, sizeof(policy_sel)); auto init_name = std::unique_ptr(duplicate_c_string(three_way_partition_init_kernel_lowered_name)); auto kernel_name = std::unique_ptr(duplicate_c_string(three_way_partition_kernel_lowered_name)); build_ptr->cc = cc.get(); // Zero-init fields set by _load, not _compile. build_ptr->library = nullptr; build_ptr->three_way_partition_init_kernel = nullptr; build_ptr->three_way_partition_kernel = nullptr; // All potentially-throwing operations come before any release() calls so that // unique_ptrs automatically clean up on exception. if (kernel_only) { auto [ltoir_size, ltoir_data] = post_build->get_program_ltoir(); build_ptr->payload = ltoir_data.release(); build_ptr->payload_size = ltoir_size; build_ptr->payload_kind = CCCL_PAYLOAD_LTOIR; } else { nvrtc_link_result result = post_build->link_program()->add_link_list(linkable_list)->finalize_program(); build_ptr->payload = (void*) result.data.release(); build_ptr->payload_size = result.size; build_ptr->payload_kind = CCCL_PAYLOAD_CUBIN; } build_ptr->runtime_policy = policy_ptr.release(); build_ptr->runtime_policy_size = policy_size; build_ptr->three_way_partition_init_kernel_lowered_name = init_name.release(); build_ptr->three_way_partition_kernel_lowered_name = kernel_name.release(); return CUDA_SUCCESS; } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_three_way_partition_compile(): %s\n", exc.what()); fflush(stdout); return CUDA_ERROR_UNKNOWN; } CUresult cccl_device_three_way_partition_load(cccl_device_three_way_partition_build_result_t* build) try { auto invalid_name = [](const char* n) { return n == nullptr || n[0] == '\0'; }; if (build == nullptr || build->payload == nullptr || build->payload_size == 0 || build->payload_kind != CCCL_PAYLOAD_CUBIN || invalid_name(build->three_way_partition_init_kernel_lowered_name) || invalid_name(build->three_way_partition_kernel_lowered_name)) { return CUDA_ERROR_INVALID_VALUE; } CUresult status = cuLibraryLoadData(&build->library, build->payload, nullptr, nullptr, 0, nullptr, nullptr, 0); if (status != CUDA_SUCCESS) { return status; } try { check(cuLibraryGetKernel( &build->three_way_partition_init_kernel, build->library, build->three_way_partition_init_kernel_lowered_name)); check(cuLibraryGetKernel( &build->three_way_partition_kernel, build->library, build->three_way_partition_kernel_lowered_name)); } catch (...) { cuLibraryUnload(build->library); build->library = nullptr; throw; } return CUDA_SUCCESS; } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_three_way_partition_load(): %s\n", exc.what()); fflush(stdout); return CUDA_ERROR_UNKNOWN; } CUresult cccl_device_three_way_partition_build_ex( cccl_device_three_way_partition_build_result_t* build_ptr, cccl_iterator_t d_in, cccl_iterator_t d_first_part_out, cccl_iterator_t d_second_part_out, cccl_iterator_t d_unselected_out, cccl_iterator_t d_num_selected_out, cccl_op_t select_first_part_op, cccl_op_t select_second_part_op, int cc_major, int cc_minor, const char* cub_path, const char* thrust_path, const char* libcudacxx_path, const char* ctk_path, cccl_build_config* config) { CUresult result = cccl_device_three_way_partition_compile( build_ptr, d_in, d_first_part_out, d_second_part_out, d_unselected_out, d_num_selected_out, select_first_part_op, select_second_part_op, cc_major, cc_minor, cub_path, thrust_path, libcudacxx_path, ctk_path, config); if (result != CUDA_SUCCESS) { return result; } CUresult load_r = cccl_device_three_way_partition_load(build_ptr); if (load_r != CUDA_SUCCESS) { cccl_device_three_way_partition_cleanup(build_ptr); } return load_r; } CUresult cccl_device_three_way_partition( cccl_device_three_way_partition_build_result_t build, void* d_temp_storage, size_t* temp_storage_bytes, cccl_iterator_t d_in, cccl_iterator_t d_first_part_out, cccl_iterator_t d_second_part_out, cccl_iterator_t d_unselected_out, cccl_iterator_t d_num_selected_out, cccl_op_t select_first_part_op, cccl_op_t select_second_part_op, uint64_t num_items, CUstream stream) { bool pushed = false; CUresult error = CUDA_SUCCESS; try { pushed = try_push_context(); CUdevice cu_device; check(cuCtxGetDevice(&cu_device)); auto exec_status = cub::detail::three_way_partition::dispatch< indirect_arg_t, // InputIteratorT indirect_arg_t, // FirstOutputIteratorT indirect_arg_t, // SecondOutputIteratorT indirect_arg_t, // UnselectedOutputIteratorT indirect_arg_t, // NumSelectedIteratorT indirect_arg_t, // SelectFirstPartOp indirect_arg_t, // SelectSecondPartOp OffsetT, // OffsetT cub::detail::three_way_partition::policy_selector, // PolicySelector three_way_partition::three_way_partition_kernel_source, // KernelSource cub::detail::CudaDriverLauncherFactory>( d_temp_storage, *temp_storage_bytes, d_in, d_first_part_out, d_second_part_out, d_unselected_out, d_num_selected_out, select_first_part_op, select_second_part_op, static_cast(num_items), stream, /* policy_selector */ *static_cast(build.runtime_policy), /* kernel_source */ {build}, /* launcher_factory */ cub::detail::CudaDriverLauncherFactory{cu_device, build.cc}); error = static_cast(exec_status); } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_three_way_partition(): %s\n", exc.what()); fflush(stdout); error = CUDA_ERROR_UNKNOWN; } if (pushed) { CUcontext dummy; cuCtxPopCurrent(&dummy); } return error; } CUresult cccl_device_three_way_partition_cleanup(cccl_device_three_way_partition_build_result_t* bld_ptr) try { if (bld_ptr == nullptr) { return CUDA_ERROR_INVALID_VALUE; } std::unique_ptr payload(reinterpret_cast(bld_ptr->payload)); std::free(bld_ptr->runtime_policy); std::unique_ptr init_name(bld_ptr->three_way_partition_init_kernel_lowered_name); std::unique_ptr kernel_name(bld_ptr->three_way_partition_kernel_lowered_name); if (bld_ptr->library != nullptr) { check(cuLibraryUnload(bld_ptr->library)); } return CUDA_SUCCESS; } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_three_way_partition_cleanup(): %s\n", exc.what()); fflush(stdout); return CUDA_ERROR_UNKNOWN; } CUresult cccl_device_three_way_partition_build( cccl_device_three_way_partition_build_result_t* build_ptr, cccl_iterator_t d_in, cccl_iterator_t d_first_part_out, cccl_iterator_t d_second_part_out, cccl_iterator_t d_unselected_out, cccl_iterator_t d_num_selected_out, cccl_op_t select_first_part_op, cccl_op_t select_second_part_op, int cc_major, int cc_minor, const char* cub_path, const char* thrust_path, const char* libcudacxx_path, const char* ctk_path) { return cccl_device_three_way_partition_build_ex( build_ptr, d_in, d_first_part_out, d_second_part_out, d_unselected_out, d_num_selected_out, select_first_part_op, select_second_part_op, cc_major, cc_minor, cub_path, thrust_path, libcudacxx_path, ctk_path, nullptr); } CUresult cccl_device_three_way_partition_link_ltoir( cccl_device_three_way_partition_build_result_t* build_ptr, const void** input_blobs, const size_t* input_sizes, size_t num_inputs) try { if (build_ptr == nullptr || build_ptr->payload == nullptr || build_ptr->payload_size == 0 || build_ptr->payload_kind != CCCL_PAYLOAD_LTOIR) { return CUDA_ERROR_INVALID_VALUE; } const int cc_major = build_ptr->cc / 10; const int cc_minor = build_ptr->cc % 10; std::vector all_blobs; std::vector all_sizes; all_blobs.push_back(build_ptr->payload); all_sizes.push_back(build_ptr->payload_size); if (num_inputs > 0 && (input_blobs == nullptr || input_sizes == nullptr)) { return CUDA_ERROR_INVALID_VALUE; } for (size_t i = 0; i < num_inputs; ++i) { if (input_blobs[i] == nullptr || input_sizes[i] == 0) { return CUDA_ERROR_INVALID_VALUE; } all_blobs.push_back(input_blobs[i]); all_sizes.push_back(input_sizes[i]); } auto [cubin, cubin_size] = nvjitlink_link(all_blobs.data(), all_sizes.data(), all_blobs.size(), cc_major, cc_minor); delete[] static_cast(build_ptr->payload); build_ptr->payload = (void*) cubin.release(); build_ptr->payload_size = cubin_size; build_ptr->payload_kind = CCCL_PAYLOAD_CUBIN; return CUDA_SUCCESS; } catch (const std::exception& exc) { printf("\nEXCEPTION in cccl_device_three_way_partition_link_ltoir(): %s\n", exc.what()); return CUDA_ERROR_UNKNOWN; } CUresult cccl_device_three_way_partition_serialize( const cccl_device_three_way_partition_build_result_t* build_ptr, void** out_buf, size_t* out_size) try { if (build_ptr == nullptr || out_buf == nullptr || out_size == nullptr) { return CUDA_ERROR_INVALID_VALUE; } if (build_ptr->payload == nullptr || build_ptr->payload_size == 0 || build_ptr->runtime_policy == nullptr || build_ptr->runtime_policy_size == 0) { *out_buf = nullptr; *out_size = 0; return CUDA_ERROR_INVALID_VALUE; } *out_buf = nullptr; *out_size = 0; using namespace cccl::serialization; buffer_writer w; write_header(w, CCCL_SERIALIZATION_ALGO_THREE_WAY_PARTITION, build_ptr->payload_kind, build_ptr->cc); w.write_blob(build_ptr->payload, build_ptr->payload_size); w.write_blob(build_ptr->runtime_policy, build_ptr->runtime_policy_size); w.write_cstring(build_ptr->three_way_partition_init_kernel_lowered_name); w.write_cstring(build_ptr->three_way_partition_kernel_lowered_name); w.release(out_buf, out_size); return CUDA_SUCCESS; } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_three_way_partition_serialize(): %s\n", exc.what()); fflush(stdout); return CUDA_ERROR_UNKNOWN; } CUresult cccl_device_three_way_partition_deserialize( cccl_device_three_way_partition_build_result_t* build_ptr, const void* buf, size_t size) try { if (build_ptr == nullptr || buf == nullptr || size == 0) { return CUDA_ERROR_INVALID_VALUE; } using namespace cccl::serialization; buffer_reader r{buf, size}; const auto h = read_and_validate_header(r, CCCL_SERIALIZATION_ALGO_THREE_WAY_PARTITION); std::unique_ptr payload_owner; size_t payload_size = 0; { void* p = nullptr; r.read_blob_new(&p, &payload_size); payload_owner.reset(static_cast(p)); } if (payload_size == 0) { throw std::runtime_error("serialization blob: empty payload"); } std::unique_ptr policy( static_cast( std::malloc(sizeof(cub::detail::three_way_partition::policy_selector))), std::free); if (!policy) { return CUDA_ERROR_OUT_OF_MEMORY; } r.read_into(policy.get(), sizeof(cub::detail::three_way_partition::policy_selector)); std::unique_ptr n_init{r.read_cstring_dup()}; std::unique_ptr n_kernel{r.read_cstring_dup()}; cccl_device_three_way_partition_build_result_t result{}; result.cc = static_cast(h.cc); result.payload_kind = static_cast(h.payload_kind); result.payload = payload_owner.release(); result.payload_size = payload_size; result.runtime_policy = policy.release(); result.runtime_policy_size = sizeof(cub::detail::three_way_partition::policy_selector); result.three_way_partition_init_kernel_lowered_name = n_init.release(); result.three_way_partition_kernel_lowered_name = n_kernel.release(); *build_ptr = result; return CUDA_SUCCESS; } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_three_way_partition_deserialize(): %s\n", exc.what()); fflush(stdout); return CUDA_ERROR_UNKNOWN; }