cloud-fan commented on code in PR #57592:
URL: https://github.com/apache/spark/pull/57592#discussion_r3665856406


##########
sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/collect.scala:
##########
@@ -316,6 +316,150 @@ case class CollectSet(
     copy(child = newChild)
 }
 
+/**
+ * Collect the distinct union of the elements of an array-typed input across 
rows.
+ *
+ * Unlike collect_set, whose input is a scalar and whose output is the set of 
those scalars,
+ * collect_union's input is itself an array and its output is the set of the 
array's
+ * *elements* unioned across all rows. It is equivalent to
+ * `array_distinct(flatten(collect_list(arr)))`, but the aggregation buffer 
holds only the
+ * distinct elements (a set), so its size is bounded by the element universe 
rather than by
+ * the number of input rows.
+ */
+@ExpressionDescription(
+  usage =
+    "_FUNC_(expr) - Collects and returns the distinct union of the elements of 
array `expr`.",
+  examples = """
+    Examples:
+      > SELECT _FUNC_(col) FROM VALUES (array(1, 2)), (array(2, 3)), 
(array(1)) AS tab(col);
+       [1,2,3]
+  """,
+  note = """
+    The function is non-deterministic because the order of collected results 
depends
+    on the order of the rows which may be non-deterministic after a shuffle.
+  """,
+  group = "agg_funcs",
+  since = "4.3.0")
+case class CollectUnion(
+    child: Expression,
+    mutableAggBufferOffset: Int = 0,
+    inputAggBufferOffset: Int = 0)
+  extends Collect[mutable.HashSet[Any]] with QueryErrorsBase with 
UnaryLike[Expression] {
+
+  def this(child: Expression) = this(child, 0, 0)
+
+  // The element type of the input array. Guarded by checkInputDataTypes to be 
an ArrayType.
+  private lazy val elementType: DataType = child.dataType match {
+    case ArrayType(et, _) => et
+    case other => other
+  }
+
+  // Result is array<elementType>; NULL elements are dropped, so it never 
contains nulls.
+  override def dataType: DataType = ArrayType(elementType, containsNull = 
false)
+
+  // The buffer stores distinct elements. Mirror CollectSet's keying so 
equality is correct
+  // for float/double (bit pattern) and binary (byte array) element types.
+  override lazy val bufferElementType: DataType = elementType match {
+    case BinaryType => ArrayType(ByteType)
+    case DoubleType => LongType
+    case FloatType => IntegerType
+    case other => other
+  }
+
+  @transient private lazy val complexNormalizer: Any => Any = {
+    val ref = BoundReference(0, elementType, nullable = true)
+    val proj = UnsafeProjection.create(NormalizeFloatingNumbers.normalize(ref))
+    (value: Any) => InternalRow.copyValue(proj(InternalRow(value)).get(0, 
elementType))
+  }
+
+  override def convertToBufferElement(value: Any): Any = elementType match {
+    // See CollectSet.convertToBufferElement for why binary/float/double are 
keyed specially.
+    case BinaryType => 
UnsafeArrayData.fromPrimitiveArray(value.asInstanceOf[Array[Byte]])
+    case DoubleType =>
+      java.lang.Double.doubleToLongBits(
+        NormalizeFloatingNumbers.DOUBLE_NORMALIZER(value).asInstanceOf[Double])
+    case FloatType =>
+      java.lang.Float.floatToIntBits(
+        NormalizeFloatingNumbers.FLOAT_NORMALIZER(value).asInstanceOf[Float])
+    case dt if NormalizeFloatingNumbers.needNormalize(dt) => 
complexNormalizer(value)
+    case _ => InternalRow.copyValue(value)
+  }
+
+  // Iterate the input array and add each non-null element to the set. NULL 
input arrays and
+  // NULL elements are skipped (following collect_set's ignore-null semantics).
+  override def update(
+      buffer: mutable.HashSet[Any],
+      input: InternalRow): mutable.HashSet[Any] = {
+    val arr = child.eval(input)
+    if (arr != null) {
+      arr.asInstanceOf[ArrayData].foreach(elementType, (_, element: Any) =>
+        if (element != null) {

Review Comment:
   The documented equivalent expression preserves one null element, so this 
aggregate should do the same. For `array(1, NULL)`, 
`array_distinct(flatten(collect_list(...)))` returns `[1, NULL]`, but this 
branch drops the null and the result metadata declares `containsNull = false`. 
Please retain one null in the buffer, update the result nullability, and change 
the null-element test accordingly.



##########
sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/collect.scala:
##########
@@ -316,6 +316,150 @@ case class CollectSet(
     copy(child = newChild)
 }
 
+/**
+ * Collect the distinct union of the elements of an array-typed input across 
rows.
+ *
+ * Unlike collect_set, whose input is a scalar and whose output is the set of 
those scalars,
+ * collect_union's input is itself an array and its output is the set of the 
array's
+ * *elements* unioned across all rows. It is equivalent to
+ * `array_distinct(flatten(collect_list(arr)))`, but the aggregation buffer 
holds only the
+ * distinct elements (a set), so its size is bounded by the element universe 
rather than by
+ * the number of input rows.
+ */
+@ExpressionDescription(
+  usage =
+    "_FUNC_(expr) - Collects and returns the distinct union of the elements of 
array `expr`.",
+  examples = """
+    Examples:
+      > SELECT _FUNC_(col) FROM VALUES (array(1, 2)), (array(2, 3)), 
(array(1)) AS tab(col);
+       [1,2,3]
+  """,
+  note = """
+    The function is non-deterministic because the order of collected results 
depends
+    on the order of the rows which may be non-deterministic after a shuffle.
+  """,
+  group = "agg_funcs",
+  since = "4.3.0")
+case class CollectUnion(
+    child: Expression,
+    mutableAggBufferOffset: Int = 0,
+    inputAggBufferOffset: Int = 0)
+  extends Collect[mutable.HashSet[Any]] with QueryErrorsBase with 
UnaryLike[Expression] {
+
+  def this(child: Expression) = this(child, 0, 0)
+
+  // The element type of the input array. Guarded by checkInputDataTypes to be 
an ArrayType.

Review Comment:
   The second sentence currently makes the element type sound like the value 
guarded to be an `ArrayType`.
   
   ```suggestion
     // The input is guarded by checkInputDataTypes to be an ArrayType; this is 
its element type.
   ```



-- 
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]

Reply via email to