This is an automated email from the ASF dual-hosted git repository.

yiguolei pushed a commit to branch branch-4.1
in repository https://gitbox.apache.org/repos/asf/doris.git


The following commit(s) were added to refs/heads/branch-4.1 by this push:
     new 5dbd0f6fdcc branch-4.1: [improvement](be) Optimize pow squares with 
exact integer results #68258 (#68443)
5dbd0f6fdcc is described below

commit 5dbd0f6fdcc7ae991c7936cb11558d0f031cffd8
Author: github-actions[bot] 
<41898282+github-actions[bot]@users.noreply.github.com>
AuthorDate: Thu Sep 24 17:30:25 2026 +0800

    branch-4.1: [improvement](be) Optimize pow squares with exact integer 
results #68258 (#68443)
    
    Cherry-picked from #68258
    
    Co-authored-by: Jerry Hu <[email protected]>
---
 be/src/exprs/function/math.cpp                     |  23 +++
 be/test/exprs/function/function_math_test.cpp      | 193 ++++++++++++++++++++-
 .../data/function_p0/test_pow_square.out           |  79 +++++++++
 .../suites/function_p0/test_pow_square.groovy      | 114 ++++++++++++
 4 files changed, 402 insertions(+), 7 deletions(-)

diff --git a/be/src/exprs/function/math.cpp b/be/src/exprs/function/math.cpp
index d25cca3a901..8bd849549ca 100644
--- a/be/src/exprs/function/math.cpp
+++ b/be/src/exprs/function/math.cpp
@@ -15,6 +15,8 @@
 // specific language governing permissions and limitations
 // under the License.
 
+#include <algorithm>
+#include <cfenv>
 #include <cstddef>
 #include <cstdint>
 #include <cstring>
@@ -643,6 +645,27 @@ private:
             auto& a = column_left_ptr->get_data();
             auto& c = column_result->get_data();
             size_t size = a.size();
+            if constexpr (std::is_same_v<Impl, PowImpl>) {
+                if (column_right_ptr->template get_value<Impl::type>() == 2.0 
&&
+                    std::fegetround() == FE_TONEAREST) {
+                    // Integer bases up to 2^26 have exact binary64 squares 
(at most 2^52).
+                    // Other bases can differ by one ULP between 
multiplication and libm.
+                    // Stop checking at the first unsafe value rather than 
checking the whole block.
+                    const auto exact_end = std::ranges::find_if_not(a, 
[](double value) {
+                        return std::abs(value) <= 0x1p26 && value == 
std::trunc(value);
+                    });
+                    const auto exact_rows = static_cast<size_t>(exact_end - 
a.begin());
+                    for (size_t i = 0; i < exact_rows; ++i) {
+                        c[i] = a[i] * a[i];
+                    }
+                    // Keep the remaining rows on libm, not compiler-folded 
pow(x, 2).
+                    volatile double exponent = 2.0;
+                    for (size_t i = exact_rows; i < size; ++i) {
+                        c[i] = Impl::apply(a[i], exponent);
+                    }
+                    return column_result;
+                }
+            }
             for (size_t i = 0; i < size; ++i) {
                 c[i] = Impl::apply(a[i], column_right_ptr->template 
get_value<Impl::type>());
             }
diff --git a/be/test/exprs/function/function_math_test.cpp 
b/be/test/exprs/function/function_math_test.cpp
index 318ca61dc72..2e2196f226e 100644
--- a/be/test/exprs/function/function_math_test.cpp
+++ b/be/test/exprs/function/function_math_test.cpp
@@ -15,20 +15,37 @@
 // specific language governing permissions and limitations
 // under the License.
 
+#include <array>
+#include <bit>
+#include <cfenv>
 #include <climits>
+#include <cmath>
 #include <cstdint>
 #include <limits>
+#include <memory>
+#include <numbers>
 #include <random>
+#include <span>
 #include <string>
+#include <utility>
+#include <vector>
 
+#include "core/block/block.h"
 #include "core/column/column_const.h"
+#include "core/column/column_nullable.h"
+#include "core/column/column_vector.h"
 #include "core/data_type/data_type_decimal.h"
+#include "core/data_type/data_type_nullable.h"
 #include "core/data_type/data_type_number.h"
 #include "core/data_type/data_type_string.h"
+#include "core/field.h"
 #include "core/types.h"
 #include "exprs/function/function_test_util.h"
