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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions include/spla.h
Original file line number Diff line number Diff line change
Expand Up @@ -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);

//////////////////////////////////////////////////////////////////////////////////////

Expand Down
30 changes: 30 additions & 0 deletions include/spla/exec.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -518,6 +518,36 @@ namespace spla {
ref_ptr<Descriptor> desc = ref_ptr<Descriptor>(),
ref_ptr<ScheduleTask>* 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<Vector> a_keys,
ref_ptr<Vector> a_vals,
ref_ptr<Vector> b_keys,
ref_ptr<Vector> b_vals,
ref_ptr<Vector> r_keys,
ref_ptr<Vector> r_vals,
ref_ptr<OpBinary> op,
ref_ptr<Descriptor> desc = ref_ptr<Descriptor>(),
ref_ptr<ScheduleTask>* task_hnd = nullptr);

}// namespace spla

#endif//SPLA_EXEC_HPP
3 changes: 3 additions & 0 deletions src/binding/c_exec.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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));
}
48 changes: 27 additions & 21 deletions src/cpu/cpu_algo_registry.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
#include <core/top.hpp>

#include <cpu/cpu_algo_callback.hpp>
#include <cpu/cpu_intersect.hpp>
#include <cpu/cpu_kron.hpp>
#include <cpu/cpu_m_eadd.hpp>
#include <cpu/cpu_m_emult.hpp>
Expand All @@ -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<Algo_callback_cpu>());

// algorthm v_count_mf
// algorithm v_count_mf
g_registry->add(MAKE_KEY_CPU_0("v_count_mf", INT), std::make_shared<Algo_v_count_mf_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("v_count_mf", UINT), std::make_shared<Algo_v_count_mf_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("v_count_mf", FLOAT), std::make_shared<Algo_v_count_mf_cpu<T_FLOAT>>());

// algorthm v_map
// algorithm v_map
g_registry->add(MAKE_KEY_CPU_0("v_map", INT), std::make_shared<Algo_v_map_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("v_map", UINT), std::make_shared<Algo_v_map_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("v_map", FLOAT), std::make_shared<Algo_v_map_cpu<T_FLOAT>>());

// algorthm v_reduce
// algorithm v_reduce
g_registry->add(MAKE_KEY_CPU_0("v_reduce", INT), std::make_shared<Algo_v_reduce_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("v_reduce", UINT), std::make_shared<Algo_v_reduce_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("v_reduce", FLOAT), std::make_shared<Algo_v_reduce_cpu<T_FLOAT>>());

// algorthm v_eadd
// algorithm v_eadd
g_registry->add(MAKE_KEY_CPU_0("v_eadd", INT), std::make_shared<Algo_v_eadd_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("v_eadd", UINT), std::make_shared<Algo_v_eadd_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("v_eadd", FLOAT), std::make_shared<Algo_v_eadd_cpu<T_FLOAT>>());

// algorthm v_emult
// algorithm v_emult
g_registry->add(MAKE_KEY_CPU_0("v_emult", INT), std::make_shared<Algo_v_emult_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("v_emult", UINT), std::make_shared<Algo_v_emult_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("v_emult", FLOAT), std::make_shared<Algo_v_emult_cpu<T_FLOAT>>());

// algorthm v_eadd_fdb
// algorithm v_eadd_fdb
g_registry->add(MAKE_KEY_CPU_0("v_eadd_fdb", INT), std::make_shared<Algo_v_eadd_fdb_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("v_eadd_fdb", UINT), std::make_shared<Algo_v_eadd_fdb_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("v_eadd_fdb", FLOAT), std::make_shared<Algo_v_eadd_fdb_cpu<T_FLOAT>>());

// algorthm v_assign_masked
// algorithm v_assign_masked
g_registry->add(MAKE_KEY_CPU_0("v_assign_masked", INT), std::make_shared<Algo_v_assign_masked_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("v_assign_masked", UINT), std::make_shared<Algo_v_assign_masked_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("v_assign_masked", FLOAT), std::make_shared<Algo_v_assign_masked_cpu<T_FLOAT>>());

