Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
33 commits
Select commit Hold shift + click to select a range
37fd9fc
fix: round for float/double
kazuyukitanimura Apr 25, 2026
9707af3
fix: round for float/double
kazuyukitanimura Apr 25, 2026
62b910d
fix: round for float/double
kazuyukitanimura Apr 25, 2026
1cbf7b4
fix: round for float/double
kazuyukitanimura Apr 25, 2026
dbc7019
fix: round for float/double
kazuyukitanimura Apr 25, 2026
ce10f8b
fix: round for float/double
kazuyukitanimura Apr 25, 2026
643a38e
fix: round for float/double
kazuyukitanimura Apr 25, 2026
e31e825
fix: round for float/double
kazuyukitanimura Apr 25, 2026
57e0287
fix: round for float/double
kazuyukitanimura Apr 25, 2026
3a7eea6
fix: round for float/double
kazuyukitanimura Apr 25, 2026
17eac11
fix: round for float/double
kazuyukitanimura Apr 25, 2026
300e403
fix: round for float/double
kazuyukitanimura Apr 25, 2026
3e39898
fix: round for float/double
kazuyukitanimura Apr 25, 2026
e045cc6
fix: round for float/double
kazuyukitanimura Apr 25, 2026
a3ab1f0
fix: round for float/double
kazuyukitanimura Apr 25, 2026
9adb5fd
fix: round for float/double
kazuyukitanimura Apr 25, 2026
60080a9
fix: round for float/double
kazuyukitanimura Apr 25, 2026
49a0209
fix: round for float/double
kazuyukitanimura Apr 25, 2026
fa786aa
fix: round for float/double
kazuyukitanimura Apr 25, 2026
7731338
fix: round for float/double
kazuyukitanimura Apr 26, 2026
f88a04c
fix: round for float/double
kazuyukitanimura Apr 27, 2026
8204171
Merge remote-tracking branch 'upstream/main' into fix-flaot-round
kazuyukitanimura May 12, 2026
2dd4d5a
Merge remote-tracking branch 'upstream/main' into fix-flaot-round
kazuyukitanimura May 16, 2026
203c319
fix: round for float/double
kazuyukitanimura May 16, 2026
3470e5f
fix: round for float/double
kazuyukitanimura May 16, 2026
f9fdfae
fix: round for float/double
kazuyukitanimura May 16, 2026
3c2af76
Merge remote-tracking branch 'upstream/main' into fix-flaot-round
kazuyukitanimura May 22, 2026
b69469b
address review comments
kazuyukitanimura May 22, 2026
28b0544
Merge remote-tracking branch 'upstream/main' into fix-flaot-round
kazuyukitanimura Jun 26, 2026
11412bf
address review comments
kazuyukitanimura Jun 27, 2026
0dff7cd
address review comments
kazuyukitanimura Jun 27, 2026
2c1145d
address review comments
kazuyukitanimura Jun 27, 2026
b296143
address review comments
kazuyukitanimura Jun 29, 2026
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
24 changes: 8 additions & 16 deletions spark/src/main/scala/org/apache/comet/serde/arithmetic.scala
Original file line number Diff line number Diff line change
Expand Up @@ -285,22 +285,6 @@ object CometRound extends CometExpressionSerde[Round] {
override def getSupportLevel(expr: Round): SupportLevel = expr.child.dataType match {
case t: DecimalType if t.scale < 0 => // Spark disallows negative scale SPARK-30252
Unsupported(Some("Decimal type has negative scale"))
case _: FloatType | DoubleType =>
// We cannot properly match with the Spark behavior for floating-point numbers.
// Spark uses BigDecimal for rounding float/double, and BigDecimal fist converts a
// double to string internally in order to create its own internal representation.
// The problem is BigDecimal uses java.lang.Double.toString() and it has complicated
// rounding algorithm. E.g. -5.81855622136895E8 is actually
// -581855622.13689494132995605468750. Note the 5th fractional digit is 4 instead of
// 5. Java(Scala)'s toString() rounds it up to -581855622.136895. This makes a
// difference when rounding at 5th digit, I.e. round(-5.81855622136895E8, 5) should be
// -5.818556221369E8, instead of -5.8185562213689E8. There is also an example that
// toString() does NOT round up. 6.1317116247283497E18 is 6131711624728349696. It can
// be rounded up to 6.13171162472835E18 that still represents the same double number.
// I.e. 6.13171162472835E18 == 6.1317116247283497E18. However, toString() does not.
// That results in round(6.1317116247283497E18, -5) == 6.1317116247282995E18 instead
// of 6.1317116247283999E18.
Unsupported(Some("Comet does not support Spark's BigDecimal rounding"))
case _ =>
Compatible()
}
Expand All @@ -319,6 +303,14 @@ object CometRound extends CometExpressionSerde[Round] {
exprToProtoInternal(Literal(null), inputs, binding)
case _: ByteType | ShortType | IntegerType | LongType if _scale >= 0 =>
childExpr // _scale(I.e. decimal place) >= 0 is a no-op for integer types in Spark
case _: FloatType | _: DoubleType =>
// Spark rounds floats/doubles by widening to double, building a BigDecimal via
// java.lang.Double.toString, applying HALF_UP, and (for floats) narrowing back. The
// toString algorithm differs between JDKs (notably 17 vs 21), so a native implementation
// can't match every JDK. Route the expression through the JVM codegen dispatcher, which
// Janino-compiles Spark's own RoundBase.doGenCode and runs it inside the Comet pipeline
// on the executor's JDK, so the result matches Spark exactly.
CometScalaUDF.emitJvmCodegenDispatch(r, inputs, binding)
case _ =>
// `scale` must be Int64 type in DataFusion
val scaleExpr = exprToProtoInternal(Literal(_scale.toLong, LongType), inputs, binding)
Expand Down
78 changes: 73 additions & 5 deletions spark/src/test/resources/sql-tests/expressions/math/round.sql
Original file line number Diff line number Diff line change
Expand Up @@ -21,18 +21,86 @@ CREATE TABLE test_round(d double, i int) USING parquet
statement
INSERT INTO test_round VALUES (2.5, 0), (3.5, 0), (-2.5, 0), (123.456, 2), (123.456, -1), (NULL, 0), (cast('NaN' as double), 0), (cast('Infinity' as double), 0), (0.0, 0)

query expect_fallback(BigDecimal rounding)
query
SELECT round(d, 0) FROM test_round WHERE i = 0

query expect_fallback(BigDecimal rounding)
query
SELECT round(d, 2) FROM test_round WHERE i = 2

query expect_fallback(BigDecimal rounding)
query
SELECT round(d, -1) FROM test_round WHERE i = -1

query expect_fallback(BigDecimal rounding)
query
SELECT round(d) FROM test_round

-- literal + literal
query expect_fallback(BigDecimal rounding)
query
SELECT round(123.456, 2), round(2.5, 0), round(3.5, 0), round(-2.5, 0), round(NULL, 0)

-- HALF_UP semantics: .5 always rounds away from zero
statement
CREATE TABLE test_round_half_up(d double) USING parquet

statement
INSERT INTO test_round_half_up VALUES (0.5), (1.5), (2.5), (-0.5), (-1.5), (-2.5)

query
SELECT d, round(d, 0) FROM test_round_half_up

-- various scales on a single value
query
SELECT round(123.456, 0), round(123.456, 1), round(123.456, 2), round(123.456, 3), round(123.456, 5)

query
SELECT round(123.456, -1), round(123.456, -2), round(123.456, -3)

-- special values
query
SELECT round(cast('NaN' as double), 2), round(cast('Infinity' as double), 2), round(cast('-Infinity' as double), 2)

query
SELECT round(0.0, 5), round(-0.0, 5)

-- very small values
query
SELECT round(1.0E-10, 15), round(1.0E-10, 10), round(1.0E-10, 5)

-- negative scale on doubles
query
SELECT round(9999.9, -1), round(9999.9, -2), round(9999.9, -3), round(9999.9, -4)

query
SELECT round(-9999.9, -1), round(-9999.9, -2), round(-9999.9, -3), round(-9999.9, -4)

-- float type
statement
CREATE TABLE test_round_float(f float) USING parquet

statement
INSERT INTO test_round_float VALUES (cast(2.5 as float)), (cast(3.5 as float)), (cast(-2.5 as float)), (cast(0.125 as float)), (cast(0.785 as float)), (cast(123.456 as float)), (cast('NaN' as float)), (cast('Infinity' as float)), (NULL)

query
SELECT round(f, 0) FROM test_round_float

query
SELECT round(f, 2) FROM test_round_float

query
SELECT round(f, -1) FROM test_round_float

-- BigDecimal rounding edge case from Spark
statement
CREATE TABLE test_round_edge(d double) USING parquet

statement
INSERT INTO test_round_edge VALUES (-5.81855622136895E8), (6.1317116247283497E18), (6.13171162472835E18)

query
SELECT round(d, 4), round(d, 5), round(d, 6) FROM test_round_edge

query
SELECT round('-8316362075006449156', -5)

-- round with column from table (not literals)
query
SELECT d, round(d, 0), round(d, 2), round(d, -1) FROM test_round
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
TakeOrderedAndProject
+- Project [COMET: Comet does not support Spark's BigDecimal rounding]
+- CometNativeColumnarToRow
CometNativeColumnarToRow
+- CometTakeOrderedAndProject
+- CometProject
+- CometSortMergeJoin
:- CometProject
: +- CometSortMergeJoin
Expand Down Expand Up @@ -76,4 +76,4 @@ TakeOrderedAndProject
+- CometFilter
+- CometNativeScan parquet spark_catalog.default.date_dim

Comet accelerated 71 out of 76 eligible operators (93%). Final plan contains 1 transitions between Spark and Comet.
Comet accelerated 73 out of 76 eligible operators (96%). Final plan contains 1 transitions between Spark and Comet.
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
TakeOrderedAndProject
+- Project [COMET: Comet does not support Spark's BigDecimal rounding]
+- CometNativeColumnarToRow
CometNativeColumnarToRow
+- CometTakeOrderedAndProject
+- CometProject
+- CometSortMergeJoin
:- CometProject
: +- CometSortMergeJoin
Expand Down Expand Up @@ -76,4 +76,4 @@ TakeOrderedAndProject
+- CometFilter
+- CometNativeScan parquet spark_catalog.default.date_dim

Comet accelerated 71 out of 76 eligible operators (93%). Final plan contains 1 transitions between Spark and Comet.
Comet accelerated 73 out of 76 eligible operators (96%). Final plan contains 1 transitions between Spark and Comet.
66 changes: 64 additions & 2 deletions spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -2962,10 +2962,24 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper {
Byte.MinValue,
Byte.MaxValue,
Short.MinValue,
Short.MaxValue)).foreach { value =>
Short.MaxValue,
Float.MinPositiveValue,
Float.MaxValue,
Float.NaN,
Float.MinValue,
Float.NegativeInfinity,
Float.PositiveInfinity,
Double.MinPositiveValue,
Double.MaxValue,
Double.NaN,
Double.MinValue,
Double.NegativeInfinity,
Double.PositiveInfinity,
-5.81855622136895e8,
6.1317116247283497e18)).foreach { value =>
val data = Seq(value)
withParquetTable(data, "tbl") {
Seq(-1000, -100, -10, -1, 0, 1, 10, 100, 1000).foreach { scale =>
Seq(-1000, -100, -10, -5, -1, 0, 1, 5, 10, 100, 1000).foreach { scale =>
Seq(true, false).foreach { ansi =>
withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi.toString) {
val res = spark.sql(s"SELECT round(_1, $scale) from tbl")
Expand All @@ -2987,6 +3001,54 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper {
}
}

test("round") {
Seq(true, false).foreach { dictionaryEnabled =>
withTempDir { dir =>
val path = new Path(dir.toURI.toString, "test.parquet")
makeParquetFileAllPrimitiveTypes(
path,
dictionaryEnabled = dictionaryEnabled,
-128,
128,
randomSize = 100)
withParquetTable(path.toString, "tbl") {
for (s <- Seq(-5, -1, 0, 1, 5, -1000, 1000, -323, -308, 308, -15, 15, -16, 16, null)) {
// array tests
// TODO: enable test for unsigned ints (_9, _10, _11, _12)
for (c <- Seq(2, 3, 4, 5, 6, 7, 8, 13, 15, 16, 17)) {
checkSparkAnswerAndOperator(s"select _${c}, round(_${c}, ${s}) FROM tbl")
}
// scalar tests
// Exclude the constant folding optimizer in order to actually execute the native round
// operations for scalar (literal) values.
withSQLConf(
"spark.sql.optimizer.excludedRules" -> "org.apache.spark.sql.catalyst.optimizer.ConstantFolding") {
for (n <- Seq("0.0", "-0.0", "0.5", "-0.5", "1.2", "-1.2")) {
checkSparkAnswerAndOperator(s"select round(cast(${n} as tinyint), ${s}) FROM tbl")
checkSparkAnswerAndOperator(s"select round(cast(${n} as float), ${s}) FROM tbl")
checkSparkAnswerAndOperator(
s"select round(cast(${n} as decimal(38, 18)), ${s}) FROM tbl")
checkSparkAnswerAndOperator(
s"select round(cast(${n} as decimal(20, 0)), ${s}) FROM tbl")
}
checkSparkAnswerAndOperator(s"select round(double('infinity'), ${s}) FROM tbl")
checkSparkAnswerAndOperator(s"select round(double('-infinity'), ${s}) FROM tbl")
checkSparkAnswerAndOperator(s"select round(double('NaN'), ${s}) FROM tbl")
checkSparkAnswerAndOperator(
s"select round(double('0.000000000000000000000000000000000001'), ${s}) FROM tbl")
}
}
}
}
}
}

test("round double from large integer string") {
withParquetTable(Seq(Tuple1("-8316362075006449156")), "tbl") {
checkSparkAnswerAndOperator("SELECT round(cast(_1 as double), -5) FROM tbl")
}
}

test("test integral divide overflow for decimal") {
// All inserted values produce a quotient > Decimal(38,0).max (~1e38), so they overflow
// the intermediate decimal result type. In legacy/try mode both Spark and Comet return
Expand Down
4 changes: 2 additions & 2 deletions spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala
Original file line number Diff line number Diff line change
Expand Up @@ -795,7 +795,7 @@ abstract class CometTestBase
record.add(4, i.toLong)
record.add(5, i.toFloat)
record.add(6, i.toDouble)
record.add(7, i.toString * 48)
record.add(7, i.toString + (i.abs.toString * 47))
record.add(8, (-i).toByte)
record.add(9, (-i).toShort)
record.add(10, -i)
Expand Down Expand Up @@ -824,7 +824,7 @@ abstract class CometTestBase
record.add(4, i)
record.add(5, java.lang.Float.intBitsToFloat(i.toInt))
record.add(6, java.lang.Double.longBitsToDouble(i))
record.add(7, i.toString * 24)
record.add(7, i.toString + (i.abs.toString * 23))
record.add(8, (-i).toByte)
record.add(9, (-i).toShort)
record.add(10, (-i).toInt)
Expand Down
Loading