+#include "exprs/function/simple_function_factory.h"
+#include "exprs/function_context.h"
 #include "testutil/any_type.h"
 #include "testutil/column_helper.h"
+#include "util/defer_op.h"
 
 namespace doris {
 
@@ -297,16 +314,178 @@ TEST(MathFunctionTest, log_test) {
 }
 
 TEST(MathFunctionTest, pow_test) {
-    std::string func_name = "pow"; // pow(x,y)
+    const InputTypeSet input_types = {TYPE_DOUBLE, TYPE_DOUBLE};
+    const double inf = std::numeric_limits<double>::infinity();
+    const double nan = std::numeric_limits<double>::quiet_NaN();
+    const DataSet data_set = {{{10.0, 1.0}, 10.0},     {{10.0, 10.0}, 
10000000000.0},
+                              {{100.0, -2.0}, 0.0001}, {{2.0, 0.5}, 
std::numbers::sqrt2},
+                              {{-2.0, 3.0}, -8.0},     {{-2.0, 0.5}, nan},
+                              {{nan, 0.0}, 1.0},       {{1.0, nan}, 1.0},
+                              {{0.0, -2.0}, inf},      {{-0.0, -3.0}, -inf},
+                              {{Null(), 2.0}, Null()}, {{2.0, Null()}, 
Null()}};
+    check_function_all_arg_comb<DataTypeFloat64, true>("pow", input_types, 
data_set);
+}
 
-    InputTypeSet input_types = {PrimitiveType::TYPE_DOUBLE, 
PrimitiveType::TYPE_DOUBLE};
+static void check_pow_square_result(const IColumn& result, std::span<const 
double> values,
+                                    bool nullable, bool const_base) {
+    ASSERT_EQ(result.size(), values.size());
+    const auto& nested =
+            nullable ? assert_cast<const 
ColumnNullable&>(result).get_nested_column() : result;
+    const auto& data = assert_cast<const ColumnFloat64&>(nested).get_data();
+    // Keep the reference call on libm rather than letting the compiler turn 
pow(x, 2) into x * x.
+    volatile double exponent = 2.0;
+    for (size_t i = 0; i < values.size(); ++i) {
+        const bool expect_null = nullable && !const_base && i == values.size() 
- 1;
+        if (nullable) {
+            EXPECT_EQ(result.is_null_at(i), expect_null);
+        }
+        if (expect_null) {
+            continue;
+        }
+        const double base = values[const_base ? 0 : i];
+        const double expected = std::pow(base, exponent);
+        if (std::isnan(expected)) {
+            EXPECT_TRUE(std::isnan(data[i]));
+        } else {
+            EXPECT_EQ(std::bit_cast<uint64_t>(data[i]), 
std::bit_cast<uint64_t>(expected))
+                    << "row=" << i << " base=" << base;
+        }
+    }
+}
+
+static void check_pow_square_column_shapes(const std::string& name, bool 
nullable, int const_mask,
+                                           std::span<const double> values) {
+    SCOPED_TRACE(testing::Message()
+                 << name << " nullable=" << nullable << " const_mask=" << 
const_mask);
+    const size_t rows = values.size();
+    DataTypePtr type = std::make_shared<DataTypeFloat64>();
+    if (nullable) {
+        type = make_nullable(type);
+    }
+    auto bases = type->create_column();
+    for (double value : values) {
+        bases->insert(Field::create_field<TYPE_DOUBLE>(value));
+    }
+    if (nullable) {
+        bases->pop_back(1);
+        bases->insert_default();
+    }
+    auto exponents = type->create_column();
+    exponents->insert(Field::create_field<TYPE_DOUBLE>(2.0));
+    ColumnPtr left = std::move(bases);
+    if (const_mask & 1) {
+        left = ColumnConst::create(left->clone_resized(1), rows);
+    }
+    ColumnPtr right = ColumnConst::create(exponents->get_ptr(), rows);
+    if (!(const_mask & 2)) {
+        right = right->convert_to_full_column_if_const();
+    }
+    Block block({{left, type, "base"}, {right, type, "exponent"}});
+    auto function = SimpleFunctionFactory::instance().get_function(
+            name, block.get_columns_with_type_and_name(), type);
+    ASSERT_NE(function, nullptr);
+    block.insert({nullptr, type, "result"});
+    FunctionUtils fn_utils(type, {type, type}, false);
+    auto* context = fn_utils.get_fn_ctx();
+    ASSERT_TRUE(function->open(context, FunctionContext::FRAGMENT_LOCAL).ok());
+    ASSERT_TRUE(function->open(context, FunctionContext::THREAD_LOCAL).ok());
+    const auto status = function->execute(context, block, {0, 1}, 2, rows);
+    EXPECT_TRUE(function->close(context, FunctionContext::THREAD_LOCAL).ok());
+    EXPECT_TRUE(function->close(context, 
FunctionContext::FRAGMENT_LOCAL).ok());
+    ASSERT_TRUE(status.ok()) << status.to_string();
+    auto result = 
block.get_by_position(2).column->convert_to_full_column_if_const();
+    check_pow_square_result(*result, values, nullable, const_mask & 1);
+}
 
-    DataSet data_set = {{{10.0, 1.0}, 10.0},
-                        {{10.0, 10.0}, 10000000000.0},
-                        {{100.0, -2.0}, 0.0001},
-                        {{2.0, 0.5}, 1.4142135623730951}};
+static void check_pow_square_all_shapes(std::span<const double> values) {
+    for (const auto* name : {"pow", "power", "dpow", "fpow"}) {
+        for (int const_mask = 0; const_mask < 4; ++const_mask) {
+            check_pow_square_column_shapes(name, false, const_mask, values);
+            check_pow_square_column_shapes(name, true, const_mask, values);
+        }
+    }
+}
 
-    static_cast<void>(check_function<DataTypeFloat64, true>(func_name, 
input_types, data_set));
+TEST(MathFunctionTest, pow_square_column_shapes) {
+    const double inf = std::numeric_limits<double>::infinity();
+    // The first value differs by one ULP between libm pow(x, 2) and x * x. 
Keep it first
+    // so that the constant-base cases also exercise it; approximate equality 
would miss this.
+    const std::array values = {1.1500729535343723e-17,
+                               -1.5,
+                               0.0,
+                               -0.0,
+                               1.0,
+                               -2.0,
+                               0.5,
+                               12345.125,
+                               1e154,
+                               1e-154,
+                               std::numeric_limits<double>::max(),
+                               std::numeric_limits<double>::min(),
+                               std::numeric_limits<double>::denorm_min(),
+                               inf,
+                               -inf,
+                               std::numeric_limits<double>::quiet_NaN(),
+                               3.0};
+    check_pow_square_all_shapes(values);
+}
+
+TEST(MathFunctionTest, pow_square_exact_integers) {
+    std::vector<double> values = {3.0, -3.0, 0x1p26, -0x1p26};
+    for (int value = -4096; value <= 4096; ++value) {
+        values.push_back(value);
+    }
+    std::mt19937_64 random(0);
+    for (size_t i = 0; i < 4096; ++i) {
+        values.push_back(static_cast<double>(random() % ((1ULL << 27) + 1)) - 
0x1p26);
+    }
+    check_pow_square_all_shapes(values);
+}
+
+TEST(MathFunctionTest, pow_square_integer_boundaries) {
+    // A safe prefix followed by out-of-range integers and fractional 
neighbours.
+    // 94906297 has a one-ULP square difference between libm and 
multiplication on some platforms.
+    std::vector<double> values = {3.0,     -3.0,       0x1p26,
+                                  -0x1p26, 94906297.0, -94906297.0,
+                                  0x1p27,  -0x1p27,    1.1500729535343723e-17};
+    for (int offset = -32; offset <= 32; ++offset) {
+        const double value = 0x1p26 + offset;
+        values.insert(values.end(),
+                      {value, -value, std::nextafter(value, 0.0),
+                       std::nextafter(value, 
std::numeric_limits<double>::infinity())});
+    }
+    check_pow_square_all_shapes(values);
+}
+
+TEST(MathFunctionTest, pow_square_empty_block) {
+    for (const auto* name : {"pow", "power", "dpow", "fpow"}) {
+        for (int const_mask = 0; const_mask < 4; ++const_mask) {
+            check_pow_square_column_shapes(name, false, const_mask, {});
+        }
+    }
+}
+
+TEST(MathFunctionTest, pow_square_random_bits) {
+    std::mt19937_64 random(1);
+    std::vector<double> values;
+    values.reserve(4096);
+    for (size_t i = 0; i < 4096; ++i) {
+        values.push_back(std::bit_cast<double>(random()));
+    }
+    check_pow_square_all_shapes(values);
+}
+
+TEST(MathFunctionTest, pow_square_rounding_modes) {
+    std::fenv_t environment;
+    ASSERT_EQ(std::fegetenv(&environment), 0);
+    Defer restore_environment([&] { EXPECT_EQ(std::fesetenv(&environment), 0); 
});
+    // All values are in the fast domain, so only the rounding-mode guard can 
disable it.
+    const std::array values = {3.0, -3.0, 0.0, -0.0, 0x1p26, -0x1p26, 0x1p26 - 
1};
+    for (int mode : {FE_TONEAREST, FE_UPWARD, FE_DOWNWARD, FE_TOWARDZERO}) {
+        SCOPED_TRACE(testing::Message() << "rounding_mode=" << mode);
+        ASSERT_EQ(std::fesetround(mode), 0);
+        check_pow_square_all_shapes(values);
+    }
 }
 
 TEST(MathFunctionTest, ceil_test) {
diff --git a/regression-test/data/function_p0/test_pow_square.out 
b/regression-test/data/function_p0/test_pow_square.out
new file mode 100644
index 00000000000..27da9ba2ff1
--- /dev/null
+++ b/regression-test/data/function_p0/test_pow_square.out
@@ -0,0 +1,79 @@
+-- This file is automatically generated. You should know what you did if you 
want to edit this
+-- !square_aliases --
+0      4       4       4       4       4
+1      2.25    2.25    2.25    2.25    2.25
+2      1       1       1       1       1
+3      0.25    0.25    0.25    0.25    0.25
+4      0       0       0       0       0
+5      0.25    0.25    0.25    0.25    0.25
+6      1       1       1       1       1
+7      2.25    2.25    2.25    2.25    2.25
+8      4       4       4       4       4
+
+-- !nullable_square --
+0      \N      \N      \N
+1      9       \N      9
+2      4       \N      4
+3      1       \N      1
+4      0       \N      0
+5      1       \N      1
+6      4       \N      4
+7      9       \N      9
+8      16      \N      16
+
+-- !column_shapes --
+0      1       1       4
+1      2       1       4
+2      4       4       4
+3      8       27      4
+4      16      256     4
+
+-- !square_shape_equality --
+0      true    true    true    true
+1      true    true    true    true
+2      false   false   false   false
+3      false   false   false   false
+
+-- !exact_integer_square --
+0      0       0       0       0       0
+1      100000000000000 100000000000000 100000000000000 100000000000000 
100000000000000
+2      400000000000000 400000000000000 400000000000000 400000000000000 
400000000000000
+3      900000000000000 900000000000000 900000000000000 900000000000000 
900000000000000
+4      1600000000000000        1600000000000000        1600000000000000        
1600000000000000        1600000000000000
+5      2500000000000000        2500000000000000        2500000000000000        
2500000000000000        2500000000000000
+6      3600000000000000        3600000000000000        3600000000000000        
3600000000000000        3600000000000000
+
+-- !integer_square_boundaries --
+0      4503599090499600        4503599090499600        4503599090499600        
4503599090499600        true
+1      4503599224717321        4503599224717321        4503599224717321        
4503599224717321        true
+2      4503599358935044        4503599358935044        4503599358935044        
4503599358935044        true
+3      4503599493152769        4503599493152769        4503599493152769        
4503599493152769        true
+4      4503599627370496        4503599627370496        4503599627370496        
4503599627370496        true
+5      4503599761588225        4503599761588225        4503599761588225        
4503599761588225        true
+6      4503599895805956        4503599895805956        4503599895805956        
4503599895805956        true
+7      4503600030023689        4503600030023689        4503600030023689        
4503600030023689        true
+8      4503600164241424        4503600164241424        4503600164241424        
4503600164241424        false
+9      4503600298459161        4503600298459161        4503600298459161        
4503600298459161        false
+
+-- !out_of_range_square_shapes --
+0      true    true    true    true
+1      true    true    true    true
+2      false   false   false   false
+3      false   false   false   false
+
+-- !other_exponents --
+0      1       -2      0.25    -8      NaN
+1      1       -1      1       -1      NaN
+2      1       0       Infinity        0       0
+3      1       1       1       1       1
+4      1       2       0.25    8       1.4142135623730951
+
+-- !square_boundaries --
+0      NaN     NaN
+1      Infinity        Infinity
+2      Infinity        Infinity
+3      0       0
+4      Infinity        Infinity
+5      0       0
+6      1e+308  1e+308
+7      1e-308  1e-308
diff --git a/regression-test/suites/function_p0/test_pow_square.groovy 
b/regression-test/suites/function_p0/test_pow_square.groovy
new file mode 100644
index 00000000000..1ede6381b22
--- /dev/null
+++ b/regression-test/suites/function_p0/test_pow_square.groovy
@@ -0,0 +1,114 @@
+// 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.
+
+suite("test_pow_square") {
+    qt_square_aliases """
+        select number, pow(x, 2.0), power(x, 2.0), dpow(x, 2.0), fpow(x, 2.0), 
x * x
+        from (select number, cast(number - 4 as double) / 2 as x
+              from numbers("number" = "9")) t
+        order by number
+    """
+
+    qt_nullable_square """
+        select number, pow(x, 2.0), pow(x, cast(null as double)), x * x
+        from (select number, if(number = 0, null, cast(number - 4 as double)) 
as x
+              from numbers("number" = "9")) t
+        order by number
+    """
+
+    qt_column_shapes """
+        select number, pow(2.0, x), pow(x, x), pow(-2.0, 2.0)
+        from (select number, cast(number as double) as x
+              from numbers("number" = "5")) t
+        order by number
+    """
+
+    // The first two rows have an exponent of 2, but y remains a vector 
because later rows
+    // have an exponent of 3. These bases expose a one-ULP difference between 
pow(x, 2) and x * x.
+    // Derive x from number to keep alias evaluation in BE rather than FE 
constant folding.
+    // Do not guard the equality with y = 2: predicate inference can simplify 
it away.
+    qt_square_shape_equality """
+        select number,
+               pow(x, 2.0) = pow(x, y),
+               power(x, 2.0) = power(x, y),
+               dpow(x, 2.0) = dpow(x, y),
+               fpow(x, 2.0) = fpow(x, y)
+        from (
+            select number,
+                   cast(2 * number - 1 as double)
+                       * cast('1.1500729535343723e-17' as double) as x,
+                   if(number < 2, 2.0, 3.0) as y
+            from numbers("number" = "4")
+        ) t order by number
+    """
+
+    qt_exact_integer_square """
+        select number, pow(x, 2.0), power(x, 2.0), dpow(x, 2.0), fpow(x, 2.0), 
x * x
+        from (
+            select number,
+                   cast(number * 10000000 as double) * if(number % 2 = 0, 1, 
-1) as x
+            from numbers("number" = "7")
+        ) t order by number
+    """
+
+    qt_integer_square_boundaries """
+        select number, pow(x, 2.0), power(x, 2.0), dpow(x, 2.0), fpow(x, 2.0),
+               pow(x, 2.0) = pow(x, y)
+        from (
+            select number,
+                   cast(number + 67108860 as double) * if(number % 2 = 0, 1, 
-1) as x,
+                   if(number < 8, 2.0, 3.0) as y
+            from numbers("number" = "10")
+        ) t order by number
+    """
+
+    // Integer bases alone are not sufficient: these squares are not exactly 
representable.
+    qt_out_of_range_square_shapes """
+        select number,
+               pow(x, 2.0) = pow(x, y), power(x, 2.0) = power(x, y),
+               dpow(x, 2.0) = dpow(x, y), fpow(x, 2.0) = fpow(x, y)
+        from (
+            select number, cast(2 * number - 1 as double) * 94906297.0 as x,
+                   if(number < 2, 2.0, 3.0) as y
+            from numbers("number" = "4")
+        ) t order by number
+    """
+
+    qt_other_exponents """
+        select number, pow(x, 0.0), pow(x, 1.0), pow(x, -2.0), pow(x, 3.0), 
pow(x, 0.5)
+        from (select number, cast(number - 2 as double) as x
+              from numbers("number" = "5")) t
+        order by number
+    """
+
+    qt_square_boundaries """
+        select number, pow(x, 2.0), x * x
+        from (
+            select number, case number
+                when 0 then cast('nan' as double)
+                when 1 then cast('inf' as double)
+                when 2 then cast('-inf' as double)
+                when 3 then cast('-0.0' as double)
+                when 4 then cast('1e308' as double)
+                when 5 then cast('1e-308' as double)
+                when 6 then cast('1e154' as double)
+                when 7 then cast('1e-154' as double)
+                end as x
+            from numbers("number" = "8")
+        ) t order by number
+    """
+}


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to