// 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<Algo_m_reduce_by_row_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_row", UINT), std::make_shared<Algo_m_reduce_by_row_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_row", FLOAT), std::make_shared<Algo_m_reduce_by_row_cpu<T_FLOAT>>());

// 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<Algo_m_reduce_by_column_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_column", UINT), std::make_shared<Algo_m_reduce_by_column_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_reduce_by_column", FLOAT), std::make_shared<Algo_m_reduce_by_column_cpu<T_FLOAT>>());

// algorthm m_reduce
// algorithm m_reduce
g_registry->add(MAKE_KEY_CPU_0("m_reduce", INT), std::make_shared<Algo_m_reduce_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_reduce", UINT), std::make_shared<Algo_m_reduce_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_reduce", FLOAT), std::make_shared<Algo_m_reduce_cpu<T_FLOAT>>());

// algorthm m_eadd
// algorithm m_eadd
g_registry->add(MAKE_KEY_CPU_0("m_eadd", INT), std::make_shared<Algo_m_eadd_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_eadd", UINT), std::make_shared<Algo_m_eadd_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_eadd", FLOAT), std::make_shared<Algo_m_eadd_cpu<T_FLOAT>>());

// algorthm m_emult
// algorithm m_emult
g_registry->add(MAKE_KEY_CPU_0("m_emult", INT), std::make_shared<Algo_m_emult_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_emult", UINT), std::make_shared<Algo_m_emult_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_emult", FLOAT), std::make_shared<Algo_m_emult_cpu<T_FLOAT>>());

// algorthm m_transpose
// algorithm m_transpose
g_registry->add(MAKE_KEY_CPU_0("m_transpose", INT), std::make_shared<Algo_m_transpose_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_transpose", UINT), std::make_shared<Algo_m_transpose_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_transpose", FLOAT), std::make_shared<Algo_m_transpose_cpu<T_FLOAT>>());

// algorthm m_extract_row
// algorithm m_extract_row
g_registry->add(MAKE_KEY_CPU_0("m_extract_row", INT), std::make_shared<Algo_m_extract_row_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_extract_row", UINT), std::make_shared<Algo_m_extract_row_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_extract_row", FLOAT), std::make_shared<Algo_m_extract_row_cpu<T_FLOAT>>());

// algorthm m_extract_column
// algorithm m_extract_column
g_registry->add(MAKE_KEY_CPU_0("m_extract_column", INT), std::make_shared<Algo_m_extract_column_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("m_extract_column", UINT), std::make_shared<Algo_m_extract_column_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("m_extract_column", FLOAT), std::make_shared<Algo_m_extract_column_cpu<T_FLOAT>>());

// algorthm mxv_masked
// algorithm mxv_masked
g_registry->add(MAKE_KEY_CPU_0("mxv_masked", INT), std::make_shared<Algo_mxv_masked_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("mxv_masked", UINT), std::make_shared<Algo_mxv_masked_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("mxv_masked", FLOAT), std::make_shared<Algo_mxv_masked_cpu<T_FLOAT>>());

// algorthm vxm_masked
// algorithm vxm_masked
g_registry->add(MAKE_KEY_CPU_0("vxm_masked", INT), std::make_shared<Algo_vxm_masked_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("vxm_masked", UINT), std::make_shared<Algo_vxm_masked_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("vxm_masked", FLOAT), std::make_shared<Algo_vxm_masked_cpu<T_FLOAT>>());

// algorthm kron
// algorithm kron
g_registry->add(MAKE_KEY_CPU_0("kron", INT), std::make_shared<Algo_kron_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("kron", UINT), std::make_shared<Algo_kron_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("kron", FLOAT), std::make_shared<Algo_kron_cpu<T_FLOAT>>());

