Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
Original file line number Diff line number Diff line change
Expand Up @@ -96,17 +96,16 @@ object DeequFunctions {

/** Data type detection with state */
def stateful_datatype(column: Column): Column = {
val statefulDataType = new StatefulDataType()
statefulDataType(column)
functions.udaf(new StatefulDataType(), Encoders.STRING)(column)
}

def stateful_kll(
column: Column,
sketchSize: Int,
shrinkingFactor: Double): Column = {
val statefulKLL = new StatefulKLLSketch(sketchSize, shrinkingFactor)
statefulKLL(column)
functions.udaf(
new StatefulKLLSketch(sketchSize, shrinkingFactor),
Encoders.DOUBLE)(column)
}
}


Original file line number Diff line number Diff line change
Expand Up @@ -17,67 +17,70 @@
package org.apache.spark.sql

import com.amazon.deequ.analyzers.DataTypeHistogram
import org.apache.spark.sql.expressions.{MutableAggregationBuffer, UserDefinedAggregateFunction}
import org.apache.spark.sql.types._
import org.apache.spark.sql.expressions.Aggregator

import scala.util.matching.Regex

private[sql] final case class DataTypeAggregationBuffer(
var numNull: Long,
var numFractional: Long,
var numIntegral: Long,
var numBoolean: Long,
var numString: Long)

