@@ -22,81 +22,97 @@ import com.amazon.deequ.analyzers.QuantileNonSample
2222import com .amazon .deequ .analyzers .catalyst .KLLSketchSerializer
2323import 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