// algorthm mxmT_masked
// algorithm mxmT_masked
g_registry->add(MAKE_KEY_CPU_0("mxmT_masked", INT), std::make_shared<Algo_mxmT_masked_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("mxmT_masked", UINT), std::make_shared<Algo_mxmT_masked_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("mxmT_masked", FLOAT), std::make_shared<Algo_mxmT_masked_cpu<T_FLOAT>>());

// algorthm mxm
// algorithm mxm
g_registry->add(MAKE_KEY_CPU_0("mxm", INT), std::make_shared<Algo_mxm_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("mxm", UINT), std::make_shared<Algo_mxm_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("mxm", FLOAT), std::make_shared<Algo_mxm_cpu<T_FLOAT>>());

// algorithm intersect
g_registry->add(MAKE_KEY_CPU_0("intersect", INT), std::make_shared<Algo_intersect_cpu<T_INT>>());
g_registry->add(MAKE_KEY_CPU_0("intersect", UINT), std::make_shared<Algo_intersect_cpu<T_UINT>>());
g_registry->add(MAKE_KEY_CPU_0("intersect", FLOAT), std::make_shared<Algo_intersect_cpu<T_FLOAT>>());
}

}// namespace spla
175 changes: 175 additions & 0 deletions src/cpu/cpu_intersect.hpp
Original file line number Diff line number Diff line change
@@ -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 <schedule/schedule_tasks.hpp>

#include <core/dispatcher.hpp>
#include <core/logger.hpp>
#include <core/registry.hpp>
#include <core/top.hpp>
#include <core/tscalar.hpp>
#include <core/ttype.hpp>
#include <core/tvector.hpp>

#include <cstdint>
#include <vector>

namespace spla {

template<typename T>
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<ScheduleTask_intersect>();
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<TVector<uint32_t>>();
if (!a_keys_vec) {
LOG_MSG(Status::InvalidArgument, "Failed to cast a_keys to TVector<uint32_t>");
return Status::InvalidArgument;
}

auto a_vals_vec = t->a_vals.template cast_safe<TVector<T>>();
if (!a_vals_vec) {
LOG_MSG(Status::InvalidArgument, "Failed to cast a_vals to TVector<T>");
return Status::InvalidArgument;
}

auto b_keys_vec = t->b_keys.template cast_safe<TVector<uint32_t>>();
if (!b_keys_vec) {
LOG_MSG(Status::InvalidArgument, "Failed to cast b_keys to TVector<uint32_t>");
return Status::InvalidArgument;
}

auto b_vals_vec = t->b_vals.template cast_safe<TVector<T>>();
if (!b_vals_vec) {
LOG_MSG(Status::InvalidArgument, "Failed to cast b_vals to TVector<T>");
return Status::InvalidArgument;
}

auto r_keys_vec = t->r_keys.template cast_safe<TVector<uint32_t>>();
if (!r_keys_vec) {
LOG_MSG(Status::InvalidArgument, "Failed to cast r_keys to TVector<uint32_t>");
return Status::InvalidArgument;
}

auto r_vals_vec = t->r_vals.template cast_safe<TVector<T>>();
if (!r_vals_vec) {
LOG_MSG(Status::InvalidArgument, "Failed to cast r_vals to TVector<T>");
return Status::InvalidArgument;
}

auto op = t->op.template cast_safe<TOpBinary<T, T, T>>();
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<CpuCooVec<uint32_t>>();
auto* p_r_vals = r_vals_vec->template get<CpuCooVec<T>>();
const auto* p_a_keys = a_keys_vec->template get<CpuDenseVec<uint32_t>>();
const auto* p_a_vals = a_vals_vec->template get<CpuDenseVec<T>>();
const auto* p_b_keys = b_keys_vec->template get<CpuDenseVec<uint32_t>>();
const auto* p_b_vals = b_vals_vec->template get<CpuDenseVec<T>>();

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
Loading
Loading