Skip to content

Commit 8de27d2

Browse files
authored
Replace deprecated stateful UDAFs with typed aggregators (#760)
1 parent 38ba5bb commit 8de27d2

5 files changed

Lines changed: 559 additions & 98 deletions

File tree

src/main/scala/com/amazon/deequ/analyzers/catalyst/DeequFunctions.scala

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -96,17 +96,16 @@ object DeequFunctions {
9696

9797
/** Data type detection with state */
9898
def stateful_datatype(column: Column): Column = {
99-
val statefulDataType = new StatefulDataType()
100-
statefulDataType(column)
99+
functions.udaf(new StatefulDataType(), Encoders.STRING)(column)
101100
}
102101

103102
def stateful_kll(
104103
column: Column,
105104
sketchSize: Int,
106105
shrinkingFactor: Double): Column = {
107-
val statefulKLL = new StatefulKLLSketch(sketchSize, shrinkingFactor)
108-
statefulKLL(column)
106+
functions.udaf(
107+
new StatefulKLLSketch(sketchSize, shrinkingFactor),
108+
Encoders.DOUBLE)(column)
109109
}
110110
}
111111

112-

src/main/scala/com/amazon/deequ/analyzers/catalyst/StatefulDataType.scala

Lines changed: 47 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -17,67 +17,70 @@
1717
package org.apache.spark.sql
1818

1919
import com.amazon.deequ.analyzers.DataTypeHistogram
20-
import org.apache.spark.sql.expressions.{MutableAggregationBuffer, UserDefinedAggregateFunction}
21-
import org.apache.spark.sql.types._
20+
import org.apache.spark.sql.expressions.Aggregator
2221

2322
import scala.util.matching.Regex
2423

24+
private[sql] final case class DataTypeAggregationBuffer(
25+
var numNull: Long,
26+
var numFractional: Long,
27+
var numIntegral: Long,
28+
var numBoolean: Long,
29+
var numString: Long)
2530

26-
private[sql] class StatefulDataType extends UserDefinedAggregateFunction {
27-
28-
val SIZE_IN_BYTES = 40
29-
30-
val NULL_POS = 0
31-
val FRACTIONAL_POS = 1
32-
val INTEGRAL_POS = 2
33-
val BOOLEAN_POS = 3
34-
val STRING_POS = 4
31+
private[sql] class StatefulDataType
32+
extends Aggregator[String, DataTypeAggregationBuffer, Array[Byte]] {
3533

3634
val FRACTIONAL: Regex = """^(-|\+)? ?\d+((\.\d+)|((?:\.\d+)?[Ee][-+]?\d+))$""".r
3735
val INTEGRAL: Regex = """^(-|\+)? ?\d+$""".r
3836
val BOOLEAN: Regex = """^(true|false)$""".r
3937

40-
override def inputSchema: StructType = StructType(StructField("value", StringType) :: Nil)
41-
42-
override def bufferSchema: StructType = StructType(StructField("null", LongType) ::
43-
StructField("fractional", LongType) :: StructField("integral", LongType) ::
44-
StructField("boolean", LongType) :: StructField("string", LongType) :: Nil)
45-
46-
override def dataType: types.DataType = BinaryType
47-
48-
override def deterministic: Boolean = true
49-
50-
override def initialize(buffer: MutableAggregationBuffer): Unit = {
51-
buffer(NULL_POS) = 0L
52-
buffer(FRACTIONAL_POS) = 0L
53-
buffer(INTEGRAL_POS) = 0L
54-
buffer(BOOLEAN_POS) = 0L
55-
buffer(STRING_POS) = 0L
38+
override def zero: DataTypeAggregationBuffer = {
39+
DataTypeAggregationBuffer(0L, 0L, 0L, 0L, 0L)
5640
}
5741

58-
override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
59-
if (input.isNullAt(0)) {
60-
buffer(NULL_POS) = buffer.getLong(NULL_POS) + 1L
42+
override def reduce(
43+
buffer: DataTypeAggregationBuffer,
44+
input: String)
45+
: DataTypeAggregationBuffer = {
46+
47+
if (input == null) {
48+
buffer.numNull += 1L
6149
} else {
62-
input.getString(0) match {
63-
case FRACTIONAL(_*) => buffer(FRACTIONAL_POS) = buffer.getLong(FRACTIONAL_POS) + 1L
64-
case INTEGRAL(_*) => buffer(INTEGRAL_POS) = buffer.getLong(INTEGRAL_POS) + 1L
65-
case BOOLEAN(_*) => buffer(BOOLEAN_POS) = buffer.getLong(BOOLEAN_POS) + 1L
66-
case _ => buffer(STRING_POS) = buffer.getLong(STRING_POS) + 1L
50+
input match {
51+
case FRACTIONAL(_*) => buffer.numFractional += 1L
52+
case INTEGRAL(_*) => buffer.numIntegral += 1L
53+
case BOOLEAN(_*) => buffer.numBoolean += 1L
54+
case _ => buffer.numString += 1L
6755
}
6856
}
57+
buffer
6958
}
7059

71-
override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
72-
buffer1(NULL_POS) = buffer1.getLong(NULL_POS) + buffer2.getLong(NULL_POS)
73-
buffer1(FRACTIONAL_POS) = buffer1.getLong(FRACTIONAL_POS) + buffer2.getLong(FRACTIONAL_POS)
74-
buffer1(INTEGRAL_POS) = buffer1.getLong(INTEGRAL_POS) + buffer2.getLong(INTEGRAL_POS)
75-
buffer1(BOOLEAN_POS) = buffer1.getLong(BOOLEAN_POS) + buffer2.getLong(BOOLEAN_POS)
76-
buffer1(STRING_POS) = buffer1.getLong(STRING_POS) + buffer2.getLong(STRING_POS)
60+
override def merge(
61+
buffer1: DataTypeAggregationBuffer,
62+
buffer2: DataTypeAggregationBuffer)
63+
: DataTypeAggregationBuffer = {
64+
65+
buffer1.numNull += buffer2.numNull
66+
buffer1.numFractional += buffer2.numFractional
67+
buffer1.numIntegral += buffer2.numIntegral
68+
buffer1.numBoolean += buffer2.numBoolean
69+
buffer1.numString += buffer2.numString
70+
buffer1
7771
}
7872

79-
override def evaluate(buffer: Row): Any = {
80-
DataTypeHistogram.toBytes(buffer.getLong(NULL_POS), buffer.getLong(FRACTIONAL_POS),
81-
buffer.getLong(INTEGRAL_POS), buffer.getLong(BOOLEAN_POS), buffer.getLong(STRING_POS))
73+
override def finish(buffer: DataTypeAggregationBuffer): Array[Byte] = {
74+
DataTypeHistogram.toBytes(
75+
buffer.numNull,
76+
buffer.numFractional,
77+
buffer.numIntegral,
78+
buffer.numBoolean,
79+
buffer.numString)
8280
}
81+
82+
override def bufferEncoder: Encoder[DataTypeAggregationBuffer] =
83+
Encoders.product[DataTypeAggregationBuffer]
84+
85+
override def outputEncoder: Encoder[Array[Byte]] = Encoders.BINARY
8386
}

src/main/scala/com/amazon/deequ/analyzers/catalyst/StatefulKLLSketch.scala

Lines changed: 65 additions & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -22,81 +22,97 @@ import com.amazon.deequ.analyzers.QuantileNonSample
2222
import com.amazon.deequ.analyzers.catalyst.KLLSketchSerializer
2323
import com.google.common.primitives.Doubles
2424

25-
import org.apache.spark.sql.expressions.{MutableAggregationBuffer, UserDefinedAggregateFunction}
26-
import org.apache.spark.sql.types._
25+
import org.apache.spark.sql.expressions.Aggregator
2726

27+
private[sql] final class KLLAggregationBuffer() extends Serializable {
2828

29-
private [sql] class StatefulKLLSketch(
30-
sketchSize: Int,
31-
shrinkingFactor: Double)
32-
extends UserDefinedAggregateFunction{
29+
private var qSketch: QuantileNonSample[Double] = _
30+
private var minimum: Double = _
31+
private var maximum: Double = _
3332

34-
val OBJECT_POS = 0
35-
val MIN_POS = 1
36-
val MAX_POS = 2
33+
def this(qSketch: QuantileNonSample[Double], minimum: Double, maximum: Double) = {
34+
this()
35+
this.qSketch = qSketch
36+
this.minimum = minimum
37+
this.maximum = maximum
38+
}
3739

38-
override def inputSchema: StructType = StructType(StructField("value", DoubleType) :: Nil)
40+
private[sql] def sketch: QuantileNonSample[Double] = qSketch
3941

40-
override def bufferSchema: StructType = StructType(StructField("data", BinaryType) ::
41-
StructField("minimum", DoubleType) :: StructField("maximum", DoubleType) :: Nil)
42+
def getSerializedSketch: Array[Byte] = KLLSketchSerializer.serializer.serialize(qSketch)
4243

43-
override def dataType: DataType = BinaryType
44+
def setSerializedSketch(bytes: Array[Byte]): Unit = {
45+
qSketch = KLLSketchSerializer.serializer.deserialize(bytes)
46+
}
4447

45-
override def deterministic: Boolean = true
48+
def getMinimum: Double = minimum
4649

47-
override def initialize(buffer: MutableAggregationBuffer): Unit = {
48-
val qsketch = new QuantileNonSample[Double](sketchSize, shrinkingFactor)
49-
buffer(OBJECT_POS) = serialize(qsketch)
50-
buffer(MIN_POS) = Int.MaxValue.toDouble
51-
buffer(MAX_POS) = Int.MinValue.toDouble
50+
def setMinimum(value: Double): Unit = {
51+
minimum = value
5252
}
5353

54-
override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
55-
if (input.isNullAt(OBJECT_POS)) {
56-
return
57-
}
54+
def getMaximum: Double = maximum
5855

59-
val tmp = input.getDouble(OBJECT_POS)
60-
val kll = deserialize(buffer.getAs[Array[Byte]](OBJECT_POS))
61-
kll.update(tmp)
62-
buffer(OBJECT_POS) = serialize(kll)
63-
buffer(MIN_POS) = Math.min(buffer.getDouble(MIN_POS), tmp)
64-
buffer(MAX_POS) = Math.max(buffer.getDouble(MAX_POS), tmp)
56+
def setMaximum(value: Double): Unit = {
57+
maximum = value
6558
}
59+
}
6660

67-
override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
68-
if (buffer2.isNullAt(OBJECT_POS)) {
69-
return
61+
private[sql] class StatefulKLLSketch(
62+
sketchSize: Int,
63+
shrinkingFactor: Double)
64+
extends Aggregator[java.lang.Double, KLLAggregationBuffer, Array[Byte]] {
65+
66+
override def zero: KLLAggregationBuffer = {
67+
new KLLAggregationBuffer(
68+
new QuantileNonSample[Double](sketchSize, shrinkingFactor),
69+
Int.MaxValue.toDouble,
70+
Int.MinValue.toDouble)
71+
}
72+
73+
override def reduce(
74+
buffer: KLLAggregationBuffer,
75+
input: java.lang.Double)
76+
: KLLAggregationBuffer = {
77+
78+
if (input != null) {
79+
val value = input.doubleValue()
80+
buffer.sketch.update(value)
81+
buffer.setMinimum(Math.min(buffer.getMinimum, value))
82+
buffer.setMaximum(Math.max(buffer.getMaximum, value))
7083
}
84+
buffer
85+
}
7186

72-
val kll_this = deserialize(buffer1.getAs[Array[Byte]](OBJECT_POS))
73-
val kll_other = deserialize(buffer2.getAs[Array[Byte]](OBJECT_POS))
74-
val kll_ret = kll_this.merge(kll_other)
75-
buffer1(OBJECT_POS) = serialize(kll_ret)
76-
buffer1(MIN_POS) = Math.min(buffer1.getDouble(MIN_POS), buffer2.getDouble(MIN_POS))
77-
buffer1(MAX_POS) = Math.max(buffer1.getDouble(MAX_POS), buffer2.getDouble(MAX_POS))
87+
override def merge(
88+
buffer1: KLLAggregationBuffer,
89+
buffer2: KLLAggregationBuffer)
90+
: KLLAggregationBuffer = {
91+
92+
buffer1.sketch.merge(buffer2.sketch)
93+
buffer1.setMinimum(Math.min(buffer1.getMinimum, buffer2.getMinimum))
94+
buffer1.setMaximum(Math.max(buffer1.getMaximum, buffer2.getMaximum))
95+
buffer1
7896
}
7997

80-
override def evaluate(buffer: Row): Any = {
81-
toBytes(buffer.getDouble(MIN_POS),
82-
buffer.getDouble(MAX_POS),
83-
buffer.getAs[Array[Byte]](OBJECT_POS))
98+
override def finish(buffer: KLLAggregationBuffer): Array[Byte] = {
99+
toBytes(buffer.getMinimum, buffer.getMaximum, serialize(buffer.sketch))
84100
}
85101

86-
def toBytes(min: Double, max: Double, obj: Array[Byte]): Array[Byte] = {
102+
override def bufferEncoder: Encoder[KLLAggregationBuffer] =
103+
Encoders.bean(classOf[KLLAggregationBuffer])
104+
105+
override def outputEncoder: Encoder[Array[Byte]] = Encoders.BINARY
106+
107+
private def toBytes(min: Double, max: Double, obj: Array[Byte]): Array[Byte] = {
87108
val buffer2 = ByteBuffer.wrap(new Array(Doubles.BYTES + Doubles.BYTES + obj.length))
88109
buffer2.putDouble(min)
89110
buffer2.putDouble(max)
90111
buffer2.put(obj)
91112
buffer2.array()
92113
}
93114

94-
def serialize(obj: QuantileNonSample[Double]): Array[Byte] = {
115+
private def serialize(obj: QuantileNonSample[Double]): Array[Byte] = {
95116
KLLSketchSerializer.serializer.serialize(obj)
96117
}
97-
98-
def deserialize(bytes: Array[Byte]): QuantileNonSample[Double] = {
99-
KLLSketchSerializer.serializer.deserialize(bytes)
100-
}
101118
}
102-

0 commit comments

Comments
 (0)