cloud-fan commented on code in PR #57815:
URL: https://github.com/apache/spark/pull/57815#discussion_r3758102672
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/window/WindowEvaluatorFactoryBase.scala:
##########
@@ -219,13 +236,70 @@ trait WindowEvaluatorFactoryBase {
MutableProjection.create(expressions, schema),
aggFilters)
}
+ lazy val distinctAggregateExpressions =
+ functions.map(_.asInstanceOf[AggregateExpression])
+ lazy val originalDistinctExpressions =
+ distinctChildren(distinctAggregateExpressions.head)
+ lazy val normalizedDistinctExpressions =
+ originalDistinctExpressions.map(NormalizeFloatingNumbers.normalize)
+ lazy val distinctInputAttributes =
+ normalizedDistinctExpressions.zipWithIndex.map { case (expression,
index) =>
+ AttributeReference(
+ s"windowDistinctValue$index",
+ expression.dataType,
+ expression.nullable)()
+ }
+ lazy val rewrittenDistinctFunctions: Array[Expression] =
+ distinctAggregateExpressions.map { ae =>
+ val distinctColumnAttributeLookup = Utils.toMap(
+ distinctChildren(ae).map(_.canonicalized),
+ distinctInputAttributes)
+ ae.aggregateFunction.transformDown {
+ case expression: Expression
+ if
distinctColumnAttributeLookup.contains(expression.canonicalized) =>
+ distinctColumnAttributeLookup(expression.canonicalized)
Review Comment:
**Non-blocking:**
This canonicalizes every matching expression twice: once for `contains` and
again for this lookup. Since canonicalization traverses the subtree, compute it
once and use a single `get` lookup in the transform callback.
##########
sql/core/src/test/scala/org/apache/spark/sql/execution/SQLWindowFunctionSuite.scala:
##########
@@ -196,11 +197,380 @@ class SQLWindowFunctionSuite extends SharedSparkSession {
val e = intercept[AnalysisException] {
sql(
"""
- |select month, area, product, sum(distinct product + 1) over
(partition by 1 order by 2)
+ |select month, area, product, sum(distinct product + 1) over (
+ | partition by 1 order by 2 rows between current row and current
row)
|from windowData
""".stripMargin)
}
- assert(e.getMessage.contains("Distinct window functions are not
supported"))
+ assert(e.getMessage.contains("Unsupported DISTINCT window function"))
+ }
+ }
+
+ test("window function: distinct rejects unorderable inputs") {
+ val e = intercept[AnalysisException] {
+ sql("SELECT count(DISTINCT map('key', id)) OVER () FROM range(1)")
+ }
+ assert(e.getCondition === "DISTINCT_WINDOW_FUNCTION_UNSUPPORTED")
+ }
+
+ test("window function: distinct aggregates with an unbounded preceding
frame") {
+ val data = Seq(
+ (1, 0, 10, "a", 10),
+ (1, 1, 20, "a", 10),
+ (1, 2, 20, "b", 20),
+ (1, 3, 20, null.asInstanceOf[String], 30),
+ (1, 4, 30, "c", 30),
+ (2, 5, 5, "b", 5),
+ (2, 6, 5, "b", 5),
+ (2, 7, 6, "a", 6)
+ ).toDF("k", "id", "v", "x", "amount")
+
+ withTempView("distinctWindowData") {
+ data.createOrReplaceTempView("distinctWindowData")
+
+ checkAnswer(
+ sql(
+ """
+ |SELECT k, id,
+ | count(DISTINCT x) OVER (PARTITION BY k ORDER BY v) AS
range_count,
+ | count(DISTINCT x) OVER (
+ | PARTITION BY k ORDER BY v, id
+ | ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW) AS
rows_count,
+ | count(DISTINCT x) OVER (
+ | PARTITION BY k ORDER BY v, id
+ | ROWS BETWEEN UNBOUNDED PRECEDING AND 1 PRECEDING) AS
preceding_count,
+ | count(DISTINCT x) OVER (
+ | PARTITION BY k ORDER BY v, id
+ | ROWS BETWEEN UNBOUNDED PRECEDING AND 1 FOLLOWING) AS
following_count,
+ | count(DISTINCT x) OVER (PARTITION BY k) AS partition_count,
+ | sum(DISTINCT amount) OVER (PARTITION BY k ORDER BY v) AS
range_sum,
+ | avg(DISTINCT amount) OVER (PARTITION BY k ORDER BY v) AS
range_avg,
+ | sort_array(collect_list(DISTINCT amount) OVER (
+ | PARTITION BY k ORDER BY v)) AS range_values
+ |FROM distinctWindowData
+ """.stripMargin),
+ Seq(
+ Row(1, 0, 1L, 1L, 0L, 1L, 3L, 10L, 10.0, Seq(10)),
+ Row(1, 1, 2L, 1L, 1L, 2L, 3L, 60L, 20.0, Seq(10, 20, 30)),
+ Row(1, 2, 2L, 2L, 1L, 2L, 3L, 60L, 20.0, Seq(10, 20, 30)),
+ Row(1, 3, 2L, 2L, 2L, 3L, 3L, 60L, 20.0, Seq(10, 20, 30)),
+ Row(1, 4, 3L, 3L, 2L, 3L, 3L, 60L, 20.0, Seq(10, 20, 30)),
+ Row(2, 5, 1L, 1L, 0L, 1L, 2L, 5L, 5.0, Seq(5)),
+ Row(2, 6, 1L, 1L, 1L, 2L, 2L, 5L, 5.0, Seq(5)),
+ Row(2, 7, 2L, 2L, 1L, 2L, 2L, 11L, 5.5, Seq(5, 6))
+ ))
+ }
+ }
+
+ test("window function: count distinct with a range offset, filter, and
multiple columns") {
+ val data = Seq(
+ (0, 10, "a", 1, true),
+ (1, 20, "a", 1, true),
+ (2, 20, "b", 1, false),
+ (3, 20, "b", 2, true),
+ (4, 30, "c", 3, true)
+ ).toDF("id", "v", "x", "y", "selected")
+
+ withTempView("distinctWindowData") {
+ data.createOrReplaceTempView("distinctWindowData")
+
+ checkAnswer(
+ sql(
+ """
+ |SELECT id,
+ | count(DISTINCT x) OVER (
+ | ORDER BY v RANGE BETWEEN UNBOUNDED PRECEDING AND 5 PRECEDING)
AS preceding_count,
+ | count(DISTINCT x) FILTER (WHERE selected) OVER (
+ | ORDER BY v RANGE BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW)
AS filtered_count,
+ | count(DISTINCT x, y) OVER (
+ | ORDER BY v RANGE BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW)
AS tuple_count
+ |FROM distinctWindowData
+ """.stripMargin),
+ Seq(
+ Row(0, 0L, 1L, 1L),
+ Row(1, 1L, 2L, 3L),
+ Row(2, 1L, 2L, 3L),
+ Row(3, 1L, 2L, 3L),
+ Row(4, 2L, 3L, 4L)
+ ))
+ }
+ }
+
+ test("window function: count distinct falls back from hash and sorter
spills") {
+ withSQLConf(
+ WINDOW_EXEC_BUFFER_IN_MEMORY_THRESHOLD.key -> "1000",
+ WINDOW_EXEC_BUFFER_SPILL_THRESHOLD.key -> "5",
+ WINDOW_EXEC_DISTINCT_HASH_FALLBACK_THRESHOLD.key -> "2") {
+ val result = sql(
+ """
+ |SELECT max(distinct_count)
+ |FROM (
+ | SELECT count(DISTINCT id % 3) OVER (
+ | ORDER BY id ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW)
AS distinct_count
+ | FROM range(100)
+ |)
+ """.stripMargin)
+ assertSpilled(sparkContext, "count distinct window hash fallback") {
+ checkAnswer(result, Row(3L))
+ }
+ }
+ }
+
+ test("window function: unbounded distinct frame skips the event sorter") {
+ withSQLConf(
+ ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+ WINDOW_EXEC_BUFFER_IN_MEMORY_THRESHOLD.key -> "1000",
+ WINDOW_EXEC_BUFFER_SPILL_THRESHOLD.key -> "5",
+ WINDOW_EXEC_DISTINCT_HASH_FALLBACK_THRESHOLD.key ->
Int.MaxValue.toString) {
+ val result = sql(
+ """
+ |SELECT max(distinct_count)
+ |FROM (
+ | SELECT count(DISTINCT id) OVER () AS distinct_count
+ | FROM range(100)
+ |)
+ """.stripMargin)
+ assertNotSpilled(sparkContext, "unbounded distinct window without an
event sorter") {
+ checkAnswer(result, Row(100L))
+ }
+ val window = result.queryExecution.executedPlan.collectFirst {
+ case window: WindowExec => window
+ }.get
+ assert(window.metrics("spillSize").value == 0)
+ }
+ }
+
+ test("window function: unbounded distinct frame falls back by size") {
+ withSQLConf(
+ WINDOW_EXEC_BUFFER_IN_MEMORY_THRESHOLD.key -> "1000",
+ WINDOW_EXEC_BUFFER_SPILL_THRESHOLD.key -> Int.MaxValue.toString,
+ WINDOW_EXEC_BUFFER_SIZE_SPILL_THRESHOLD.key -> "1",
+ WINDOW_EXEC_DISTINCT_HASH_FALLBACK_THRESHOLD.key ->
Int.MaxValue.toString) {
+ val result = sql(
+ """
+ |SELECT id,
+ | count(DISTINCT id % 3) OVER () AS distinct_count,
+ | sum(DISTINCT id % 3) OVER () AS distinct_sum,
+ | sort_array(collect_list(DISTINCT id % 3) OVER ()) AS
distinct_values
+ |FROM range(20)
+ """.stripMargin)
+ // An entire-partition frame has no event sorter, and the input buffer
stays in memory, so
Review Comment:
**Non-blocking:**
This test does not isolate the distinct-key sorter spill: the 1-byte
`WINDOW_EXEC_BUFFER_SIZE_SPILL_THRESHOLD` also applies to the window input
buffer, so that buffer can spill and make `assertSpilled` pass. Keep the
input-buffer size threshold high and lower only the distinct-key fallback
threshold, or assert the specific sorter path directly.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]