-
Notifications
You must be signed in to change notification settings - Fork 9k
Expand file tree
/
Copy patharrayJaccardIndex.cpp
More file actions
201 lines (173 loc) · 8.79 KB
/
Copy patharrayJaccardIndex.cpp
File metadata and controls
201 lines (173 loc) · 8.79 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
#include <Columns/ColumnArray.h>
#include <Columns/IColumn.h>
#include <Columns/ColumnsNumber.h>
#include <DataTypes/DataTypeArray.h>
#include <DataTypes/DataTypesNumber.h>
#include <DataTypes/IDataType.h>
#include <Functions/FunctionFactory.h>
#include <Functions/FunctionHelpers.h>
#include <DataTypes/DataTypeNothing.h>
#include <Core/ColumnWithTypeAndName.h>
namespace DB
{
namespace ErrorCodes
{
extern const int ILLEGAL_COLUMN;
extern const int ILLEGAL_TYPE_OF_ARGUMENT;
extern const int LOGICAL_ERROR;
}
class FunctionArrayJaccardIndex final : public IFunction
{
private:
using ResultType = Float64;
struct LeftAndRightSizes
{
size_t left_size;
size_t right_size;
};
template <bool left_is_const, bool right_is_const>
static LeftAndRightSizes getArraySizes(const ColumnArray::Offsets & left_offsets, const ColumnArray::Offsets & right_offsets, size_t i)
{
size_t left_size = 0;
size_t right_size = 0;
if constexpr (left_is_const)
left_size = left_offsets[0];
else
left_size = left_offsets[i] - left_offsets[i - 1];
if constexpr (right_is_const)
right_size = right_offsets[0];
else
right_size = right_offsets[i] - right_offsets[i - 1];
return {left_size, right_size};
}
static void vector(
const ColumnArray::Offsets & intersect_offsets,
const ColumnUInt32::Container & left_unique_sizes,
const ColumnUInt32::Container & right_unique_sizes,
PaddedPODArray<ResultType> & res)
{
for (size_t i = 0; i < res.size(); ++i)
{
size_t intersect_size = intersect_offsets[i] - intersect_offsets[i - 1];
size_t union_size = static_cast<size_t>(left_unique_sizes[i])
+ static_cast<size_t>(right_unique_sizes[i]) - intersect_size;
res[i] = static_cast<ResultType>(intersect_size) / static_cast<ResultType>(union_size);
}
}
template <bool left_is_const, bool right_is_const>
static void vectorWithEmptyIntersect(const ColumnArray::Offsets & left_offsets, const ColumnArray::Offsets & right_offsets, PaddedPODArray<ResultType> & res)
{
for (size_t i = 0; i < res.size(); ++i)
{
LeftAndRightSizes sizes = getArraySizes<left_is_const, right_is_const>(left_offsets, right_offsets, i);
if (sizes.left_size == 0 && sizes.right_size == 0)
throw Exception(ErrorCodes::ILLEGAL_TYPE_OF_ARGUMENT, "array aggregate functions cannot be performed on two empty arrays");
res[i] = 0;
}
}
public:
static constexpr auto name = "arrayJaccardIndex";
String getName() const override { return name; }
static FunctionPtr create(ContextPtr context_) { return std::make_shared<FunctionArrayJaccardIndex>(context_); }
explicit FunctionArrayJaccardIndex(ContextPtr context_)
: array_intersect(FunctionFactory::instance().get("arrayIntersect", context_))
, array_uniq(FunctionFactory::instance().get("arrayUniq", context_))
{
}
size_t getNumberOfArguments() const override { return 2; }
bool isSuitableForShortCircuitArgumentsExecution(const DataTypesWithConstInfo &) const override { return true; }
bool useDefaultImplementationForConstants() const override { return true; }
DataTypePtr getReturnTypeImpl(const ColumnsWithTypeAndName & arguments) const override
{
FunctionArgumentDescriptors args{
{"array_1", static_cast<FunctionArgumentDescriptor::TypeValidator>(&isArray), nullptr, "Array"},
{"array_2", static_cast<FunctionArgumentDescriptor::TypeValidator>(&isArray), nullptr, "Array"},
};
validateFunctionArguments(*this, arguments, args);
return std::make_shared<DataTypeNumber<ResultType>>();
}
ColumnPtr executeImpl(const ColumnsWithTypeAndName & arguments, const DataTypePtr &, size_t input_rows_count) const override
{
auto cast_to_array = [&](const ColumnWithTypeAndName & col) -> std::pair<const ColumnArray *, bool>
{
if (const ColumnConst * col_const = typeid_cast<const ColumnConst *>(col.column.get()))
{
const ColumnArray & col_const_array = checkAndGetColumn<ColumnArray>(*col_const->getDataColumnPtr());
return {&col_const_array, true};
}
if (const ColumnArray * col_non_const_array = checkAndGetColumn<ColumnArray>(col.column.get()))
return {col_non_const_array, false};
throw Exception(
ErrorCodes::ILLEGAL_COLUMN, "Argument for function {} must be array but it has type {}.", col.column->getName(), getName());
};
const auto & [left_array, left_is_const] = cast_to_array(arguments[0]);
const auto & [right_array, right_is_const] = cast_to_array(arguments[1]);
auto intersect_array = array_intersect->build(arguments);
ColumnWithTypeAndName intersect_column;
intersect_column.type = intersect_array->getResultType();
intersect_column.column = intersect_array->execute(arguments, intersect_column.type, input_rows_count, /* dry_run = */ false);
const auto * intersect_column_type = checkAndGetDataType<DataTypeArray>(intersect_column.type.get());
if (!intersect_column_type)
throw Exception(ErrorCodes::LOGICAL_ERROR, "Unexpected return type for function arrayIntersect");
ColumnPtr left_unique_column;
ColumnPtr right_unique_column;
const ColumnUInt32 * left_unique_sizes = nullptr;
const ColumnUInt32 * right_unique_sizes = nullptr;
if (!typeid_cast<const DataTypeNothing *>(intersect_column_type->getNestedType().get()))
{
auto execute_array_uniq = [&](const ColumnWithTypeAndName & argument)
{
ColumnsWithTypeAndName single_argument{argument};
auto uniq_function = array_uniq->build(single_argument);
return uniq_function->execute(single_argument, uniq_function->getResultType(), input_rows_count, /* dry_run = */ false)
->convertToFullColumnIfConst();
};
left_unique_column = execute_array_uniq(arguments[0]);
right_unique_column = execute_array_uniq(arguments[1]);
left_unique_sizes = checkAndGetColumn<ColumnUInt32>(left_unique_column.get());
right_unique_sizes = checkAndGetColumn<ColumnUInt32>(right_unique_column.get());
if (!left_unique_sizes || !right_unique_sizes)
throw Exception(ErrorCodes::LOGICAL_ERROR, "Unexpected return type for function arrayUniq");
}
auto col_res = ColumnVector<ResultType>::create();
typename ColumnVector<ResultType>::Container & vec_res = col_res->getData();
vec_res.resize(input_rows_count);
#define EXECUTE_VECTOR(left_is_const, right_is_const) \
if (typeid_cast<const DataTypeNothing *>(intersect_column_type->getNestedType().get())) \
vectorWithEmptyIntersect<left_is_const, right_is_const>(left_array->getOffsets(), right_array->getOffsets(), vec_res); \
else \
{ \
const ColumnArray & intersect_column_array = checkAndGetColumn<ColumnArray>(*intersect_column.column); \
vector(intersect_column_array.getOffsets(), left_unique_sizes->getData(), right_unique_sizes->getData(), vec_res); \
}
if (!left_is_const && !right_is_const)
EXECUTE_VECTOR(false, false)
else if (!left_is_const && right_is_const)
EXECUTE_VECTOR(false, true)
else if (left_is_const && !right_is_const)
EXECUTE_VECTOR(true, false)
else
EXECUTE_VECTOR(true, true)
#undef EXECUTE_VECTOR
return col_res;
}
private:
FunctionOverloadResolverPtr array_intersect;
FunctionOverloadResolverPtr array_uniq;
};
REGISTER_FUNCTION(ArrayJaccardIndex)
{
FunctionDocumentation::Description description = "Returns the [Jaccard index](https://en.wikipedia.org/wiki/Jaccard_index) of two arrays.";
FunctionDocumentation::Syntax syntax = "arrayJaccardIndex(arr_x, arr_y)";
FunctionDocumentation::Arguments arguments = {
{"arr_x", "First array.", {"Array(T)"}},
{"arr_y", "Second array.", {"Array(T)"}},
};
FunctionDocumentation::ReturnedValue returned_value = {"Returns the Jaccard index of `arr_x` and `arr_y`", {"Float64"}};
FunctionDocumentation::Examples examples = {{"Usage example", "SELECT arrayJaccardIndex([1, 2], [2, 3]) AS res", "0.3333333333333333"}};
FunctionDocumentation::IntroducedIn introduced_in = {23, 7};
FunctionDocumentation::Category category = FunctionDocumentation::Category::Array;
FunctionDocumentation documentation = {description, syntax, arguments, {}, returned_value, examples, introduced_in, category};
factory.registerFunction<FunctionArrayJaccardIndex>(documentation);
}
}