private[sql] class StatefulDataType extends UserDefinedAggregateFunction {

val SIZE_IN_BYTES = 40

val NULL_POS = 0
val FRACTIONAL_POS = 1
val INTEGRAL_POS = 2
val BOOLEAN_POS = 3
val STRING_POS = 4
private[sql] class StatefulDataType
extends Aggregator[String, DataTypeAggregationBuffer, Array[Byte]] {

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

override def inputSchema: StructType = StructType(StructField("value", StringType) :: Nil)

override def bufferSchema: StructType = StructType(StructField("null", LongType) ::
StructField("fractional", LongType) :: StructField("integral", LongType) ::
StructField("boolean", LongType) :: StructField("string", LongType) :: Nil)

override def dataType: types.DataType = BinaryType

override def deterministic: Boolean = true

override def initialize(buffer: MutableAggregationBuffer): Unit = {
buffer(NULL_POS) = 0L
buffer(FRACTIONAL_POS) = 0L
buffer(INTEGRAL_POS) = 0L
buffer(BOOLEAN_POS) = 0L
buffer(STRING_POS) = 0L
override def zero: DataTypeAggregationBuffer = {
DataTypeAggregationBuffer(0L, 0L, 0L, 0L, 0L)
}

override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
if (input.isNullAt(0)) {
buffer(NULL_POS) = buffer.getLong(NULL_POS) + 1L
override def reduce(
buffer: DataTypeAggregationBuffer,
input: String)
: DataTypeAggregationBuffer = {

if (input == null) {
buffer.numNull += 1L
} else {
input.getString(0) match {
case FRACTIONAL(_*) => buffer(FRACTIONAL_POS) = buffer.getLong(FRACTIONAL_POS) + 1L
case INTEGRAL(_*) => buffer(INTEGRAL_POS) = buffer.getLong(INTEGRAL_POS) + 1L
case BOOLEAN(_*) => buffer(BOOLEAN_POS) = buffer.getLong(BOOLEAN_POS) + 1L
case _ => buffer(STRING_POS) = buffer.getLong(STRING_POS) + 1L
input match {
case FRACTIONAL(_*) => buffer.numFractional += 1L
case INTEGRAL(_*) => buffer.numIntegral += 1L
case BOOLEAN(_*) => buffer.numBoolean += 1L
case _ => buffer.numString += 1L
}
}
buffer
}

override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
buffer1(NULL_POS) = buffer1.getLong(NULL_POS) + buffer2.getLong(NULL_POS)
buffer1(FRACTIONAL_POS) = buffer1.getLong(FRACTIONAL_POS) + buffer2.getLong(FRACTIONAL_POS)
buffer1(INTEGRAL_POS) = buffer1.getLong(INTEGRAL_POS) + buffer2.getLong(INTEGRAL_POS)
buffer1(BOOLEAN_POS) = buffer1.getLong(BOOLEAN_POS) + buffer2.getLong(BOOLEAN_POS)
buffer1(STRING_POS) = buffer1.getLong(STRING_POS) + buffer2.getLong(STRING_POS)
override def merge(
buffer1: DataTypeAggregationBuffer,
buffer2: DataTypeAggregationBuffer)
: DataTypeAggregationBuffer = {

buffer1.numNull += buffer2.numNull
buffer1.numFractional += buffer2.numFractional
buffer1.numIntegral += buffer2.numIntegral
buffer1.numBoolean += buffer2.numBoolean
buffer1.numString += buffer2.numString
buffer1
}

override def evaluate(buffer: Row): Any = {
DataTypeHistogram.toBytes(buffer.getLong(NULL_POS), buffer.getLong(FRACTIONAL_POS),
buffer.getLong(INTEGRAL_POS), buffer.getLong(BOOLEAN_POS), buffer.getLong(STRING_POS))
override def finish(buffer: DataTypeAggregationBuffer): Array[Byte] = {
DataTypeHistogram.toBytes(
buffer.numNull,
buffer.numFractional,
buffer.numIntegral,
buffer.numBoolean,
buffer.numString)
}

override def bufferEncoder: Encoder[DataTypeAggregationBuffer] =
Encoders.product[DataTypeAggregationBuffer]

override def outputEncoder: Encoder[Array[Byte]] = Encoders.BINARY
}
Original file line number Diff line number Diff line change
Expand Up @@ -22,81 +22,97 @@ import com.amazon.deequ.analyzers.QuantileNonSample
import com.amazon.deequ.analyzers.catalyst.KLLSketchSerializer
import com.google.common.primitives.Doubles

import org.apache.spark.sql.expressions.{MutableAggregationBuffer, UserDefinedAggregateFunction}
import org.apache.spark.sql.types._
import org.apache.spark.sql.expressions.Aggregator

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

private [sql] class StatefulKLLSketch(
sketchSize: Int,
shrinkingFactor: Double)
extends UserDefinedAggregateFunction{
private var qSketch: QuantileNonSample[Double] = _
private var minimum: Double = _
private var maximum: Double = _

val OBJECT_POS = 0
val MIN_POS = 1
val MAX_POS = 2
def this(qSketch: QuantileNonSample[Double], minimum: Double, maximum: Double) = {
this()
this.qSketch = qSketch
this.minimum = minimum
this.maximum = maximum
}

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

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

override def dataType: DataType = BinaryType
def setSerializedSketch(bytes: Array[Byte]): Unit = {
qSketch = KLLSketchSerializer.serializer.deserialize(bytes)
}

override def deterministic: Boolean = true
def getMinimum: Double = minimum

override def initialize(buffer: MutableAggregationBuffer): Unit = {
val qsketch = new QuantileNonSample[Double](sketchSize, shrinkingFactor)
buffer(OBJECT_POS) = serialize(qsketch)
buffer(MIN_POS) = Int.MaxValue.toDouble
buffer(MAX_POS) = Int.MinValue.toDouble
def setMinimum(value: Double): Unit = {
minimum = value
}

override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
if (input.isNullAt(OBJECT_POS)) {
return
}
def getMaximum: Double = maximum

val tmp = input.getDouble(OBJECT_POS)
val kll = deserialize(buffer.getAs[Array[Byte]](OBJECT_POS))
kll.update(tmp)
buffer(OBJECT_POS) = serialize(kll)
buffer(MIN_POS) = Math.min(buffer.getDouble(MIN_POS), tmp)
buffer(MAX_POS) = Math.max(buffer.getDouble(MAX_POS), tmp)
def setMaximum(value: Double): Unit = {
maximum = value
}
}

override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
if (buffer2.isNullAt(OBJECT_POS)) {
return
private[sql] class StatefulKLLSketch(
sketchSize: Int,
shrinkingFactor: Double)
extends Aggregator[java.lang.Double, KLLAggregationBuffer, Array[Byte]] {

override def zero: KLLAggregationBuffer = {
new KLLAggregationBuffer(
new QuantileNonSample[Double](sketchSize, shrinkingFactor),
Int.MaxValue.toDouble,
Int.MinValue.toDouble)
}

override def reduce(
buffer: KLLAggregationBuffer,
input: java.lang.Double)
: KLLAggregationBuffer = {

if (input != null) {
val value = input.doubleValue()
buffer.sketch.update(value)
buffer.setMinimum(Math.min(buffer.getMinimum, value))
buffer.setMaximum(Math.max(buffer.getMaximum, value))
}
buffer
}

val kll_this = deserialize(buffer1.getAs[Array[Byte]](OBJECT_POS))
val kll_other = deserialize(buffer2.getAs[Array[Byte]](OBJECT_POS))
val kll_ret = kll_this.merge(kll_other)
buffer1(OBJECT_POS) = serialize(kll_ret)
buffer1(MIN_POS) = Math.min(buffer1.getDouble(MIN_POS), buffer2.getDouble(MIN_POS))
buffer1(MAX_POS) = Math.max(buffer1.getDouble(MAX_POS), buffer2.getDouble(MAX_POS))
override def merge(
buffer1: KLLAggregationBuffer,
buffer2: KLLAggregationBuffer)
: KLLAggregationBuffer = {

buffer1.sketch.merge(buffer2.sketch)
buffer1.setMinimum(Math.min(buffer1.getMinimum, buffer2.getMinimum))
buffer1.setMaximum(Math.max(buffer1.getMaximum, buffer2.getMaximum))
buffer1
}

override def evaluate(buffer: Row): Any = {
toBytes(buffer.getDouble(MIN_POS),
buffer.getDouble(MAX_POS),
buffer.getAs[Array[Byte]](OBJECT_POS))
override def finish(buffer: KLLAggregationBuffer): Array[Byte] = {
toBytes(buffer.getMinimum, buffer.getMaximum, serialize(buffer.sketch))
}

def toBytes(min: Double, max: Double, obj: Array[Byte]): Array[Byte] = {
override def bufferEncoder: Encoder[KLLAggregationBuffer] =
Encoders.bean(classOf[KLLAggregationBuffer])

override def outputEncoder: Encoder[Array[Byte]] = Encoders.BINARY

private def toBytes(min: Double, max: Double, obj: Array[Byte]): Array[Byte] = {
val buffer2 = ByteBuffer.wrap(new Array(Doubles.BYTES + Doubles.BYTES + obj.length))
buffer2.putDouble(min)
buffer2.putDouble(max)
buffer2.put(obj)
buffer2.array()
}

def serialize(obj: QuantileNonSample[Double]): Array[Byte] = {
private def serialize(obj: QuantileNonSample[Double]): Array[Byte] = {
KLLSketchSerializer.serializer.serialize(obj)
}

def deserialize(bytes: Array[Byte]): QuantileNonSample[Double] = {
KLLSketchSerializer.serializer.deserialize(bytes)
}
}

Loading
Loading