diff --git a/be/src/exprs/aggregate/aggregate_function.h b/be/src/exprs/aggregate/aggregate_function.h index 94eaaa9ad72403..ae44e5b1187740 100644 --- a/be/src/exprs/aggregate/aggregate_function.h +++ b/be/src/exprs/aggregate/aggregate_function.h @@ -47,6 +47,7 @@ struct AggregateFunctionAttr { bool is_window_function {false}; bool is_foreach {false}; bool enable_aggregate_function_null_v2 {false}; + bool new_version_percentile {false}; std::vector column_names; }; diff --git a/be/src/exprs/aggregate/aggregate_function_percentile.cpp b/be/src/exprs/aggregate/aggregate_function_percentile.cpp index 61d8e78d7b8db2..86457929c29a0a 100644 --- a/be/src/exprs/aggregate/aggregate_function_percentile.cpp +++ b/be/src/exprs/aggregate/aggregate_function_percentile.cpp @@ -64,9 +64,13 @@ void register_aggregate_function_percentile(AggregateFunctionSimpleFactory& fact using creator = creator_with_type_list; factory.register_function_both("percentile", creator::creator); - factory.register_alias("percentile", "percentile_cont"); factory.register_function_both("percentile_array", creator::creator); + factory.register_function_both("percentile_v2", + creator::creator); + factory.register_function_both("percentile_array_v2", + creator::creator); + factory.register_alias("percentile", "percentile_cont"); } void register_percentile_approx_old_function(AggregateFunctionSimpleFactory& factory) { @@ -89,4 +93,4 @@ void register_aggregate_function_percentile_approx(AggregateFunctionSimpleFactor } #include "common/compile_check_end.h" -} // namespace doris \ No newline at end of file +} // namespace doris diff --git a/be/src/exprs/aggregate/aggregate_function_percentile.h b/be/src/exprs/aggregate/aggregate_function_percentile.h index e36e29a478fa48..2f88f3a2550ea3 100644 --- a/be/src/exprs/aggregate/aggregate_function_percentile.h +++ b/be/src/exprs/aggregate/aggregate_function_percentile.h @@ -36,10 +36,11 @@ #include "core/data_type/data_type_array.h" #include "core/data_type/data_type_nullable.h" #include "core/data_type/data_type_number.h" +#include "core/pod_array.h" #include "core/pod_array_fwd.h" #include "core/types.h" #include "exprs/aggregate/aggregate_function.h" -#include "util/counts.h" +#include "util/percentile_util.h" #include "util/tdigest.h" namespace doris { @@ -48,14 +49,6 @@ namespace doris { class Arena; class BufferReadable; -inline void check_quantile(double quantile) { - if (quantile < 0 || quantile > 1) { - throw Exception(ErrorCode::INVALID_ARGUMENT, - "quantile in func percentile should in [0, 1], but real data is:" + - std::to_string(quantile)); - } -} - struct PercentileApproxState { static constexpr double INIT_QUANTILE = -1.0; PercentileApproxState() = default; @@ -437,6 +430,193 @@ struct PercentileState { } }; +template +struct PercentileExactState { + using ValueType = typename PrimitiveTypeTraits::CppType; + static constexpr size_t bytes_in_arena = 64 - sizeof(PODArray); + using Array = PODArrayWithStackMemory; + + void add_single_range(const ValueType* data, size_t count, double quantile) { + if (!inited_flag) { + _set_single_level(quantile); + inited_flag = true; + } + _append(data, count); + } + + void add_many_range(const ValueType* data, size_t count, + const PaddedPODArray& quantiles_data, const NullMap& null_maps, + size_t start, int64_t arg_size) { + if (!inited_flag) { + _set_many_levels(quantiles_data, null_maps, start, arg_size); + inited_flag = true; + } + if (levels.empty()) { + return; + } + _append(data, count); + } + + void write(BufferWritable& buf) const { + buf.write_binary(inited_flag); + if (!inited_flag) { + return; + } + + levels.write(buf); + size_t size = values.size(); + buf.write_binary(size); + if (size > 0) { + buf.write(reinterpret_cast(values.data()), sizeof(ValueType) * size); + } + } + + void read(BufferReadable& buf) { + reset(); + buf.read_binary(inited_flag); + if (!inited_flag) { + return; + } + + levels.read(buf); + size_t size = 0; + buf.read_binary(size); + values.resize(size); + if (size > 0) { + auto raw = buf.read(sizeof(ValueType) * size); + memcpy(values.data(), raw.data, raw.size); + } + } + + void merge(const PercentileExactState& rhs) { + if (!rhs.inited_flag) { + return; + } + + if (!inited_flag) { + levels = rhs.levels; + inited_flag = true; + } else { + levels.merge(rhs.levels); + } + _append(rhs.values.data(), rhs.values.size()); + } + + void reset() { + values.clear(); + levels.clear(); + inited_flag = false; + } + + double get() const { + if (!inited_flag || levels.empty() || values.empty()) { + return 0.0; + } + + DCHECK_EQ(levels.quantiles.size(), 1); + return _get_result(levels.quantiles[0]); + } + + void insert_result_into(IColumn& to) const { + auto& column_data = assert_cast(to).get_data(); + if (!inited_flag || levels.empty() || values.empty()) { + return; + } + + size_t old_size = column_data.size(); + size_t size = levels.quantiles.size(); + column_data.resize(old_size + size); + auto* result = column_data.data() + old_size; + + if (values.size() == 1) { + for (size_t i = 0; i < size; ++i) { + result[i] = static_cast(values.front()); + } + return; + } + + size_t prev_index = 0; + const auto& quantiles = levels.quantiles; + const auto& permutation = levels.get_permutation(); + for (size_t i = 0; i < size; ++i) { + auto level_index = permutation[i]; + auto level = quantiles[level_index]; + double u = static_cast(values.size() - 1) * level; + auto index = static_cast(u); + + if (index + 1 >= values.size()) { + result[level_index] = + static_cast(*std::max_element(values.begin(), values.end())); + } else { + std::nth_element(values.begin() + prev_index, values.begin() + index, values.end()); + auto* nth_elem = std::min_element(values.begin() + index + 1, values.end()); + result[level_index] = + static_cast(values[index]) + + (u - static_cast(index)) * (static_cast(*nth_elem) - + static_cast(values[index])); + prev_index = index; + } + } + } + +private: + void _set_single_level(double quantile) { + DCHECK(levels.empty()); + check_quantile(quantile); + levels.quantiles.push_back(quantile); + levels.permutation.push_back(0); + } + + void _set_many_levels(const PaddedPODArray& quantiles_data, const NullMap& null_maps, + size_t start, int64_t arg_size) { + DCHECK(levels.empty()); + size_t size = cast_set(arg_size); + levels.quantiles.resize(size); + levels.permutation.resize(size); + for (size_t i = 0; i < size; ++i) { + if (null_maps[start + i]) { + throw Exception(ErrorCode::INVALID_ARGUMENT, + "quantiles in func percentile_array should not have null"); + } + check_quantile(quantiles_data[start + i]); + levels.quantiles[i] = quantiles_data[start + i]; + levels.permutation[i] = i; + } + } + + void _append(const ValueType* data, size_t count) { + if (count == 0) { + return; + } + values.reserve(values.size() + count); + values.insert_assume_reserved(data, data + count); + } + + double _get_result(double quantile) const { + if (values.size() == 1) { + return static_cast(values.front()); + } + + double u = static_cast(values.size() - 1) * quantile; + auto index = static_cast(u); + + if (index + 1 >= values.size()) { + return static_cast(*std::max_element(values.begin(), values.end())); + } + + std::nth_element(values.begin(), values.begin() + index, values.end()); + auto* nth_elem = std::min_element(values.begin() + index + 1, values.end()); + + return static_cast(values[index]) + + (u - static_cast(index)) * + (static_cast(*nth_elem) - static_cast(values[index])); + } + + mutable Array values; + mutable PercentileLevels levels; + bool inited_flag = false; +}; + template class AggregateFunctionPercentile final : public IAggregateFunctionDataHelper, AggregateFunctionPercentile>, @@ -569,5 +749,194 @@ class AggregateFunctionPercentileArray final } }; +template +class AggregateFunctionPercentileV2 final + : public IAggregateFunctionDataHelper, + AggregateFunctionPercentileV2>, + MultiExpression, + NullableAggregateFunction { +public: + using ColVecType = typename PrimitiveTypeTraits::ColumnType; + using Base = + IAggregateFunctionDataHelper, AggregateFunctionPercentileV2>; + AggregateFunctionPercentileV2(const DataTypes& argument_types_) : Base(argument_types_) {} + + String get_name() const override { return "percentile_v2"; } + + DataTypePtr get_return_type() const override { return std::make_shared(); } + + void add(AggregateDataPtr __restrict place, const IColumn** columns, ssize_t row_num, + Arena&) const override { + const auto& sources = + assert_cast(*columns[0]); + const auto& quantile = + assert_cast(*columns[1]); + AggregateFunctionPercentileV2::data(place).add_single_range(&sources.get_data()[row_num], 1, + quantile.get_data()[0]); + } + + void add_batch_single_place(size_t batch_size, AggregateDataPtr place, const IColumn** columns, + Arena&) const override { + const auto& sources = + assert_cast(*columns[0]); + const auto& quantile = + assert_cast(*columns[1]); + DCHECK_EQ(sources.get_data().size(), batch_size); + AggregateFunctionPercentileV2::data(place).add_single_range( + sources.get_data().data(), batch_size, quantile.get_data()[0]); + } + + void add_batch_range(size_t batch_begin, size_t batch_end, AggregateDataPtr place, + const IColumn** columns, Arena&, bool has_null) override { + const auto& sources = + assert_cast(*columns[0]); + const auto& quantile = + assert_cast(*columns[1]); + DCHECK(!has_null); + AggregateFunctionPercentileV2::data(place).add_single_range( + sources.get_data().data() + batch_begin, batch_end - batch_begin + 1, + quantile.get_data()[0]); + } + + void reset(AggregateDataPtr __restrict place) const override { + AggregateFunctionPercentileV2::data(place).reset(); + } + + void merge(AggregateDataPtr __restrict place, ConstAggregateDataPtr rhs, + Arena&) const override { + AggregateFunctionPercentileV2::data(place).merge(AggregateFunctionPercentileV2::data(rhs)); + } + + void serialize(ConstAggregateDataPtr __restrict place, BufferWritable& buf) const override { + AggregateFunctionPercentileV2::data(place).write(buf); + } + + void deserialize(AggregateDataPtr __restrict place, BufferReadable& buf, + Arena&) const override { + AggregateFunctionPercentileV2::data(place).read(buf); + } + + void insert_result_into(ConstAggregateDataPtr __restrict place, IColumn& to) const override { + auto& col = assert_cast(to); + col.insert_value(AggregateFunctionPercentileV2::data(place).get()); + } +}; + +template +class AggregateFunctionPercentileArrayV2 final + : public IAggregateFunctionDataHelper, + AggregateFunctionPercentileArrayV2>, + MultiExpression, + NotNullableAggregateFunction { +public: + using ColVecType = typename PrimitiveTypeTraits::ColumnType; + using Base = IAggregateFunctionDataHelper, + AggregateFunctionPercentileArrayV2>; + AggregateFunctionPercentileArrayV2(const DataTypes& argument_types_) : Base(argument_types_) {} + + String get_name() const override { return "percentile_array_v2"; } + + DataTypePtr get_return_type() const override { + return std::make_shared(make_nullable(std::make_shared())); + } + + void add(AggregateDataPtr __restrict place, const IColumn** columns, ssize_t row_num, + Arena&) const override { + const auto& sources = + assert_cast(*columns[0]); + const auto& quantile_array = + assert_cast(*columns[1]); + const auto& offset_column_data = quantile_array.get_offsets(); + const auto& null_maps = assert_cast( + quantile_array.get_data()) + .get_null_map_data(); + const auto& nested_column = assert_cast( + quantile_array.get_data()) + .get_nested_column(); + const auto& nested_column_data = + assert_cast(nested_column); + size_t start = row_num == 0 ? 0 : offset_column_data[row_num - 1]; + AggregateFunctionPercentileArrayV2::data(place).add_many_range( + &sources.get_data()[row_num], 1, nested_column_data.get_data(), null_maps, start, + cast_set(offset_column_data[row_num] - start)); + } + + void add_batch_single_place(size_t batch_size, AggregateDataPtr place, const IColumn** columns, + Arena&) const override { + const auto& sources = + assert_cast(*columns[0]); + const auto& quantile_array = + assert_cast(*columns[1]); + const auto& offset_column_data = quantile_array.get_offsets(); + const auto& null_maps = assert_cast( + quantile_array.get_data()) + .get_null_map_data(); + const auto& nested_column = assert_cast( + quantile_array.get_data()) + .get_nested_column(); + const auto& nested_column_data = + assert_cast(nested_column); + DCHECK_EQ(sources.get_data().size(), batch_size); + AggregateFunctionPercentileArrayV2::data(place).add_many_range( + sources.get_data().data(), batch_size, nested_column_data.get_data(), null_maps, 0, + cast_set(offset_column_data[0])); + } + + void add_batch_range(size_t batch_begin, size_t batch_end, AggregateDataPtr place, + const IColumn** columns, Arena&, bool has_null) override { + const auto& sources = + assert_cast(*columns[0]); + const auto& quantile_array = + assert_cast(*columns[1]); + const auto& offset_column_data = quantile_array.get_offsets(); + const auto& null_maps = assert_cast( + quantile_array.get_data()) + .get_null_map_data(); + const auto& nested_column = assert_cast( + quantile_array.get_data()) + .get_nested_column(); + const auto& nested_column_data = + assert_cast(nested_column); + DCHECK(!has_null); + size_t start = batch_begin == 0 ? 0 : offset_column_data[batch_begin - 1]; + AggregateFunctionPercentileArrayV2::data(place).add_many_range( + sources.get_data().data() + batch_begin, batch_end - batch_begin + 1, + nested_column_data.get_data(), null_maps, start, + cast_set(offset_column_data[batch_begin] - start)); + } + + void reset(AggregateDataPtr __restrict place) const override { + AggregateFunctionPercentileArrayV2::data(place).reset(); + } + + void merge(AggregateDataPtr __restrict place, ConstAggregateDataPtr rhs, + Arena&) const override { + AggregateFunctionPercentileArrayV2::data(place).merge( + AggregateFunctionPercentileArrayV2::data(rhs)); + } + + void serialize(ConstAggregateDataPtr __restrict place, BufferWritable& buf) const override { + AggregateFunctionPercentileArrayV2::data(place).write(buf); + } + + void deserialize(AggregateDataPtr __restrict place, BufferReadable& buf, + Arena&) const override { + AggregateFunctionPercentileArrayV2::data(place).read(buf); + } + + void insert_result_into(ConstAggregateDataPtr __restrict place, IColumn& to) const override { + auto& to_arr = assert_cast(to); + auto& to_nested_col = to_arr.get_data(); + if (to_nested_col.is_nullable()) { + auto* col_null = reinterpret_cast(&to_nested_col); + AggregateFunctionPercentileArrayV2::data(place).insert_result_into( + col_null->get_nested_column()); + col_null->get_null_map_data().resize_fill(col_null->get_nested_column().size(), 0); + } else { + AggregateFunctionPercentileArrayV2::data(place).insert_result_into(to_nested_col); + } + to_arr.get_offsets().push_back(to_nested_col.size()); + } +}; #include "common/compile_check_end.h" -} // namespace doris \ No newline at end of file +} // namespace doris diff --git a/be/src/exprs/aggregate/aggregate_function_simple_factory.h b/be/src/exprs/aggregate/aggregate_function_simple_factory.h index 16f331aa6db60d..48e9edba5dde8c 100644 --- a/be/src/exprs/aggregate/aggregate_function_simple_factory.h +++ b/be/src/exprs/aggregate/aggregate_function_simple_factory.h @@ -136,6 +136,15 @@ class AggregateFunctionSimpleFactory { if (function_alias.contains(name)) { name_str = function_alias[name]; } + + if (attr.new_version_percentile) { + if (name_str == "percentile" || name_str == "percentile_cont") { + name_str = "percentile_v2"; + } else if (name_str == "percentile_array") { + name_str = "percentile_array_v2"; + } + } + if (nullable) { return nullable_aggregate_functions.find(name_str) == nullable_aggregate_functions.end() ? nullptr diff --git a/be/src/exprs/vectorized_agg_fn.cpp b/be/src/exprs/vectorized_agg_fn.cpp index 4eac3830c71a91..a093693dbc74a7 100644 --- a/be/src/exprs/vectorized_agg_fn.cpp +++ b/be/src/exprs/vectorized_agg_fn.cpp @@ -230,6 +230,9 @@ Status AggFnEvaluator::prepare(RuntimeState* state, const RowDescriptor& desc, .is_foreach = is_foreach, .enable_aggregate_function_null_v2 = state->enable_aggregate_function_null_v2(), + .new_version_percentile = + state->query_options().__isset.new_version_percentile && + state->query_options().new_version_percentile, .column_names = std::move(column_names)}); } else { _function = AggregateFunctionSimpleFactory::instance().get( @@ -239,6 +242,9 @@ Status AggFnEvaluator::prepare(RuntimeState* state, const RowDescriptor& desc, .is_foreach = is_foreach, .enable_aggregate_function_null_v2 = state->enable_aggregate_function_null_v2(), + .new_version_percentile = + state->query_options().__isset.new_version_percentile && + state->query_options().new_version_percentile, .column_names = std::move(column_names)}); } } diff --git a/be/src/util/counts.h b/be/src/util/percentile_util.h similarity index 79% rename from be/src/util/counts.h rename to be/src/util/percentile_util.h index a0299f6318b0eb..05003e5cb636d0 100644 --- a/be/src/util/counts.h +++ b/be/src/util/percentile_util.h @@ -21,13 +21,27 @@ #include #include +#include #include +#include +#include +#include "common/cast_set.h" +#include "common/exception.h" +#include "core/column/column_nullable.h" #include "core/pod_array.h" #include "core/string_buffer.hpp" namespace doris { +inline void check_quantile(double quantile) { + if (quantile < 0 || quantile > 1) { + throw Exception(ErrorCode::INVALID_ARGUMENT, + "quantile in func percentile should in [0, 1], but real data is:" + + std::to_string(quantile)); + } +} + template class Counts { public: @@ -223,4 +237,64 @@ class Counts { std::vector> _sorted_nums_vec; }; +class PercentileLevels { +public: + void merge(const PercentileLevels& rhs) { + if (rhs.empty()) { + return; + } + + if (empty()) { + quantiles = rhs.quantiles; + permutation = rhs.permutation; + return; + } + + DCHECK_EQ(quantiles.size(), rhs.quantiles.size()); + for (size_t i = 0; i < quantiles.size(); ++i) { + DCHECK_EQ(quantiles[i], rhs.quantiles[i]); + } + } + + void write(BufferWritable& buf) const { + int size_num = cast_set(quantiles.size()); + buf.write_binary(size_num); + for (const auto& quantile : quantiles) { + buf.write_binary(quantile); + } + } + + void read(BufferReadable& buf) { + int size_num = 0; + buf.read_binary(size_num); + + quantiles.resize(size_num); + permutation.resize(size_num); + for (int i = 0; i < size_num; ++i) { + buf.read_binary(quantiles[i]); + permutation[i] = cast_set(i); + } + } + + void clear() { + quantiles.clear(); + permutation.clear(); + } + + bool empty() const { return quantiles.empty(); } + + const std::vector& get_permutation() const { + sort_permutation(); + return permutation; + } + + void sort_permutation() const { + pdqsort(permutation.begin(), permutation.end(), + [this](size_t lhs, size_t rhs) { return quantiles[lhs] < quantiles[rhs]; }); + } + + std::vector quantiles; + mutable std::vector permutation; +}; + } // namespace doris diff --git a/be/test/exprs/aggregate/agg_test.cpp b/be/test/exprs/aggregate/agg_test.cpp index 44116c752a38d4..459c24cdb6af05 100644 --- a/be/test/exprs/aggregate/agg_test.cpp +++ b/be/test/exprs/aggregate/agg_test.cpp @@ -28,6 +28,7 @@ #include "core/column/column_nullable.h" #include "core/column/column_string.h" #include "core/column/column_vector.h" +#include "core/data_type/data_type_array.h" #include "core/data_type/data_type_date.h" #include "core/data_type/data_type_date_time.h" #include "core/data_type/data_type_nullable.h" @@ -49,6 +50,7 @@ void register_aggregate_function_sum(AggregateFunctionSimpleFactory& factory); void register_aggregate_function_topn(AggregateFunctionSimpleFactory& factory); void register_aggregate_function_bit(AggregateFunctionSimpleFactory& factory); void register_aggregate_function_minmax(AggregateFunctionSimpleFactory& factory); +void register_aggregate_function_percentile(AggregateFunctionSimpleFactory& factory); void register_aggregate_function_replace_reader_load(AggregateFunctionSimpleFactory& factory); TEST(AggTest, basic_test) { @@ -136,6 +138,53 @@ TEST(AggTest, window_function_test) { EXPECT_EQ(size2, 1); } +TEST(AggTest, percentile_query_option_routes_default_names_to_v2) { + AggregateFunctionSimpleFactory factory; + register_aggregate_function_percentile(factory); + int be_version = BeExecVersionManager::get_newest_version(); + + DataTypes percentile_types = {std::make_shared(), + std::make_shared()}; + auto percentile_result_type = std::make_shared(); + auto percentile_v1 = + factory.get("percentile", percentile_types, percentile_result_type, false, be_version); + ASSERT_NE(percentile_v1, nullptr); + EXPECT_EQ(percentile_v1->get_name(), "percentile"); + + auto percentile_v2 = + factory.get("percentile", percentile_types, percentile_result_type, false, be_version, + {.new_version_percentile = true, .column_names = {}}); + ASSERT_NE(percentile_v2, nullptr); + EXPECT_EQ(percentile_v2->get_name(), "percentile_v2"); + + auto percentile_cont_v1 = factory.get("percentile_cont", percentile_types, + percentile_result_type, false, be_version); + ASSERT_NE(percentile_cont_v1, nullptr); + EXPECT_EQ(percentile_cont_v1->get_name(), "percentile"); + + auto percentile_cont_v2 = + factory.get("percentile_cont", percentile_types, percentile_result_type, false, + be_version, {.new_version_percentile = true, .column_names = {}}); + ASSERT_NE(percentile_cont_v2, nullptr); + EXPECT_EQ(percentile_cont_v2->get_name(), "percentile_v2"); + + DataTypes percentile_array_types = { + std::make_shared(), + std::make_shared(make_nullable(std::make_shared()))}; + auto percentile_array_result_type = + std::make_shared(make_nullable(std::make_shared())); + auto percentile_array_v1 = factory.get("percentile_array", percentile_array_types, + percentile_array_result_type, false, be_version); + ASSERT_NE(percentile_array_v1, nullptr); + EXPECT_EQ(percentile_array_v1->get_name(), "percentile_array"); + + auto percentile_array_v2 = + factory.get("percentile_array", percentile_array_types, percentile_array_result_type, + false, be_version, {.new_version_percentile = true, .column_names = {}}); + ASSERT_NE(percentile_array_v2, nullptr); + EXPECT_EQ(percentile_array_v2->get_name(), "percentile_array_v2"); +} + TEST(AggTest, window_function_test2) { AggregateFunctionSimpleFactory factory; register_aggregate_function_sum(factory); diff --git a/be/test/util/counts_test.cpp b/be/test/util/counts_test.cpp deleted file mode 100644 index 72f0fc122dfcf2..00000000000000 --- a/be/test/util/counts_test.cpp +++ /dev/null @@ -1,82 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -#include "util/counts.h" - -#include -#include - -#include - -#include "gtest/gtest_pred_impl.h" - -namespace doris { - -class TCountsTest : public testing::Test {}; - -TEST_F(TCountsTest, TotalTest) { - Counts counts; - // 1 1 1 2 5 7 7 9 9 19 - // >>> import numpy as np - // >>> a = np.array([1,1,1,2,5,7,7,9,9,19]) - // >>> p = np.percentile(a, 20) - counts.increment(1, 3); - counts.increment(5, 1); - counts.increment(2, 1); - counts.increment(9, 1); - counts.increment(9, 1); - counts.increment(19, 1); - counts.increment(7, 2); - - double result = counts.terminate(0.2); - EXPECT_EQ(1, result); - - auto cs = ColumnString::create(); - BufferWritable bw(*cs); - counts.serialize(bw); - bw.commit(); - - Counts other; - StringRef res(cs->get_chars().data(), cs->get_chars().size()); - BufferReadable br(res); - other.unserialize(br); - double result1 = other.terminate(0.2); - EXPECT_EQ(result, result1); - - Counts other1; - other1.increment(1, 1); - other1.increment(100, 3); - other1.increment(50, 3); - other1.increment(10, 1); - other1.increment(99, 2); - - // deserialize other1 - cs->clear(); - other1.serialize(bw); - bw.commit(); - Counts other1_deserialized; - BufferReadable br1(res); - other1_deserialized.unserialize(br1); - - Counts merge_res; - merge_res.merge(&other); - merge_res.merge(&other1_deserialized); - // 1 1 1 1 2 5 7 7 9 9 10 19 50 50 50 99 99 100 100 100 - EXPECT_EQ(merge_res.terminate(0.3), 6.4); -} - -} // namespace doris diff --git a/be/test/util/percentile_util_test.cpp b/be/test/util/percentile_util_test.cpp new file mode 100644 index 00000000000000..5f6126978e6fff --- /dev/null +++ b/be/test/util/percentile_util_test.cpp @@ -0,0 +1,290 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +#include "util/percentile_util.h" + +#include + +#include +#include +#include + +namespace doris { + +class PercentileUtilTest : public testing::Test {}; + +TEST_F(PercentileUtilTest, CountsTotalTest) { + Counts counts; + // 1 1 1 2 5 7 7 9 9 19 + // >>> import numpy as np + // >>> a = np.array([1,1,1,2,5,7,7,9,9,19]) + // >>> p = np.percentile(a, 20) + counts.increment(1, 3); + counts.increment(5, 1); + counts.increment(2, 1); + counts.increment(9, 1); + counts.increment(9, 1); + counts.increment(19, 1); + counts.increment(7, 2); + + double result = counts.terminate(0.2); + EXPECT_EQ(1, result); + + auto cs = ColumnString::create(); + BufferWritable bw(*cs); + counts.serialize(bw); + bw.commit(); + + Counts other; + StringRef res(cs->get_chars().data(), cs->get_chars().size()); + BufferReadable br(res); + other.unserialize(br); + double result1 = other.terminate(0.2); + EXPECT_EQ(result, result1); + + Counts other1; + other1.increment(1, 1); + other1.increment(100, 3); + other1.increment(50, 3); + other1.increment(10, 1); + other1.increment(99, 2); + + // deserialize other1 + cs->clear(); + other1.serialize(bw); + bw.commit(); + Counts other1_deserialized; + BufferReadable br1(res); + other1_deserialized.unserialize(br1); + + Counts merge_res; + merge_res.merge(&other); + merge_res.merge(&other1_deserialized); + // 1 1 1 1 2 5 7 7 9 9 10 19 50 50 50 99 99 100 100 100 + EXPECT_EQ(merge_res.terminate(0.3), 6.4); +} + +TEST_F(PercentileUtilTest, CountsBoundaryBehavior) { + Counts empty_counts; + EXPECT_DOUBLE_EQ(0.0, empty_counts.terminate(0.5)); + + Counts single_counts; + single_counts.increment(42); + EXPECT_DOUBLE_EQ(42.0, single_counts.terminate(0.0)); + EXPECT_DOUBLE_EQ(42.0, single_counts.terminate(0.5)); + EXPECT_DOUBLE_EQ(42.0, single_counts.terminate(1.0)); +} + +TEST_F(PercentileUtilTest, CountsSerializeMergedState) { + Counts left; + left.increment(5); + left.increment(1); + + auto col = ColumnString::create(); + BufferWritable writer(*col); + left.serialize(writer); + writer.commit(); + + StringRef left_data(col->get_chars().data(), col->get_chars().size()); + BufferReadable left_reader(left_data); + Counts left_sorted; + left_sorted.unserialize(left_reader); + + Counts right; + right.increment(9); + right.increment(7); + + col->clear(); + right.serialize(writer); + writer.commit(); + + StringRef right_data(col->get_chars().data(), col->get_chars().size()); + BufferReadable right_reader(right_data); + Counts right_sorted; + right_sorted.unserialize(right_reader); + + Counts merged; + merged.merge(&left_sorted); + merged.merge(&right_sorted); + + col->clear(); + merged.serialize(writer); + writer.commit(); + + StringRef merged_data(col->get_chars().data(), col->get_chars().size()); + BufferReadable merged_reader(merged_data); + Counts restored; + restored.unserialize(merged_reader); + + EXPECT_DOUBLE_EQ(1.0, restored.terminate(0.0)); + EXPECT_DOUBLE_EQ(6.0, restored.terminate(0.5)); + EXPECT_DOUBLE_EQ(9.0, restored.terminate(1.0)); +} + +TEST_F(PercentileUtilTest, CheckQuantileBoundary) { + EXPECT_NO_THROW(check_quantile(0.0)); + EXPECT_NO_THROW(check_quantile(0.5)); + EXPECT_NO_THROW(check_quantile(1.0)); + EXPECT_THROW(check_quantile(-0.0001), Exception); + EXPECT_THROW(check_quantile(1.0001), Exception); +} + +TEST_F(PercentileUtilTest, EmptyLevelsState) { + PercentileLevels levels; + EXPECT_TRUE(levels.empty()); + EXPECT_TRUE(levels.quantiles.empty()); + EXPECT_TRUE(levels.permutation.empty()); + + auto col = ColumnString::create(); + BufferWritable writer(*col); + levels.write(writer); + writer.commit(); + + StringRef data(col->get_chars().data(), col->get_chars().size()); + BufferReadable reader(data); + + PercentileLevels restored; + restored.read(reader); + EXPECT_TRUE(restored.empty()); + EXPECT_TRUE(restored.quantiles.empty()); + EXPECT_TRUE(restored.permutation.empty()); +} + +TEST_F(PercentileUtilTest, SortPermutationKeepsQuantiles) { + PercentileLevels levels; + levels.quantiles = {0.56, 0.12, 0.45, 0.23}; + levels.permutation = {0, 1, 2, 3}; + + const auto& permutation = levels.get_permutation(); + + EXPECT_EQ((std::vector {1, 3, 2, 0}), permutation); + EXPECT_EQ((std::vector {1, 3, 2, 0}), levels.permutation); + EXPECT_EQ((std::vector {0.56, 0.12, 0.45, 0.23}), levels.quantiles); +} + +TEST_F(PercentileUtilTest, SortPermutationWithDuplicateQuantiles) { + PercentileLevels levels; + levels.quantiles = {0.7, 0.2, 0.2, 0.9, 0.7}; + levels.permutation = {0, 1, 2, 3, 4}; + + const auto& permutation = levels.get_permutation(); + + ASSERT_EQ(levels.quantiles.size(), permutation.size()); + std::vector sorted_perm = permutation; + std::sort(sorted_perm.begin(), sorted_perm.end()); + EXPECT_EQ((std::vector {0, 1, 2, 3, 4}), sorted_perm); + for (size_t i = 1; i < permutation.size(); ++i) { + EXPECT_LE(levels.quantiles[permutation[i - 1]], levels.quantiles[permutation[i]]); + } +} + +TEST_F(PercentileUtilTest, WriteReadDefersPermutationSort) { + PercentileLevels levels; + levels.quantiles = {0.56, 0.12, 0.45, 0.23}; + levels.permutation = {0, 1, 2, 3}; + + auto col = ColumnString::create(); + BufferWritable writer(*col); + levels.write(writer); + writer.commit(); + + StringRef data(col->get_chars().data(), col->get_chars().size()); + BufferReadable reader(data); + + PercentileLevels restored; + restored.read(reader); + + EXPECT_EQ(levels.quantiles, restored.quantiles); + EXPECT_EQ((std::vector {0, 1, 2, 3}), restored.permutation); + EXPECT_EQ((std::vector {1, 3, 2, 0}), restored.get_permutation()); +} + +TEST_F(PercentileUtilTest, ReadBuildsPermutationFromUnsortedQuantiles) { + auto col = ColumnString::create(); + BufferWritable writer(*col); + writer.write_binary(4); + writer.write_binary(0.56); + writer.write_binary(0.12); + writer.write_binary(0.45); + writer.write_binary(0.23); + writer.commit(); + + StringRef data(col->get_chars().data(), col->get_chars().size()); + BufferReadable reader(data); + + PercentileLevels levels; + levels.read(reader); + + EXPECT_EQ((std::vector {0.56, 0.12, 0.45, 0.23}), levels.quantiles); + EXPECT_EQ((std::vector {0, 1, 2, 3}), levels.permutation); + EXPECT_EQ((std::vector {1, 3, 2, 0}), levels.get_permutation()); +} + +TEST_F(PercentileUtilTest, MergeFromEmptyCopiesState) { + PercentileLevels src; + src.quantiles = {0.56, 0.12, 0.45, 0.23}; + src.permutation = {0, 1, 2, 3}; + + PercentileLevels dst; + dst.merge(src); + + EXPECT_EQ(src.quantiles, dst.quantiles); + EXPECT_EQ(src.permutation, dst.permutation); +} + +TEST_F(PercentileUtilTest, MergeEmptyRightKeepsLeft) { + PercentileLevels lhs; + lhs.quantiles = {0.56, 0.12, 0.45, 0.23}; + lhs.permutation = {0, 1, 2, 3}; + + PercentileLevels rhs; + lhs.merge(rhs); + + EXPECT_EQ((std::vector {0.56, 0.12, 0.45, 0.23}), lhs.quantiles); + EXPECT_EQ((std::vector {0, 1, 2, 3}), lhs.permutation); + EXPECT_EQ((std::vector {1, 3, 2, 0}), lhs.get_permutation()); +} + +TEST_F(PercentileUtilTest, MergeSameStateKeepsState) { + PercentileLevels lhs; + lhs.quantiles = {0.56, 0.12, 0.45, 0.23}; + lhs.permutation = {0, 1, 2, 3}; + + PercentileLevels rhs; + rhs.quantiles = lhs.quantiles; + rhs.permutation = lhs.permutation; + + lhs.merge(rhs); + + EXPECT_EQ((std::vector {0.56, 0.12, 0.45, 0.23}), lhs.quantiles); + EXPECT_EQ((std::vector {0, 1, 2, 3}), lhs.permutation); + EXPECT_EQ((std::vector {1, 3, 2, 0}), lhs.get_permutation()); +} + +TEST_F(PercentileUtilTest, ClearResetsState) { + PercentileLevels levels; + levels.quantiles = {0.56, 0.12, 0.45, 0.23}; + levels.permutation = {1, 3, 2, 0}; + + levels.clear(); + + EXPECT_TRUE(levels.empty()); + EXPECT_TRUE(levels.quantiles.empty()); + EXPECT_TRUE(levels.permutation.empty()); +} + +} // namespace doris diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/load/NereidsStreamLoadPlanner.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/load/NereidsStreamLoadPlanner.java index ad8fa4475e5aa8..9f368f6dacd423 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/load/NereidsStreamLoadPlanner.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/load/NereidsStreamLoadPlanner.java @@ -325,6 +325,7 @@ public TPipelineFragmentParams plan(TUniqueId loadId, int fragmentInstanceIdInde : false; queryOptions.setEnableMemtableOnSinkNode(enableMemtableOnSinkNode); queryOptions.setNewVersionUnixTimestamp(true); + queryOptions.setNewVersionPercentile(true); params.setQueryOptions(queryOptions); TQueryGlobals queryGlobals = new TQueryGlobals(); queryGlobals.setNowString(TimeUtils.getDatetimeFormatWithTimeZone().format(LocalDateTime.now())); diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/FoldConstantRuleOnBE.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/FoldConstantRuleOnBE.java index 331d70aac8ed88..76b57eece78a48 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/FoldConstantRuleOnBE.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/FoldConstantRuleOnBE.java @@ -320,6 +320,7 @@ private static Map evalOnBE(Map> tQueryOptions.setBeExecVersion(Config.be_exec_version); tQueryOptions.setEnableDecimal256(context.getSessionVariable().isEnableDecimal256()); tQueryOptions.setNewVersionUnixTimestamp(true); + tQueryOptions.setNewVersionPercentile(true); tQueryOptions.setEnableStrictCast(SessionVariable.enableStrictCast()); TFoldConstantParams tParams = new TFoldConstantParams(paramMap, queryGlobals); diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/CoordinatorContext.java b/fe/fe-core/src/main/java/org/apache/doris/qe/CoordinatorContext.java index 945cf66bbfe6d0..61aee063452984 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/CoordinatorContext.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/CoordinatorContext.java @@ -308,6 +308,7 @@ public static CoordinatorContext buildForLoad( queryOptions.setProfileLevel(2); queryOptions.setBeExecVersion(Config.be_exec_version); queryOptions.setNewVersionUnixTimestamp(true); + queryOptions.setNewVersionPercentile(true); TQueryGlobals queryGlobals = new TQueryGlobals(); queryGlobals.setNowString(TimeUtils.getDatetimeFormatWithTimeZone().format(LocalDateTime.now())); diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/NereidsCoordinator.java b/fe/fe-core/src/main/java/org/apache/doris/qe/NereidsCoordinator.java index 4084e4a84e8a2a..f8c7509f102678 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/NereidsCoordinator.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/NereidsCoordinator.java @@ -492,6 +492,7 @@ private void setForInsert(long jobId) { // Set this field to true to avoid data entering the normal cache LRU queue this.coordinatorContext.queryOptions.setDisableFileCache(true); this.coordinatorContext.queryOptions.setNewVersionUnixTimestamp(true); + this.coordinatorContext.queryOptions.setNewVersionPercentile(true); } private void setForQuery() { diff --git a/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java b/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java index c9020b6c9dc37d..c140ead2369565 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java +++ b/fe/fe-core/src/main/java/org/apache/doris/qe/SessionVariable.java @@ -5637,6 +5637,7 @@ public TQueryOptions toThrift() { tResult.setEnableStrictCast(enableStrictCast()); tResult.setEnableInsertStrict(enableInsertStrict); tResult.setNewVersionUnixTimestamp(true); // once FE upgraded, always use new version + tResult.setNewVersionPercentile(true); tResult.setHnswEfSearch(hnswEFSearch); tResult.setHnswCheckRelativeDistance(hnswCheckRelativeDistance);