//===----------------------------------------------------------------------===// // // 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) 2024 NVIDIA CORPORATION & AFFILIATES. // //===----------------------------------------------------------------------===// #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include using OffsetT = unsigned long long; static_assert(std::is_same_v, OffsetT>, "OffsetT must be size_t"); namespace binary_search { struct op_state_header_t { void* data; OffsetT num_data; }; struct op_state_t { std::unique_ptr storage; void* get() { return storage.get(); } }; static size_t align_up(size_t offset, size_t alignment) { const size_t remainder = offset % alignment; return remainder == 0 ? offset : offset + alignment - remainder; } static size_t comparator_state_offset(cccl_op_t op) { return op.type == CCCL_STATEFUL ? align_up(sizeof(op_state_header_t), op.alignment) : sizeof(op_state_header_t); } static size_t op_state_alignment(cccl_op_t op) { return op.type == CCCL_STATEFUL ? std::max(alignof(op_state_header_t), op.alignment) : alignof(op_state_header_t); } static size_t op_state_size(cccl_op_t op) { const size_t unaligned_size = op.type == CCCL_STATEFUL ? comparator_state_offset(op) + op.size : sizeof(op_state_header_t); return align_up(unaligned_size, op_state_alignment(op)); } static op_state_t make_op_state(cccl_iterator_t data, OffsetT num_data, cccl_op_t op) { op_state_t result{std::make_unique(op_state_size(op))}; char* raw = static_cast(result.get()); auto* header = reinterpret_cast(raw); header->data = data.state; header->num_data = num_data; if (op.type == CCCL_STATEFUL) { std::memcpy(raw + comparator_state_offset(op), op.state, op.size); } return result; } } // namespace binary_search static CUresult Invoke( cccl_iterator_t d_in, uint64_t num_items, cccl_iterator_t d_values, uint64_t num_values, cccl_iterator_t d_out, cccl_op_t op, cccl_device_binary_search_build_result_t build, CUstream stream) { auto state = binary_search::make_op_state(d_in, static_cast(num_items), op); cccl_op_t transform_op = op; transform_op.type = CCCL_STATEFUL; transform_op.state = state.get(); transform_op.size = build.op_state_size; transform_op.alignment = build.op_state_alignment; return cccl_device_unary_transform(build.transform, d_values, d_out, num_values, transform_op, stream); } struct binary_search_data_iterator_tag; struct binary_search_values_iterator_tag; struct binary_search_op_tag; CUresult cccl_device_binary_search_compile( cccl_device_binary_search_build_result_t* build_ptr, cccl_binary_search_mode_t mode, cccl_iterator_t d_data, cccl_iterator_t d_values, cccl_iterator_t d_out, cccl_op_t 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 { if (d_data.type == cccl_iterator_kind_t::CCCL_ITERATOR) { throw std::runtime_error(std::string("Iterators are unsupported in for_each currently")); } auto [d_data_it_name, d_data_it_src] = get_specialization(template_id(), d_data); auto [op_name, op_src] = get_specialization( template_id(), op, d_data.value_type, d_data.value_type); const std::string mode_t = [&] { switch (mode) { case CCCL_BINARY_SEARCH_LOWER_BOUND: return "cub::detail::find::lower_bound"; case CCCL_BINARY_SEARCH_UPPER_BOUND: return "cub::detail::find::upper_bound"; } throw std::runtime_error(std::format("Invalid binary search mode ({})", static_cast(mode))); }(); const bool user_defined_comparator = op.type == CCCL_STATEFUL || op.type == CCCL_STATELESS; const bool comparator_stateful = op.type == CCCL_STATEFUL; const auto comparator_offset = binary_search::comparator_state_offset(op); const std::string output_t = cccl_type_enum_to_name(d_out.value_type.type); const size_t storage_size = std::max({d_data.value_type.size, d_values.value_type.size, d_out.value_type.size}); const size_t storage_alignment = std::max({d_data.value_type.alignment, d_values.value_type.alignment, d_out.value_type.alignment}); const std::string transform_op_src = std::format( R"XXX( #include #include {8} struct __align__({7}) storage_t {{ char data[{6}]; }}; {0} {2} using OffsetT = cuda::std::size_t; struct binary_search_op_state {{ {1} data; OffsetT num_data; }}; extern "C" __device__ void binary_search_transform_op(void* state, const void* value, void* result) {{ auto* header = static_cast(state); const auto& item = *static_cast(value); {4} comparator{{}}; if constexpr ({9}) {{ ::cuda::std::memcpy(&comparator, static_cast(state) + {10}, sizeof(comparator)); }} const auto search_op = cub::detail::find::make_binary_search_transform_op<{11}>(header->data, header->num_data, comparator); *static_cast<{5}*>(result) = static_cast<{5}>(search_op(item)); }} )XXX", d_data_it_src, d_data_it_name, op_src, cccl_type_enum_to_name(d_values.value_type.type), op_name, output_t, storage_size, storage_alignment, jit_template_header_contents, comparator_stateful ? "true" : "false", comparator_offset, mode_t); if (is_custom_op(op)) { // kernel-only (link_ltoir) mode is not supported for binary_search: the // binary_search_transform_op wrapper cannot be decoupled from the comparator type. return CUDA_ERROR_INVALID_VALUE; } cccl_op_t transform_op = op; transform_op.type = CCCL_STATEFUL; transform_op.name = "binary_search_transform_op"; transform_op.code = transform_op_src.c_str(); transform_op.code_size = transform_op_src.size(); transform_op.code_type = CCCL_OP_CPP_SOURCE; transform_op.size = binary_search::op_state_size(op); transform_op.alignment = binary_search::op_state_alignment(op); std::vector extra_ltoirs; std::vector extra_ltoir_sizes; std::unique_ptr comparator_ltoir; 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", "-default-device", "-DCUB_DISABLE_CDP", "-std=c++20"}; cccl::detail::extend_args_with_build_config(args, config); constexpr size_t num_lto_args = 2; const char* lopts[num_lto_args] = {"-lto", arch.c_str()}; if (user_defined_comparator && op.code_type == CCCL_OP_CPP_SOURCE && op.code_size != 0) { auto [lto_size, lto_buf] = begin_linking_nvrtc_program(num_lto_args, lopts) ->add_program(nvrtc_translation_unit{op.code, op.name}) ->compile_program({args.data(), args.size()}) ->get_program_ltoir(); comparator_ltoir = std::move(lto_buf); extra_ltoirs.push_back(comparator_ltoir.get()); extra_ltoir_sizes.push_back(lto_size); } else if (user_defined_comparator && op.code_type == CCCL_OP_LTOIR && op.code_size != 0) { extra_ltoirs.push_back(op.code); extra_ltoir_sizes.push_back(op.code_size); } for (size_t i = 0; user_defined_comparator && i < op.num_extra_ltoirs; ++i) { extra_ltoirs.push_back(op.extra_ltoirs[i]); extra_ltoir_sizes.push_back(op.extra_ltoir_sizes[i]); } transform_op.extra_ltoirs = extra_ltoirs.data(); transform_op.extra_ltoir_sizes = extra_ltoir_sizes.data(); transform_op.num_extra_ltoirs = extra_ltoirs.size(); check(cccl_device_unary_transform_compile( &build_ptr->transform, d_values, d_out, transform_op, cc_major, cc_minor, cub_path, thrust_path, libcudacxx_path, ctk_path, config)); build_ptr->op_state_size = transform_op.size; build_ptr->op_state_alignment = transform_op.alignment; return CUDA_SUCCESS; } catch (...) { return CUDA_ERROR_UNKNOWN; } CUresult cccl_device_binary_search_load(cccl_device_binary_search_build_result_t* build_ptr) { if (build_ptr == nullptr) { return CUDA_ERROR_INVALID_VALUE; } return cccl_device_transform_load(&build_ptr->transform); } CUresult cccl_device_binary_search_build_ex( cccl_device_binary_search_build_result_t* build_ptr, cccl_binary_search_mode_t mode, cccl_iterator_t d_data, cccl_iterator_t d_values, cccl_iterator_t d_out, cccl_op_t 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 r = cccl_device_binary_search_compile( build_ptr, mode, d_data, d_values, d_out, op, cc_major, cc_minor, cub_path, thrust_path, libcudacxx_path, ctk_path, config); if (r != CUDA_SUCCESS) { return r; } CUresult load_r = cccl_device_binary_search_load(build_ptr); if (load_r != CUDA_SUCCESS) { cccl_device_binary_search_cleanup(build_ptr); } return load_r; } CUresult cccl_device_binary_search( cccl_device_binary_search_build_result_t build, cccl_iterator_t d_data, uint64_t num_items, cccl_iterator_t d_values, uint64_t num_values, cccl_iterator_t d_out, cccl_op_t op, CUstream stream) { bool pushed = false; CUresult error = CUDA_SUCCESS; try { pushed = try_push_context(); error = Invoke(d_data, num_items, d_values, num_values, d_out, op, build, stream); } catch (...) { error = CUDA_ERROR_UNKNOWN; } if (pushed) { CUcontext dummy; cuCtxPopCurrent(&dummy); } return error; } CUresult cccl_device_binary_search_build( cccl_device_binary_search_build_result_t* build, cccl_binary_search_mode_t mode, cccl_iterator_t d_data, cccl_iterator_t d_values, cccl_iterator_t d_out, cccl_op_t 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_binary_search_build_ex( build, mode, d_data, d_values, d_out, op, cc_major, cc_minor, cub_path, thrust_path, libcudacxx_path, ctk_path, nullptr); } CUresult cccl_device_binary_search_link_ltoir( cccl_device_binary_search_build_result_t* build, const void** input_blobs, const size_t* input_sizes, size_t num_inputs, const char* /* kernel_lowered_name */, size_t /* values_value_size */, size_t /* output_value_size */, size_t op_state_size, size_t op_state_alignment, int /* cc_major */, int /* cc_minor */) { if (build == nullptr) { return CUDA_ERROR_INVALID_VALUE; } build->op_state_size = op_state_size; build->op_state_alignment = op_state_alignment; return cccl_device_transform_link_ltoir(&build->transform, input_blobs, input_sizes, num_inputs); } CUresult cccl_device_binary_search_cleanup(cccl_device_binary_search_build_result_t* build_ptr) try { if (build_ptr == nullptr) { return CUDA_ERROR_INVALID_VALUE; } return cccl_device_transform_cleanup(&build_ptr->transform); } catch (...) { return CUDA_ERROR_UNKNOWN; } // binary_search delegates the bulk of its build_result_t to transform; we // reuse the transform serializer for the inner blob and add a small outer // header carrying the binary_search-specific scalar fields. CUresult cccl_device_binary_search_serialize( const cccl_device_binary_search_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; } void* inner_buf = nullptr; size_t inner_buf_size = 0; CUresult rc = cccl_device_transform_serialize(&build_ptr->transform, &inner_buf, &inner_buf_size); if (rc != CUDA_SUCCESS) { *out_buf = nullptr; *out_size = 0; return rc; } // Ensure the inner buffer is freed even if writing the outer blob throws. std::unique_ptr inner_owner(static_cast(inner_buf)); using namespace cccl::serialization; buffer_writer w; // Outer header re-uses the transform's payload_kind / cc fields, since the // binary_search build_result_t doesn't carry its own copies. write_header(w, CCCL_SERIALIZATION_ALGO_BINARY_SEARCH, build_ptr->transform.payload_kind, build_ptr->transform.cc); w.write_pod(static_cast(build_ptr->op_state_size)); w.write_pod(static_cast(build_ptr->op_state_alignment)); w.write_blob(inner_owner.get(), inner_buf_size); w.release(out_buf, out_size); return CUDA_SUCCESS; } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_binary_search_serialize(): %s\n", exc.what()); fflush(stdout); return CUDA_ERROR_UNKNOWN; } CUresult cccl_device_binary_search_deserialize(cccl_device_binary_search_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}; read_and_validate_header(r, CCCL_SERIALIZATION_ALGO_BINARY_SEARCH); const uint64_t state_size = r.read_pod(); const uint64_t state_align = r.read_pod(); if (state_size > 0) { if (state_align == 0 || (state_align & (state_align - 1)) != 0) { throw std::runtime_error( std::format("serialization blob: invalid binary_search state alignment ({})", state_align)); } } // Pull the inner transform blob and hand it to the transform deserializer. std::unique_ptr inner_owner; size_t inner_size = 0; { void* p = nullptr; r.read_blob_new(&p, &inner_size); inner_owner.reset(static_cast(p)); } if (inner_size == 0) { throw std::runtime_error("serialization blob: empty payload"); } cccl_device_binary_search_build_result_t result{}; CUresult rc = cccl_device_transform_deserialize(&result.transform, inner_owner.get(), inner_size); if (rc != CUDA_SUCCESS) { return rc; } result.op_state_size = static_cast(state_size); result.op_state_alignment = static_cast(state_align); *build_ptr = result; return CUDA_SUCCESS; } catch (const std::exception& exc) { fflush(stderr); printf("\nEXCEPTION in cccl_device_binary_search_deserialize(): %s\n", exc.what()); fflush(stdout); return CUDA_ERROR_UNKNOWN; }