From c176a88695810101da0b2da87622a107466ddf90 Mon Sep 17 00:00:00 2001 From: Anton Smirnov Date: Mon, 21 Sep 2026 14:39:32 +0300 Subject: [PATCH] gh-216: add intersect --- include/spla.h | 1 + include/spla/exec.hpp | 30 ++ src/binding/c_exec.cpp | 3 + src/cpu/cpu_algo_registry.cpp | 48 +- src/cpu/cpu_intersect.hpp | 175 +++++++ src/exec.cpp | 23 + src/opencl/cl_algo_registry.cpp | 26 +- src/opencl/cl_intersect.hpp | 188 +++++++ src/opencl/generated/auto_intersect.hpp | 102 ++++ src/opencl/kernels/intersect.cl | 62 +++ src/schedule/schedule_tasks.cpp | 22 + src/schedule/schedule_tasks.hpp | 22 + tests/CMakeLists.txt | 4 +- tests/test_intersect.cpp | 370 +++++++++++++ tests/test_opencl_intersect.cpp | 669 ++++++++++++++++++++++++ 15 files changed, 1713 insertions(+), 32 deletions(-) create mode 100644 src/cpu/cpu_intersect.hpp create mode 100644 src/opencl/cl_intersect.hpp create mode 100644 src/opencl/generated/auto_intersect.hpp create mode 100644 src/opencl/kernels/intersect.cl create mode 100644 tests/test_intersect.cpp create mode 100644 tests/test_opencl_intersect.cpp diff --git a/include/spla.h b/include/spla.h index 907865fab..f3d4409a8 100644 --- a/include/spla.h +++ b/include/spla.h @@ -386,6 +386,7 @@ SPLA_API spla_Status spla_Exec_v_assign_masked(spla_Vector r, spla_Vector mask, SPLA_API spla_Status spla_Exec_v_map(spla_Vector r, spla_Vector v, spla_OpUnary op, spla_Descriptor desc, spla_ScheduleTask* task); SPLA_API spla_Status spla_Exec_v_reduce(spla_Scalar r, spla_Scalar s, spla_Vector v, spla_OpBinary op_reduce, spla_Descriptor desc, spla_ScheduleTask* task); SPLA_API spla_Status spla_Exec_v_count_mf(spla_Scalar r, spla_Vector v, spla_Descriptor desc, spla_ScheduleTask* task); +SPLA_API spla_Status spla_Exec_intersect(spla_Vector a_keys, spla_Vector a_vals, spla_Vector b_keys, spla_Vector b_vals, spla_Vector r_keys, spla_Vector r_vals, spla_OpBinary op, spla_Descriptor desc, spla_ScheduleTask* task); ////////////////////////////////////////////////////////////////////////////////////// diff --git a/include/spla/exec.hpp b/include/spla/exec.hpp index d24029ffb..3c7f76c1f 100644 --- a/include/spla/exec.hpp +++ b/include/spla/exec.hpp @@ -518,6 +518,36 @@ namespace spla { ref_ptr desc = ref_ptr(), ref_ptr* task_hnd = nullptr); + /** + * @brief Execute (schedule) set intersection of two sorted key-value arrays + * + * Finds intersection of keys from two arrays and applies binary operation + * to corresponding values. + * + * @param a_keys Keys of first array (sorted, unique) + * @param a_vals Values of first array + * @param b_keys Keys of second array (sorted, unique) + * @param b_vals Values of second array + * @param r_keys Result keys (intersection) + * @param r_vals Result values (f(a_val, b_val)) + * @param op Binary operation f(a_val, b_val) -> T + * @param desc Scheduled task descriptor; default is null + * @param task_hnd Optional task hnd; pass not-null pointer to store task + * + * @return Status on task execution or status on hnd creation + */ + + SPLA_API Status exec_intersect( + ref_ptr a_keys, + ref_ptr a_vals, + ref_ptr b_keys, + ref_ptr b_vals, + ref_ptr r_keys, + ref_ptr r_vals, + ref_ptr op, + ref_ptr desc = ref_ptr(), + ref_ptr* task_hnd = nullptr); + }// namespace spla #endif//SPLA_EXEC_HPP diff --git a/src/binding/c_exec.cpp b/src/binding/c_exec.cpp index cda0ffe27..ddb2100d5 100644 --- a/src/binding/c_exec.cpp +++ b/src/binding/c_exec.cpp @@ -102,4 +102,7 @@ spla_Status spla_Exec_v_reduce(spla_Scalar r, spla_Scalar s, spla_Vector v, spla } spla_Status spla_Exec_v_count_mf(spla_Scalar r, spla_Vector v, spla_Descriptor desc, spla_ScheduleTask* task) { SPLA_WRAP_EXEC(exec_v_count_mf, AS_S(r), AS_V(v)); +} +spla_Status spla_Exec_intersect(spla_Vector a_keys, spla_Vector a_vals, spla_Vector b_keys, spla_Vector b_vals, spla_Vector r_keys, spla_Vector r_vals, spla_OpBinary op, spla_Descriptor desc, spla_ScheduleTask* task) { + SPLA_WRAP_EXEC(exec_intersect, AS_V(a_keys), AS_V(a_vals), AS_V(b_keys), AS_V(b_vals), AS_V(r_keys), AS_V(r_vals), AS_OB(op)); } \ No newline at end of file diff --git a/src/cpu/cpu_algo_registry.cpp b/src/cpu/cpu_algo_registry.cpp index 4944eb756..f7a74eca4 100644 --- a/src/cpu/cpu_algo_registry.cpp +++ b/src/cpu/cpu_algo_registry.cpp @@ -31,6 +31,7 @@ #include #include +#include #include #include #include @@ -55,108 +56,113 @@ namespace spla { void register_algo_cpu(Registry* g_registry) { - // algorthm callback + // algorithm callback g_registry->add("callback" CPU_SUFFIX, std::make_shared()); - // algorthm v_count_mf + // algorithm v_count_mf g_registry->add(MAKE_KEY_CPU_0("v_count_mf", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_count_mf", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_count_mf", FLOAT), std::make_shared>()); - // algorthm v_map + // algorithm v_map g_registry->add(MAKE_KEY_CPU_0("v_map", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_map", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_map", FLOAT), std::make_shared>()); - // algorthm v_reduce + // algorithm v_reduce g_registry->add(MAKE_KEY_CPU_0("v_reduce", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_reduce", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_reduce", FLOAT), std::make_shared>()); - // algorthm v_eadd + // algorithm v_eadd g_registry->add(MAKE_KEY_CPU_0("v_eadd", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_eadd", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_eadd", FLOAT), std::make_shared>()); - // algorthm v_emult + // algorithm v_emult g_registry->add(MAKE_KEY_CPU_0("v_emult", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_emult", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_emult", FLOAT), std::make_shared>()); - // algorthm v_eadd_fdb + // algorithm v_eadd_fdb g_registry->add(MAKE_KEY_CPU_0("v_eadd_fdb", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_eadd_fdb", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_eadd_fdb", FLOAT), std::make_shared>()); - // algorthm v_assign_masked + // algorithm v_assign_masked g_registry->add(MAKE_KEY_CPU_0("v_assign_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_assign_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("v_assign_masked", FLOAT), std::make_shared>()); - // algorthm m_reduce_by_row + // algorithm m_reduce_by_row g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_row", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_row", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_row", FLOAT), std::make_shared>()); - // algorthm m_reduce_by_column + // algorithm m_reduce_by_column g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_column", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_column", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_column", FLOAT), std::make_shared>()); - // algorthm m_reduce + // algorithm m_reduce g_registry->add(MAKE_KEY_CPU_0("m_reduce", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_reduce", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_reduce", FLOAT), std::make_shared>()); - // algorthm m_eadd + // algorithm m_eadd g_registry->add(MAKE_KEY_CPU_0("m_eadd", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_eadd", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_eadd", FLOAT), std::make_shared>()); - // algorthm m_emult + // algorithm m_emult g_registry->add(MAKE_KEY_CPU_0("m_emult", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_emult", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_emult", FLOAT), std::make_shared>()); - // algorthm m_transpose + // algorithm m_transpose g_registry->add(MAKE_KEY_CPU_0("m_transpose", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_transpose", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_transpose", FLOAT), std::make_shared>()); - // algorthm m_extract_row + // algorithm m_extract_row g_registry->add(MAKE_KEY_CPU_0("m_extract_row", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_extract_row", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_extract_row", FLOAT), std::make_shared>()); - // algorthm m_extract_column + // algorithm m_extract_column g_registry->add(MAKE_KEY_CPU_0("m_extract_column", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_extract_column", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("m_extract_column", FLOAT), std::make_shared>()); - // algorthm mxv_masked + // algorithm mxv_masked g_registry->add(MAKE_KEY_CPU_0("mxv_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("mxv_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("mxv_masked", FLOAT), std::make_shared>()); - // algorthm vxm_masked + // algorithm vxm_masked g_registry->add(MAKE_KEY_CPU_0("vxm_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("vxm_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("vxm_masked", FLOAT), std::make_shared>()); - // algorthm kron + // algorithm kron g_registry->add(MAKE_KEY_CPU_0("kron", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("kron", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("kron", FLOAT), std::make_shared>()); - // algorthm mxmT_masked + // algorithm mxmT_masked g_registry->add(MAKE_KEY_CPU_0("mxmT_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("mxmT_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("mxmT_masked", FLOAT), std::make_shared>()); - // algorthm mxm + // algorithm mxm g_registry->add(MAKE_KEY_CPU_0("mxm", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("mxm", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CPU_0("mxm", FLOAT), std::make_shared>()); + + // algorithm intersect + g_registry->add(MAKE_KEY_CPU_0("intersect", INT), std::make_shared>()); + g_registry->add(MAKE_KEY_CPU_0("intersect", UINT), std::make_shared>()); + g_registry->add(MAKE_KEY_CPU_0("intersect", FLOAT), std::make_shared>()); } }// namespace spla \ No newline at end of file diff --git a/src/cpu/cpu_intersect.hpp b/src/cpu/cpu_intersect.hpp new file mode 100644 index 000000000..02568b7a5 --- /dev/null +++ b/src/cpu/cpu_intersect.hpp @@ -0,0 +1,175 @@ +/**********************************************************************************/ +/* This file is part of spla project */ +/* https://github.com/SparseLinearAlgebra/spla */ +/**********************************************************************************/ +/* MIT License */ +/* */ +/* Copyright (c) 2023 SparseLinearAlgebra */ +/* */ +/* Permission is hereby granted, free of charge, to any person obtaining a copy */ +/* of this software and associated documentation files (the "Software"), to deal */ +/* in the Software without restriction, including without limitation the rights */ +/* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell */ +/* copies of the Software, and to permit persons to whom the Software is */ +/* furnished to do so, subject to the following conditions: */ +/* */ +/* The above copyright notice and this permission notice shall be included in all */ +/* copies or substantial portions of the Software. */ +/* */ +/* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR */ +/* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, */ +/* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE */ +/* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER */ +/* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, */ +/* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE */ +/* SOFTWARE. */ +/**********************************************************************************/ + +#ifndef SPLA_CPU_INTERSECT_HPP +#define SPLA_CPU_INTERSECT_HPP + +#include + +#include +#include +#include +#include +#include +#include +#include + +#include +#include + +namespace spla { + + template + class Algo_intersect_cpu final : public RegistryAlgo { + public: + ~Algo_intersect_cpu() override = default; + + std::string get_name() override { + return "intersect"; + } + + std::string get_description() override { + return "sequential set intersection of two key-value arrays"; + } + + Status execute(const DispatchContext& ctx) override { + auto t = ctx.task.template cast_safe(); + if (!t) { + LOG_MSG(Status::InvalidArgument, "invalid task type for cpu_intersect"); + return Status::InvalidArgument; + } + + if (!t->a_keys || !t->a_vals || !t->b_keys || !t->b_vals || !t->r_keys || !t->r_vals) { + LOG_MSG(Status::InvalidArgument, "null vector argument for intersect"); + return Status::InvalidArgument; + } + + const uint a_size = t->a_keys->get_n_rows(); + const uint b_size = t->b_keys->get_n_rows(); + + if (a_size == 0 || b_size == 0) { + if (t->r_keys) t->r_keys->clear(); + if (t->r_vals) t->r_vals->clear(); + return Status::Ok; + } + + auto a_keys_vec = t->a_keys.template cast_safe>(); + if (!a_keys_vec) { + LOG_MSG(Status::InvalidArgument, "Failed to cast a_keys to TVector"); + return Status::InvalidArgument; + } + + auto a_vals_vec = t->a_vals.template cast_safe>(); + if (!a_vals_vec) { + LOG_MSG(Status::InvalidArgument, "Failed to cast a_vals to TVector"); + return Status::InvalidArgument; + } + + auto b_keys_vec = t->b_keys.template cast_safe>(); + if (!b_keys_vec) { + LOG_MSG(Status::InvalidArgument, "Failed to cast b_keys to TVector"); + return Status::InvalidArgument; + } + + auto b_vals_vec = t->b_vals.template cast_safe>(); + if (!b_vals_vec) { + LOG_MSG(Status::InvalidArgument, "Failed to cast b_vals to TVector"); + return Status::InvalidArgument; + } + + auto r_keys_vec = t->r_keys.template cast_safe>(); + if (!r_keys_vec) { + LOG_MSG(Status::InvalidArgument, "Failed to cast r_keys to TVector"); + return Status::InvalidArgument; + } + + auto r_vals_vec = t->r_vals.template cast_safe>(); + if (!r_vals_vec) { + LOG_MSG(Status::InvalidArgument, "Failed to cast r_vals to TVector"); + return Status::InvalidArgument; + } + + auto op = t->op.template cast_safe>(); + if (!op) { + LOG_MSG(Status::Error, "Failed to cast binary operation"); + return Status::Error; + } + + r_keys_vec->validate_wd(FormatVector::CpuCoo); + r_vals_vec->validate_wd(FormatVector::CpuCoo); + a_keys_vec->validate_rw(FormatVector::CpuDense); + a_vals_vec->validate_rw(FormatVector::CpuDense); + b_keys_vec->validate_rw(FormatVector::CpuDense); + b_vals_vec->validate_rw(FormatVector::CpuDense); + + auto* p_r_keys = r_keys_vec->template get>(); + auto* p_r_vals = r_vals_vec->template get>(); + const auto* p_a_keys = a_keys_vec->template get>(); + const auto* p_a_vals = a_vals_vec->template get>(); + const auto* p_b_keys = b_keys_vec->template get>(); + const auto* p_b_vals = b_vals_vec->template get>(); + + p_r_keys->Ai.clear(); + p_r_keys->Ax.clear(); + p_r_vals->Ai.clear(); + p_r_vals->Ax.clear(); + p_r_keys->values = 0; + p_r_vals->values = 0; + + const auto& function = op->function; + + // Two-pointer scan over sorted arrays + uint i = 0, j = 0; + while (i < a_size && j < b_size) { + const uint32_t key_a = p_a_keys->Ax[i]; + const uint32_t key_b = p_b_keys->Ax[j]; + + if (key_a < key_b) { + ++i; + } else if (key_b < key_a) { + ++j; + } else { + // Match found + p_r_keys->Ai.push_back(p_r_keys->values); + p_r_keys->Ax.push_back(key_a); + p_r_vals->Ai.push_back(p_r_vals->values); + p_r_vals->Ax.push_back(function(p_a_vals->Ax[i], p_b_vals->Ax[j])); + p_r_keys->values++; + p_r_vals->values++; + ++i; + ++j; + } + } + + LOG_MSG(Status::Ok, "Found " << p_r_keys->values << " matches"); + return Status::Ok; + } + }; + +}// namespace spla + +#endif//SPLA_CPU_INTERSECT_HPP \ No newline at end of file diff --git a/src/exec.cpp b/src/exec.cpp index d62647fba..87f1a6df4 100644 --- a/src/exec.cpp +++ b/src/exec.cpp @@ -405,4 +405,27 @@ namespace spla { EXEC_OR_MAKE_TASK } + Status exec_intersect( + ref_ptr a_keys, + ref_ptr a_vals, + ref_ptr b_keys, + ref_ptr b_vals, + ref_ptr r_keys, + ref_ptr r_vals, + ref_ptr op, + ref_ptr desc, + ref_ptr* task_hnd) { + auto task = make_ref(); + task->a_keys = std::move(a_keys); + task->a_vals = std::move(a_vals); + task->b_keys = std::move(b_keys); + task->b_vals = std::move(b_vals); + task->r_keys = std::move(r_keys); + task->r_vals = std::move(r_vals); + task->op = std::move(op); + task->desc = std::move(desc); + + EXEC_OR_MAKE_TASK + } + }// namespace spla \ No newline at end of file diff --git a/src/opencl/cl_algo_registry.cpp b/src/opencl/cl_algo_registry.cpp index 65726579b..8714970ee 100644 --- a/src/opencl/cl_algo_registry.cpp +++ b/src/opencl/cl_algo_registry.cpp @@ -30,6 +30,7 @@ #include #include +#include #include #include #include @@ -44,55 +45,60 @@ namespace spla { void register_algo_cl(class Registry* g_registry) { - // algorthm v_count_mf + // algorithm v_count_mf g_registry->add(MAKE_KEY_CL_0("v_count_mf", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_count_mf", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_count_mf", FLOAT), std::make_shared>()); - // algorthm v_map + // algorithm v_map g_registry->add(MAKE_KEY_CL_0("v_map", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_map", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_map", FLOAT), std::make_shared>()); - // algorthm v_reduce + // algorithm v_reduce g_registry->add(MAKE_KEY_CL_0("v_reduce", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_reduce", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_reduce", FLOAT), std::make_shared>()); - // algorthm v_eadd + // algorithm v_eadd g_registry->add(MAKE_KEY_CL_0("v_eadd", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_eadd", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_eadd", FLOAT), std::make_shared>()); - // algorthm v_eadd_fdb + // algorithm v_eadd_fdb g_registry->add(MAKE_KEY_CL_0("v_eadd_fdb", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_eadd_fdb", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_eadd_fdb", FLOAT), std::make_shared>()); - // algorthm v_assign_masked + // algorithm v_assign_masked g_registry->add(MAKE_KEY_CL_0("v_assign_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_assign_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("v_assign_masked", FLOAT), std::make_shared>()); - // algorthm m_reduce + // algorithm m_reduce g_registry->add(MAKE_KEY_CL_0("m_reduce", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("m_reduce", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("m_reduce", FLOAT), std::make_shared>()); - // algorthm mxv_masked + // algorithm mxv_masked g_registry->add(MAKE_KEY_CL_0("mxv_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("mxv_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("mxv_masked", FLOAT), std::make_shared>()); - // algorthm vxm_masked + // algorithm vxm_masked g_registry->add(MAKE_KEY_CL_0("vxm_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("vxm_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("vxm_masked", FLOAT), std::make_shared>()); - // algorthm mxmT_masked + // algorithm mxmT_masked g_registry->add(MAKE_KEY_CL_0("mxmT_masked", INT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("mxmT_masked", UINT), std::make_shared>()); g_registry->add(MAKE_KEY_CL_0("mxmT_masked", FLOAT), std::make_shared>()); + + // algorithm intersect + g_registry->add(MAKE_KEY_CL_0("intersect", INT), std::make_shared>()); + g_registry->add(MAKE_KEY_CL_0("intersect", UINT), std::make_shared>()); + g_registry->add(MAKE_KEY_CL_0("intersect", FLOAT), std::make_shared>()); } }// namespace spla diff --git a/src/opencl/cl_intersect.hpp b/src/opencl/cl_intersect.hpp new file mode 100644 index 000000000..1cf841945 --- /dev/null +++ b/src/opencl/cl_intersect.hpp @@ -0,0 +1,188 @@ +/**********************************************************************************/ +/* This file is part of spla project */ +/* https://github.com/SparseLinearAlgebra/spla */ +/**********************************************************************************/ +/* MIT License */ +/* */ +/* Copyright (c) 2023 SparseLinearAlgebra */ +/* */ +/* Permission is hereby granted, free of charge, to any person obtaining a copy */ +/* of this software and associated documentation files (the "Software"), to deal */ +/* in the Software without restriction, including without limitation the rights */ +/* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell */ +/* copies of the Software, and to permit persons to whom the Software is */ +/* furnished to do so, subject to the following conditions: */ +/* */ +/* The above copyright notice and this permission notice shall be included in all */ +/* copies or substantial portions of the Software. */ +/* */ +/* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR */ +/* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, */ +/* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE */ +/* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER */ +/* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, */ +/* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE */ +/* SOFTWARE. */ +/**********************************************************************************/ + +#ifndef SPLA_CL_INTERSECT_HPP +#define SPLA_CL_INTERSECT_HPP + +#include + +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +namespace spla { + + template + class Algo_intersect_cl final : public RegistryAlgo { + public: + ~Algo_intersect_cl() override = default; + + std::string get_name() override { + return "intersect"; + } + + std::string get_description() override { + return "parallel set intersection of key-value arrays on opencl device"; + } + + Status execute(const DispatchContext& ctx) override { + auto t = ctx.task.template cast_safe(); + if (!t) { + LOG_MSG(Status::InvalidArgument, "invalid task type for cl_intersect"); + return Status::InvalidArgument; + } + + return execute_impl(ctx); + } + + private: + Status execute_impl(const DispatchContext& ctx) { + TIME_PROFILE_SCOPE("opencl/intersect"); + + auto t = ctx.task.template cast_safe(); + + if (!t->a_keys || !t->a_vals || !t->b_keys || !t->b_vals || !t->r_keys || !t->r_vals) { + LOG_MSG(Status::InvalidArgument, "null vector argument for cl_intersect"); + return Status::InvalidArgument; + } + + ref_ptr> a_keys_vec = t->a_keys.template cast_safe>(); + ref_ptr> a_vals_vec = t->a_vals.template cast_safe>(); + ref_ptr> b_keys_vec = t->b_keys.template cast_safe>(); + ref_ptr> b_vals_vec = t->b_vals.template cast_safe>(); + + ref_ptr> r_keys_vec = t->r_keys.template cast_safe>(); + ref_ptr> r_vals_vec = t->r_vals.template cast_safe>(); + + ref_ptr> op = t->op.template cast_safe>(); + if (!op) { + LOG_MSG(Status::Error, "Failed to cast binary operation"); + return Status::Error; + } + + r_keys_vec->validate_wd(FormatVector::AccCoo); + r_vals_vec->validate_wd(FormatVector::AccCoo); + a_keys_vec->validate_rw(FormatVector::AccDense); + a_vals_vec->validate_rw(FormatVector::AccDense); + b_keys_vec->validate_rw(FormatVector::AccDense); + b_vals_vec->validate_rw(FormatVector::AccDense); + + auto* p_cl_r_keys = r_keys_vec->template get>(); + auto* p_cl_r_vals = r_vals_vec->template get>(); + const auto* p_cl_a_keys = a_keys_vec->template get>(); + const auto* p_cl_a_vals = a_vals_vec->template get>(); + const auto* p_cl_b_keys = b_keys_vec->template get>(); + const auto* p_cl_b_vals = b_vals_vec->template get>(); + + auto* p_cl_acc = get_acc_cl(); + auto& queue = p_cl_acc->get_queue_default(); + + const uint a_size = a_keys_vec->get_n_rows(); + const uint b_size = b_keys_vec->get_n_rows(); + + if (a_size == 0 || b_size == 0) { + r_keys_vec->clear(); + r_vals_vec->clear(); + return Status::Ok; + } + + const uint max_result = std::min(a_size, b_size); + + cl_coo_vec_resize(max_result, *p_cl_r_keys); + cl_coo_vec_resize(max_result, *p_cl_r_vals); + + // Atomic counter for safe parallel writes to output arrays + CLCounterWrapper cl_result_count; + cl_result_count.set(queue, 0); + + std::shared_ptr program; + if (!ensure_kernel(op, program)) { + return Status::CompilationError; + } + + auto kernel = program->make_kernel("intersect"); + kernel.setArg(0, p_cl_a_keys->Ax); + kernel.setArg(1, p_cl_a_vals->Ax); + kernel.setArg(2, a_size); + kernel.setArg(3, p_cl_b_keys->Ax); + kernel.setArg(4, p_cl_b_vals->Ax); + kernel.setArg(5, b_size); + kernel.setArg(6, p_cl_r_keys->Ai); + kernel.setArg(7, p_cl_r_keys->Ax); + kernel.setArg(8, p_cl_r_vals->Ai); + kernel.setArg(9, p_cl_r_vals->Ax); + kernel.setArg(10, cl_result_count.buffer()); + + const uint wgs = p_cl_acc->get_default_wgs(); + const uint n_groups = div_up_clamp(max_result, wgs, 1, 1024); + const uint global_size = n_groups * wgs; + + cl::NDRange global(global_size); + cl::NDRange local(wgs); + + CL_DISPATCH_PROFILED("exec", queue, kernel, cl::NDRange(), global, local); + + // Read total number of matches from atomic counter + const uint result_count = cl_result_count.get(queue); + LOG_MSG(Status::Ok, "Found " << result_count << " matches"); + + p_cl_r_keys->values = result_count; + p_cl_r_vals->values = result_count; + + return Status::Ok; + } + + bool ensure_kernel(const ref_ptr>& op, std::shared_ptr& program) { + // Build OpenCL kernel with type-specific macros at runtime + CLProgramBuilder builder; + builder.set_name("intersect") + .add_type("TYPE", get_ttype().template as()) + .add_define("WARP_SIZE", get_acc_cl()->get_wave_size()) + .add_op("OP_BINARY", op.template as()) + .set_source(source_intersect) + .acquire(); + + program = builder.get_program(); + return true; + } + }; + +}// namespace spla + +#endif//SPLA_CL_INTERSECT_HPP \ No newline at end of file diff --git a/src/opencl/generated/auto_intersect.hpp b/src/opencl/generated/auto_intersect.hpp new file mode 100644 index 000000000..35c1fb9c7 --- /dev/null +++ b/src/opencl/generated/auto_intersect.hpp @@ -0,0 +1,102 @@ +//////////////////////////////////////////////////////////////////// +// Copyright (c) 2021 - 2026 SparseLinearAlgebra +// Autogenerated file, do not modify +//////////////////////////////////////////////////////////////////// + +#pragma once + +static const char source_intersect[] = R"( + + + +// memory bank conflict-free address and local buffer size +#ifdef LM_NUM_MEM_BANKS + #define LM_ADDR(address) (address + ((address) / LM_NUM_MEM_BANKS)) + #define LM_SIZE(size) (size + (size) / LM_NUM_MEM_BANKS) +#endif + +#define SWAP_KEYS(x, y) \ + uint tmp1 = x; \ + x = y; \ + y = tmp1; + +#define SWAP_VALUES(x, y) \ + TYPE tmp2 = x; \ + x = y; \ + y = tmp2; + +// nearest power of two number greater equals n +uint ceil_to_pow2(uint n) { + uint r = 1; + while (r < n) r *= 2; + return r; +} + +// find first element in a sorted array such x <= element +uint lower_bound(const uint x, + uint first, + uint size, + __global const uint* array) { + while (size > 0) { + int step = size / 2; + + if (array[first + step] < x) { + first = first + step + 1; + size -= step + 1; + } else { + size = step; + } + } + return first; +} + +// find first element in a sorted array such x <= element +uint lower_bound_local(const uint x, + uint first, + uint size, + __local const uint* array) { + while (size > 0) { + int step = size / 2; + + if (array[first + step] < x) { + first = first + step + 1; + size -= step + 1; + } else { + size = step; + } + } + return first; +} +__kernel void intersect( + __global const uint* a_keys, + __global const TYPE* a_vals, + const uint a_size, + __global const uint* b_keys, + __global const TYPE* b_vals, + const uint b_size, + __global uint* r_keys_ai, + __global uint* r_keys_ax, + __global uint* r_vals_ai, + __global TYPE* r_vals_ax, + __global uint* r_size) { + + const uint gid = get_global_id(0); + const uint gsize = get_global_size(0); + + for (uint i = gid; i < a_size; i += gsize) { + const uint key = a_keys[i]; + const TYPE val_a = a_vals[i]; + + // Binary search for matching key in B + const uint pos = lower_bound(key, 0, b_size, b_keys); + + if (pos < b_size && b_keys[pos] == key) { + const uint idx = atomic_add(r_size, 1); + r_keys_ai[idx] = idx; + r_keys_ax[idx] = key; + r_vals_ai[idx] = idx; + r_vals_ax[idx] = OP_BINARY(val_a, b_vals[pos]); + } + } +} +)"; \ No newline at end of file diff --git a/src/opencl/kernels/intersect.cl b/src/opencl/kernels/intersect.cl new file mode 100644 index 000000000..d42f51c80 --- /dev/null +++ b/src/opencl/kernels/intersect.cl @@ -0,0 +1,62 @@ +/**********************************************************************************/ +/* This file is part of spla project */ +/* https://github.com/SparseLinearAlgebra/spla */ +/**********************************************************************************/ +/* MIT License */ +/* */ +/* Copyright (c) 2023 SparseLinearAlgebra */ +/* */ +/* Permission is hereby granted, free of charge, to any person obtaining a copy */ +/* of this software and associated documentation files (the "Software"), to deal */ +/* in the Software without restriction, including without limitation the rights */ +/* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell */ +/* copies of the Software, and to permit persons to whom the Software is */ +/* furnished to do so, subject to the following conditions: */ +/* */ +/* The above copyright notice and this permission notice shall be included in all */ +/* copies or substantial portions of the Software. */ +/* */ +/* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR */ +/* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, */ +/* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE */ +/* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER */ +/* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, */ +/* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE */ +/* SOFTWARE. */ +/**********************************************************************************/ + +#include "common_def.cl" +#include "common_func.cl" + +__kernel void intersect( + __global const uint* a_keys, + __global const TYPE* a_vals, + const uint a_size, + __global const uint* b_keys, + __global const TYPE* b_vals, + const uint b_size, + __global uint* r_keys_ai, + __global uint* r_keys_ax, + __global uint* r_vals_ai, + __global TYPE* r_vals_ax, + __global uint* r_size) { + + const uint gid = get_global_id(0); + const uint gsize = get_global_size(0); + + for (uint i = gid; i < a_size; i += gsize) { + const uint key = a_keys[i]; + const TYPE val_a = a_vals[i]; + + // Binary search for matching key in B + const uint pos = lower_bound(key, 0, b_size, b_keys); + + if (pos < b_size && b_keys[pos] == key) { + const uint idx = atomic_add(r_size, 1); + r_keys_ai[idx] = idx; + r_keys_ax[idx] = key; + r_vals_ai[idx] = idx; + r_vals_ax[idx] = OP_BINARY(val_a, b_vals[pos]); + } + } +} \ No newline at end of file diff --git a/src/schedule/schedule_tasks.cpp b/src/schedule/schedule_tasks.cpp index a65ef2934..2e0c40cff 100644 --- a/src/schedule/schedule_tasks.cpp +++ b/src/schedule/schedule_tasks.cpp @@ -490,4 +490,26 @@ namespace spla { return {r.as(), v.as()}; } + std::string ScheduleTask_intersect::get_name() { + return "intersect"; + } + std::string ScheduleTask_intersect::get_key() { + std::stringstream key; + key << get_name() + << TYPE_KEY(r_vals->get_type()); + return key.str(); + } + std::string ScheduleTask_intersect::get_key_full() { + std::stringstream key; + key << get_name() + << TYPE_KEY(r_keys->get_type()) + << OP_KEY(op); + return key.str(); + } + std::vector> ScheduleTask_intersect::get_args() { + return {a_keys.as(), a_vals.as(), + b_keys.as(), b_vals.as(), + r_keys.as(), r_vals.as(), + op.as()}; + } }// namespace spla diff --git a/src/schedule/schedule_tasks.hpp b/src/schedule/schedule_tasks.hpp index 49b0d87a5..c95393b8b 100644 --- a/src/schedule/schedule_tasks.hpp +++ b/src/schedule/schedule_tasks.hpp @@ -464,6 +464,28 @@ namespace spla { ref_ptr v; }; + /** + * @class ScheduleTask_intersect + * @brief Set intersection of two sorted key-value arrays + */ + class ScheduleTask_intersect final : public ScheduleTaskBase { + public: + ~ScheduleTask_intersect() override = default; + + std::string get_name() override; + std::string get_key() override; + std::string get_key_full() override; + std::vector> get_args() override; + + ref_ptr a_keys; + ref_ptr a_vals; + ref_ptr b_keys; + ref_ptr b_vals; + ref_ptr r_keys; + ref_ptr r_vals; + ref_ptr op; + }; + /** * @} */ diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index e20657a89..d217dccb7 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -18,6 +18,8 @@ spla_test_target(test_op) if (SPLA_BUILD_OPENCL) spla_test_target(test_opencl) spla_test_target(test_opencl_merge) + spla_test_target(test_opencl_intersect) endif () spla_test_target(test_schedule) -spla_test_target(test_vector) \ No newline at end of file +spla_test_target(test_vector) +spla_test_target(test_intersect) \ No newline at end of file diff --git a/tests/test_intersect.cpp b/tests/test_intersect.cpp new file mode 100644 index 000000000..c6865d7d9 --- /dev/null +++ b/tests/test_intersect.cpp @@ -0,0 +1,370 @@ +/**********************************************************************************/ +/* This file is part of spla project */ +/* https://github.com/SparseLinearAlgebra/spla */ +/**********************************************************************************/ +/* MIT License */ +/* */ +/* Copyright (c) 2023 SparseLinearAlgebra */ +/* */ +/* Permission is hereby granted, free of charge, to any person obtaining a copy */ +/* of this software and associated documentation files (the "Software"), to deal */ +/* in the Software without restriction, including without limitation the rights */ +/* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell */ +/* copies of the Software, and to permit persons to whom the Software is */ +/* furnished to do so, subject to the following conditions: */ +/* */ +/* The above copyright notice and this permission notice shall be included in all */ +/* copies or substantial portions of the Software. */ +/* */ +/* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR */ +/* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, */ +/* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE */ +/* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER */ +/* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, */ +/* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE */ +/* SOFTWARE. */ +/**********************************************************************************/ + +#include "test_common.hpp" + +#include +#include + +TEST(intersect, simple_case_float) { + spla::Library::get()->set_accelerator(spla::AcceleratorType::None); + const spla::uint N_A = 4, N_B = 5; + + const spla::uint a_keys_data[N_A] = {0, 1, 4, 7}; + const float a_vals_data[N_A] = {10, 10, 5, 30}; + const spla::uint b_keys_data[N_B] = {2, 4, 5, 7, 8}; + const float b_vals_data[N_B] = {20, 2, 2, 2, 30}; + + auto a_keys = spla::Vector::make(N_A, spla::UINT); + auto a_vals = spla::Vector::make(N_A, spla::FLOAT); + auto b_keys = spla::Vector::make(N_B, spla::UINT); + auto b_vals = spla::Vector::make(N_B, spla::FLOAT); + + for (spla::uint i = 0; i < N_A; ++i) { + a_keys->set_uint(i, a_keys_data[i]); + a_vals->set_float(i, a_vals_data[i]); + } + for (spla::uint i = 0; i < N_B; ++i) { + b_keys->set_uint(i, b_keys_data[i]); + b_vals->set_float(i, b_vals_data[i]); + } + + auto r_keys = spla::Vector::make(std::min(N_A, N_B), spla::UINT); + auto r_vals = spla::Vector::make(std::min(N_A, N_B), spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto count_keys = spla::Scalar::make_int(0); + auto count_vals = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count_keys, r_keys); + spla::exec_v_count_mf(count_vals, r_vals); + EXPECT_EQ(count_keys->as_int(), 2); + EXPECT_EQ(count_vals->as_int(), 2); + + const spla::uint expected_keys[2] = {4, 7}; + const float expected_vals[2] = {7, 32}; + + spla::uint key; + float val; + for (spla::uint i = 0; i < 2; ++i) { + r_keys->get_uint(i, key); + EXPECT_EQ(key, expected_keys[i]); + r_vals->get_float(i, val); + EXPECT_FLOAT_EQ(val, expected_vals[i]); + } +} + +TEST(intersect, no_keys_match_float) { + spla::Library::get()->set_accelerator(spla::AcceleratorType::None); + const spla::uint N = 4; + const spla::uint a_keys_data[N] = {1, 3, 5, 7}; + const float a_vals_data[N] = {10, 30, 50, 70}; + const spla::uint b_keys_data[N] = {2, 4, 6, 8}; + const float b_vals_data[N] = {20, 40, 60, 80}; + + auto a_keys = spla::Vector::make(N, spla::UINT); + auto a_vals = spla::Vector::make(N, spla::FLOAT); + auto b_keys = spla::Vector::make(N, spla::UINT); + auto b_vals = spla::Vector::make(N, spla::FLOAT); + + for (spla::uint i = 0; i < N; ++i) { + a_keys->set_uint(i, a_keys_data[i]); + a_vals->set_float(i, a_vals_data[i]); + b_keys->set_uint(i, b_keys_data[i]); + b_vals->set_float(i, b_vals_data[i]); + } + + auto r_keys = spla::Vector::make(N, spla::UINT); + auto r_vals = spla::Vector::make(N, spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto count_keys = spla::Scalar::make_int(0); + auto count_vals = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count_keys, r_keys); + spla::exec_v_count_mf(count_vals, r_vals); + EXPECT_EQ(count_keys->as_int(), 0); + EXPECT_EQ(count_vals->as_int(), 0); +} + +TEST(intersect, simple_case_int) { + spla::Library::get()->set_accelerator(spla::AcceleratorType::None); + const spla::uint N_A = 4, N_B = 5; + const spla::uint a_keys_data[N_A] = {0, 1, 4, 7}; + const int a_vals_data[N_A] = {10, 10, 5, 30}; + const spla::uint b_keys_data[N_B] = {2, 4, 5, 7, 8}; + const int b_vals_data[N_B] = {20, 2, 2, 2, 30}; + + auto a_keys = spla::Vector::make(N_A, spla::UINT); + auto a_vals = spla::Vector::make(N_A, spla::INT); + auto b_keys = spla::Vector::make(N_B, spla::UINT); + auto b_vals = spla::Vector::make(N_B, spla::INT); + + for (spla::uint i = 0; i < N_A; ++i) { + a_keys->set_uint(i, a_keys_data[i]); + a_vals->set_int(i, a_vals_data[i]); + } + for (spla::uint i = 0; i < N_B; ++i) { + b_keys->set_uint(i, b_keys_data[i]); + b_vals->set_int(i, b_vals_data[i]); + } + + auto r_keys = spla::Vector::make(std::min(N_A, N_B), spla::UINT); + auto r_vals = spla::Vector::make(std::min(N_A, N_B), spla::INT); + + auto op = spla::PLUS_INT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto count_keys = spla::Scalar::make_int(0); + auto count_vals = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count_keys, r_keys); + spla::exec_v_count_mf(count_vals, r_vals); + EXPECT_EQ(count_keys->as_int(), 2); + EXPECT_EQ(count_vals->as_int(), 2); + + const spla::uint expected_keys[2] = {4, 7}; + const int expected_vals[2] = {7, 32}; + + spla::uint key; + int val; + for (spla::uint i = 0; i < 2; ++i) { + r_keys->get_uint(i, key); + EXPECT_EQ(key, expected_keys[i]); + r_vals->get_int(i, val); + EXPECT_EQ(val, expected_vals[i]); + } +} + +TEST(intersect, all_keys_match_float) { + spla::Library::get()->set_accelerator(spla::AcceleratorType::None); + const spla::uint N = 4; + const spla::uint a_keys_data[N] = {1, 2, 3, 4}; + const float a_vals_data[N] = {10, 20, 30, 40}; + const spla::uint b_keys_data[N] = {1, 2, 3, 4}; + const float b_vals_data[N] = {1, 2, 3, 4}; + + auto a_keys = spla::Vector::make(N, spla::UINT); + auto a_vals = spla::Vector::make(N, spla::FLOAT); + auto b_keys = spla::Vector::make(N, spla::UINT); + auto b_vals = spla::Vector::make(N, spla::FLOAT); + + for (spla::uint i = 0; i < N; ++i) { + a_keys->set_uint(i, a_keys_data[i]); + a_vals->set_float(i, a_vals_data[i]); + b_keys->set_uint(i, b_keys_data[i]); + b_vals->set_float(i, b_vals_data[i]); + } + + auto r_keys = spla::Vector::make(N, spla::UINT); + auto r_vals = spla::Vector::make(N, spla::FLOAT); + + auto op = spla::MULT_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto count_keys = spla::Scalar::make_int(0); + auto count_vals = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count_keys, r_keys); + spla::exec_v_count_mf(count_vals, r_vals); + EXPECT_EQ(count_keys->as_int(), 4); + EXPECT_EQ(count_vals->as_int(), 4); + + const spla::uint expected_keys[N] = {1, 2, 3, 4}; + const float expected_vals[N] = {10, 40, 90, 160}; + + spla::uint key; + float val; + for (spla::uint i = 0; i < N; ++i) { + r_keys->get_uint(i, key); + EXPECT_EQ(key, expected_keys[i]); + r_vals->get_float(i, val); + EXPECT_FLOAT_EQ(val, expected_vals[i]); + } +} + +TEST(intersect, different_sizes_float) { + spla::Library::get()->set_accelerator(spla::AcceleratorType::None); + const spla::uint N_A = 3, N_B = 11; + const spla::uint a_keys_data[N_A] = {1, 10, 100}; + const float a_vals_data[N_A] = {100, 200, 300}; + const spla::uint b_keys_data[N_B] = {1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 100}; + const float b_vals_data[N_B] = {1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 1000}; + + auto a_keys = spla::Vector::make(N_A, spla::UINT); + auto a_vals = spla::Vector::make(N_A, spla::FLOAT); + auto b_keys = spla::Vector::make(N_B, spla::UINT); + auto b_vals = spla::Vector::make(N_B, spla::FLOAT); + + for (spla::uint i = 0; i < N_A; ++i) { + a_keys->set_uint(i, a_keys_data[i]); + a_vals->set_float(i, a_vals_data[i]); + } + for (spla::uint i = 0; i < N_B; ++i) { + b_keys->set_uint(i, b_keys_data[i]); + b_vals->set_float(i, b_vals_data[i]); + } + + auto r_keys = spla::Vector::make(N_A, spla::UINT); + auto r_vals = spla::Vector::make(N_A, spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto count_keys = spla::Scalar::make_int(0); + auto count_vals = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count_keys, r_keys); + spla::exec_v_count_mf(count_vals, r_vals); + EXPECT_EQ(count_keys->as_int(), 3); + EXPECT_EQ(count_vals->as_int(), 3); + + const spla::uint expected_keys[3] = {1, 10, 100}; + const float expected_vals[3] = {101, 210, 1300}; + + spla::uint key; + float val; + for (spla::uint i = 0; i < 3; ++i) { + r_keys->get_uint(i, key); + EXPECT_EQ(key, expected_keys[i]); + r_vals->get_float(i, val); + EXPECT_FLOAT_EQ(val, expected_vals[i]); + } +} + +TEST(intersect, subtraction_float) { + spla::Library::get()->set_accelerator(spla::AcceleratorType::None); + const spla::uint N_A = 4, N_B = 5; + const spla::uint a_keys_data[N_A] = {0, 1, 4, 7}; + const float a_vals_data[N_A] = {10, 10, 5, 30}; + const spla::uint b_keys_data[N_B] = {2, 4, 5, 7, 8}; + const float b_vals_data[N_B] = {20, 2, 2, 2, 30}; + + auto a_keys = spla::Vector::make(N_A, spla::UINT); + auto a_vals = spla::Vector::make(N_A, spla::FLOAT); + auto b_keys = spla::Vector::make(N_B, spla::UINT); + auto b_vals = spla::Vector::make(N_B, spla::FLOAT); + + for (spla::uint i = 0; i < N_A; ++i) { + a_keys->set_uint(i, a_keys_data[i]); + a_vals->set_float(i, a_vals_data[i]); + } + for (spla::uint i = 0; i < N_B; ++i) { + b_keys->set_uint(i, b_keys_data[i]); + b_vals->set_float(i, b_vals_data[i]); + } + + auto r_keys = spla::Vector::make(std::min(N_A, N_B), spla::UINT); + auto r_vals = spla::Vector::make(std::min(N_A, N_B), spla::FLOAT); + + auto op = spla::OpBinary::make_float("sub", + "(float a, float b) { return a - b; }", + [](float a, float b) { return a - b; }); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto count_keys = spla::Scalar::make_int(0); + auto count_vals = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count_keys, r_keys); + spla::exec_v_count_mf(count_vals, r_vals); + EXPECT_EQ(count_keys->as_int(), 2); + EXPECT_EQ(count_vals->as_int(), 2); + + const spla::uint expected_keys[2] = {4, 7}; + const float expected_vals[2] = {3, 28}; + + spla::uint key; + float val; + for (spla::uint i = 0; i < 2; ++i) { + r_keys->get_uint(i, key); + EXPECT_EQ(key, expected_keys[i]); + r_vals->get_float(i, val); + EXPECT_FLOAT_EQ(val, expected_vals[i]); + } +} + +TEST(intersect, edge_keys_uint) { + spla::Library::get()->set_accelerator(spla::AcceleratorType::None); + const spla::uint N_A = 3, N_B = 3; + const spla::uint a_keys_data[N_A] = {0, 1, UINT_MAX}; + const spla::uint a_vals_data[N_A] = {10, 20, 30}; + const spla::uint b_keys_data[N_B] = {0, 5, UINT_MAX}; + const spla::uint b_vals_data[N_B] = {1, 2, 3}; + + auto a_keys = spla::Vector::make(N_A, spla::UINT); + auto a_vals = spla::Vector::make(N_A, spla::UINT); + auto b_keys = spla::Vector::make(N_B, spla::UINT); + auto b_vals = spla::Vector::make(N_B, spla::UINT); + + for (spla::uint i = 0; i < N_A; ++i) { + a_keys->set_uint(i, a_keys_data[i]); + a_vals->set_uint(i, a_vals_data[i]); + } + for (spla::uint i = 0; i < N_B; ++i) { + b_keys->set_uint(i, b_keys_data[i]); + b_vals->set_uint(i, b_vals_data[i]); + } + + auto r_keys = spla::Vector::make(std::min(N_A, N_B), spla::UINT); + auto r_vals = spla::Vector::make(std::min(N_A, N_B), spla::UINT); + + auto op = spla::PLUS_UINT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto count_keys = spla::Scalar::make_int(0); + auto count_vals = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count_keys, r_keys); + spla::exec_v_count_mf(count_vals, r_vals); + EXPECT_EQ(count_keys->as_int(), 2); + EXPECT_EQ(count_vals->as_int(), 2); + + const spla::uint expected_keys[2] = {0, UINT_MAX}; + const spla::uint expected_vals[2] = {11, 33}; + + spla::uint key; + spla::uint val; + for (spla::uint i = 0; i < 2; ++i) { + r_keys->get_uint(i, key); + EXPECT_EQ(key, expected_keys[i]); + r_vals->get_uint(i, val); + EXPECT_EQ(val, expected_vals[i]); + } +} + +SPLA_GTEST_MAIN_WITH_FINALIZE \ No newline at end of file diff --git a/tests/test_opencl_intersect.cpp b/tests/test_opencl_intersect.cpp new file mode 100644 index 000000000..6d996d9c4 --- /dev/null +++ b/tests/test_opencl_intersect.cpp @@ -0,0 +1,669 @@ +/**********************************************************************************/ +/* This file is part of spla project */ +/* https://github.com/SparseLinearAlgebra/spla */ +/**********************************************************************************/ +/* MIT License */ +/* */ +/* Copyright (c) 2023 SparseLinearAlgebra */ +/* */ +/* Permission is hereby granted, free of charge, to any person obtaining a copy */ +/* of this software and associated documentation files (the "Software"), to deal */ +/* in the Software without restriction, including without limitation the rights */ +/* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell */ +/* copies of the Software, and to permit persons to whom the Software is */ +/* furnished to do so, subject to the following conditions: */ +/* */ +/* The above copyright notice and this permission notice shall be included in all */ +/* copies or substantial portions of the Software. */ +/* */ +/* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR */ +/* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, */ +/* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE */ +/* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER */ +/* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, */ +/* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE */ +/* SOFTWARE. */ +/**********************************************************************************/ + +#include "test_common.hpp" + +#include + +#include +#include +#include +#include +#include +#include +#include + +using uint = spla::uint; + +inline std::mt19937& generator() { + static thread_local std::mt19937 gen(std::random_device{}()); + return gen; +} + +template +T my_rand(T min, T max) { + std::uniform_int_distribution dist(min, max); + return dist(generator()); +} + +template<> +float my_rand(float min, float max) { + std::uniform_real_distribution dist(min, max); + return dist(generator()); +} + +inline bool ensure_opencl() { + auto* lib = spla::Library::get(); + if (lib->get_accelerator() != nullptr) { + return true; + } + return lib->set_accelerator(spla::AcceleratorType::OpenCL) == spla::Status::Ok && + lib->get_accelerator() != nullptr; +} + +template +void fill_vector(spla::ref_ptr& vec, const std::vector& data) { + for (size_t i = 0; i < data.size(); ++i) { + if constexpr (std::is_same_v) { + vec->set_float(i, data[i]); + } else if constexpr (std::is_same_v) { + vec->set_int(i, data[i]); + } else if constexpr (std::is_same_v) { + vec->set_uint(i, data[i]); + } + } +} + +template +std::vector> collect_sorted(spla::ref_ptr& r_keys, + spla::ref_ptr& r_vals) { + auto count = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count, r_keys); + + std::vector> res; + res.reserve(count->as_int()); + for (int i = 0; i < count->as_int(); ++i) { + uint key; + T val; + r_keys->get_uint(static_cast(i), key); + if constexpr (std::is_same_v) { + r_vals->get_float(static_cast(i), val); + } else if constexpr (std::is_same_v) { + r_vals->get_int(static_cast(i), val); + } else { + r_vals->get_uint(static_cast(i), val); + } + res.emplace_back(key, val); + } + + std::sort(res.begin(), res.end(), + [](const std::pair& a, const std::pair& b) { return a.first < b.first; }); + return res; +} + +template +std::vector> reference_intersect(const std::vector& a_keys, + const std::vector& a_vals, + const std::vector& b_keys, + const std::vector& b_vals, + const std::function& op) { + std::vector> res; + size_t i = 0, j = 0; + while (i < a_keys.size() && j < b_keys.size()) { + if (a_keys[i] < b_keys[j]) { + ++i; + } else if (b_keys[j] < a_keys[i]) { + ++j; + } else { + res.emplace_back(a_keys[i], op(a_vals[i], b_vals[j])); + ++i; + ++j; + } + } + return res; +} + +template +void expect_result(const std::vector>& got, + const std::vector>& expected) { + ASSERT_EQ(got.size(), expected.size()); + for (size_t i = 0; i < expected.size(); ++i) { + EXPECT_EQ(got[i].first, expected[i].first); + if constexpr (std::is_same_v) { + EXPECT_FLOAT_EQ(got[i].second, expected[i].second); + } else { + EXPECT_EQ(got[i].second, expected[i].second); + } + } +} + +std::vector random_sorted_unique_keys(uint n, uint key_range) { + std::vector pool(key_range); + std::iota(pool.begin(), pool.end(), 0u); + std::shuffle(pool.begin(), pool.end(), generator()); + + std::vector keys(pool.begin(), pool.begin() + n); + std::sort(keys.begin(), keys.end()); + return keys; +} + +static int count_keys(const spla::ref_ptr& keys) { + auto count = spla::Scalar::make_int(0); + spla::exec_v_count_mf(count, keys); + return count->as_int(); +} + +static spla::ref_ptr make_dense_uint(const std::vector& data) { + auto v = spla::Vector::make(static_cast(data.size()), spla::UINT); + v->set_format(spla::FormatVector::CpuDense); + for (std::size_t i = 0; i < data.size(); ++i) { + v->set_uint(static_cast(i), data[i]); + } + return v; +} + +static spla::ref_ptr make_dense_float(const std::vector& data) { + auto v = spla::Vector::make(static_cast(data.size()), spla::FLOAT); + v->set_format(spla::FormatVector::CpuDense); + for (std::size_t i = 0; i < data.size(); ++i) { + v->set_float(static_cast(i), data[i]); + } + return v; +} + +TEST(opencl_intersect, simple_case_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 1, 4, 7}; + std::vector a_vals = {10, 10, 5, 30}; + std::vector b_keys = {2, 4, 5, 7, 8}; + std::vector b_vals = {20, 2, 2, 2, 30}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{4, 7}, {7, 32}}); +} + +TEST(opencl_intersect, simple_case_int) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 1, 4, 7}; + std::vector a_vals = {10, 10, 5, 30}; + std::vector b_keys = {2, 4, 5, 7, 8}; + std::vector b_vals = {20, 2, 2, 2, 30}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::INT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::INT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::INT); + + auto op = spla::PLUS_INT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{4, 7}, {7, 32}}); +} + +TEST(opencl_intersect, no_keys_match_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {1, 3, 5, 7}; + std::vector a_vals = {10, 30, 50, 70}; + std::vector b_keys = {2, 4, 6, 8}; + std::vector b_vals = {20, 40, 60, 80}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(a_keys.size(), spla::UINT); + auto r_vals = spla::Vector::make(a_keys.size(), spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), std::vector>{}); +} + +TEST(opencl_intersect, all_keys_match_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {1, 2, 3, 4}; + std::vector a_vals = {10, 20, 30, 40}; + std::vector b_keys = {1, 2, 3, 4}; + std::vector b_vals = {1, 2, 3, 4}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(a_keys.size(), spla::UINT); + auto r_vals = spla::Vector::make(a_keys.size(), spla::FLOAT); + + auto op = spla::MULT_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{1, 10}, {2, 40}, {3, 90}, {4, 160}}); +} + +TEST(opencl_intersect, different_sizes_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {1, 10, 100}; + std::vector a_vals = {100, 200, 300}; + std::vector b_keys = {1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 100}; + std::vector b_vals = {1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 1000}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(a_keys.size(), spla::UINT); + auto r_vals = spla::Vector::make(a_keys.size(), spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{1, 101}, {10, 210}, {100, 1300}}); +} + +TEST(opencl_intersect, subtraction_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 1, 4, 7}; + std::vector a_vals = {10, 10, 5, 30}; + std::vector b_keys = {2, 4, 5, 7, 8}; + std::vector b_vals = {20, 2, 2, 2, 30}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::FLOAT); + + auto op = spla::OpBinary::make_float("sub", + "(float a, float b) { return a - b; }", + [](float a, float b) { return a - b; }); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{4, 3}, {7, 28}}); +} + +TEST(opencl_intersect, zero_key_and_value_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 3}; + std::vector a_vals = {5, 7}; + std::vector b_keys = {0, 3}; + std::vector b_vals = {5, 7}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::FLOAT); + + auto op = spla::OpBinary::make_float("sub", + "(float a, float b) { return a - b; }", + [](float a, float b) { return a - b; }); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{0, 0}, {3, 0}}); +} + +TEST(opencl_intersect, uint_type) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 1, 4, 7}; + std::vector a_vals = {10, 10, 5, 30}; + std::vector b_keys = {2, 4, 5, 7, 8}; + std::vector b_vals = {20, 2, 2, 2, 30}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::UINT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::UINT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + + auto op = spla::PLUS_UINT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{4, 7}, {7, 32}}); +} + +TEST(opencl_intersect, edge_keys_uint) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 1, UINT_MAX}; + std::vector a_vals = {10, 20, 30}; + std::vector b_keys = {0, 5, UINT_MAX}; + std::vector b_vals = {1, 2, 3}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::UINT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::UINT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + + auto op = spla::PLUS_UINT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{0, 11}, {UINT_MAX, 33}}); +} + +TEST(opencl_intersect, null_input_returns_error) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + auto a_vals = make_dense_float({1.0f, 2.0f}); + auto b_keys = make_dense_uint({0u, 1u}); + auto b_vals = make_dense_float({1.0f, 2.0f}); + auto r_keys = spla::Vector::make(2, spla::UINT); + auto r_vals = spla::Vector::make(2, spla::FLOAT); + + spla::ref_ptr a_keys; + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + EXPECT_EQ(status, spla::Status::InvalidArgument); +} + +TEST(opencl_intersect, null_op_returns_error) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + auto a_keys = make_dense_uint({0u, 1u}); + auto a_vals = make_dense_float({1.0f, 2.0f}); + auto b_keys = make_dense_uint({0u, 1u}); + auto b_vals = make_dense_float({1.0f, 2.0f}); + auto r_keys = spla::Vector::make(2, spla::UINT); + auto r_vals = spla::Vector::make(2, spla::FLOAT); + + spla::ref_ptr op; + + auto status = spla::exec_intersect(a_keys, a_vals, b_keys, b_vals, r_keys, r_vals, op); + EXPECT_EQ(status, spla::Status::Error); +} + +TEST(opencl_intersect, deferred_task) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 1, 4, 7}; + std::vector a_vals = {10, 10, 5, 30}; + std::vector b_keys = {2, 4, 5, 7, 8}; + std::vector b_vals = {20, 2, 2, 2, 30}; + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + spla::ref_ptr task; + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op, + spla::ref_ptr(), &task); + ASSERT_EQ(status, spla::Status::Ok); + ASSERT_TRUE(task); + + auto schedule = spla::make_schedule(); + ASSERT_EQ(schedule->step_task(task), spla::Status::Ok); + ASSERT_EQ(schedule->submit(), spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{4, 7}, {7, 32}}); +} + +TEST(opencl_intersect, large_disjoint_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + const uint N = 500000; + + std::vector a_keys(N); + std::vector b_keys(N); + std::vector a_vals(N, 1.0f); + std::vector b_vals(N, 2.0f); + + for (uint i = 0; i < N; ++i) { + a_keys[i] = 2 * i; + b_keys[i] = 2 * i + 1; + } + + auto a_keys_vec = make_dense_uint(a_keys); + auto a_vals_vec = make_dense_float(a_vals); + auto b_keys_vec = make_dense_uint(b_keys); + auto b_vals_vec = make_dense_float(b_vals); + auto r_keys = spla::Vector::make(N, spla::UINT); + auto r_vals = spla::Vector::make(N, spla::FLOAT); + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + EXPECT_EQ(count_keys(r_keys), 0); +} + +TEST(opencl_intersect, large_full_overlap_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + const uint N = 500000; + + std::vector a_keys(N); + std::vector b_keys(N); + std::vector a_vals(N, 3.0f); + std::vector b_vals(N, 4.0f); + + for (uint i = 0; i < N; ++i) { + a_keys[i] = i; + b_keys[i] = i; + } + + auto a_keys_vec = make_dense_uint(a_keys); + auto a_vals_vec = make_dense_float(a_vals); + auto b_keys_vec = make_dense_uint(b_keys); + auto b_vals_vec = make_dense_float(b_vals); + auto r_keys = spla::Vector::make(N, spla::UINT); + auto r_vals = spla::Vector::make(N, spla::FLOAT); + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + EXPECT_EQ(count_keys(r_keys), static_cast(N)); +} + +TEST(opencl_intersect, asymmetric_sizes_large_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + const uint N = 500000; + + std::vector a_keys = {5, 500, N - 1}; + std::vector a_vals = {1.0f, 2.0f, 3.0f}; + std::vector b_keys(N); + std::vector b_vals(N, 1.0f); + + for (uint i = 0; i < N; ++i) { + b_keys[i] = i; + } + + auto a_keys_vec = make_dense_uint(a_keys); + auto a_vals_vec = make_dense_float(a_vals); + auto b_keys_vec = make_dense_uint(b_keys); + auto b_vals_vec = make_dense_float(b_vals); + auto r_keys = spla::Vector::make(a_keys.size(), spla::UINT); + auto r_vals = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + EXPECT_EQ(count_keys(r_keys), 3); +} + +TEST(opencl_intersect, large_values_uint) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + std::vector a_keys = {0, 1, 2}; + std::vector a_vals = {4000000000u, 3000000000u, 1u}; + std::vector b_keys = {0, 1, 2}; + std::vector b_vals = {1u, 2u, 3u}; + + auto a_keys_vec = make_dense_uint(a_keys); + auto a_vals_vec = make_dense_uint(a_vals); + auto b_keys_vec = make_dense_uint(b_keys); + auto b_vals_vec = make_dense_uint(b_vals); + auto r_keys = spla::Vector::make(3, spla::UINT); + auto r_vals = spla::Vector::make(3, spla::UINT); + auto op = spla::PLUS_UINT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + expect_result(collect_sorted(r_keys, r_vals), + std::vector>{{0, 4000000001u}, {1, 3000000002u}, {2, 4u}}); +} + +TEST(opencl_intersect, random_large_float) { + if (!ensure_opencl()) GTEST_SKIP() << "OpenCL accelerator is not available"; + + const uint N = 200000; + const uint KEY_RANGE = 400000; + + std::vector a_keys = random_sorted_unique_keys(N, KEY_RANGE); + std::vector b_keys = random_sorted_unique_keys(N, KEY_RANGE); + std::vector a_vals(N); + std::vector b_vals(N); + + for (uint i = 0; i < N; ++i) { + a_vals[i] = my_rand(0.0f, 1000.0f); + b_vals[i] = my_rand(0.0f, 1000.0f); + } + + auto a_keys_vec = spla::Vector::make(a_keys.size(), spla::UINT); + auto a_vals_vec = spla::Vector::make(a_vals.size(), spla::FLOAT); + auto b_keys_vec = spla::Vector::make(b_keys.size(), spla::UINT); + auto b_vals_vec = spla::Vector::make(b_vals.size(), spla::FLOAT); + + fill_vector(a_keys_vec, a_keys); + fill_vector(a_vals_vec, a_vals); + fill_vector(b_keys_vec, b_keys); + fill_vector(b_vals_vec, b_vals); + + auto r_keys = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::UINT); + auto r_vals = spla::Vector::make(std::min(a_keys.size(), b_keys.size()), spla::FLOAT); + + auto op = spla::PLUS_FLOAT.template cast_safe(); + + auto status = spla::exec_intersect(a_keys_vec, a_vals_vec, b_keys_vec, b_vals_vec, r_keys, r_vals, op); + ASSERT_EQ(status, spla::Status::Ok); + + auto got = collect_sorted(r_keys, r_vals); + auto expected = reference_intersect(a_keys, a_vals, b_keys, b_vals, + [](float a, float b) { return a + b; }); + + EXPECT_GT(got.size(), 0u); + expect_result(got, expected); +} + +SPLA_GTEST_MAIN_WITH_FINALIZE