From 145cec36a61dc53af22552830a459c66fd6b1b3c Mon Sep 17 00:00:00 2001 From: Zhongjun Jin Date: Mon, 30 Mar 2026 20:33:31 +0000 Subject: [PATCH] Disable registration of gluten-side implemented "round" function, use Bolt-side "round" function instead * fix 'round' test result for round(3.1415f, 3) * Enable bolt for negative scale * Fix unit test * Revert "Add unit tests" * [fix] disable gluten internal 'round' function and rely on Bolt's RoundFunction implementation * Add unit tests See merge request: !1752 --- .../functions/RegistrationAllFunctions.cc | 15 ++++++++------- .../substrait/SubstraitToBoltPlanValidator.cc | 11 ----------- cpp/bolt/tests/SparkFunctionTest.cc | 12 ++++++++++-- .../expressions/GlutenMathExpressionsSuite.scala | 2 +- .../expressions/GlutenMathExpressionsSuite.scala | 2 +- .../expressions/GlutenMathExpressionsSuite.scala | 2 +- .../expressions/GlutenMathExpressionsSuite.scala | 2 +- 7 files changed, 22 insertions(+), 24 deletions(-) diff --git a/cpp/bolt/operators/functions/RegistrationAllFunctions.cc b/cpp/bolt/operators/functions/RegistrationAllFunctions.cc index 1ceec1c8676..21bbfaa52d2 100644 --- a/cpp/bolt/operators/functions/RegistrationAllFunctions.cc +++ b/cpp/bolt/operators/functions/RegistrationAllFunctions.cc @@ -49,13 +49,14 @@ namespace gluten { namespace { void registerFunctionOverwrite() { - bolt::functions::registerUnaryNumeric({"round"}); - bolt::registerFunction({"round"}); - bolt::registerFunction({"round"}); - bolt::registerFunction({"round"}); - bolt::registerFunction({"round"}); - bolt::registerFunction({"round"}); - bolt::registerFunction({"round"}); + // Disable gluten internal round function + // bolt::functions::registerUnaryNumeric({"round"}); + // bolt::registerFunction({"round"}); + // bolt::registerFunction({"round"}); + // bolt::registerFunction({"round"}); + // bolt::registerFunction({"round"}); + // bolt::registerFunction({"round"}); + // bolt::registerFunction({"round"}); auto kRowConstructorWithNull = RowConstructorWithNullCallToSpecialForm::kRowConstructorWithNull; bolt::exec::registerVectorFunction( diff --git a/cpp/bolt/substrait/SubstraitToBoltPlanValidator.cc b/cpp/bolt/substrait/SubstraitToBoltPlanValidator.cc index bd997c99947..140ed99ef75 100644 --- a/cpp/bolt/substrait/SubstraitToBoltPlanValidator.cc +++ b/cpp/bolt/substrait/SubstraitToBoltPlanValidator.cc @@ -130,23 +130,12 @@ bool SubstraitToBoltPlanValidator::validateRound( return false; } - // Bolt has different result with Spark on negative scale. auto typeCase = arguments[1].value().literal().literal_type_case(); switch (typeCase) { case ::substrait::Expression_Literal::LiteralTypeCase::kI32: { - int32_t scale = arguments[1].value().literal().i32(); - if (scale < 0) { - LOG_VALIDATION_MSG("Round scale validation failed: scale " + std::to_string(scale) + " is negative."); - return false; - } return true; } case ::substrait::Expression_Literal::LiteralTypeCase::kI64: { - int64_t scale = arguments[1].value().literal().i64(); - if (scale < 0) { - LOG_VALIDATION_MSG("Round scale validation failed: scale " + std::to_string(scale) + " is negative."); - return false; - } return true; } default: diff --git a/cpp/bolt/tests/SparkFunctionTest.cc b/cpp/bolt/tests/SparkFunctionTest.cc index 62f0dbf4c63..40d0c443c2c 100644 --- a/cpp/bolt/tests/SparkFunctionTest.cc +++ b/cpp/bolt/tests/SparkFunctionTest.cc @@ -67,13 +67,21 @@ class SparkFunctionTest : public SparkFunctionBaseTest { template std::vector> testRoundWithDecFloatAndDoubleData() { - return {{1.122112, 0, 1}, {1.129, 1, 1.1}, {1.129, 2, 1.13}, {1.0 / 3, 0, 0.0}, + std::vector> data = + {{1.122112, 0, 1}, {1.129, 1, 1.1}, {1.129, 2, 1.13}, {1.0 / 3, 0, 0.0}, {1.0 / 3, 1, 0.3}, {1.0 / 3, 2, 0.33}, {1.0 / 3, 6, 0.333333}, {-1.122112, 0, -1}, {-1.129, 1, -1.1}, {-1.129, 2, -1.13}, {-1.129, 2, -1.13}, {-1.0 / 3, 0, 0.0}, {-1.0 / 3, 1, -0.3}, {-1.0 / 3, 2, -0.33}, {-1.0 / 3, 6, -0.333333}, {1.0, -1, 0.0}, {0.0, -2, 0.0}, {-1.0, -3, 0.0}, {11111.0, -1, 11110.0}, {11111.0, -2, 11100.0}, {11111.0, -3, 11000.0}, {11111.0, -4, 10000.0}, {0.575, 2, 0.58}, {0.574, 2, 0.57}, - {-0.575, 2, -0.58}, {-0.574, 2, -0.57}}; + {-0.575, 2, -0.58}, {-0.574, 2, -0.57}, {102.825291, 6, 102.825291}, {10.424817, 6, 10.424817}, + {-82.737209, 6, -82.737209}, {-123.106541, 6, -123.106541}}; + // To be consistent with Spark + if constexpr (std::is_same_v) { + data[22] = {0.575, 2, 0.57}; + data[24] = {-0.575, 2, -0.57}; + } + return data; } template diff --git a/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala b/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala index 1f343e1bbff..998fac6d3b2 100644 --- a/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala +++ b/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala @@ -41,7 +41,7 @@ class GlutenMathExpressionsSuite extends MathExpressionsSuite with GlutenTestsTr Seq(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 3.0, 3.1, 3.14, 3.142, 3.1416, 3.14159, 3.141593) val floatResults: Seq[Float] = - Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.142f, 3.1415f, 3.1415f, 3.1415f) + Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f) val bRoundFloatResults: Seq[Float] = Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f) diff --git a/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala b/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala index e4c59095eea..bfdc406d337 100644 --- a/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala +++ b/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala @@ -69,7 +69,7 @@ class GlutenMathExpressionsSuite extends MathExpressionsSuite with GlutenTestsTr Seq(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 3.0, 3.1, 3.14, 3.142, 3.1416, 3.14159, 3.141593) val floatResults: Seq[Float] = - Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.142f, 3.1415f, 3.1415f, 3.1415f) + Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f) val bRoundFloatResults: Seq[Float] = Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f) diff --git a/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala b/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala index 826176334c1..dbadfda90e8 100644 --- a/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala +++ b/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala @@ -34,7 +34,7 @@ class GlutenMathExpressionsSuite extends MathExpressionsSuite with GlutenTestsTr Seq(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 3.0, 3.1, 3.14, 3.142, 3.1416, 3.14159, 3.141593) val floatResults: Seq[Float] = - Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.142f, 3.1415f, 3.1415f, 3.1415f) + Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f) val bRoundFloatResults: Seq[Float] = Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f) diff --git a/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala b/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala index d49bbd3555e..2a3bc9b1157 100644 --- a/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala +++ b/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/catalyst/expressions/GlutenMathExpressionsSuite.scala @@ -34,7 +34,7 @@ class GlutenMathExpressionsSuite extends MathExpressionsSuite with GlutenTestsTr Seq(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 3.0, 3.1, 3.14, 3.142, 3.1416, 3.14159, 3.141593) val floatResults: Seq[Float] = - Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.142f, 3.1415f, 3.1415f, 3.1415f) + Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f) val bRoundFloatResults: Seq[Float] = Seq(0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 3.0f, 3.1f, 3.14f, 3.141f, 3.1415f, 3.1415f, 3.1415f)